[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:
@@ -55,7 +55,10 @@ class LayerwiseOffloadManager:
|
|||||||
# Armed on the first denoise forward, so that the load-time prefetch below
|
# 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.
|
# does not pin the whole resident set before the DiT is the active component.
|
||||||
self._residency_active = False
|
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())
|
self.enabled = bool(enabled and torch.get_device_module().is_available())
|
||||||
if not self.enabled:
|
if not self.enabled:
|
||||||
return
|
return
|
||||||
@@ -254,6 +257,7 @@ class LayerwiseOffloadManager:
|
|||||||
self.prepare_for_next_req(non_blocking=False)
|
self.prepare_for_next_req(non_blocking=False)
|
||||||
|
|
||||||
self.register_forward_hooks()
|
self.register_forward_hooks()
|
||||||
|
self._configured = True
|
||||||
logger.info(
|
logger.info(
|
||||||
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, num resident layers: {self.resident_layers}, total num layers: {self.num_layers}"
|
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)
|
manager.prepare_for_next_req(non_blocking=True)
|
||||||
|
|
||||||
def disable_offload(self) -> None:
|
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:
|
if self.layerwise_offload_managers is None:
|
||||||
return
|
return
|
||||||
for manager in self.layerwise_offload_managers:
|
for manager in self.layerwise_offload_managers:
|
||||||
if manager.enabled:
|
if manager.enabled:
|
||||||
manager.remove_forward_hooks()
|
manager.remove_forward_hooks()
|
||||||
manager.load_all_layers()
|
manager.load_all_layers()
|
||||||
|
manager.enabled = False
|
||||||
|
|
||||||
def enable_offload(self) -> None:
|
def enable_offload(self) -> None:
|
||||||
"""Re-enable layerwise offload: sync weights to CPU, release layers, and restore hooks."""
|
"""Re-enable layerwise offload: sync weights to CPU, release layers, and restore hooks."""
|
||||||
if self.layerwise_offload_managers is None:
|
if self.layerwise_offload_managers is None:
|
||||||
return
|
return
|
||||||
for manager in self.layerwise_offload_managers:
|
for manager in self.layerwise_offload_managers:
|
||||||
if manager.enabled:
|
if manager._configured:
|
||||||
|
manager.enabled = True
|
||||||
manager.sync_all_layers_to_cpu()
|
manager.sync_all_layers_to_cpu()
|
||||||
manager.release_all()
|
manager.release_all()
|
||||||
manager.register_forward_hooks()
|
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))
|
comp.configure_layerwise_offload(_server_args(dit_layerwise_resident_layers=0.5))
|
||||||
# 0.5 * 8 = 4 leading layers resident
|
# 0.5 * 8 = 4 leading layers resident
|
||||||
assert comp.layerwise_offload_managers[0].resident_layers == 4
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user