[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:
co-authored by
Mick Qian
Claude Opus 5
parent
a1b2b976fe
commit
5f9c6b9eb0
+3
-1
@@ -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.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user