Support fastsafetensors no-GDS loading and page-cache release (#31859)
This commit is contained in:
@@ -145,6 +145,12 @@ python -m sglang.launch_server \
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Filename pattern for per-rank shards.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>model-rank-{rank}-part-{part}.safetensors</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>fastsafetensors</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>enable_gds</code> (bool)</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Use GPU Direct Storage. Set to <code>false</code> when the host does not provide the NVIDIA GPUDirect Storage kernel driver, such as in a gVisor sandbox.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>true</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>bitsandbytes</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>qlora_adapter_name_or_path</code> (str)</td>
|
||||
@@ -200,7 +206,7 @@ Top-level arguments that tune how safetensors weights are read, independent of `
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--weight-loader-drop-cache-after-load</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Call <code>posix_fadvise(DONTNEED)</code> on each safetensors shard after loading it, freeing page cache.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Call <code>posix_fadvise(DONTNEED)</code> after successfully loading each shard, freeing page cache. Supported by the standard safetensors and <code>fastsafetensors</code> loaders.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>off</td>
|
||||
</tr>
|
||||
<tr>
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user