[diffusion] fix: fix multi-group layerwise offload startup memory (#35509)
This commit is contained in:
+55
-10
@@ -205,6 +205,16 @@ class LayerwiseOffloadManager:
|
|||||||
except Exception:
|
except Exception:
|
||||||
return None
|
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:
|
def _get_shared_empty_tensor(self, dtype: torch.dtype) -> torch.Tensor:
|
||||||
placeholder = self._offload_placeholders.get(dtype)
|
placeholder = self._offload_placeholders.get(dtype)
|
||||||
if placeholder is None:
|
if placeholder is None:
|
||||||
@@ -252,13 +262,26 @@ class LayerwiseOffloadManager:
|
|||||||
if not self.enabled:
|
if not self.enabled:
|
||||||
return
|
return
|
||||||
|
|
||||||
self._named_parameters = dict(self.model.named_parameters())
|
|
||||||
self._named_buffers = dict(self.model.named_buffers())
|
|
||||||
|
|
||||||
if self._synchronous_mps:
|
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()
|
self._initialize_mps_cpu_weights()
|
||||||
return
|
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:
|
# 1. collect and group layer parameters by dtype. Keep buffers resident:
|
||||||
# shared buffers such as RoPE caches may be referenced by many layers.
|
# shared buffers such as RoPE caches may be referenced by many layers.
|
||||||
layer_groups: Dict[int, Dict[torch.dtype, List[Tuple[str, torch.Tensor]]]] = {}
|
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
|
self._consolidated_cpu_weights[layer_idx][dtype] = cpu_buffer
|
||||||
|
|
||||||
# Keep non-layer parameters resident on GPU. Layer tensors have already
|
def _finalize_initialization(self) -> None:
|
||||||
# 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)
|
|
||||||
|
|
||||||
# prefetch the head of the stream for warm-up; residency is not armed
|
# prefetch the head of the stream for warm-up; residency is not armed
|
||||||
# yet, so this is layer 0 regardless of policy
|
# yet, so this is layer 0 regardless of policy
|
||||||
self.prepare_for_next_req(non_blocking=False)
|
self.prepare_for_next_req(non_blocking=False)
|
||||||
@@ -1012,7 +1030,7 @@ class LayerwiseOffloadableModuleMixin:
|
|||||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
pin_cpu_memory=server_args.pin_cpu_memory,
|
||||||
prefetch_size=prefetch_size,
|
prefetch_size=prefetch_size,
|
||||||
resident_layers=resident_layers,
|
resident_layers=resident_layers,
|
||||||
initialize=not current_platform.is_mps(),
|
initialize=False,
|
||||||
residency_policy=(
|
residency_policy=(
|
||||||
server_args.dit_layerwise_residency_policy
|
server_args.dit_layerwise_residency_policy
|
||||||
if dit_tuning_enabled
|
if dit_tuning_enabled
|
||||||
@@ -1026,6 +1044,33 @@ class LayerwiseOffloadableModuleMixin:
|
|||||||
for manager in self.layerwise_offload_managers:
|
for manager in self.layerwise_offload_managers:
|
||||||
manager.initialize()
|
manager.initialize()
|
||||||
self._capture_mps_cpu_non_layer_weights()
|
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:
|
if configured_layer_names:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
|
|||||||
@@ -749,6 +749,25 @@ class _AuxiliaryResidentComponent(_ResidentComponent):
|
|||||||
layerwise_offload_dit_group_enabled = False
|
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):
|
def _patch_fake_device(monkeypatch):
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
|
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):
|
def test_holds_residents_reflects_configuration(monkeypatch):
|
||||||
_patch_fake_device(monkeypatch)
|
_patch_fake_device(monkeypatch)
|
||||||
resident = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=2)
|
resident = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=2)
|
||||||
|
|||||||
Reference in New Issue
Block a user