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:
co-authored by
Brayden Zhong
Mohammad Miadh Angkad
parent
0645398a32
commit
92a4d8b5ee
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user