Fix Bug in Group Offload when diffusers version is low. (#474)

This commit is contained in:
Bubbliiiing
2026-03-10 15:02:16 +08:00
committed by GitHub
parent 5202421e7c
commit ad72867c0f
16 changed files with 136 additions and 27 deletions
+2 -1
View File
@@ -504,8 +504,8 @@ class CombineZImagePipeline:
)
pipeline.remove_all_hooks()
safe_remove_group_offloading(pipeline)
undo_convert_weight_dtype_wrapper(transformer)
pipeline.to(device=offload_device)
transformer = transformer.to(weight_dtype)
if GPU_memory_mode == "sequential_cpu_offload":
@@ -747,6 +747,7 @@ class LoadZImageControlNetInPipeline:
# Remove hooks
funmodels["pipeline"].remove_all_hooks()
safe_remove_group_offloading(funmodels["pipeline"])
# Load config
config_path = f"{script_directory}/config/{config}"
+1 -1
View File
@@ -19,7 +19,7 @@ accelerate launch --mixed_precision="bf16" scripts/z_image/train_distill_lora.py
--learning_rate=1e-04 \
--learning_rate_critic=1e-05 \
--seed=42 \
--output_dir="output_dir_z_image_distill" \
--output_dir="output_dir_z_image_distill_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
+1 -1
View File
@@ -20,7 +20,7 @@ accelerate launch --mixed_precision="bf16" scripts/z_image/train.py \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_z_image" \
--output_dir="output_dir_z_image_turbo" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
+1 -1
View File
@@ -18,7 +18,7 @@ accelerate launch --mixed_precision="bf16" scripts/z_image/train_lora.py \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_z_image_lora" \
--output_dir="output_dir_z_image_turbo_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
+4 -3
View File
@@ -1,7 +1,8 @@
from .dataset_image import CC15M, ImageEditDataset
from .dataset_image_video import (ImageVideoControlDataset, ImageVideoDataset, TextDataset,
ImageVideoSampler)
from .dataset_video import VideoDataset, VideoSpeechDataset, VideoAnimateDataset, WebVid10M
from .dataset_image_video import (ImageVideoControlDataset, ImageVideoDataset,
ImageVideoSampler, TextDataset)
from .dataset_video import (VideoAnimateDataset, VideoDataset,
VideoSpeechDataset, WebVid10M)
from .utils import (VIDEO_READER_TIMEOUT, Camera, VideoReader_contextmanager,
custom_meshgrid, get_random_mask, get_relative_pose,
get_video_reader_batch, padding_image, process_pose_file,
-1
View File
@@ -9,7 +9,6 @@ from contextlib import contextmanager
from random import shuffle
from threading import Thread
import albumentations
import cv2
import numpy as np
import torch
-1
View File
@@ -8,7 +8,6 @@ import random
from contextlib import contextmanager
from threading import Thread
import albumentations
import cv2
import librosa
import numpy as np
-1
View File
@@ -9,7 +9,6 @@ from contextlib import contextmanager
from random import shuffle
from threading import Thread
import albumentations
import cv2
import numpy as np
import torch
+8 -2
View File
@@ -28,8 +28,8 @@ from diffusers.utils import (BaseOutput, is_torch_xla_available, logging,
from diffusers.utils.torch_utils import randn_tensor
from ..models import (AutoencoderKLQwenImage,
Qwen2_5_VLForConditionalGeneration,
Qwen2Tokenizer, QwenImageTransformer2DModel)
Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer,
QwenImageTransformer2DModel)
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
@@ -629,6 +629,12 @@ class QwenImagePipeline(DiffusionPipeline):
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
)
if num_inference_steps == 2 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 2/3
elif num_inference_steps <= 4 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 1/2
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -699,6 +699,12 @@ class QwenImageControlPipeline(DiffusionPipeline):
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
)
if num_inference_steps == 2 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 2/3
elif num_inference_steps <= 4 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 1/2
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -14,11 +14,11 @@
# limitations under the License.
import inspect
import math
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Union
import numpy as np
import math
import PIL.Image
import torch
from diffusers.image_processor import VaeImageProcessor
@@ -29,8 +29,8 @@ from diffusers.utils import (BaseOutput, is_torch_xla_available, logging,
from diffusers.utils.torch_utils import randn_tensor
from ..models import (AutoencoderKLQwenImage,
Qwen2_5_VLForConditionalGeneration, Qwen2VLProcessor,
Qwen2Tokenizer, QwenImageTransformer2DModel)
Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer,
Qwen2VLProcessor, QwenImageTransformer2DModel)
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
@@ -801,6 +801,12 @@ class QwenImageEditPipeline(DiffusionPipeline):
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
)
if num_inference_steps == 2 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 2/3
elif num_inference_steps <= 4 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 1/2
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -14,11 +14,11 @@
# limitations under the License.
import inspect
import math
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Union
import numpy as np
import math
import PIL.Image
import torch
from diffusers.image_processor import VaeImageProcessor
@@ -29,8 +29,8 @@ from diffusers.utils import (BaseOutput, is_torch_xla_available, logging,
from diffusers.utils.torch_utils import randn_tensor
from ..models import (AutoencoderKLQwenImage,
Qwen2_5_VLForConditionalGeneration, Qwen2VLProcessor,
Qwen2Tokenizer, QwenImageTransformer2DModel)
Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer,
Qwen2VLProcessor, QwenImageTransformer2DModel)
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
@@ -786,6 +786,12 @@ class QwenImageEditPlusPipeline(DiffusionPipeline):
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
)
if num_inference_steps == 2 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 2/3
elif num_inference_steps <= 4 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 1/2
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -764,6 +764,12 @@ class QwenImageControlNetPipeline(DiffusionPipeline):
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
)
if num_inference_steps == 2 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 2/3
elif num_inference_steps <= 4 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 1/2
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -770,6 +770,12 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
image_seq_len = latents.shape[1]
base_seqlen = 256 * 256 / 16 / 16
mu = (image_latents.shape[1] / base_seqlen) ** 0.5
if num_inference_steps == 2 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 2/3
elif num_inference_steps <= 4 and hasattr(self.scheduler, "config") and \
hasattr(self.scheduler.config, "shift_terminal"):
self.scheduler.config.shift_terminal = 1/2
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
+41 -6
View File
@@ -1219,8 +1219,14 @@ def remove_group_offloading(
def safe_remove_group_offloading(obj, *args, **kwargs):
"""Safely call remove_group_offloading"""
return remove_group_offloading(obj, *args, **kwargs)
"""Safely call remove_group_offloading, and restore _execution_device if it was patched."""
result = remove_group_offloading(obj, *args, **kwargs)
# Restore the original _execution_device from the MRO if we had patched it.
if hasattr(obj, 'components') and hasattr(obj.__class__, '_execution_device_original'):
obj.__class__._execution_device = obj.__class__._execution_device_original
del obj.__class__._execution_device_original
logger.debug("Restored original _execution_device after removing group offload.")
return result
def enable_group_offload(
@@ -1347,12 +1353,41 @@ def enable_group_offload(
def safe_enable_group_offload(obj, *args, **kwargs):
"""Safely call enable_group_offload, register default implementation if not exists"""
"""Safely call enable_group_offload, register default implementation if not exists.
Also patches obj._execution_device so that pipelines using group offload (which does
not use Accelerate _hf_hook) can still return the correct onload device instead of
falling back to self.device (which may be CPU after offloading).
"""
if not hasattr(obj, 'enable_group_offload'):
obj.enable_group_offload = types.MethodType(enable_group_offload, obj)
return obj.enable_group_offload(*args, **kwargs)
result = obj.enable_group_offload(*args, **kwargs)
# Patch _execution_device on the pipeline so it correctly returns the
# onload (GPU) device instead of self.device (CPU) when group offload is active.
onload_device = kwargs.get('onload_device') or (args[0] if args else None)
if onload_device is not None and hasattr(obj, 'components'):
onload_device = torch.device(onload_device) if isinstance(onload_device, str) else onload_device
# Save the original _execution_device before patching so safe_remove can restore it.
if not hasattr(obj.__class__, '_execution_device_original'):
obj.__class__._execution_device_original = obj.__class__._execution_device
@property
def _execution_device(self):
# Dynamically check: if any component still has group offload active,
# return the onload device; otherwise fall through to the original impl.
for _, component in self.components.items():
if isinstance(component, torch.nn.Module) and _is_group_offload_enabled(component):
return onload_device
# Group offload has been removed, delegate to the saved original.
return self.__class__._execution_device_original.fget(self)
obj.__class__._execution_device = _execution_device
logger.debug(f"Patched _execution_device to return {onload_device} for group offload.")
return result
def register_auto_device_hook(model):
+42 -3
View File
@@ -18,6 +18,12 @@ from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear
from safetensors.torch import load_file
from transformers import T5EncoderModel
from videox_fun.utils.group_offload import (_get_top_level_group_offload_hook,
_is_group_offload_enabled,
register_auto_device_hook,
safe_enable_group_offload,
safe_remove_group_offloading)
class LoRAModule(torch.nn.Module):
"""
@@ -461,7 +467,6 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
error_count = 0
for layer, elems in updates.items():
if "lora_te" in layer:
if transformer_only:
skipped_count += 1
@@ -554,7 +559,25 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
if sequential_cpu_offload_flag:
print(f"[LoRA Merge] Re-enabling sequential CPU offload...")
pipeline.enable_sequential_cpu_offload(device=offload_device)
else:
# When group offload is active, remove and re-apply it on the whole pipeline
# so that all ModuleGroup.cpu_param_dict references are rebuilt from the
# freshly-merged weights. A simple _maybe_remove_and_reapply is not enough
# because .to() during merge creates new tensor objects that break cached refs.
try:
local_transformer = getattr(pipeline, sub_transformer_name)
if _is_group_offload_enabled(local_transformer):
print(f"[LoRA Merge] Removing group offload hooks from pipeline...")
safe_remove_group_offloading(pipeline)
print(f"[LoRA Merge] Re-applying group offload hooks to pipeline...")
register_auto_device_hook(getattr(pipeline, sub_transformer_name))
safe_enable_group_offload(
pipeline,
onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True
)
print(pipeline._execution_device)
except Exception as e:
print(f"[LoRA Merge] Warning: Failed to refresh group offload: {e}")
print(f"[LoRA Merge] ✓ LoRA merge finished successfully")
return pipeline
@@ -702,6 +725,22 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
if sequential_cpu_offload_flag:
print(f"[LoRA Unmerge] Re-enabling sequential CPU offload...")
pipeline.enable_sequential_cpu_offload(device=device)
else:
# Same as merge_lora: remove and re-apply group offload on the whole pipeline
# to rebuild cpu_param_dict with the post-unmerge weights.
try:
local_transformer = getattr(pipeline, sub_transformer_name)
if _is_group_offload_enabled(local_transformer):
print(f"[LoRA Unmerge] Removing group offload hooks from pipeline...")
safe_remove_group_offloading(pipeline)
print(f"[LoRA Unmerge] Re-applying group offload hooks to pipeline...")
register_auto_device_hook(getattr(pipeline, sub_transformer_name))
safe_enable_group_offload(
pipeline,
onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True
)
except Exception as e:
print(f"[LoRA Unmerge] Warning: Failed to refresh group offload: {e}")
print(f"[LoRA Unmerge] ✓ LoRA unmerge finished successfully")
return pipeline