[Model Loading] Overlap checkpoint staging with CUDA graph capture during startup (#32017)

Co-authored-by: Wenhui Zhu <wzhu59@asu.edu>
Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
Han Yu
2026-08-13 12:26:25 -07:00
committed by GitHub
co-authored by Wenhui Zhu Alex Nails
parent 0772e79ee7
commit 6b94d39f13
16 changed files with 2020 additions and 79 deletions
@@ -0,0 +1,69 @@
"""End-to-end parity test for post-capture startup weight loading."""
import unittest
import sglang as sgl
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
CustomTestCase,
)
register_cuda_ci(est_time=120, stage="base-b", runner_config="1-gpu-small")
class TestStartupWeightLoad(CustomTestCase):
@staticmethod
def _generate(startup_weight_load_mode=None):
kwargs = dict(
model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
dtype="bfloat16",
random_seed=42,
cuda_graph_max_bs_decode=1,
max_total_tokens=256,
)
if startup_weight_load_mode is not None:
kwargs["startup_weight_load_mode"] = startup_weight_load_mode
with sgl.Engine(**kwargs) as engine:
return engine.generate(
"The capital of France is",
sampling_params={
"temperature": 0,
"max_new_tokens": 8,
"ignore_eos": True,
},
return_logprob=True,
logprob_start_len=0,
)
def test_overlap_matches_default_serial_startup(self):
# Omitting the flag is intentional: it pins the merge-safe default path.
serial = self._generate()
overlap = self._generate("overlap")
self.assertEqual(serial["output_ids"], overlap["output_ids"])
self.assertEqual(serial["text"], overlap["text"])
serial_logprobs = serial["meta_info"]["output_token_logprobs"]
overlap_logprobs = overlap["meta_info"]["output_token_logprobs"]
self.assertEqual(len(serial_logprobs), len(overlap_logprobs))
self.assertGreater(len(serial_logprobs), 0)
for index, (serial_item, overlap_item) in enumerate(
zip(serial_logprobs, overlap_logprobs)
):
self.assertEqual(
serial_item[1],
overlap_item[1],
f"token id differs at output position {index}",
)
self.assertAlmostEqual(
serial_item[0],
overlap_item[0],
delta=1e-5,
msg=f"logprob differs at output position {index}",
)
if __name__ == "__main__":
unittest.main()
@@ -220,6 +220,20 @@ class TestMlxExtendRouting(CustomTestCase):
worker._mlx_pool_initialized = True
return worker
def test_startup_weight_overlap_is_rejected_before_mlx_model_load(self):
from sglang.srt.hardware_backend.mlx.model_runner_stub import (
MlxModelRunnerStub,
)
from sglang.srt.hardware_backend.mlx.tp_worker import MlxTpModelWorker
worker = MlxTpModelWorker.__new__(MlxTpModelWorker)
worker.server_args = SimpleNamespace(is_startup_weight_load_overlap=True)
with self.assertRaisesRegex(ValueError, "CUDA only"):
MlxModelRunnerStub.validate_startup_weight_load_mode(worker.server_args)
with self.assertRaisesRegex(ValueError, "CUDA only"):
worker._init_model_runner()
# ---------- the shared decision helper ----------
# The helper takes no seq_len: length cannot distinguish a 1-token
# continuation from a genuine decode -- request state does.
@@ -0,0 +1,732 @@
"""Unit tests for the post-capture startup weight-loading component."""
import dataclasses
import re
import unittest
from types import SimpleNamespace
from unittest.mock import call, patch
import torch
from torch import nn
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.configs.device_config import DeviceConfig
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
from sglang.srt.configs.model_config import ModelImpl
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.model_executor.model_runner_components.startup_weight_load import (
ModelStorageManifest,
StartupWeightLoadManager,
StartupWeightLoadOptions,
StartupWeightLoadState,
)
from sglang.srt.model_loader.loader import DefaultModelLoader
from sglang.srt.model_loader.weight_utils import initialize_capture_safe_weights
from sglang.srt.runtime_context import get_context
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
_STARTUP_MODULE = (
"sglang.srt.model_executor.model_runner_components.startup_weight_load"
)
class _CanonicalModel:
pass
class _ExternalModel:
pass
def _make_options(**overrides):
options = StartupWeightLoadOptions(
device="cuda",
is_cuda_platform=True,
cuda_graph_enabled=True,
prefill_cuda_graph_backend=Backend.FULL,
is_draft_worker=False,
speculative_algorithm=None,
tp_size=1,
attn_cp_size=1,
dcp_size=1,
pp_size=1,
dp_size=1,
ep_size=1,
cpu_offload_gb=0,
offload_group_size=-1,
enable_memory_saver=False,
enable_weights_cpu_backup=False,
torchao_config="",
enable_lora=False,
has_lora_paths=False,
weight_loader_disable_mmap=False,
weight_loader_drop_cache_after_load=False,
has_custom_weight_loader=False,
enable_torch_compile=False,
prefetch_num_threads=4,
)
return dataclasses.replace(options, **overrides)
def _make_model_config(**overrides):
values = dict(
hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]),
dtype=torch.bfloat16,
quantization=None,
modelopt_quant=None,
is_multimodal=False,
is_generation=True,
model_impl=ModelImpl.SGLANG,
_resolved_model_impl=ModelImpl.SGLANG,
)
values.update(overrides)
return SimpleNamespace(**values)
class _RecordingPrefetchHandle:
def __init__(self, trace, *, done=False, errors=()):
self._trace = trace
self.done = done
self.errors = errors
@property
def failed(self):
return bool(self.errors)
def wait(self, timeout=None):
self._trace.append("wait_prefetch")
def stop(self, timeout=None):
self._trace.append("stop_prefetch")
self.wait()
self.done = True
class _RecordingLoader:
def __init__(self, model, trace):
self._model = model
self._trace = trace
self.prefetch_handle = _RecordingPrefetchHandle(trace)
def initialize_model_for_startup(self, *, model_config, device_config):
self._trace.append("initialize")
return self._model
def resolve_model_weights(self, model_config, model):
self._trace.append("resolve")
return (object(),)
def start_checkpoint_prefetch(self, resolved_sources, *, num_threads):
self._trace.append("start_prefetch")
return self.prefetch_handle
def prepare_model_for_capture(self, *, model, model_config):
self._trace.append("prepare_capture")
return model
def commit_model_weights(
self,
*,
model,
model_config,
resolved_sources,
target_device,
startup_prefetch_active,
):
self._trace.append("commit")
self.startup_prefetch_active = startup_prefetch_active
with torch.no_grad():
for parameter in model.parameters():
parameter.fill_(3)
class _TiedWeightModel(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.ones(2, 2))
self.tied_weight = self.weight
self.register_buffer("scale", torch.ones(2))
class TestStartupWeightLoadSelector(CustomTestCase):
def setUp(self):
self.load_config = LoadConfig(load_format=LoadFormat.SAFETENSORS)
self.loader = DefaultModelLoader(self.load_config)
self.device_config = DeviceConfig("cuda", 0)
def _create(
self,
*,
options=None,
model_config=None,
load_config=None,
loader=None,
resolved_model_class=None,
):
model_config = _make_model_config() if model_config is None else model_config
architecture = model_config.hf_config.architectures[0]
with (
patch(
f"{_STARTUP_MODULE}.get_model_architecture",
return_value=(
resolved_model_class or _CanonicalModel,
architecture,
),
),
patch(
f"{_STARTUP_MODULE}._get_canonical_model_class",
return_value=_CanonicalModel,
),
):
return StartupWeightLoadManager.create(
loader=self.loader if loader is None else loader,
model_config=model_config,
load_config=self.load_config if load_config is None else load_config,
device_config=self.device_config,
options=_make_options() if options is None else options,
)
def test_supported_overlap_creates_a_manager(self):
self.assertIsInstance(self._create(), StartupWeightLoadManager)
self.assertIsInstance(
self._create(options=_make_options(tp_size=2)),
StartupWeightLoadManager,
)
def test_unsupported_overlap_is_rejected_instead_of_falling_back(self):
cases = (
(
"non_cuda",
dict(options=_make_options(device="cpu", is_cuda_platform=False)),
"CUDA only",
),
(
"graphs_disabled",
dict(options=_make_options(cuda_graph_enabled=False)),
"CUDA graph capture is disabled",
),
(
"tc_piecewise_prefill",
dict(
options=_make_options(
prefill_cuda_graph_backend=Backend.TC_PIECEWISE
)
),
"tc_piecewise prefill CUDA graphs are not supported",
),
(
"pt_checkpoint",
dict(load_config=LoadConfig(load_format=LoadFormat.PT)),
"load format must be auto or safetensors",
),
(
"draft_worker",
dict(options=_make_options(is_draft_worker=True)),
"draft workers are not supported",
),
(
"draft_model_checkpoint",
dict(
load_config=LoadConfig(
load_format=LoadFormat.SAFETENSORS,
draft_model_idx=0,
)
),
"draft model loading is unsupported",
),
(
"speculative_decoding",
dict(options=_make_options(speculative_algorithm="EAGLE")),
"speculative decoding is not supported",
),
(
"tp3",
dict(options=_make_options(tp_size=3)),
"only TP1 and TP2 are supported",
),
(
"attention_context_parallel",
dict(options=_make_options(tp_size=2, attn_cp_size=2)),
"attention context parallelism is not supported",
),
(
"decode_context_parallel",
dict(options=_make_options(tp_size=2, dcp_size=2)),
"decode context parallelism is not supported",
),
(
"quantized_model",
dict(model_config=_make_model_config(quantization="fp8")),
"quantization is not supported",
),
(
"layer_group_offload",
dict(options=_make_options(offload_group_size=1)),
"layer-group offloading is not supported",
),
(
"torch_compile",
dict(options=_make_options(enable_torch_compile=True)),
"torch.compile is not supported",
),
(
"transformers_model_impl",
dict(
model_config=_make_model_config(
model_impl=ModelImpl.TRANSFORMERS,
_resolved_model_impl=ModelImpl.TRANSFORMERS,
),
resolved_model_class=_ExternalModel,
),
"the native SGLang model implementation is required",
),
(
"external_model_implementation",
dict(resolved_model_class=_ExternalModel),
"the native SGLang model implementation is required",
),
(
"unknown_architecture",
dict(
model_config=_make_model_config(
hf_config=SimpleNamespace(architectures=["OtherForCausalLM"])
)
),
"model architecture is not in the startup-overlap allowlist",
),
)
for name, kwargs, reason in cases:
with self.subTest(name=name):
with self.assertRaisesRegex(ValueError, re.escape(reason)):
self._create(**kwargs)
class TestStartupWeightLoadManager(CustomTestCase):
def _manager(self, loader):
return StartupWeightLoadManager(
loader=loader,
model_config=_make_model_config(),
device_config=DeviceConfig("cpu", 0),
options=_make_options(),
)
def test_prepare_capture_finalize_state_and_order(self):
trace = []
model = _TiedWeightModel()
manager = self._manager(_RecordingLoader(model, trace))
self.assertEqual(manager.state, StartupWeightLoadState.CREATED)
self.assertIs(manager.prepare(), model)
self.assertEqual(manager.state, StartupWeightLoadState.CAPTURE_READY)
manager.start_prefetch()
self.assertEqual(manager.state, StartupWeightLoadState.PREFETCHING)
# CUDA graph capture is owned by Scheduler and occurs between these calls.
trace.append("capture")
with (
patch(
f"{_STARTUP_MODULE}.monkey_patch_vllm_parallel_state"
) as parallel_state_patch,
patch(f"{_STARTUP_MODULE}.torch.cuda.synchronize"),
patch(f"{_STARTUP_MODULE}.logger.info") as log_info,
):
manager.finalize()
self.assertEqual(manager.state, StartupWeightLoadState.READY)
self.assertEqual(
trace,
[
"initialize",
"resolve",
"prepare_capture",
"start_prefetch",
"capture",
"commit",
"stop_prefetch",
"wait_prefetch",
],
)
# Finalization is idempotent after a successful commit.
manager.finalize()
self.assertEqual(trace.count("commit"), 1)
self.assertIs(model.weight, model.tied_weight)
torch.testing.assert_close(model.weight, torch.full_like(model.weight, 3))
self.assertTrue(log_info.call_args.args[0].startswith("Load weight end."))
self.assertTrue(manager._loader.startup_prefetch_active)
self.assertEqual(
parallel_state_patch.call_args_list,
[call(), call(reverse=True)],
)
def test_finalize_rejects_graph_visible_storage_rebind(self):
trace = []
model = _TiedWeightModel()
loader = _RecordingLoader(model, trace)
def rebind_tied_weight(**kwargs):
trace.append("commit")
model.tied_weight = nn.Parameter(model.tied_weight.detach().clone())
loader.commit_model_weights = rebind_tied_weight
manager = self._manager(loader)
manager.prepare()
manager.start_prefetch()
with (
patch(f"{_STARTUP_MODULE}.monkey_patch_vllm_parallel_state"),
patch(f"{_STARTUP_MODULE}.torch.cuda.synchronize"),
self.assertRaisesRegex(
RuntimeError,
"changed graph-visible tensor storage: parameter:tied_weight",
),
):
manager.finalize()
def test_finalize_rejects_parameter_left_at_capture_sentinel(self):
trace = []
model = _TiedWeightModel()
loader = _RecordingLoader(model, trace)
def skip_commit(**kwargs):
trace.append("commit")
loader.commit_model_weights = skip_commit
manager = self._manager(loader)
manager.prepare()
with torch.no_grad():
model.weight.fill_(1e-3)
manager.start_prefetch()
with (
patch(f"{_STARTUP_MODULE}.monkey_patch_vllm_parallel_state"),
patch(f"{_STARTUP_MODULE}.torch.cuda.synchronize"),
self.assertRaisesRegex(
RuntimeError,
"did not replace capture-safe dummy values: parameter:tied_weight",
),
):
manager.finalize()
def test_completed_prefetch_restores_normal_loader(self):
trace = []
model = _TiedWeightModel()
loader = _RecordingLoader(model, trace)
loader.prefetch_handle.done = True
manager = self._manager(loader)
manager.prepare()
manager.start_prefetch()
with (
patch(f"{_STARTUP_MODULE}.monkey_patch_vllm_parallel_state"),
patch(f"{_STARTUP_MODULE}.torch.cuda.synchronize"),
):
manager.finalize()
self.assertFalse(loader.startup_prefetch_active)
self.assertIn("wait_prefetch", trace)
self.assertNotIn("stop_prefetch", trace)
def test_failed_prefetch_falls_back_and_logs_summary(self):
trace = []
model = _TiedWeightModel()
loader = _RecordingLoader(model, trace)
loader.prefetch_handle.errors = (("bad.safetensors", OSError("failed")),)
manager = self._manager(loader)
manager.prepare()
manager.start_prefetch()
with (
patch(f"{_STARTUP_MODULE}.monkey_patch_vllm_parallel_state"),
patch(f"{_STARTUP_MODULE}.torch.cuda.synchronize"),
patch(f"{_STARTUP_MODULE}.logger.warning") as warning,
):
manager.finalize()
self.assertFalse(loader.startup_prefetch_active)
warning.assert_called_once()
self.assertIn("falling back", warning.call_args.args[2])
def test_stop_timeout_after_commit_does_not_fail_startup(self):
trace = []
model = _TiedWeightModel()
loader = _RecordingLoader(model, trace)
def _stop_times_out(timeout=None):
trace.append("stop_prefetch")
raise TimeoutError("Timed out waiting for checkpoint prefetching")
loader.prefetch_handle.stop = _stop_times_out
manager = self._manager(loader)
manager.prepare()
manager.start_prefetch()
with (
patch(f"{_STARTUP_MODULE}.monkey_patch_vllm_parallel_state"),
patch(f"{_STARTUP_MODULE}.torch.cuda.synchronize"),
patch(f"{_STARTUP_MODULE}.logger.warning") as warning,
):
manager.finalize()
self.assertEqual(manager.state, StartupWeightLoadState.READY)
self.assertIn("stop_prefetch", trace)
warning.assert_called_once()
self.assertIn("did not stop within its timeout", warning.call_args.args[0])
def test_start_prefetch_requires_capture_ready_and_starts_once(self):
trace = []
manager = self._manager(_RecordingLoader(nn.Linear(2, 2), trace))
with self.assertRaisesRegex(RuntimeError, "from state"):
manager.start_prefetch()
manager.prepare()
manager.start_prefetch()
self.assertEqual(manager.state, StartupWeightLoadState.PREFETCHING)
with self.assertRaisesRegex(RuntimeError, "from state"):
manager.start_prefetch()
self.assertEqual(trace.count("start_prefetch"), 1)
class TestModelStorageManifest(CustomTestCase):
def test_in_place_updates_preserve_the_manifest(self):
model = _TiedWeightModel()
manifest = ModelStorageManifest.capture(model)
with torch.no_grad():
model.weight.fill_(2)
model.scale.fill_(3)
self.assertEqual(manifest.changed_names(model), ())
def test_manifest_keeps_strong_tensor_references(self):
model = _TiedWeightModel()
manifest = ModelStorageManifest.capture(model)
metadata = dict(manifest.tensors)["parameter:weight"]
self.assertIs(metadata.tensor, model.weight)
def test_capture_sentinel_check_ignores_buffers(self):
model = _TiedWeightModel()
with torch.no_grad():
model.weight.fill_(1e-3)
model.scale.fill_(1e-3)
manifest = ModelStorageManifest.capture(model)
self.assertEqual(
manifest.unchanged_parameter_names(1e-3),
("parameter:tied_weight",),
)
def test_parameter_rebind_and_alias_break_are_detected(self):
model = _TiedWeightModel()
manifest = ModelStorageManifest.capture(model)
model.tied_weight = nn.Parameter(model.tied_weight.detach().clone())
self.assertEqual(
manifest.changed_names(model),
("parameter:tied_weight",),
)
class TestCaptureSafeWeightInitialization(CustomTestCase):
def test_only_parameters_are_filled(self):
model = _TiedWeightModel()
initialize_capture_safe_weights(model, value=0.125)
torch.testing.assert_close(model.weight, torch.full_like(model.weight, 0.125))
torch.testing.assert_close(model.scale, torch.ones_like(model.scale))
class _LifecycleRunner:
def __init__(self, name, trace):
self._name = name
self._trace = trace
def start_startup_weight_load(self):
self._trace.append(f"start:{self._name}")
def finalize_startup_weight_load(self):
self._trace.append(f"finalize:{self._name}")
class TestStartupWeightLoadFanout(CustomTestCase):
def test_primary_and_multi_runner_extras_are_started_once(self):
trace = []
primary = _LifecycleRunner("primary", trace)
extra_1 = _LifecycleRunner("extra_1", trace)
extra_2 = _LifecycleRunner("extra_2", trace)
worker = TpModelWorker.__new__(TpModelWorker)
worker._model_runner = primary
worker.model_runner_list = [primary, extra_1, extra_2]
worker.start_startup_weight_load()
self.assertEqual(
trace,
["start:primary", "start:extra_1", "start:extra_2"],
)
def test_primary_and_multi_runner_extras_are_finalized_once(self):
for multi_runner in (False, True):
with self.subTest(multi_runner=multi_runner):
trace = []
primary = _LifecycleRunner("primary", trace)
extra_1 = _LifecycleRunner("extra_1", trace)
extra_2 = _LifecycleRunner("extra_2", trace)
worker = TpModelWorker.__new__(TpModelWorker)
worker._model_runner = primary
worker.model_runner_list = (
[primary, extra_1, extra_2] if multi_runner else []
)
worker.finalize_startup_weight_load()
self.assertEqual(
trace,
(
["finalize:primary", "finalize:extra_1", "finalize:extra_2"]
if multi_runner
else ["finalize:primary"]
),
)
class _RunnerStartupManager:
def __init__(self, trace):
self._trace = trace
def start_prefetch(self):
self._trace.append("start_prefetch")
def finalize(self):
self._trace.append("finalize")
class TestModelRunnerStartupWeightLoadOwnership(CustomTestCase):
@staticmethod
def _runner(manager):
runner = ModelRunner.__new__(ModelRunner)
runner.startup_weight_load = manager
runner.server_args = SimpleNamespace(
elastic_ep_backend=None,
is_ep_joiner=False,
)
runner.ps = SimpleNamespace(tp_rank=0)
return runner
def test_start_delegates_to_the_manager(self):
trace = []
runner = self._runner(_RunnerStartupManager(trace))
runner.start_startup_weight_load()
self.assertEqual(trace, ["start_prefetch"])
def test_success_releases_ownership_after_the_barrier(self):
trace = []
manager = _RunnerStartupManager(trace)
runner = self._runner(manager)
def barrier(**kwargs):
self.assertIs(runner.startup_weight_load, manager)
trace.append("barrier")
with (
patch(
"sglang.srt.model_executor.model_runner.dist_barrier_after_load",
side_effect=barrier,
),
get_context().override_server_args(),
):
runner.finalize_startup_weight_load()
self.assertEqual(trace, ["finalize", "barrier"])
self.assertIsNone(runner.startup_weight_load)
class _SchedulerWorker:
def __init__(self, trace, *, post_capture_active=False):
self._trace = trace
self.model_runner = SimpleNamespace(
token_to_kv_pool=SimpleNamespace(post_capture_active=post_capture_active),
post_capture_resize_kv_pool=lambda: trace.append("resize"),
)
def start_startup_weight_load(self):
self._trace.append("start")
def finalize_startup_weight_load(self):
self._trace.append("finalize")
class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
@staticmethod
def _scheduler(worker, trace, *, mode):
from sglang.srt.managers.scheduler import Scheduler
scheduler = Scheduler.__new__(Scheduler)
scheduler.server_args = SimpleNamespace(
is_startup_weight_load_overlap=mode == "overlap"
)
scheduler.init_tp_model_worker = lambda: setattr(scheduler, "tp_worker", worker)
scheduler.maybe_init_draft_worker = lambda: setattr(
scheduler, "draft_worker", None
)
scheduler.init_memory_pools = lambda: trace.append("memory_pool")
scheduler.init_all_attention_backends = lambda: trace.append("attention")
scheduler.init_all_cuda_graphs = lambda: trace.append("capture")
return scheduler
def _run_startup(self, mode):
trace = []
worker = _SchedulerWorker(trace, post_capture_active=True)
scheduler = self._scheduler(worker, trace, mode=mode)
def stop_after_startup():
raise RuntimeError("stop after startup")
scheduler.spec_algorithm = SimpleNamespace(is_none=stop_after_startup)
with (
patch(
"sglang.srt.managers.scheduler.get_exec",
return_value=SimpleNamespace(
moe=SimpleNamespace(
elastic_ep_backend=None,
ep_join_mode=None,
)
),
),
self.assertRaisesRegex(RuntimeError, "stop after startup"),
):
scheduler.init_model_worker()
return trace
def test_serial_path_skips_overlap_hooks(self):
self.assertEqual(
self._run_startup("serial"),
["memory_pool", "attention", "capture", "resize"],
)
def test_overlap_starts_before_capture_and_finalizes_after(self):
self.assertEqual(
self._run_startup("overlap"),
["start", "memory_pool", "attention", "capture", "resize", "finalize"],
)
if __name__ == "__main__":
unittest.main()
@@ -7,10 +7,11 @@ to weights loaded without prefetch.
import os
import tempfile
import threading
import unittest
from concurrent.futures import Future
from types import SimpleNamespace
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import safetensors.torch
import torch
@@ -18,6 +19,7 @@ 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 (
CheckpointFilePrefetchHandle,
_prefetch_all_checkpoints,
buffered_multi_thread_safetensors_weights_iterator,
fastsafetensors_weights_iterator,
@@ -37,6 +39,12 @@ class _InlineThread:
def start(self):
self.target()
def join(self, timeout=None):
pass
def is_alive(self):
return False
class _InlineExecutor:
def __init__(self, max_workers):
@@ -99,6 +107,45 @@ class TestPrefetchCheckpoints(CustomTestCase):
with self.assertRaisesRegex(ValueError, "num_threads"):
_prefetch_all_checkpoints(["dummy.safetensors"], num_threads=0)
@patch("torch.distributed.is_initialized", return_value=False)
def test_wait_returns_after_worker_thread_failure(self, _):
worker_errors = []
with (
patch(
"concurrent.futures.ThreadPoolExecutor",
side_effect=RuntimeError("worker failed"),
),
patch(
"threading.excepthook",
side_effect=lambda args: worker_errors.append(args.exc_value),
),
):
handle = _prefetch_all_checkpoints(["dummy.safetensors"], num_threads=1)
handle.wait(timeout=5)
self.assertTrue(handle.done)
self.assertTrue(handle.failed)
self.assertEqual(handle.errors, ())
self.assertEqual(len(worker_errors), 1)
self.assertIsInstance(worker_errors[0], RuntimeError)
def test_prefetch_stop_has_a_bounded_default_wait(self):
thread = MagicMock()
thread.is_alive.return_value = True
cancel_event = threading.Event()
handle = CheckpointFilePrefetchHandle(
thread=thread,
cancel_event=cancel_event,
succeeded_event=threading.Event(),
errors=[],
)
with self.assertRaisesRegex(TimeoutError, "checkpoint prefetching"):
handle.stop()
self.assertTrue(cancel_event.is_set())
thread.join.assert_called_once_with(60.0)
@patch("torch.distributed.is_initialized", return_value=False)
def test_prefetch_keeps_bounded_pending_window(self, _):
paths = [f"model-{i:05d}.safetensors" for i in range(20)]
@@ -106,9 +153,9 @@ class TestPrefetchCheckpoints(CustomTestCase):
submitted_paths = []
class RecordingExecutor(_InlineExecutor):
def submit(self, fn, path):
def submit(self, fn, path, *args):
submitted_paths.append(path)
return super().submit(fn, path)
return super().submit(fn, path, *args)
def record_pending_size(fs, return_when):
pending_sizes.append(len(fs))
@@ -129,7 +176,7 @@ class TestPrefetchCheckpoints(CustomTestCase):
def test_prefetch_logs_failed_futures(self, _):
paths = ["bad.safetensors"]
def fail_prefetch(path):
def fail_prefetch(path, cancel_event):
raise OSError(f"failed {path}")
with (
@@ -142,15 +189,18 @@ class TestPrefetchCheckpoints(CustomTestCase):
),
patch("sglang.srt.model_loader.weight_utils.logger.warning") as warning,
):
_prefetch_all_checkpoints(paths, num_threads=1)
handle = _prefetch_all_checkpoints(paths, num_threads=1)
handle.wait()
self.assertEqual(handle.errors[0][0], paths[0])
self.assertIsInstance(handle.errors[0][1], OSError)
warning.assert_called_once()
self.assertEqual(
warning.call_args.args[0],
"Failed to prefetch checkpoint file %r.",
"Failed to prefetch checkpoint file %r: %s",
)
self.assertEqual(warning.call_args.args[1], paths[0])
self.assertTrue(warning.call_args.kwargs["exc_info"])
self.assertIsInstance(warning.call_args.args[2], OSError)
@patch("torch.distributed.is_initialized", return_value=False)
def test_prefetch_progress_logs_all_crossed_buckets(self, _):
@@ -193,13 +243,40 @@ class TestPrefetchCheckpoints(CustomTestCase):
),
patch(
"sglang.srt.model_loader.weight_utils._prefetch_checkpoint_file",
side_effect=loaded_paths.append,
side_effect=lambda path, cancel_event: loaded_paths.append(path),
),
):
_prefetch_all_checkpoints(paths, num_threads=2)
self.assertEqual(sorted(loaded_paths), sorted(paths[1::3]))
@patch("torch.distributed.is_initialized", return_value=False)
def test_prefetch_handle_cancels_before_scheduling_next_shard(self, _):
paths = [f"model-{i:05d}.safetensors" for i in range(3)]
started = threading.Event()
release = threading.Event()
loaded_paths = []
def block_first_prefetch(path, cancel_event):
loaded_paths.append(path)
started.set()
self.assertTrue(release.wait(timeout=5))
with patch(
"sglang.srt.model_loader.weight_utils._prefetch_checkpoint_file",
side_effect=block_first_prefetch,
):
handle = _prefetch_all_checkpoints(paths, num_threads=1)
self.assertTrue(started.wait(timeout=5))
with self.assertRaisesRegex(TimeoutError, "checkpoint prefetching"):
handle.wait(timeout=0)
handle.cancel()
release.set()
handle.wait(timeout=5)
self.assertTrue(handle.cancelled)
self.assertEqual(loaded_paths, paths[:1])
@patch("torch.distributed.is_initialized", return_value=False)
def test_buffered_loader_drops_cache_after_each_loaded_shard(self, _):
with tempfile.TemporaryDirectory() as tmpdir:
@@ -319,11 +396,16 @@ class TestPrefetchDispatch(CustomTestCase):
weight_loader_drop_cache_after_load=drop_cache,
)
def _run(self, loader):
def _run(self, loader, **iterator_kwargs):
# _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()))
list(
loader._get_weights_iterator(
self._make_source(),
**iterator_kwargs,
)
)
def _patch_dispatch(self, prefetch, disable_mmap=False, drop_cache=False):
return (
@@ -332,14 +414,6 @@ class TestPrefetchDispatch(CustomTestCase):
"_prepare_weights",
return_value=("/dummy", ["f.safetensors"], True),
),
patch(
"sglang.srt.model_loader.loader.get_server_args",
return_value=self._server_args(
prefetch,
disable_mmap,
drop_cache,
),
),
patch(
"sglang.srt.model_loader.loader.get_model",
return_value=self._server_args(prefetch, disable_mmap, drop_cache),
@@ -360,12 +434,11 @@ class TestPrefetchDispatch(CustomTestCase):
"""Prefetch on + no explicit multithread config -> single-threaded,
and the opt-out warning fires once."""
loader = self._make_loader({})
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
@@ -380,12 +453,11 @@ class TestPrefetchDispatch(CustomTestCase):
"""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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
@@ -401,12 +473,11 @@ class TestPrefetchDispatch(CustomTestCase):
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
@@ -423,12 +494,11 @@ class TestPrefetchDispatch(CustomTestCase):
"""Prefetch off -> multi-threaded iterator is used (default), no
override warning."""
loader = self._make_loader({})
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=False
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
@@ -439,16 +509,117 @@ class TestPrefetchDispatch(CustomTestCase):
mock_single.assert_not_called()
mock_warning.assert_not_called()
def test_startup_prefetch_reuses_existing_background_handle(self):
"""Startup commit reuses resolved shards and the active prefetch handle."""
loader = self._make_loader({})
source = self._make_source()
resolved_source = DefaultModelLoader.ResolvedSource(
source=source,
hf_folder="/dummy",
weight_files=("f.safetensors",),
use_safetensors=True,
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=False
)
with (
p_prep as mock_prepare,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
):
list(
loader._get_weights_iterator(
source,
resolved_source=resolved_source,
startup_prefetch_started=True,
startup_prefetch_active=True,
)
)
mock_prepare.assert_not_called()
mock_single.assert_called_once()
self.assertFalse(mock_single.call_args.kwargs["prefetch"])
mock_buffered.assert_not_called()
mock_warning.assert_called_once()
def test_completed_startup_prefetch_restores_multithread_loader(self):
loader = self._make_loader({})
source = self._make_source()
resolved_source = DefaultModelLoader.ResolvedSource(
source=source,
hf_folder="/dummy",
weight_files=("f.safetensors",),
use_safetensors=True,
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=False
)
with (
p_prep as mock_prepare,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
):
list(
loader._get_weights_iterator(
source,
resolved_source=resolved_source,
startup_prefetch_started=True,
startup_prefetch_active=False,
)
)
mock_prepare.assert_not_called()
mock_buffered.assert_called_once()
self.assertFalse(mock_buffered.call_args.kwargs["prefetch"])
mock_single.assert_not_called()
mock_warning.assert_not_called()
def test_completed_startup_prefetch_is_not_started_twice(self):
loader = self._make_loader({})
source = self._make_source()
resolved_source = DefaultModelLoader.ResolvedSource(
source=source,
hf_folder="/dummy",
weight_files=("f.safetensors",),
use_safetensors=True,
)
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
):
list(
loader._get_weights_iterator(
source,
resolved_source=resolved_source,
startup_prefetch_started=True,
startup_prefetch_active=False,
)
)
mock_buffered.assert_called_once()
self.assertFalse(mock_buffered.call_args.kwargs["prefetch"])
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_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True, disable_mmap=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
@@ -463,7 +634,7 @@ class TestPrefetchDispatch(CustomTestCase):
"""FASTSAFETENSORS ignores both flags; override + warning must not
fire."""
loader = self._make_loader({}, load_format=LoadFormat.FASTSAFETENSORS)
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
@@ -472,7 +643,6 @@ class TestPrefetchDispatch(CustomTestCase):
return_value=iter([]),
) as mock_fast,
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
@@ -492,7 +662,7 @@ class TestPrefetchDispatch(CustomTestCase):
loader = self._make_loader(
{"enable_gds": False}, load_format=LoadFormat.FASTSAFETENSORS
)
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=False,
drop_cache=True,
)
@@ -502,7 +672,6 @@ class TestPrefetchDispatch(CustomTestCase):
return_value=iter([]),
) as mock_fast,
p_prep,
p_args,
p_model,
p_buffered,
p_single,
@@ -93,6 +93,25 @@ class TestServerArgsAnnotatedCli(CustomTestCase):
sa = self._parse(["--image-processor-backend", backend])
self.assertEqual(sa.image_processor_backend, backend)
def test_startup_weight_load_mode(self):
"""The startup loading mode keeps serial as the safe default."""
serial = self._parse([])
overlap = self._parse(["--startup-weight-load-mode", "overlap"])
self.assertEqual(serial.startup_weight_load_mode, "serial")
self.assertFalse(serial.is_startup_weight_load_overlap)
self.assertEqual(overlap.startup_weight_load_mode, "overlap")
self.assertTrue(overlap.is_startup_weight_load_overlap)
with self.assertRaises(SystemExit):
self.parser.parse_args(
[
"--model",
"dummy",
"--startup-weight-load-mode",
"unsupported",
]
)
def test_deprecated_flags_still_work(self):
"""Deprecated flags set the correct dest field."""
sa = self._parse(["--stream-output"])