diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py index 372a3934c..8b7734068 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py @@ -55,7 +55,10 @@ class LayerwiseOffloadManager: # Armed on the first denoise forward, so that the load-time prefetch below # does not pin the whole resident set before the DiT is the active component. self._residency_active = False - + # True once _initialize builds the CPU buffers; unlike `enabled` it + # never flips back, so disable_offload/enable_offload can toggle + # `enabled` without losing track of which managers can be re-armed. + self._configured = False self.enabled = bool(enabled and torch.get_device_module().is_available()) if not self.enabled: return @@ -254,6 +257,7 @@ class LayerwiseOffloadManager: self.prepare_for_next_req(non_blocking=False) self.register_forward_hooks() + self._configured = True logger.info( f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, num resident layers: {self.resident_layers}, total num layers: {self.num_layers}" ) @@ -663,20 +667,31 @@ class LayerwiseOffloadableModuleMixin: manager.prepare_for_next_req(non_blocking=True) def disable_offload(self) -> None: - """Disable layerwise offload: load all layers to GPU and remove hooks.""" + """Disable layerwise offload: load all layers to GPU and remove hooks. + + Also flips `manager.enabled` off so every layerwise path — + is_layerwise_offloaded_module(), release_all(), prepare_for_next_req() + — short-circuits until enable_offload() re-arms it. Without this, a + residency strategy built while the module was offloaded (e.g. the + temporary offload_during_compile window) keeps calling release_all() + on use-site switches after the hooks are gone, replacing restored + weights with (1,) placeholders that nothing swaps back in. + """ if self.layerwise_offload_managers is None: return for manager in self.layerwise_offload_managers: if manager.enabled: manager.remove_forward_hooks() manager.load_all_layers() + manager.enabled = False def enable_offload(self) -> None: """Re-enable layerwise offload: sync weights to CPU, release layers, and restore hooks.""" if self.layerwise_offload_managers is None: return for manager in self.layerwise_offload_managers: - if manager.enabled: + if manager._configured: + manager.enabled = True manager.sync_all_layers_to_cpu() manager.release_all() manager.register_forward_hooks() diff --git a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py index 0053793f0..d73487817 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -689,3 +689,71 @@ def test_configure_resolves_resident_layers_ratio(monkeypatch): comp.configure_layerwise_offload(_server_args(dit_layerwise_resident_layers=0.5)) # 0.5 * 8 = 4 leading layers resident assert comp.layerwise_offload_managers[0].resident_layers == 4 + + +class _MixinBlock(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.weight = torch.nn.Parameter( + torch.arange(9, dtype=torch.float32).reshape(3, 3) + ) + self.bias = torch.nn.Parameter(torch.arange(3, dtype=torch.float32)) + + +class _MixinModel(torch.nn.Module, LayerwiseOffloadableModuleMixin): + layer_names = ["blocks"] + + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([_MixinBlock() for _ in range(3)]) + + +def _configure_mixin_model(monkeypatch) -> _MixinModel: + _patch_fake_device(monkeypatch) + model = _MixinModel() + model.configure_layerwise_offload(_server_args()) + assert is_layerwise_offloaded_module(model) + return model + + +def test_disable_offload_short_circuits_residency_release(monkeypatch): + """disable_offload() must make later layerwise calls no-ops. + + Regression test: a ComponentResidencyManager strategy built while the + module was offloaded (the offload_during_compile window) keeps calling + release_all() on use-site switches. After disable_offload() removed the + hooks, those releases replaced restored weights with (1,) placeholders + that nothing ever swapped back in, crashing dual-DiT models (Wan2.2-A14B + boundary experts, Ideogram-4 paired towers) on the first real request. + """ + model = _configure_mixin_model(monkeypatch) + model.disable_offload() + + assert not is_layerwise_offloaded_module(model) + for name, param in model.named_parameters(): + assert tuple(param.shape) != (1,), name + + # The exact call path the residency strategy takes on use-site switches. + LayerwiseOffloadStrategy().exit(model) + model.prepare_for_next_req() + for name, param in model.named_parameters(): + assert tuple(param.shape) != (1,), name + + +def test_enable_offload_rearms_after_disable(monkeypatch): + model = _configure_mixin_model(monkeypatch) + # blocks[2] holds a placeholder right after configure; the real values are + # what _MixinBlock was constructed with. + original = torch.arange(9, dtype=torch.float32).reshape(3, 3) + + model.disable_offload() + assert not is_layerwise_offloaded_module(model) + + model.enable_offload() + assert is_layerwise_offloaded_module(model) + + manager = model.layerwise_offload_managers[0] + manager.release_layer(2) + assert tuple(model.blocks[2].weight.shape) == (1,) + manager.prefetch_layer(2, non_blocking=False) + assert torch.equal(model.blocks[2].weight.data, original)