[diffusion] fix: fix dual-DiT models crash with (1,)-placeholder weights after compile-time offload (#32743)

Co-authored-by: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-07-29 20:44:13 +08:00
committed by GitHub
co-authored by Claude Sonnet 5
parent ca6e0ff2c8
commit 67c2258906
2 changed files with 86 additions and 3 deletions
@@ -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()
@@ -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)