Disable multi-threaded load by default when prefetch is on (#30146)
This commit is contained in:
@@ -130,7 +130,7 @@ python -m sglang.launch_server \
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>auto</code> / <code>safetensors</code> / <code>pt</code> / <code>npcache</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>enable_multithread_load</code> (bool)</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Read weight shards with a thread pool instead of sequentially.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Read weight shards with a thread pool instead of sequentially. Disabled by default when <code>--weight-loader-prefetch-checkpoints</code> is set (to avoid I/O oversubscription with the prefetch threads); set this to <code>true</code> to opt back in.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>true</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -190,7 +190,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-prefetch-checkpoints</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Prefetch checkpoint files into the OS page cache before loading. Each rank prefetches a fraction of the shards, cutting total network I/O on shared filesystems (NFS/Lustre) from N×checkpoint to 1×checkpoint. Recommended for models on network storage.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Prefetch checkpoint files into the OS page cache before loading. Each rank prefetches a fraction of the shards, cutting total network I/O on shared filesystems (NFS/Lustre) from N×checkpoint to 1×checkpoint. Recommended for models on network storage. When enabled, multi-threaded safetensors loading is disabled by default to avoid I/O oversubscription with the prefetch threads; set <code>enable_multithread_load=true</code> in <code>--model-loader-extra-config</code> to keep multi-threaded loading (e.g. on local NVMe where prefetch is a no-op).</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>off</td>
|
||||
</tr>
|
||||
<tr>
|
||||
|
||||
@@ -3095,7 +3095,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--weight-loader-prefetch-checkpoints`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Prefetch checkpoint files into OS page cache before loading. Each rank prefetches a fraction of the shards in a background thread, reducing total network I/O on shared filesystems (NFS/Lustre) from N\*checkpoint to 1\*checkpoint. Recommended for models on network storage.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Prefetch checkpoint files into OS page cache before loading. Each rank prefetches a fraction of the shards in a background thread, reducing total network I/O on shared filesystems (NFS/Lustre) from N\*checkpoint to 1\*checkpoint. Recommended for models on network storage. When enabled, multi-threaded safetensors loading is disabled by default to avoid I/O oversubscription with the prefetch threads; set `enable_multithread_load=true` in `--model-loader-extra-config` to keep multi-threaded loading (e.g. on local NVMe where prefetch is a no-op).</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`False`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>bool flag (set to enable)</td>
|
||||
</tr>
|
||||
|
||||
@@ -576,6 +576,34 @@ class DefaultModelLoader(BaseModelLoader):
|
||||
server_args.weight_loader_drop_cache_after_load
|
||||
)
|
||||
|
||||
# Prefetch and multi-threaded loading both read the same shards,
|
||||
# competing for I/O on shared/network storage. When prefetch is
|
||||
# active (mmap path, not FASTSAFETENSORS) and the user didn't
|
||||
# explicitly request multi-threaded loading, fall back to the
|
||||
# single-threaded loader and let prefetch feed the page cache.
|
||||
# Setting enable_multithread_load or num_threads in
|
||||
# --model-loader-extra-config opts out (the latter is consumed
|
||||
# only by the multi-threaded iterator, so it signals intent);
|
||||
# e.g. local NVMe, where prefetch is a no-op and multi-threading
|
||||
# helps.
|
||||
if (
|
||||
weight_loader_prefetch
|
||||
and not weight_loader_disable_mmap
|
||||
and self.load_config.load_format != LoadFormat.FASTSAFETENSORS
|
||||
and use_multithread
|
||||
and not (
|
||||
{"enable_multithread_load", "num_threads"} & extra_config.keys()
|
||||
)
|
||||
):
|
||||
logger.warning(
|
||||
"--weight-loader-prefetch-checkpoints is enabled; falling "
|
||||
"back to single-threaded weight loading to avoid I/O "
|
||||
"oversubscription with the prefetch threads. Set "
|
||||
"enable_multithread_load=true in --model-loader-extra-config "
|
||||
"to keep multi-threaded loading."
|
||||
)
|
||||
use_multithread = False
|
||||
|
||||
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
|
||||
weights_iterator = fastsafetensors_weights_iterator(
|
||||
hf_weights_files,
|
||||
|
||||
@@ -2442,7 +2442,7 @@ class ServerArgs:
|
||||
] = False
|
||||
weight_loader_prefetch_checkpoints: A[
|
||||
bool,
|
||||
"Prefetch checkpoint files into OS page cache before loading. Each rank prefetches a fraction of the shards, reducing total network I/O on shared filesystems (NFS/Lustre) from N*checkpoint to 1*checkpoint. Recommended for models on network storage.",
|
||||
"Prefetch checkpoint files into OS page cache before loading. Each rank prefetches a fraction of the shards, reducing total network I/O on shared filesystems (NFS/Lustre) from N*checkpoint to 1*checkpoint. Recommended for models on network storage. When enabled, multi-threaded safetensors loading is disabled by default to avoid I/O oversubscription with the prefetch threads; set enable_multithread_load=true in --model-loader-extra-config to keep multi-threaded loading (e.g. on local NVMe where prefetch is a no-op).",
|
||||
] = False
|
||||
weight_loader_prefetch_num_threads: A[
|
||||
int,
|
||||
|
||||
@@ -9,17 +9,21 @@ import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from concurrent.futures import Future
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import safetensors.torch
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||
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,
|
||||
safetensors_weights_iterator,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
@@ -56,7 +60,7 @@ def _wait_all(fs, return_when):
|
||||
return set(fs), set()
|
||||
|
||||
|
||||
class TestPrefetchCheckpoints(unittest.TestCase):
|
||||
class TestPrefetchCheckpoints(CustomTestCase):
|
||||
"""Verify coordinated checkpoint prefetch behavior."""
|
||||
|
||||
def _create_safetensors_files(self, tmpdir, num_shards=3):
|
||||
@@ -218,5 +222,192 @@ class TestPrefetchCheckpoints(unittest.TestCase):
|
||||
)
|
||||
|
||||
|
||||
class TestPrefetchDispatch(CustomTestCase):
|
||||
"""Verify _get_weights_iterator dispatches to the right safetensors
|
||||
iterator based on prefetch / multi-thread config.
|
||||
|
||||
Prefetch + default (multi-thread on) must fall back to the
|
||||
single-threaded iterator; an explicit enable_multithread_load=true is
|
||||
honored as an opt-out; FASTSAFETENSORS and disable_mmap bypass the
|
||||
override.
|
||||
"""
|
||||
|
||||
def _make_loader(self, extra_config, load_format=LoadFormat.SAFETENSORS):
|
||||
load_config = LoadConfig(
|
||||
load_format=load_format,
|
||||
model_loader_extra_config=extra_config,
|
||||
)
|
||||
return DefaultModelLoader(load_config)
|
||||
|
||||
def _make_source(self):
|
||||
# model_config=None skips maybe_add_mtp_safetensors.
|
||||
return SimpleNamespace(
|
||||
model_or_path="/dummy",
|
||||
revision=None,
|
||||
fall_back_to_pt=False,
|
||||
model_config=None,
|
||||
prefix="",
|
||||
)
|
||||
|
||||
def _server_args(self, prefetch, disable_mmap=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,
|
||||
)
|
||||
|
||||
def _run(self, loader):
|
||||
# _get_weights_iterator returns a generator wrapping the chosen
|
||||
# iterator; consuming it forces the eager dispatch (the if/elif/else
|
||||
# that calls the iterator factory) to execute.
|
||||
list(loader._get_weights_iterator(self._make_source()))
|
||||
|
||||
def _patch_dispatch(self, prefetch, disable_mmap=False):
|
||||
return (
|
||||
patch.object(
|
||||
DefaultModelLoader,
|
||||
"_prepare_weights",
|
||||
return_value=("/dummy", ["f.safetensors"], True),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.model_loader.loader.get_global_server_args",
|
||||
return_value=self._server_args(prefetch, disable_mmap),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.model_loader.loader."
|
||||
"buffered_multi_thread_safetensors_weights_iterator",
|
||||
return_value=iter([]),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.model_loader.loader.safetensors_weights_iterator",
|
||||
return_value=iter([]),
|
||||
),
|
||||
patch("sglang.srt.model_loader.loader.logger.warning"),
|
||||
)
|
||||
|
||||
def test_prefetch_uses_single_thread_for_default_config(self):
|
||||
"""Prefetch on + no explicit multithread config -> single-threaded,
|
||||
and the opt-out warning fires once."""
|
||||
loader = self._make_loader({})
|
||||
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
p_prep,
|
||||
p_args,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_single.assert_called_once()
|
||||
mock_buffered.assert_not_called()
|
||||
mock_warning.assert_called_once()
|
||||
|
||||
def test_explicit_enable_multithread_keeps_buffered_with_prefetch(self):
|
||||
"""Explicit enable_multithread_load=true is the escape hatch; the
|
||||
override and its warning must not fire."""
|
||||
loader = self._make_loader({"enable_multithread_load": True})
|
||||
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
p_prep,
|
||||
p_args,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
|
||||
def test_num_threads_only_keeps_buffered_with_prefetch(self):
|
||||
"""num_threads alone (relying on the enable_multithread_load=True
|
||||
default) also signals multi-thread intent, so the override must not
|
||||
fire and num_threads stays live."""
|
||||
loader = self._make_loader({"num_threads": 64})
|
||||
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
p_prep,
|
||||
p_args,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
# num_threads is forwarded as max_workers to the buffered iterator.
|
||||
self.assertEqual(mock_buffered.call_args.kwargs["max_workers"], 64)
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
|
||||
def test_no_prefetch_uses_multithread(self):
|
||||
"""Prefetch off -> multi-threaded iterator is used (default), no
|
||||
override warning."""
|
||||
loader = self._make_loader({})
|
||||
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
prefetch=False
|
||||
)
|
||||
with (
|
||||
p_prep,
|
||||
p_args,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
|
||||
def test_prefetch_does_not_override_when_mmap_disabled(self):
|
||||
"""Prefetch is a no-op without mmap, so the override and its warning
|
||||
must not fire."""
|
||||
loader = self._make_loader({})
|
||||
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
prefetch=True, disable_mmap=True
|
||||
)
|
||||
with (
|
||||
p_prep,
|
||||
p_args,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
|
||||
def test_prefetch_does_not_override_for_fastsafetensors(self):
|
||||
"""FASTSAFETENSORS ignores both flags; override + warning must not
|
||||
fire."""
|
||||
loader = self._make_loader({}, load_format=LoadFormat.FASTSAFETENSORS)
|
||||
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.model_loader.loader.fastsafetensors_weights_iterator",
|
||||
return_value=iter([]),
|
||||
) as mock_fast,
|
||||
p_prep,
|
||||
p_args,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_fast.assert_called_once()
|
||||
mock_buffered.assert_not_called()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user