[Feature] support fastsafetensors (#15091)
Signed-off-by: Xuchun Shang <xuchun.shang@gmail.com> Co-authored-by: Xuchun Shang <xuchun.shang@gmail.com>
This commit is contained in:
@@ -29,6 +29,7 @@ class LoadFormat(str, enum.Enum):
|
||||
REMOTE_INSTANCE = "remote_instance"
|
||||
RDMA = "rdma"
|
||||
LOCAL_CACHED = "local_cached"
|
||||
FASTSAFETENSORS = "fastsafetensors"
|
||||
PRIVATE = "private"
|
||||
|
||||
|
||||
|
||||
@@ -89,6 +89,7 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.model_loader.weight_utils import (
|
||||
download_safetensors_index_file_from_hf,
|
||||
download_weights_from_hf,
|
||||
fastsafetensors_weights_iterator,
|
||||
filter_duplicate_safetensors_files,
|
||||
filter_files_not_needed_for_inference,
|
||||
get_gguf_extra_tensor_names,
|
||||
@@ -386,7 +387,10 @@ class DefaultModelLoader(BaseModelLoader):
|
||||
# Some quantized models use .pt files for storing the weights.
|
||||
if load_format == LoadFormat.AUTO:
|
||||
allow_patterns = ["*.safetensors", "*.bin"]
|
||||
elif load_format == LoadFormat.SAFETENSORS:
|
||||
elif (
|
||||
load_format == LoadFormat.SAFETENSORS
|
||||
or load_format == LoadFormat.FASTSAFETENSORS
|
||||
):
|
||||
use_safetensors = True
|
||||
allow_patterns = ["*.safetensors"]
|
||||
elif load_format == LoadFormat.MISTRAL:
|
||||
@@ -474,7 +478,11 @@ class DefaultModelLoader(BaseModelLoader):
|
||||
get_global_server_args().weight_loader_disable_mmap
|
||||
)
|
||||
|
||||
if extra_config.get("enable_multithread_load"):
|
||||
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
|
||||
weights_iterator = fastsafetensors_weights_iterator(
|
||||
hf_weights_files,
|
||||
)
|
||||
elif extra_config.get("enable_multithread_load"):
|
||||
weights_iterator = multi_thread_safetensors_weights_iterator(
|
||||
hf_weights_files,
|
||||
max_workers=extra_config.get(
|
||||
|
||||
@@ -49,6 +49,21 @@ from sglang.srt.model_loader.weight_validation import (
|
||||
from sglang.srt.utils import find_local_repo_dir, log_info_on_rank0, print_warning_once
|
||||
from sglang.utils import is_in_ci
|
||||
|
||||
try:
|
||||
from fastsafetensors import SafeTensorsFileLoader, SingleGroup
|
||||
except ImportError:
|
||||
|
||||
class PlaceholderModule:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
|
||||
def __getattr__(self, name):
|
||||
raise ImportError(f"Please install {self.name}")
|
||||
|
||||
fastsafetensors = PlaceholderModule("fastsafetensors")
|
||||
SafeTensorsFileLoader = None
|
||||
SingleGroup = None
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# use system-level temp directory for file locks, so that multiple users
|
||||
@@ -826,6 +841,61 @@ def safetensors_weights_iterator(
|
||||
yield name, f.get_tensor(name)
|
||||
|
||||
|
||||
def fastsafetensors_weights_iterator(
|
||||
hf_weights_files: List[str],
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""
|
||||
Iterate over the weights in the model safetensor files
|
||||
using fastsafetensor library to accelerate loading via GPU Direct Storage (if available).
|
||||
"""
|
||||
if SafeTensorsFileLoader is None:
|
||||
raise ImportError(
|
||||
"Please install fastsafetensors via `pip install fastsafetensors`"
|
||||
)
|
||||
|
||||
if torch.distributed.is_initialized():
|
||||
pg = torch.distributed.group.WORLD
|
||||
else:
|
||||
pg = SingleGroup()
|
||||
|
||||
try:
|
||||
rank = pg.rank()
|
||||
except Exception:
|
||||
rank = 0
|
||||
|
||||
device = torch.device(f"cuda:{rank}")
|
||||
|
||||
weight_files_sub_lists = [
|
||||
hf_weights_files[i : i + pg.size()]
|
||||
for i in range(0, len(hf_weights_files), pg.size())
|
||||
]
|
||||
|
||||
_BAR_FORMAT = (
|
||||
"{l_bar}{bar}| {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]"
|
||||
)
|
||||
|
||||
for f_list in tqdm(
|
||||
weight_files_sub_lists,
|
||||
desc="Loading safetensors using Fastsafetensor loader",
|
||||
disable=False,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
loader = SafeTensorsFileLoader(pg, device)
|
||||
rank_file_map = {i: [f] for i, f in enumerate(f_list)}
|
||||
loader.add_filenames(rank_file_map)
|
||||
try:
|
||||
fb = loader.copy_files_to_device()
|
||||
try:
|
||||
keys = list(fb.key_to_rank_lidx.keys())
|
||||
for k in keys:
|
||||
t = fb.get_tensor(k)
|
||||
yield k, t
|
||||
finally:
|
||||
pass
|
||||
finally:
|
||||
loader.close()
|
||||
|
||||
|
||||
def multi_thread_safetensors_weights_iterator(
|
||||
hf_weights_files: List[str],
|
||||
is_all_weights_sharded: bool = False,
|
||||
|
||||
@@ -84,6 +84,7 @@ LOAD_FORMAT_CHOICES = [
|
||||
"flash_rl",
|
||||
"remote",
|
||||
"remote_instance",
|
||||
"fastsafetensors",
|
||||
"private",
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user