[diffusion] fix: separate a use-scoped layerwise release from release_all (#40590)

Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-09-22 10:38:44 +08:00
committed by GitHub
co-authored by Mick Qian Claude Opus 5
parent a1b2b976fe
commit 5f9c6b9eb0
3 changed files with 106 additions and 6 deletions
@@ -279,7 +279,9 @@ class LayerwiseOffloadStrategy(ComponentResidencyStrategy):
if not isinstance(module, LayerwiseOffloadableModuleMixin): if not isinstance(module, LayerwiseOffloadableModuleMixin):
return return
for manager in module.layerwise_offload_managers: for manager in module.layerwise_offload_managers:
manager.release_all() # Not release_all: this is a use ending, not a reset. The default
# still drops the resident set, so behaviour is unchanged here.
manager.release_after_use()
# The layers are gone; the rest of this component is dead weight on the # The layers are gone; the rest of this component is dead weight on the
# device until it is used again, and the stage that follows may be the # device until it is used again, and the stage that follows may be the
# one that needs the room. # one that needs the room.
@@ -2121,10 +2121,30 @@ class LayerwiseOffloadManager:
torch.mps.synchronize() torch.mps.synchronize()
torch.mps.empty_cache() torch.mps.empty_cache()
@torch.compiler.disable
def release_after_use(self, *, keep_resident: bool = False) -> None:
"""This component's use has ended; release what that use was streaming.
Distinct from `release_all`, which is the literal operation and stays
that way for a full reset. A use ending asks a narrower question: the
streamed window is certainly dead, but the resident set only is if
nothing will want it before something else needs the room.
The two were the same call, and that is why `resident_layers` does
nothing for any component whose use is a single forward pass rather
than a denoise loop -- the set is prefetched at the start of the use
and dropped at the end of it, every request. `keep_resident` is how a
caller that knows the memory picture says otherwise; it defaults to the
long-standing behaviour, so nothing moves until someone asks.
"""
self._release_layers(drop_resident=not keep_resident)
@torch.compiler.disable @torch.compiler.disable
def release_all(self) -> None: def release_all(self) -> None:
"""Release every layer, including the resident ones: this ends the """Release every layer, resident ones included. A full reset."""
denoise stage that the resident set is scoped to.""" self._release_layers(drop_resident=True)
def _release_layers(self, *, drop_resident: bool) -> None:
self._log_direct_read_summary() self._log_direct_read_summary()
self._log_debug_timing() self._log_debug_timing()
if self._mapped_populator is not None: if self._mapped_populator is not None:
@@ -2140,10 +2160,13 @@ class LayerwiseOffloadManager:
self._collect_mapped_layer(layer_idx) self._collect_mapped_layer(layer_idx)
for layer_idx in list(self._gpu_layers): for layer_idx in list(self._gpu_layers):
self.release_layer(layer_idx, force=True) # `force` is what overrides release_layer's own skip of the resident
# set, so not forcing is all it takes to leave that set alone.
self.release_layer(layer_idx, force=drop_resident)
# The next use starts a new request; its first pass over the layers may # The next use starts a new request; its first pass over the layers may
# find their pages evicted and is the one worth faulting in sequentially. # find their pages evicted and is the one worth faulting in
self._first_pass = True # sequentially. Layers still on the device were never evicted.
self._first_pass = drop_resident
@torch.compiler.disable @torch.compiler.disable
def load_all_layers(self) -> None: def load_all_layers(self) -> None:
@@ -1058,6 +1058,81 @@ def test_prepare_for_next_req_repins_residents(monkeypatch):
assert {0, 1, 2} <= manager._gpu_layers assert {0, 1, 2} <= manager._gpu_layers
def test_release_after_use_defaults_to_the_old_release_all(monkeypatch):
"""The rename must not move anything: `release_after_use()` == the previous call.
`finish_use` used to call `release_all()` unconditionally. It now says
`release_after_use()`, and with the default argument that has to clear exactly the
same layers, or this refactor is a behaviour change wearing a new name.
"""
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(6), num_layers=6, prefetch_size=1, resident_layers=3
)
_arm_residency(manager)
manager.prepare_for_next_req(non_blocking=False)
assert manager._gpu_layers
manager.release_after_use()
assert not manager._gpu_layers
assert manager._first_pass is True
def test_release_all_still_drops_everything(monkeypatch):
"""`release_all` keeps its literal contract for the full-reset callers.
`enable_offload` syncs to CPU and expects nothing left on the device; it
must not inherit the resident-set exemption.
"""
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(6), num_layers=6, prefetch_size=1, resident_layers=3
)
_arm_residency(manager)
manager.prepare_for_next_req(non_blocking=False)
manager.release_all()
assert not manager._gpu_layers
assert manager._first_pass is True
def test_release_after_use_can_keep_the_resident_set(monkeypatch):
"""`keep_resident` is the whole point of naming the two calls apart.
A component whose use is one forward pass has its resident set prefetched
at the start of the use and dropped at the end, so `resident_layers` buys
it nothing. Measured on Qwen-Image-2.1 / RTX 5090:
`--layerwise-resident-layers text_encoder=0.8` logs `resident=53/66` and
moves neither memory nor latency.
"""
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(6), num_layers=6, prefetch_size=1, resident_layers=3
)
_arm_residency(manager)
manager.prepare_for_next_req(non_blocking=False)
manager.release_after_use(keep_resident=True)
assert set(manager._gpu_layers) == set(manager._retained_set)
# Those layers never left the device, so the next use must not re-do the
# sequential first pass that exists for evicted pages.
assert manager._first_pass is False
def test_release_after_use_keeps_nothing_when_no_residents_are_configured(monkeypatch):
"""`keep_resident` with an empty resident set is still a full release."""
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(6), num_layers=6, prefetch_size=1, resident_layers=0
)
manager.prefetch_layer(0, non_blocking=False)
manager.prefetch_layer(1, non_blocking=False)
assert manager._gpu_layers
manager.release_after_use(keep_resident=True)
assert not manager._gpu_layers
def _record_prepare(manager, monkeypatch): def _record_prepare(manager, monkeypatch):
"""Log the order of prefetches and stream waits inside prepare_for_next_req. """Log the order of prefetches and stream waits inside prepare_for_next_req.