[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:
co-authored by
Wenhui Zhu
Alex Nails
parent
0772e79ee7
commit
6b94d39f13
@@ -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.
|
||||
|
||||
+732
@@ -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"])
|
||||
|
||||
Reference in New Issue
Block a user