[diffusion] fix: fix LoRA weight merging when using layerwise offload (#16737)

Co-authored-by: niehen6174 <niehen.6174@gmail.com>
Co-authored-by: DavisTao <dwt614707404@163.com>
Co-authored-by: niehen6174 <nihen6174@gmail.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
WenhaoZhang
2026-01-10 20:17:35 +08:00
committed by GitHub
parent dae6a4092a
commit bdb76b34db
5 changed files with 292 additions and 83 deletions

View File

@@ -60,6 +60,8 @@ class LayerwiseOffloadManager:
self._named_parameters: Dict[str, torch.nn.Parameter] = {}
self._named_buffers: Dict[str, torch.Tensor] = {}
# Store forward hooks for removal
self._forward_hooks: List[Any] = []
self._initialize()
@@ -204,6 +206,51 @@ class LayerwiseOffloadManager:
for layer_idx in list(self._gpu_layers):
self.release_layer(layer_idx)
@torch.compiler.disable
def load_all_layers(self) -> None:
"""Load all layers from CPU to GPU."""
if not self.enabled or self.device is None:
return
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
for layer_idx in range(self.num_layers):
if layer_idx not in self._gpu_layers:
self.prefetch_layer(layer_idx, non_blocking=False)
@torch.compiler.disable
def sync_layer_to_cpu(self, layer_idx: int) -> None:
"""Sync a layer's weights from GPU back to CPU."""
if not self.enabled or layer_idx not in self._gpu_layers:
return
if layer_idx not in self._consolidated_cpu_weights:
return
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
# Collect current GPU weights and write back to CPU buffer
for name, meta in self._weight_metadata.get(layer_idx, {}).items():
target = self.get_target_with_name(name)
gpu_weight = target.data.flatten().cpu()
dtype = meta["dtype"]
cpu_buffer = self._consolidated_cpu_weights[layer_idx][dtype]
offset = meta["offset"]
numel = meta["numel"]
cpu_buffer[offset : offset + numel].copy_(gpu_weight)
@torch.compiler.disable
def sync_all_layers_to_cpu(self) -> None:
"""Sync all loaded layers' weights from GPU back to CPU."""
if not self.enabled or self.device is None:
return
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
for layer_idx in list(self._gpu_layers):
self.sync_layer_to_cpu(layer_idx)
def register_forward_hooks(self) -> None:
if not self.enabled:
return
@@ -225,9 +272,17 @@ class LayerwiseOffloadManager:
return hook
# register prefetch & release hooks for each layer
self._forward_hooks.clear()
for i, layer in enumerate(layers):
layer.register_forward_pre_hook(make_pre_hook(i))
layer.register_forward_hook(make_post_hook(i))
pre_hook_handle = layer.register_forward_pre_hook(make_pre_hook(i))
post_hook_handle = layer.register_forward_hook(make_post_hook(i))
self._forward_hooks.extend([pre_hook_handle, post_hook_handle])
def remove_forward_hooks(self) -> None:
"""Remove all registered forward hooks."""
for hook_handle in self._forward_hooks:
hook_handle.remove()
self._forward_hooks.clear()
class OffloadableDiTMixin:
@@ -266,3 +321,24 @@ class OffloadableDiTMixin:
return
for manager in self.layerwise_offload_managers:
manager.prepare_for_next_denoise(non_blocking=True)
def disable_offload(self) -> None:
"""Disable layerwise offload: load all layers to GPU and remove hooks."""
if self.layerwise_offload_managers is None:
return
for manager in self.layerwise_offload_managers:
if manager.enabled:
manager.remove_forward_hooks()
manager.load_all_layers()
def enable_offload(self) -> None:
"""Re-enable layerwise offload: sync weights to CPU, release layers, and restore hooks."""
if self.layerwise_offload_managers is None:
return
for manager in self.layerwise_offload_managers:
if manager.enabled:
manager.sync_all_layers_to_cpu()
for layer_idx in list(manager._gpu_layers):
if layer_idx > 0:
manager.release_layer(layer_idx)
manager.register_forward_hooks()