[diffusion] fix: fix multi-group layerwise offload startup memory (#35509)

This commit is contained in:
Mick
2026-08-19 20:32:31 +08:00
committed by GitHub
parent c57ada81e1
commit 29f5d1c7c3
2 changed files with 104 additions and 10 deletions
@@ -205,6 +205,16 @@ class LayerwiseOffloadManager:
except Exception:
return None
def _managed_parameter_bytes(self) -> int:
total_bytes = 0
for name, tensor in self.model.named_parameters():
layer_idx = self._match_layer_idx(name)
if layer_idx is None or layer_idx >= self.num_layers:
continue
local_tensor = self._to_local_tensor(tensor)
total_bytes += local_tensor.numel() * local_tensor.element_size()
return total_bytes
def _get_shared_empty_tensor(self, dtype: torch.dtype) -> torch.Tensor:
placeholder = self._offload_placeholders.get(dtype)
if placeholder is None:
@@ -252,13 +262,26 @@ class LayerwiseOffloadManager:
if not self.enabled:
return
self._named_parameters = dict(self.model.named_parameters())
self._named_buffers = dict(self.model.named_buffers())
if self._synchronous_mps:
self._named_parameters = dict(self.model.named_parameters())
self._named_buffers = dict(self.model.named_buffers())
self._initialize_mps_cpu_weights()
return
self._initialize_layer_weights()
# Keep non-layer parameters resident on GPU. Layer tensors have already
# been replaced by tiny device placeholders, so this does not reload the
# offloaded layer weights.
if not self._has_dtensor_weights:
self.model.to(self.device)
self._finalize_initialization()
def _initialize_layer_weights(self) -> None:
self._named_parameters = dict(self.model.named_parameters())
self._named_buffers = dict(self.model.named_buffers())
# 1. collect and group layer parameters by dtype. Keep buffers resident:
# shared buffers such as RoPE caches may be referenced by many layers.
layer_groups: Dict[int, Dict[torch.dtype, List[Tuple[str, torch.Tensor]]]] = {}
@@ -353,12 +376,7 @@ class LayerwiseOffloadManager:
self._consolidated_cpu_weights[layer_idx][dtype] = cpu_buffer
# Keep non-layer parameters resident on GPU. Layer tensors have already
# been replaced by tiny device placeholders, so this does not reload the
# offloaded layer weights.
if not self._has_dtensor_weights:
self.model.to(self.device)
def _finalize_initialization(self) -> None:
# prefetch the head of the stream for warm-up; residency is not armed
# yet, so this is layer 0 regardless of policy
self.prepare_for_next_req(non_blocking=False)
@@ -1012,7 +1030,7 @@ class LayerwiseOffloadableModuleMixin:
pin_cpu_memory=server_args.pin_cpu_memory,
prefetch_size=prefetch_size,
resident_layers=resident_layers,
initialize=not current_platform.is_mps(),
initialize=False,
residency_policy=(
server_args.dit_layerwise_residency_policy
if dit_tuning_enabled
@@ -1026,6 +1044,33 @@ class LayerwiseOffloadableModuleMixin:
for manager in self.layerwise_offload_managers:
manager.initialize()
self._capture_mps_cpu_non_layer_weights()
else:
enabled_managers = [
manager
for manager in self.layerwise_offload_managers
if manager.enabled
]
initialization_order = sorted(
enabled_managers,
key=lambda manager: manager._managed_parameter_bytes(),
reverse=True,
)
# release the largest managed groups first when checkpoint loading
# already placed weights on the accelerator; keep the stored manager
# order unchanged for prefetch and forward lifecycle semantics
for manager in initialization_order:
manager._initialize_layer_weights()
# Every managed layer group must be replaced before moving the
# remaining parameters, otherwise an earlier manager transiently
# moves later groups to the device.
if enabled_managers and not any(
manager._has_dtensor_weights for manager in enabled_managers
):
self.to(enabled_managers[0].device)
for manager in enabled_managers:
manager._finalize_initialization()
if configured_layer_names:
logger.debug(
@@ -749,6 +749,25 @@ class _AuxiliaryResidentComponent(_ResidentComponent):
layerwise_offload_dit_group_enabled = False
class _MultiGroupComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin):
layer_names = ["small_blocks", "large_blocks"]
def __init__(self) -> None:
super().__init__()
self.small_blocks = torch.nn.ModuleList([_DummyBlock()])
self.large_blocks = torch.nn.ModuleList(
[_DummyBlock(), _DummyBlock(), _DummyBlock()]
)
self.non_layer = torch.nn.Parameter(torch.ones(2))
self.to_parameter_shapes = []
def to(self, *args, **kwargs):
self.to_parameter_shapes.append(
{name: tuple(param.shape) for name, param in self.named_parameters()}
)
return super().to(*args, **kwargs)
def _patch_fake_device(monkeypatch):
monkeypatch.setattr(
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
@@ -916,6 +935,36 @@ def test_configure_resolves_residency_policy(monkeypatch):
)
def test_configure_offloads_all_layer_groups_before_moving_non_layers(monkeypatch):
_patch_fake_device(monkeypatch)
model = _MultiGroupComponent()
initialization_order = []
initialize_layer_weights = LayerwiseOffloadManager._initialize_layer_weights
def record_initialization(manager):
initialization_order.append(manager.layers_attr_str)
initialize_layer_weights(manager)
monkeypatch.setattr(
LayerwiseOffloadManager,
"_initialize_layer_weights",
record_initialization,
)
model.configure_layerwise_offload(_server_args())
assert initialization_order == ["large_blocks", "small_blocks"]
assert [
manager.layers_attr_str for manager in model.layerwise_offload_managers
] == ["small_blocks", "large_blocks"]
assert len(model.to_parameter_shapes) == 1
shapes_at_move = model.to_parameter_shapes[0]
assert shapes_at_move["non_layer"] == (2,)
for name, shape in shapes_at_move.items():
if name != "non_layer":
assert shape == (1,), name
def test_holds_residents_reflects_configuration(monkeypatch):
_patch_fake_device(monkeypatch)
resident = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=2)