Clean logging under --weight-loader-prefetch-checkpoints (#33930)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
Co-authored-by: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com>
This commit is contained in:
Brayden Zhong
2026-09-04 20:05:53 -07:00
committed by GitHub
co-authored by Brayden Zhong Mohammad Miadh Angkad
parent 0645398a32
commit 92a4d8b5ee
19 changed files with 111 additions and 78 deletions
@@ -211,13 +211,13 @@ class TestPrefetchCheckpoints(CustomTestCase):
patch("concurrent.futures.ThreadPoolExecutor", _InlineExecutor),
patch("concurrent.futures.wait", side_effect=_wait_all),
patch("sglang.srt.model_loader.weight_utils._prefetch_checkpoint_file"),
patch("sglang.srt.model_loader.weight_utils.logger.info") as log_info,
patch("sglang.srt.model_loader.weight_utils.logger.debug") as log_debug,
):
_prefetch_all_checkpoints(paths, num_threads=1)
progress_pcts = [
call.args[2]
for call in log_info.call_args_list
for call in log_debug.call_args_list
if call.args
and call.args[0] == "Rank %d: prefetching checkpoint files: %d%% (%d/%d)"
]
@@ -426,14 +426,24 @@ class TestPrefetchDispatch(CustomTestCase):
"sglang.srt.model_loader.loader.safetensors_weights_iterator",
return_value=iter([]),
),
patch("sglang.srt.model_loader.loader.logger.warning"),
patch("sglang.srt.model_loader.loader.logger.debug"),
)
@staticmethod
def _override_notices(mock_log):
"""The single-thread override notice among the captured log calls."""
return [
call
for call in mock_log.call_args_list
if call.args
and "falling back to single-threaded weight loading" in call.args[0]
]
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."""
and the opt-out notice fires once."""
loader = self._make_loader({})
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=True
)
with (
@@ -441,18 +451,18 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
self._run(loader)
mock_single.assert_called_once()
mock_buffered.assert_not_called()
mock_warning.assert_called_once()
self.assertEqual(len(self._override_notices(mock_log)), 1)
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=True
)
with (
@@ -460,19 +470,19 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
self._run(loader)
mock_buffered.assert_called_once()
mock_single.assert_not_called()
mock_warning.assert_not_called()
self.assertEqual(self._override_notices(mock_log), [])
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=True
)
with (
@@ -480,20 +490,20 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
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()
self.assertEqual(self._override_notices(mock_log), [])
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=False
)
with (
@@ -501,12 +511,12 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
self._run(loader)
mock_buffered.assert_called_once()
mock_single.assert_not_called()
mock_warning.assert_not_called()
self.assertEqual(self._override_notices(mock_log), [])
def test_startup_prefetch_reuses_existing_background_handle(self):
"""Startup commit reuses resolved shards and the active prefetch handle."""
@@ -518,7 +528,7 @@ class TestPrefetchDispatch(CustomTestCase):
weight_files=("f.safetensors",),
use_safetensors=True,
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=False
)
with (
@@ -526,7 +536,7 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
list(
loader._get_weights_iterator(
@@ -541,7 +551,7 @@ class TestPrefetchDispatch(CustomTestCase):
mock_single.assert_called_once()
self.assertFalse(mock_single.call_args.kwargs["prefetch"])
mock_buffered.assert_not_called()
mock_warning.assert_called_once()
self.assertEqual(len(self._override_notices(mock_log)), 1)
def test_completed_startup_prefetch_restores_multithread_loader(self):
loader = self._make_loader({})
@@ -552,7 +562,7 @@ class TestPrefetchDispatch(CustomTestCase):
weight_files=("f.safetensors",),
use_safetensors=True,
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=False
)
with (
@@ -560,7 +570,7 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
list(
loader._get_weights_iterator(
@@ -575,7 +585,7 @@ class TestPrefetchDispatch(CustomTestCase):
mock_buffered.assert_called_once()
self.assertFalse(mock_buffered.call_args.kwargs["prefetch"])
mock_single.assert_not_called()
mock_warning.assert_not_called()
self.assertEqual(self._override_notices(mock_log), [])
def test_completed_startup_prefetch_is_not_started_twice(self):
loader = self._make_loader({})
@@ -586,7 +596,7 @@ class TestPrefetchDispatch(CustomTestCase):
weight_files=("f.safetensors",),
use_safetensors=True,
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=True
)
with (
@@ -594,7 +604,7 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
list(
loader._get_weights_iterator(
@@ -608,13 +618,13 @@ class TestPrefetchDispatch(CustomTestCase):
mock_buffered.assert_called_once()
self.assertFalse(mock_buffered.call_args.kwargs["prefetch"])
mock_single.assert_not_called()
mock_warning.assert_not_called()
self.assertEqual(self._override_notices(mock_log), [])
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=True, disable_mmap=True
)
with (
@@ -622,18 +632,18 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
self._run(loader)
mock_buffered.assert_called_once()
mock_single.assert_not_called()
mock_warning.assert_not_called()
self.assertEqual(self._override_notices(mock_log), [])
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=True
)
with (
@@ -645,7 +655,7 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
p_log as mock_log,
):
self._run(loader)
mock_fast.assert_called_once_with(
@@ -655,13 +665,13 @@ class TestPrefetchDispatch(CustomTestCase):
)
mock_buffered.assert_not_called()
mock_single.assert_not_called()
mock_warning.assert_not_called()
self.assertEqual(self._override_notices(mock_log), [])
def test_fastsafetensors_gds_can_be_disabled(self):
loader = self._make_loader(
{"enable_gds": False}, load_format=LoadFormat.FASTSAFETENSORS
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
prefetch=False,
drop_cache=True,
)
@@ -674,7 +684,7 @@ class TestPrefetchDispatch(CustomTestCase):
p_model,
p_buffered,
p_single,
p_warn,
p_log,
):
self._run(loader)
mock_fast.assert_called_once_with(