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