diff --git a/docs_new/docs/advanced_features/model_loading.mdx b/docs_new/docs/advanced_features/model_loading.mdx
index a5b41066e..eb67e3c4b 100644
--- a/docs_new/docs/advanced_features/model_loading.mdx
+++ b/docs_new/docs/advanced_features/model_loading.mdx
@@ -145,6 +145,12 @@ python -m sglang.launch_server \
Filename pattern for per-rank shards. |
model-rank-{rank}-part-{part}.safetensors |
+
+ fastsafetensors |
+ enable_gds (bool) |
+ Use GPU Direct Storage. Set to false when the host does not provide the NVIDIA GPUDirect Storage kernel driver, such as in a gVisor sandbox. |
+ true |
+
bitsandbytes |
qlora_adapter_name_or_path (str) |
@@ -200,7 +206,7 @@ Top-level arguments that tune how safetensors weights are read, independent of `
--weight-loader-drop-cache-after-load |
- Call posix_fadvise(DONTNEED) on each safetensors shard after loading it, freeing page cache. |
+ Call posix_fadvise(DONTNEED) after successfully loading each shard, freeing page cache. Supported by the standard safetensors and fastsafetensors loaders. |
off |
diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py
index 8bafdd5cd..db18b4823 100644
--- a/python/sglang/srt/model_loader/loader.py
+++ b/python/sglang/srt/model_loader/loader.py
@@ -399,6 +399,14 @@ class DefaultModelLoader(BaseModelLoader):
super().__init__(load_config)
extra_config = load_config.model_loader_extra_config
allowed_keys = {"enable_multithread_load", "num_threads"}
+ if load_config.load_format == LoadFormat.FASTSAFETENSORS:
+ allowed_keys.add("enable_gds")
+ if "enable_gds" in extra_config and not isinstance(
+ extra_config["enable_gds"], bool
+ ):
+ raise ValueError(
+ "enable_gds in --model-loader-extra-config must be a boolean"
+ )
unexpected_keys = set(extra_config.keys()) - allowed_keys
if unexpected_keys:
@@ -609,8 +617,11 @@ class DefaultModelLoader(BaseModelLoader):
use_multithread = False
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
+ enable_gds = extra_config.get("enable_gds", True)
weights_iterator = fastsafetensors_weights_iterator(
hf_weights_files,
+ enable_gds=enable_gds,
+ drop_cache_after_load=weight_loader_drop_cache_after_load,
)
elif use_multithread:
weights_iterator = buffered_multi_thread_safetensors_weights_iterator(
diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py
index 16babccaf..506d359a4 100644
--- a/python/sglang/srt/model_loader/weight_utils.py
+++ b/python/sglang/srt/model_loader/weight_utils.py
@@ -999,6 +999,8 @@ def safetensors_weights_iterator(
def fastsafetensors_weights_iterator(
hf_weights_files: List[str],
+ enable_gds: bool = True,
+ drop_cache_after_load: bool = False,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""
Iterate over the weights in the model safetensor files
@@ -1036,7 +1038,7 @@ def fastsafetensors_weights_iterator(
disable=False,
bar_format=_BAR_FORMAT,
):
- loader = SafeTensorsFileLoader(pg, device)
+ loader = SafeTensorsFileLoader(pg, device, nogds=not enable_gds)
rank_file_map = {i: [f] for i, f in enumerate(f_list)}
loader.add_filenames(rank_file_map)
try:
@@ -1050,6 +1052,9 @@ def fastsafetensors_weights_iterator(
pass
finally:
loader.close()
+ if drop_cache_after_load:
+ for loaded_file in rank_file_map.get(rank, []):
+ _drop_file_cache_after_load(loaded_file)
def multi_thread_safetensors_weights_iterator(
diff --git a/test/registered/unit/model_loader/test_prefetch_checkpoints.py b/test/registered/unit/model_loader/test_prefetch_checkpoints.py
index dc2d4c295..367c44312 100644
--- a/test/registered/unit/model_loader/test_prefetch_checkpoints.py
+++ b/test/registered/unit/model_loader/test_prefetch_checkpoints.py
@@ -20,6 +20,7 @@ from sglang.srt.model_loader.loader import DefaultModelLoader
from sglang.srt.model_loader.weight_utils import (
_prefetch_all_checkpoints,
buffered_multi_thread_safetensors_weights_iterator,
+ fastsafetensors_weights_iterator,
safetensors_weights_iterator,
)
from sglang.test.ci.ci_register import register_cpu_ci
@@ -221,6 +222,67 @@ class TestPrefetchCheckpoints(CustomTestCase):
paths,
)
+ @patch("torch.distributed.is_initialized", return_value=False)
+ def test_fastsafetensors_drops_only_the_rank_owned_file(self, _):
+ events = []
+ test_case = self
+
+ class FakeGroup:
+ def rank(self):
+ return 0
+
+ def size(self):
+ return 1
+
+ class FakeBuffer:
+ key_to_rank_lidx = {"weight": (0, 0)}
+
+ def get_tensor(self, name):
+ test_case.assertEqual(name, "weight")
+ return torch.tensor([1.0])
+
+ class FakeLoader:
+ def __init__(self, group, device, nogds):
+ test_case.assertIsInstance(group, FakeGroup)
+ test_case.assertEqual(device.type, "cuda")
+ test_case.assertFalse(nogds)
+
+ def add_filenames(self, rank_file_map):
+ test_case.assertEqual(
+ rank_file_map,
+ {0: ["model.safetensors"]},
+ )
+
+ def copy_files_to_device(self):
+ return FakeBuffer()
+
+ def close(self):
+ events.append("close")
+
+ with (
+ patch(
+ "sglang.srt.model_loader.weight_utils.SingleGroup",
+ FakeGroup,
+ ),
+ patch(
+ "sglang.srt.model_loader.weight_utils.SafeTensorsFileLoader",
+ FakeLoader,
+ ),
+ patch(
+ "sglang.srt.model_loader.weight_utils._drop_file_cache_after_load",
+ side_effect=lambda path: events.append(f"drop:{path}"),
+ ),
+ ):
+ loaded = list(
+ fastsafetensors_weights_iterator(
+ ["model.safetensors"],
+ drop_cache_after_load=True,
+ )
+ )
+
+ torch.testing.assert_close(loaded[0][1], torch.tensor([1.0]))
+ self.assertEqual(events, ["close", "drop:model.safetensors"])
+
class TestPrefetchDispatch(CustomTestCase):
"""Verify _get_weights_iterator dispatches to the right safetensors
@@ -249,12 +311,12 @@ class TestPrefetchDispatch(CustomTestCase):
prefix="",
)
- def _server_args(self, prefetch, disable_mmap=False):
+ def _server_args(self, prefetch, disable_mmap=False, drop_cache=False):
return SimpleNamespace(
weight_loader_disable_mmap=disable_mmap,
weight_loader_prefetch_checkpoints=prefetch,
weight_loader_prefetch_num_threads=4,
- weight_loader_drop_cache_after_load=False,
+ weight_loader_drop_cache_after_load=drop_cache,
)
def _run(self, loader):
@@ -263,7 +325,7 @@ class TestPrefetchDispatch(CustomTestCase):
# that calls the iterator factory) to execute.
list(loader._get_weights_iterator(self._make_source()))
- def _patch_dispatch(self, prefetch, disable_mmap=False):
+ def _patch_dispatch(self, prefetch, disable_mmap=False, drop_cache=False):
return (
patch.object(
DefaultModelLoader,
@@ -272,7 +334,11 @@ class TestPrefetchDispatch(CustomTestCase):
),
patch(
"sglang.srt.model_loader.loader.get_server_args",
- return_value=self._server_args(prefetch, disable_mmap),
+ return_value=self._server_args(
+ prefetch,
+ disable_mmap,
+ drop_cache,
+ ),
),
patch(
"sglang.srt.model_loader.loader."
@@ -403,11 +469,47 @@ class TestPrefetchDispatch(CustomTestCase):
p_warn as mock_warning,
):
self._run(loader)
- mock_fast.assert_called_once()
+ mock_fast.assert_called_once_with(
+ ["f.safetensors"],
+ enable_gds=True,
+ drop_cache_after_load=False,
+ )
mock_buffered.assert_not_called()
mock_single.assert_not_called()
mock_warning.assert_not_called()
+ def test_fastsafetensors_gds_can_be_disabled(self):
+ loader = self._make_loader(
+ {"enable_gds": False}, load_format=LoadFormat.FASTSAFETENSORS
+ )
+ p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
+ prefetch=False,
+ drop_cache=True,
+ )
+ with (
+ patch(
+ "sglang.srt.model_loader.loader.fastsafetensors_weights_iterator",
+ return_value=iter([]),
+ ) as mock_fast,
+ p_prep,
+ p_args,
+ p_buffered,
+ p_single,
+ p_warn,
+ ):
+ self._run(loader)
+ mock_fast.assert_called_once_with(
+ ["f.safetensors"],
+ enable_gds=False,
+ drop_cache_after_load=True,
+ )
+
+ def test_fastsafetensors_enable_gds_requires_boolean(self):
+ with self.assertRaisesRegex(ValueError, "enable_gds.*must be a boolean"):
+ self._make_loader(
+ {"enable_gds": "false"}, load_format=LoadFormat.FASTSAFETENSORS
+ )
+
if __name__ == "__main__":
unittest.main()