From b5d396f08339c00af7625e870585c92b9689f902 Mon Sep 17 00:00:00 2001 From: DAI0818 Date: Wed, 12 Aug 2026 16:43:21 +0800 Subject: [PATCH 1/2] fix(sglang): avoid duplicate model allocation on retry Signed-off-by: DAI0818 --- .../python/modelexpress/adapter.py | 7 +- .../modelexpress/engines/sglang/adapter.py | 46 +++++++++-- .../modelexpress/engines/sglang/loader.py | 4 + .../modelexpress/load_strategy/__init__.py | 2 + .../python/modelexpress/load_strategy/base.py | 25 ++++++ .../load_strategy/rdma_strategy.py | 22 +++-- .../python/tests/test_sglang_loader.py | 82 ++++++++++++++++++- .../python/tests/test_source_selection.py | 61 ++++++++++++-- 8 files changed, 223 insertions(+), 26 deletions(-) diff --git a/modelexpress_client/python/modelexpress/adapter.py b/modelexpress_client/python/modelexpress/adapter.py index 8efa7980..532b326c 100644 --- a/modelexpress_client/python/modelexpress/adapter.py +++ b/modelexpress_client/python/modelexpress/adapter.py @@ -154,7 +154,12 @@ def load_via_native(self, result: LoadResult) -> LoadResult: @gated_capability def reinit_for_retry(self, result: LoadResult) -> LoadResult: - """Replace a possibly-mutated model with a fresh engine model instance.""" + """Restore a possibly-mutated model to freshly initialized state. + + Adapters may return a different model object, or preserve the root + object's identity while replacing its complete internal state when an + engine-owned caller retains the original root reference. + """ ... def get_unique_id(self) -> str: diff --git a/modelexpress_client/python/modelexpress/engines/sglang/adapter.py b/modelexpress_client/python/modelexpress/engines/sglang/adapter.py index 0446df4f..15c392ae 100644 --- a/modelexpress_client/python/modelexpress/engines/sglang/adapter.py +++ b/modelexpress_client/python/modelexpress/engines/sglang/adapter.py @@ -6,6 +6,7 @@ from __future__ import annotations import copy +import gc import logging import uuid from importlib.metadata import version as pkg_version @@ -165,14 +166,33 @@ def reinit_for_retry(self, result: LoadResult) -> LoadResult: ) from sglang.srt.model_loader.utils import set_default_torch_dtype - old_value = result.value + model = result.model + if model is None: + raise RuntimeError("SGLang retry reinitialization requires result.model") + if result.value is not model: + raise RuntimeError( + "SGLang retry reinitialization requires result.value and " + "result.model to reference the same model root" + ) + + publishable = result.publishable + metadata = result.metadata result.value = None result.model = None - del old_value + + # SGLang's RemoteInstanceModelLoader and MxModelLoader both retain the + # root model object while this hook runs. Deleting LoadResult references + # cannot release its parameters, so constructing a replacement directly + # would temporarily allocate two full models. Preserve the externally + # owned root identity, but first turn it into an empty shell so the old + # CUDA allocations can be reclaimed before initialization starts. + model.__dict__.clear() + gc.collect() self.accelerator_backend.empty_cache() logger.info( - "[Worker %s] Re-initializing SGLang model after failed strategy", + "[Worker %s] Re-initializing SGLang model state in-place after " + "failed strategy", self.get_global_rank(), ) quant_config = _get_quantization_config(self.model_config, self.load_config) @@ -180,12 +200,28 @@ def reinit_for_retry(self, result: LoadResult) -> LoadResult: # configured dtype instead of PyTorch's default float32. with set_default_torch_dtype(self.model_config.dtype): with self.target_device: - model = _initialize_model( + fresh_model = _initialize_model( self.model_config, self.load_config, quant_config, ) - return LoadResult(value=model, model=model, publishable=result.publishable) + if type(fresh_model) is not type(model): + raise RuntimeError( + "SGLang retry initialization returned a different model type: " + f"expected {type(model).__qualname__}, " + f"got {type(fresh_model).__qualname__}" + ) + + # Both roots briefly reference the same new children, so there is still + # only one set of parameter storage. The externally owned root remains + # valid after the temporary fresh root is dropped. + model.__dict__.update(fresh_model.__dict__) + del fresh_model + result.value = model + result.model = model + result.publishable = publishable + result.metadata = metadata + return result def _process_weights_after_loading(self, result: LoadResult) -> LoadResult: if result.model is None: diff --git a/modelexpress_client/python/modelexpress/engines/sglang/loader.py b/modelexpress_client/python/modelexpress/engines/sglang/loader.py index 7e58e831..9ff9fd45 100644 --- a/modelexpress_client/python/modelexpress/engines/sglang/loader.py +++ b/modelexpress_client/python/modelexpress/engines/sglang/loader.py @@ -14,6 +14,7 @@ from ... import envs, p2p_pb2 from ...load_strategy import LoadContext, LoadStrategyChain +from ...load_strategy.base import clear_exception_tracebacks from ...load_strategy.context import LoadResult from ...metadata.publisher import PublisherThread from ...metadata.payload import tensor_source_metadata, worker_tensor_descriptors @@ -197,6 +198,9 @@ def _load_model_via_transfer_engine( exc, exc_info=True, ) + registered_tensors = None + tensors = {} + clear_exception_tracebacks(exc) result = ctx.adapter.reinit_for_retry(result) result = ctx.adapter.load_via_native(result) tensors = ctx.adapter.discover_tensors(result) diff --git a/modelexpress_client/python/modelexpress/load_strategy/__init__.py b/modelexpress_client/python/modelexpress/load_strategy/__init__.py index 8b1d92cf..fedfa376 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/__init__.py +++ b/modelexpress_client/python/modelexpress/load_strategy/__init__.py @@ -22,6 +22,7 @@ LoadResult, LoadStrategy, SourceTransferError, + clear_exception_tracebacks, publish_source_if_supported, register_tensors, publish_metadata, @@ -99,6 +100,7 @@ def run(model: nn.Module, ctx: LoadContext) -> nn.Module: ) strategy.rollback(ctx) if e.mutated: + clear_exception_tracebacks(e) result = LoadStrategyChain._reinit_for_retry(result, ctx, strategy) continue except Exception as e: diff --git a/modelexpress_client/python/modelexpress/load_strategy/base.py b/modelexpress_client/python/modelexpress/load_strategy/base.py index bac7d367..22ec17da 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/base.py +++ b/modelexpress_client/python/modelexpress/load_strategy/base.py @@ -6,6 +6,7 @@ from __future__ import annotations import logging +import traceback import uuid from abc import ABC, abstractmethod from typing import TYPE_CHECKING, ClassVar @@ -25,6 +26,30 @@ logger = logging.getLogger("modelexpress.load_strategy") +def clear_exception_tracebacks(exc: BaseException) -> None: + """Drop completed failure frames before releasing a mutated model. + + Transfer failures commonly retain target tensors through traceback frame + locals (for example ``local_tensor`` in the NIXL matching loop). Clearing + only ``LoadResult`` and ``LoadContext`` therefore does not guarantee that + CUDA allocations become unreachable before retry initialization. + """ + pending: list[BaseException] = [exc] + seen: set[int] = set() + while pending: + current = pending.pop() + if id(current) in seen: + continue + seen.add(id(current)) + if current.__cause__ is not None: + pending.append(current.__cause__) + if current.__context__ is not None: + pending.append(current.__context__) + if current.__traceback__ is not None: + traceback.clear_frames(current.__traceback__) + current.__traceback__ = None + + class SourceTransferError(Exception): """Raised when a failure is demonstrably from the remote source side. diff --git a/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py b/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py index 3437c2aa..4c594383 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py +++ b/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py @@ -30,6 +30,7 @@ LoadStrategy, SourceTransferError, _as_load_result, + clear_exception_tracebacks, register_tensors, ) from .context import LoadResult @@ -156,7 +157,6 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: attempts = candidates[:MAX_SOURCE_RETRIES] policy = configured_policy_label() - needs_outer_reinit = False for attempt_index, instance in enumerate(attempts): mx_source_id = instance.mx_source_id worker_id = instance.worker_id @@ -215,8 +215,6 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: "transfer_retry" if has_next_candidate else "transfer_fallback", ) if not has_next_candidate: - if needs_outer_reinit and not e.mutated: - raise StrategyFailed(str(e), mutated=True) from e raise logger.warning( @@ -233,14 +231,24 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: ) from cleanup_error if e.mutated: try: - result = ctx.adapter.reinit_for_retry(result) + clear_exception_tracebacks(e) + reinitialized = ctx.adapter.reinit_for_retry(result) + # LoadResult is the stable envelope shared with the + # outer strategy chain. Some adapters return a new + # envelope, so copy its restored state back rather than + # leaving the outer owner with the cleared pre-retry + # object if all later candidates miss. + if reinitialized is not result: + result.value = reinitialized.value + result.model = reinitialized.model + result.publishable = reinitialized.publishable + result.metadata = reinitialized.metadata except Exception as reinit_error: raise StrategyFailed( f"Failed to reinitialize target after source worker " f"{worker_id} failed: {reinit_error}", mutated=True, ) from reinit_error - needs_outer_reinit = True continue except BaseException: selection_metrics.observe_transfer_seconds( @@ -259,11 +267,9 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: f"[Worker {ctx.global_rank}] Tried {tried} of {len(candidates)} source workers " f"(max retries={MAX_SOURCE_RETRIES}), falling through" ) - # An internal reinit returns a new result, but the outer strategy chain - # still owns the original result that the adapter cleared. raise StrategyFailed( "No RDMA source succeeded", - mutated=needs_outer_reinit, + mutated=False, ) def _find_source_instances( diff --git a/modelexpress_client/python/tests/test_sglang_loader.py b/modelexpress_client/python/tests/test_sglang_loader.py index 6223a9c2..a27f4ae7 100644 --- a/modelexpress_client/python/tests/test_sglang_loader.py +++ b/modelexpress_client/python/tests/test_sglang_loader.py @@ -5,6 +5,7 @@ import os import sys +import weakref from contextlib import contextmanager from types import ModuleType from types import SimpleNamespace @@ -20,6 +21,7 @@ build_sglang_load_context, ) from modelexpress.engines.sglang.loader import MxModelLoader +from modelexpress.load_strategy.context import LoadResult def _load_config(**overrides): @@ -381,6 +383,8 @@ def test_sglang_retry_initializes_model_with_configured_dtype(monkeypatch): loader_mod = ModuleType("sglang.srt.model_loader.loader") model_loader_utils_mod = ModuleType("sglang.srt.model_loader.utils") observed_dtypes = [] + initial_model = nn.Linear(2, 2) + initial_weight_ref = weakref.ref(initial_model.weight) @contextmanager def set_default_torch_dtype(dtype): @@ -394,6 +398,7 @@ def set_default_torch_dtype(dtype): loader_mod._get_quantization_config = lambda *_: None def initialize_model(*_): + assert initial_weight_ref() is None observed_dtypes.append(torch.get_default_dtype()) return nn.Linear(2, 2) @@ -411,16 +416,85 @@ def initialize_model(*_): model_config = _model_config(dtype=torch.bfloat16) adapter = SglangAdapter(_load_config(), model_config, _device_config()) - result = SimpleNamespace( - value=nn.Linear(2, 2), - model=nn.Linear(2, 2), + result = LoadResult( + value=initial_model, + model=initial_model, publishable=True, ) - adapter.reinit_for_retry(result) + retried = adapter.reinit_for_retry(result) assert observed_dtypes == [torch.bfloat16] assert torch.get_default_dtype() == original_dtype + assert retried.value is initial_model + assert retried.model is initial_model + assert list(initial_model.parameters()) + + +def test_sglang_retry_reuses_root_for_native_fallback(monkeypatch): + sglang_mod = ModuleType("sglang") + srt_mod = ModuleType("sglang.srt") + model_loader_mod = ModuleType("sglang.srt.model_loader") + loader_mod = ModuleType("sglang.srt.model_loader.loader") + model_loader_utils_mod = ModuleType("sglang.srt.model_loader.utils") + configs_mod = ModuleType("sglang.srt.configs") + load_config_mod = ModuleType("sglang.srt.configs.load_config") + + @contextmanager + def set_default_torch_dtype(_dtype): + yield + + initial_model = nn.Linear(2, 2) + initial_weight_ref = weakref.ref(initial_model.weight) + native_roots = [] + + loader_mod._get_quantization_config = lambda *_: None + + def initialize_model(*_): + assert initial_weight_ref() is None + return nn.Linear(2, 2) + + class DefaultModelLoader: + def __init__(self, _load_config): + pass + + def _get_all_weights(self, _model_config, model): + native_roots.append(model) + return iter([]) + + @staticmethod + def load_weights_and_postprocess(model, _weights, _target_device): + model.weight.data.fill_(7) + + loader_mod._initialize_model = initialize_model + loader_mod.DefaultModelLoader = DefaultModelLoader + model_loader_utils_mod.set_default_torch_dtype = set_default_torch_dtype + load_config_mod.LoadFormat = SimpleNamespace(AUTO="auto") + monkeypatch.setitem(sys.modules, "sglang", sglang_mod) + monkeypatch.setitem(sys.modules, "sglang.srt", srt_mod) + monkeypatch.setitem(sys.modules, "sglang.srt.model_loader", model_loader_mod) + monkeypatch.setitem(sys.modules, "sglang.srt.model_loader.loader", loader_mod) + monkeypatch.setitem( + sys.modules, + "sglang.srt.model_loader.utils", + model_loader_utils_mod, + ) + monkeypatch.setitem(sys.modules, "sglang.srt.configs", configs_mod) + monkeypatch.setitem( + sys.modules, + "sglang.srt.configs.load_config", + load_config_mod, + ) + + adapter = SglangAdapter(_load_config(), _model_config(), _device_config()) + result = LoadResult(value=initial_model, model=initial_model) + + retried = adapter.reinit_for_retry(result) + loaded = adapter.load_via_native(retried) + + assert loaded.model is initial_model + assert native_roots == [initial_model] + assert torch.all(initial_model.weight == 7) def test_mx_model_loader_delegates_to_shared_strategy_chain(): diff --git a/modelexpress_client/python/tests/test_source_selection.py b/modelexpress_client/python/tests/test_source_selection.py index 8d5583b3..b405ddec 100644 --- a/modelexpress_client/python/tests/test_source_selection.py +++ b/modelexpress_client/python/tests/test_source_selection.py @@ -10,7 +10,9 @@ from __future__ import annotations +import gc import logging +import weakref from types import SimpleNamespace from unittest.mock import MagicMock @@ -18,7 +20,7 @@ from modelexpress import p2p_pb2 from modelexpress.adapter import StrategyFailed -from modelexpress.load_strategy.base import LoadResult +from modelexpress.load_strategy.base import LoadResult, clear_exception_tracebacks from modelexpress.load_strategy.rdma_strategy import MAX_SOURCE_RETRIES, RdmaStrategy from modelexpress.source_selection import ( ENV_SELECTOR, @@ -56,6 +58,27 @@ def _sources(n, worker_rank=0): return [_ref(f"src{i:04x}aaaaaaaaaa", f"w{i}", worker_rank) for i in range(n)] +def test_clear_exception_tracebacks_releases_transfer_frame_locals(): + class Allocation: + pass + + def fail_with_local_allocation(): + allocation = Allocation() + allocation_ref = weakref.ref(allocation) + try: + raise RuntimeError("transfer failed") + except RuntimeError as exc: + return exc, allocation_ref + + exc, allocation_ref = fail_with_local_allocation() + assert allocation_ref() is not None + + clear_exception_tracebacks(exc) + gc.collect() + + assert allocation_ref() is None + + # --------------------------------------------------------------------------- # Registry / config resolution # --------------------------------------------------------------------------- @@ -605,8 +628,30 @@ def test_load_transfer_failure_reinitializes_and_tries_next_source(): assert strat._fetch_worker_metadata.call_count == 2 assert strat._load_as_target.call_count == 2 ctx.adapter.reinit_for_retry.assert_called_once() - assert ctx.adapter.reinit_for_retry.call_args.args[0].value is original_result - assert strat._load_as_target.call_args_list[1].args[0] is retry_result + retry_envelope = strat._load_as_target.call_args_list[1].args[0] + assert ctx.adapter.reinit_for_retry.call_args.args[0] is retry_envelope + assert retry_envelope.value is retry_result.value + assert retry_envelope.model is retry_result.model + + +@pytest.mark.parametrize("vmm_arena", [None, object()]) +def test_load_transfer_failure_supports_identity_preserving_retry(vmm_arena): + strat = RdmaStrategy() + strat._find_source_instances = MagicMock(return_value=_sources(2)) + strat._fetch_worker_metadata = MagicMock(return_value=MagicMock()) + strat._load_as_target = MagicMock( + side_effect=[StrategyFailed("receive failed", mutated=True), "loaded"] + ) + model = MagicMock(name="engine-owned-model-root") + result = LoadResult(value=model, model=model) + ctx = MagicMock(global_rank=0) + ctx.accelerator_backend.name = "" + ctx.vmm_arena = vmm_arena + ctx.adapter.reinit_for_retry.side_effect = lambda current: current + + assert strat.load(result, ctx) == "loaded" + assert strat._load_as_target.call_args_list[1].args[0] is result + assert ctx.vmm_arena is vmm_arena def test_load_clean_transfer_failure_tries_next_source_without_reinit(): @@ -623,7 +668,7 @@ def test_load_clean_transfer_failure_tries_next_source_without_reinit(): ctx.adapter.reinit_for_retry.assert_not_called() -def test_load_requires_outer_reinit_after_reinit_then_metadata_miss(): +def test_load_internal_reinit_is_visible_after_later_metadata_miss(): strat = RdmaStrategy() strat._find_source_instances = MagicMock(return_value=_sources(2)) strat._fetch_worker_metadata = MagicMock(side_effect=[MagicMock(), None]) @@ -646,12 +691,12 @@ def reinit(result): with pytest.raises(StrategyFailed) as exc: strat.load(original_result, ctx) - assert exc.value.mutated is True - assert original_result.model is None + assert exc.value.mutated is False + assert original_result.model is not None ctx.adapter.reinit_for_retry.assert_called_once() -def test_load_requires_outer_reinit_after_reinit_then_clean_failure(): +def test_load_internal_reinit_then_clean_failure_stays_clean(): strat = RdmaStrategy() strat._find_source_instances = MagicMock(return_value=_sources(2)) strat._fetch_worker_metadata = MagicMock(return_value=MagicMock()) @@ -668,7 +713,7 @@ def test_load_requires_outer_reinit_after_reinit_then_clean_failure(): with pytest.raises(StrategyFailed, match="clean failure") as exc: strat.load(MagicMock(), ctx) - assert exc.value.mutated is True + assert exc.value.mutated is False ctx.adapter.reinit_for_retry.assert_called_once() From 7c41dcdf30f7e4db03277ff6e8fad4499df7ef84 Mon Sep 17 00:00:00 2001 From: DAI0818 Date: Wed, 12 Aug 2026 17:04:09 +0800 Subject: [PATCH 2/2] fix(sglang): fail closed when retry recovery fails Signed-off-by: DAI0818 --- .../python/modelexpress/adapter.py | 9 ++++ .../modelexpress/engines/sglang/adapter.py | 37 +++++++++----- .../modelexpress/load_strategy/__init__.py | 8 ++- .../load_strategy/rdma_strategy.py | 10 ++-- .../python/tests/test_sglang_loader.py | 39 +++++++++++++++ .../python/tests/test_source_selection.py | 38 +++++++++++--- .../python/tests/test_vllm_loader.py | 49 ++++++++++++++++++- 7 files changed, 163 insertions(+), 27 deletions(-) diff --git a/modelexpress_client/python/modelexpress/adapter.py b/modelexpress_client/python/modelexpress/adapter.py index 532b326c..7cb76241 100644 --- a/modelexpress_client/python/modelexpress/adapter.py +++ b/modelexpress_client/python/modelexpress/adapter.py @@ -34,6 +34,15 @@ def __init__(self, message: str, *, mutated: bool = False): self.mutated = mutated +class StrategyRecoveryError(RuntimeError): + """Raised when a failed strategy cannot restore a safe model state. + + The strategy chain must stop immediately: trying another loader with a + partially cleared or otherwise unrecoverable model would hide the original + recovery failure and may publish invalid weights. + """ + + def gated_capability(method): """Create an optional adapter method that engines must override to support it. diff --git a/modelexpress_client/python/modelexpress/engines/sglang/adapter.py b/modelexpress_client/python/modelexpress/engines/sglang/adapter.py index 15c392ae..a97883e8 100644 --- a/modelexpress_client/python/modelexpress/engines/sglang/adapter.py +++ b/modelexpress_client/python/modelexpress/engines/sglang/adapter.py @@ -198,19 +198,32 @@ def reinit_for_retry(self, result: LoadResult) -> LoadResult: quant_config = _get_quantization_config(self.model_config, self.load_config) # Match SGLang's initial load path so retry parameters use the model's # configured dtype instead of PyTorch's default float32. - with set_default_torch_dtype(self.model_config.dtype): - with self.target_device: - fresh_model = _initialize_model( - self.model_config, - self.load_config, - quant_config, + try: + with set_default_torch_dtype(self.model_config.dtype): + with self.target_device: + fresh_model = _initialize_model( + self.model_config, + self.load_config, + quant_config, + ) + if type(fresh_model) is not type(model): + raise RuntimeError( + "SGLang retry initialization returned a different model type: " + f"expected {type(model).__qualname__}, " + f"got {type(fresh_model).__qualname__}" ) - if type(fresh_model) is not type(model): - raise RuntimeError( - "SGLang retry initialization returned a different model type: " - f"expected {type(model).__qualname__}, " - f"got {type(fresh_model).__qualname__}" - ) + except BaseException: + # The old parameter graph was intentionally released before fresh + # initialization and cannot be restored without retaining the HBM + # that caused duplicate-model OOM. Restore the envelope to the + # engine-owned empty root so callers do not observe None, then let + # the original failure abort startup rather than attempting another + # strategy with an invalid model. + result.value = model + result.model = model + result.publishable = publishable + result.metadata = metadata + raise # Both roots briefly reference the same new children, so there is still # only one set of parameter storage. The externally owned root remains diff --git a/modelexpress_client/python/modelexpress/load_strategy/__init__.py b/modelexpress_client/python/modelexpress/load_strategy/__init__.py index fedfa376..2e12f51d 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/__init__.py +++ b/modelexpress_client/python/modelexpress/load_strategy/__init__.py @@ -16,7 +16,7 @@ from modelexpress.tracing import tracer -from ..adapter import StrategyFailed, UnsupportedCapability +from ..adapter import StrategyFailed, StrategyRecoveryError, UnsupportedCapability from .base import ( LoadContext, LoadResult, @@ -93,6 +93,12 @@ def run(model: nn.Module, ctx: LoadContext) -> nn.Module: publish_source_if_supported(result, ctx) span.set_attribute("weight_loading_strategy", strategy.name) return result.value + except StrategyRecoveryError: + # Recovery already failed, so no later strategy can safely + # use the current model. Fail closed and retain the original + # recovery error as the exception cause. + strategy.rollback(ctx) + raise except StrategyFailed as e: logger.warning( f"[Worker {ctx.global_rank}] Strategy {strategy.name} failed, " diff --git a/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py b/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py index 4c594383..272757df 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py +++ b/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py @@ -11,7 +11,7 @@ import time from .. import envs, p2p_pb2 -from ..adapter import EngineAdapter, StrategyFailed +from ..adapter import EngineAdapter, StrategyFailed, StrategyRecoveryError from ..metadata.payload import ( accelerators_compatible, worker_tensor_count, @@ -239,15 +239,11 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: # leaving the outer owner with the cleared pre-retry # object if all later candidates miss. if reinitialized is not result: - result.value = reinitialized.value - result.model = reinitialized.model - result.publishable = reinitialized.publishable - result.metadata = reinitialized.metadata + vars(result).update(vars(reinitialized)) except Exception as reinit_error: - raise StrategyFailed( + raise StrategyRecoveryError( f"Failed to reinitialize target after source worker " f"{worker_id} failed: {reinit_error}", - mutated=True, ) from reinit_error continue except BaseException: diff --git a/modelexpress_client/python/tests/test_sglang_loader.py b/modelexpress_client/python/tests/test_sglang_loader.py index a27f4ae7..f6cce2e8 100644 --- a/modelexpress_client/python/tests/test_sglang_loader.py +++ b/modelexpress_client/python/tests/test_sglang_loader.py @@ -431,6 +431,45 @@ def initialize_model(*_): assert list(initial_model.parameters()) +def test_sglang_retry_failure_restores_envelope_and_original_error(monkeypatch): + sglang_mod = ModuleType("sglang") + srt_mod = ModuleType("sglang.srt") + model_loader_mod = ModuleType("sglang.srt.model_loader") + loader_mod = ModuleType("sglang.srt.model_loader.loader") + model_loader_utils_mod = ModuleType("sglang.srt.model_loader.utils") + + @contextmanager + def set_default_torch_dtype(_dtype): + yield + + failure = RuntimeError("fresh initialization failed") + loader_mod._get_quantization_config = lambda *_: None + loader_mod._initialize_model = MagicMock(side_effect=failure) + model_loader_utils_mod.set_default_torch_dtype = set_default_torch_dtype + monkeypatch.setitem(sys.modules, "sglang", sglang_mod) + monkeypatch.setitem(sys.modules, "sglang.srt", srt_mod) + monkeypatch.setitem(sys.modules, "sglang.srt.model_loader", model_loader_mod) + monkeypatch.setitem(sys.modules, "sglang.srt.model_loader.loader", loader_mod) + monkeypatch.setitem( + sys.modules, + "sglang.srt.model_loader.utils", + model_loader_utils_mod, + ) + + model = nn.Linear(2, 2) + adapter = SglangAdapter(_load_config(), _model_config(), _device_config()) + result = LoadResult(value=model, model=model, metadata={"attempt": 1}) + + with pytest.raises(RuntimeError) as exc: + adapter.reinit_for_retry(result) + + assert exc.value is failure + assert result.value is model + assert result.model is model + assert result.metadata == {"attempt": 1} + assert model.__dict__ == {} + + def test_sglang_retry_reuses_root_for_native_fallback(monkeypatch): sglang_mod = ModuleType("sglang") srt_mod = ModuleType("sglang.srt") diff --git a/modelexpress_client/python/tests/test_source_selection.py b/modelexpress_client/python/tests/test_source_selection.py index b405ddec..da5003eb 100644 --- a/modelexpress_client/python/tests/test_source_selection.py +++ b/modelexpress_client/python/tests/test_source_selection.py @@ -19,7 +19,7 @@ import pytest from modelexpress import p2p_pb2 -from modelexpress.adapter import StrategyFailed +from modelexpress.adapter import StrategyFailed, StrategyRecoveryError from modelexpress.load_strategy.base import LoadResult, clear_exception_tracebacks from modelexpress.load_strategy.rdma_strategy import MAX_SOURCE_RETRIES, RdmaStrategy from modelexpress.source_selection import ( @@ -611,8 +611,15 @@ def test_load_transfer_failure_reinitializes_and_tries_next_source(): strat._load_as_target = MagicMock( side_effect=[StrategyFailed("receive failed", mutated=True), "loaded"] ) - original_result = MagicMock(name="original-result") - retry_result = MagicMock(name="retry-result") + original_model = MagicMock(name="original-model") + retry_model = MagicMock(name="retry-model") + original_result = LoadResult(value=original_model, model=original_model) + retry_result = LoadResult( + value=retry_model, + model=retry_model, + publishable=False, + metadata={"retry": True}, + ) ctx = MagicMock(global_rank=0) ctx.accelerator_backend.name = "" # unknown target -> accelerator gate accepts ctx.adapter.reinit_for_retry.return_value = retry_result @@ -632,6 +639,8 @@ def test_load_transfer_failure_reinitializes_and_tries_next_source(): assert ctx.adapter.reinit_for_retry.call_args.args[0] is retry_envelope assert retry_envelope.value is retry_result.value assert retry_envelope.model is retry_result.model + assert retry_envelope.publishable is False + assert retry_envelope.metadata == {"retry": True} @pytest.mark.parametrize("vmm_arena", [None, object()]) @@ -717,7 +726,24 @@ def test_load_internal_reinit_then_clean_failure_stays_clean(): ctx.adapter.reinit_for_retry.assert_called_once() -def test_load_reinit_failure_remains_mutated(): +def test_load_last_candidate_mutated_failure_propagates_mutated(): + strat = RdmaStrategy() + strat._find_source_instances = MagicMock(return_value=_sources(1)) + strat._fetch_worker_metadata = MagicMock(return_value=MagicMock()) + strat._load_as_target = MagicMock( + side_effect=StrategyFailed("mutated failure", mutated=True) + ) + ctx = MagicMock(global_rank=0) + ctx.accelerator_backend.name = "" + + with pytest.raises(StrategyFailed, match="mutated failure") as exc: + strat.load(MagicMock(), ctx) + + assert exc.value.mutated is True + ctx.adapter.reinit_for_retry.assert_not_called() + + +def test_load_reinit_failure_is_unrecoverable(): strat = RdmaStrategy() strat._find_source_instances = MagicMock(return_value=_sources(2)) strat._fetch_worker_metadata = MagicMock(return_value=MagicMock()) @@ -728,10 +754,10 @@ def test_load_reinit_failure_remains_mutated(): ctx.accelerator_backend.name = "" ctx.adapter.reinit_for_retry.side_effect = RuntimeError("reinit failed") - with pytest.raises(StrategyFailed, match="reinit failed") as exc: + with pytest.raises(StrategyRecoveryError, match="reinit failed") as exc: strat.load(MagicMock(), ctx) - assert exc.value.mutated is True + assert isinstance(exc.value.__cause__, RuntimeError) def test_load_cleanup_failure_aborts_before_reinit(): diff --git a/modelexpress_client/python/tests/test_vllm_loader.py b/modelexpress_client/python/tests/test_vllm_loader.py index 573cf7c3..ecd9a970 100644 --- a/modelexpress_client/python/tests/test_vllm_loader.py +++ b/modelexpress_client/python/tests/test_vllm_loader.py @@ -14,7 +14,7 @@ import torch.nn as nn from modelexpress import p2p_pb2 -from modelexpress.adapter import EngineAdapter, StrategyFailed +from modelexpress.adapter import EngineAdapter, StrategyFailed, StrategyRecoveryError from modelexpress.load_strategy.context import LoadResult from modelexpress.nixl_transfer import NixlTransferManager @@ -896,6 +896,53 @@ def fallback_load(self_or_model, *args, **kwargs): assert call_order == ["failed", "rollback", "fallback"] ctx.adapter.reinit_for_retry.assert_not_called() + def test_strategy_recovery_error_aborts_without_fallback(self): + from modelexpress.load_strategy import LoadStrategyChain + + call_order = [] + + def failed_recovery(self_or_result, *_args, **_kwargs): + call_order.append("failed") + raise StrategyRecoveryError("model recovery failed") + + def rollback(self_or_ctx, *_args, **_kwargs): + call_order.append("rollback") + + def fallback_load(self_or_result, *_args, **_kwargs): + call_order.append("fallback") + return self_or_result + + ctx = _make_load_context() + with patch( + "modelexpress.load_strategy.rdma_strategy.RdmaStrategy.is_available", + return_value=False, + ), patch( + "modelexpress.load_strategy.model_streamer_strategy." + "ModelStreamerStrategy.is_available", + return_value=True, + ), patch( + "modelexpress.load_strategy.model_streamer_strategy." + "ModelStreamerStrategy.load", + failed_recovery, + ), patch( + "modelexpress.load_strategy.model_streamer_strategy." + "ModelStreamerStrategy.rollback", + rollback, + ), patch( + "modelexpress.load_strategy.gds_strategy.GdsStrategy.is_available", + return_value=False, + ), patch( + "modelexpress.load_strategy.default_strategy.DefaultStrategy.is_available", + return_value=True, + ), patch( + "modelexpress.load_strategy.default_strategy.DefaultStrategy.load", + fallback_load, + ): + with pytest.raises(StrategyRecoveryError, match="model recovery failed"): + LoadStrategyChain.run(MagicMock(), ctx) + + assert call_order == ["failed", "rollback"] + def test_strategy_failed_runs_rollback_and_reinit_when_mutated(self): from modelexpress.load_strategy import LoadStrategyChain