[diffusion] feat: release a layerwise component's non-layer weights between uses (#35734)

This commit is contained in:
Mick
2026-08-23 12:05:24 +08:00
committed by GitHub
parent 70319a0881
commit 97ae27893a
3 changed files with 232 additions and 0 deletions
@@ -210,6 +210,7 @@ class LayerwiseOffloadStrategy(ComponentResidencyStrategy):
return
_module_to_local_device(module, dtype=use.target_dtype)
return
module.restore_non_layer_weights()
module.prepare_for_next_req()
def finish_use(
@@ -222,6 +223,10 @@ class LayerwiseOffloadStrategy(ComponentResidencyStrategy):
return
for manager in module.layerwise_offload_managers:
manager.release_all()
# 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
# one that needs the room.
module.park_non_layer_weights()
if current_platform.is_mps():
torch.mps.synchronize()
module.restore_mps_cpu_non_layer_weights()
@@ -97,6 +97,13 @@ def compute_streamed_layers(
# of table size to rows actually read makes residency clearly wasteful.
HOST_RESIDENT_TABLE_MIN_BYTES = 256 * 1024**2
# Parking a component's non-layer weights frees device memory at the cost of two
# transfers per use and a host copy that competes with the page cache. It is
# worth that only when what it frees is a meaningful share of the headroom
# actually available; on a card with room it is pure loss. Below this share of
# free device memory, the component stays where it is.
PARK_SIGNIFICANCE = 0.1
def _resolve_submodule(root: torch.nn.Module, path: str) -> torch.nn.Module | None:
current: Any = root
@@ -1430,6 +1437,106 @@ class LayerwiseOffloadableModuleMixin:
host_resident_table_names: List[str] = []
layerwise_offload_managers: list[LayerwiseOffloadManager] = []
# Whether to park non-layer parameters on the host between uses. Costs a
# transfer per request and is worth it only when device memory is the
# binding constraint, so it follows --performance-mode memory.
park_non_layer_weights_between_uses: bool = False
def _managed_layer_parameter_names(self) -> set:
"""Parameter names some layerwise manager already streams."""
return {
name
for manager in self.layerwise_offload_managers
for names in manager._weight_metadata.values()
for name in names
}
def park_non_layer_weights(self) -> None:
"""Move the parameters no manager streams back to the host.
A layerwise component holds its non-layer parameters on the device for
the whole request. That is right while it is the component being used
and pure cost afterwards. Measured on H3 at 864x480 / 124 frames: the
DiT keeps 2.09 GB and the text encoder 1.40 GB through a VAE decode
that touches neither, and the decode is exactly where the budget runs
out -- with the VAE's blocks held resident it needs 11.86 GiB against a
12 GiB card, and fails for want of 20 MiB.
Buffers are left where they are. Layerwise offload keeps them resident
on purpose, because a shared buffer such as a RoPE cache is referenced
by many layers.
"""
if not self.park_non_layer_weights_between_uses:
return
if current_platform.is_mps():
# MPS parks its own non-layer weights, scoped to subphases
return
managed = self._managed_layer_parameter_names()
resident = [
(name, parameter)
for name, parameter in self.named_parameters()
if name not in managed and parameter.device.type != "cpu"
]
holds = sum(p.numel() * p.element_size() for _, p in resident)
if holds <= self._device_headroom_bytes() * PARK_SIGNIFICANCE:
# There is room. Give back any host copies rather than hold them.
self._parked_non_layer_weights.clear()
return
parked = self._parked_non_layer_weights
with torch.inference_mode(False), torch.no_grad():
for name, parameter in resident:
if name not in parked:
parked[name] = parameter.detach().to("cpu", copy=True)
parameter.data = self._park_placeholder(parameter)
def _device_headroom_bytes(self) -> int:
"""What an allocation could get without the allocator growing its pool.
`get_available_gpu_memory` reports driver-level free memory, which
excludes blocks the caching allocator has already reserved and not
handed out. On a warm process that undercounts the real headroom badly,
so the allocator's own unused reserve is added back.
"""
free = int(
current_platform.get_available_gpu_memory(empty_cache=False) * (1 << 30)
)
device_module = torch.get_device_module()
unused_reserve = (
device_module.memory_reserved() - device_module.memory_allocated()
)
return free + max(0, unused_reserve)
def _park_placeholder(self, parameter: torch.Tensor) -> torch.Tensor:
"""One shared stand-in per (device, dtype), not one per parked weight."""
key = (parameter.device, parameter.dtype)
placeholder = self._park_placeholders.get(key)
if placeholder is None:
placeholder = torch.empty(
(1,), dtype=parameter.dtype, device=parameter.device
)
self._park_placeholders[key] = placeholder
return placeholder
def restore_non_layer_weights(self) -> None:
"""Bring parked parameters back before this component is used again."""
parked = self._parked_non_layer_weights
if not parked:
return
device = current_platform.get_local_torch_device()
parameters = dict(self.named_parameters())
with torch.inference_mode(False), torch.no_grad():
for name, host_tensor in parked.items():
parameter = parameters.get(name)
if parameter is None:
continue
# The parked copy is pageable, so this transfer stages through
# the driver's own pinned buffer and is synchronous whatever is
# asked for. Pinning it instead would make the copy async, at
# the price of host memory the kernel can never reclaim -- the
# wrong trade on the hosts this path exists for.
parameter.data = host_tensor.to(device)
def _capture_mps_cpu_non_layer_weights(self) -> None:
managed_names = {
name
@@ -1522,6 +1629,22 @@ class LayerwiseOffloadableModuleMixin:
for name, tensor in self._mps_cpu_buffers.items():
buffers[name].data = tensor
@property
def _parked_non_layer_weights(self) -> dict:
store = self.__dict__.get("_parked_non_layer_weight_store")
if store is None:
store = {}
self.__dict__["_parked_non_layer_weight_store"] = store
return store
@property
def _park_placeholders(self) -> dict:
store = self.__dict__.get("_park_placeholder_store")
if store is None:
store = {}
self.__dict__["_park_placeholder_store"] = store
return store
def configure_layerwise_offload(
self,
server_args: ServerArgs,
@@ -1529,6 +1652,9 @@ class LayerwiseOffloadableModuleMixin:
pin_budget: HostPinBudget | None = None,
component_name: str | None = None,
):
self.park_non_layer_weights_between_uses = (
server_args.performance_mode == "memory"
)
self.layerwise_offload_managers = []
named_modules = dict(self.named_modules())
configured_layer_names = []
@@ -205,6 +205,7 @@ def _server_args(**kwargs):
defaults = dict(
component_residency=None,
disagg_role=RoleType.MONOLITHIC,
performance_mode="auto",
_required_resident_components=set(),
_component_layerwise_capabilities={},
_explicit_arg_names=set(),
@@ -797,6 +798,12 @@ class _ResidentComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin):
self.blocks = torch.nn.ModuleList([_DummyBlock() for _ in range(n)])
class _ParkableResidentComponent(_ResidentComponent):
def __init__(self, n: int) -> None:
super().__init__(n)
self.non_layer = torch.nn.Parameter(torch.ones(2))
class _AuxiliaryResidentComponent(_ResidentComponent):
layerwise_offload_dit_group_enabled = False
@@ -1684,3 +1691,97 @@ def test_layerwise_tuning_accepts_json_and_pair_forms():
assert pair.layerwise_tuning_for("text_encoder", dit_group=False)[1] == 2.0
as_json = _server_args(layerwise_resident_layers='{"vae": 6}')
assert as_json.layerwise_tuning_for("vae", dit_group=False)[1] == 6.0
def test_non_layer_parking_follows_memory_performance_mode(monkeypatch):
"""The extra transfer per request only pays for itself under memory mode."""
_patch_fake_device(monkeypatch)
tight = _ResidentComponent(4)
tight.configure_layerwise_offload(_server_args(performance_mode="memory"))
assert tight.park_non_layer_weights_between_uses
relaxed = _ResidentComponent(4)
relaxed.configure_layerwise_offload(_server_args(performance_mode="speed"))
assert not relaxed.park_non_layer_weights_between_uses
def test_parking_leaves_streamed_layer_weights_alone(monkeypatch):
"""Only the parameters no manager streams are moved to the host."""
comp = _ParkableResidentComponent(4)
comp.configure_layerwise_offload(_server_args(performance_mode="memory"))
_headroom(monkeypatch, 0)
managed = comp._managed_layer_parameter_names()
assert managed, "the managers should own the block parameters"
comp.park_non_layer_weights()
parked = comp._parked_non_layer_weights
assert not (set(parked) & managed), "a streamed layer weight was parked"
for name, host_tensor in parked.items():
assert host_tensor.device.type == "cpu", name
comp.restore_non_layer_weights()
restored = dict(comp.named_parameters())
for name, host_tensor in parked.items():
assert restored[name].shape == host_tensor.shape
def test_parking_is_a_no_op_outside_memory_mode(monkeypatch):
comp = _ParkableResidentComponent(4)
comp.configure_layerwise_offload(_server_args(performance_mode="speed"))
comp.park_non_layer_weights()
assert not comp._parked_non_layer_weights
def _headroom(monkeypatch, gib):
monkeypatch.setattr(
layerwise_offload_mod.current_platform,
"get_available_gpu_memory",
lambda **_: float(gib),
)
module = layerwise_offload_mod.torch.get_device_module()
monkeypatch.setattr(module, "memory_reserved", lambda *_: 0, raising=False)
monkeypatch.setattr(module, "memory_allocated", lambda *_: 0, raising=False)
def test_parking_is_skipped_when_the_card_has_room(monkeypatch):
"""A component holding a sliver of a large headroom is left alone."""
comp = _ParkableResidentComponent(4)
comp.configure_layerwise_offload(_server_args(performance_mode="memory"))
_headroom(monkeypatch, 400)
comp.park_non_layer_weights()
assert not comp._parked_non_layer_weights
def test_parking_happens_when_the_headroom_is_small(monkeypatch):
comp = _ParkableResidentComponent(4)
comp.configure_layerwise_offload(_server_args(performance_mode="memory"))
_headroom(monkeypatch, 0)
comp.park_non_layer_weights()
assert comp._parked_non_layer_weights
def test_host_copies_are_given_back_when_room_appears(monkeypatch):
"""Skipping must not leave host memory held for a park that will not happen."""
comp = _ParkableResidentComponent(4)
comp.configure_layerwise_offload(_server_args(performance_mode="memory"))
_headroom(monkeypatch, 0)
comp.park_non_layer_weights()
assert comp._parked_non_layer_weights
comp.restore_non_layer_weights()
_headroom(monkeypatch, 400)
comp.park_non_layer_weights()
assert not comp._parked_non_layer_weights, "host copies should be released"
def test_park_placeholders_are_shared(monkeypatch):
"""One stand-in per (device, dtype), not one allocation per parked weight."""
comp = _ParkableResidentComponent(4)
comp.configure_layerwise_offload(_server_args(performance_mode="memory"))
_headroom(monkeypatch, 0)
comp.park_non_layer_weights()
managed = comp._managed_layer_parameter_names()
stand_ins = {
id(p) for n, p in comp.named_parameters() if n not in managed and p.numel() == 1
}
assert len(stand_ins) <= len(comp._park_placeholders)