Fix Bug in Group Offload when diffusers version is low. (#474)
This commit is contained in:
@@ -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}"
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,7 +8,6 @@ import random
|
||||
from contextlib import contextmanager
|
||||
from threading import Thread
|
||||
|
||||
import albumentations
|
||||
import cv2
|
||||
import librosa
|
||||
import numpy as np
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user