Disable multi-threaded load by default when prefetch is on (#30146)

This commit is contained in:
Mohammad Miadh Angkad
2026-07-09 00:28:53 -07:00
committed by GitHub
parent 64e2a73c80
commit 666a09fe2a
5 changed files with 224 additions and 5 deletions
@@ -130,7 +130,7 @@ python -m sglang.launch_server \
<tr> <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", 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.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> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>true</code></td>
</tr> </tr>
<tr> <tr>
@@ -190,7 +190,7 @@ Top-level arguments that tune how safetensors weights are read, independent of `
</tr> </tr>
<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", 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> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>off</td>
</tr> </tr>
<tr> <tr>
@@ -3095,7 +3095,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
</tr> </tr>
<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", 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.02)"}}>`False`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>bool flag (set to enable)</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>bool flag (set to enable)</td>
</tr> </tr>
+28
View File
@@ -576,6 +576,34 @@ class DefaultModelLoader(BaseModelLoader):
server_args.weight_loader_drop_cache_after_load 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: if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
weights_iterator = fastsafetensors_weights_iterator( weights_iterator = fastsafetensors_weights_iterator(
hf_weights_files, hf_weights_files,
+1 -1
View File
@@ -2442,7 +2442,7 @@ class ServerArgs:
] = False ] = False
weight_loader_prefetch_checkpoints: A[ weight_loader_prefetch_checkpoints: A[
bool, 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 ] = False
weight_loader_prefetch_num_threads: A[ weight_loader_prefetch_num_threads: A[
int, int,
@@ -9,17 +9,21 @@ import os
import tempfile import tempfile
import unittest import unittest
from concurrent.futures import Future from concurrent.futures import Future
from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import safetensors.torch import safetensors.torch
import 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 ( from sglang.srt.model_loader.weight_utils import (
_prefetch_all_checkpoints, _prefetch_all_checkpoints,
buffered_multi_thread_safetensors_weights_iterator, buffered_multi_thread_safetensors_weights_iterator,
safetensors_weights_iterator, safetensors_weights_iterator,
) )
from sglang.test.ci.ci_register import register_cpu_ci 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") 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() return set(fs), set()
class TestPrefetchCheckpoints(unittest.TestCase): class TestPrefetchCheckpoints(CustomTestCase):
"""Verify coordinated checkpoint prefetch behavior.""" """Verify coordinated checkpoint prefetch behavior."""
def _create_safetensors_files(self, tmpdir, num_shards=3): 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__": if __name__ == "__main__":
unittest.main() unittest.main()