[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:
Teng Ma
2025-12-23 22:33:56 +08:00
committed by GitHub
parent 758b9067a0
commit d7301c89ba
4 changed files with 82 additions and 2 deletions

View File

@@ -29,6 +29,7 @@ class LoadFormat(str, enum.Enum):
REMOTE_INSTANCE = "remote_instance"
RDMA = "rdma"
LOCAL_CACHED = "local_cached"
FASTSAFETENSORS = "fastsafetensors"
PRIVATE = "private"

View File

@@ -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(

View File

@@ -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,

View File

@@ -84,6 +84,7 @@ LOAD_FORMAT_CHOICES = [
"flash_rl",
"remote",
"remote_instance",
"fastsafetensors",
"private",
]