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>
<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>
+28
View File
@@ -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,
+1 -1
View File
@@ -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()