diff --git a/python/sglang/srt/configs/load_config.py b/python/sglang/srt/configs/load_config.py index c2fd4332c..ddf8d2967 100644 --- a/python/sglang/srt/configs/load_config.py +++ b/python/sglang/srt/configs/load_config.py @@ -29,6 +29,7 @@ class LoadFormat(str, enum.Enum): REMOTE_INSTANCE = "remote_instance" RDMA = "rdma" LOCAL_CACHED = "local_cached" + FASTSAFETENSORS = "fastsafetensors" PRIVATE = "private" diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 7590b3c10..3f624aca2 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -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( diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py index b474780e1..a0eba7ef1 100644 --- a/python/sglang/srt/model_loader/weight_utils.py +++ b/python/sglang/srt/model_loader/weight_utils.py @@ -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, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8817e1690..3064d6e9c 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -84,6 +84,7 @@ LOAD_FORMAT_CHOICES = [ "flash_rl", "remote", "remote_instance", + "fastsafetensors", "private", ]