diff --git a/comfyui/z_image/nodes.py b/comfyui/z_image/nodes.py index d06d650..f8a7426 100644 --- a/comfyui/z_image/nodes.py +++ b/comfyui/z_image/nodes.py @@ -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}" diff --git a/scripts/z_image/train_distill_lora.sh b/scripts/z_image/train_distill_lora.sh index bdd7e2f..7469e70 100644 --- a/scripts/z_image/train_distill_lora.sh +++ b/scripts/z_image/train_distill_lora.sh @@ -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 \ diff --git a/scripts/z_image/train_turbo.sh b/scripts/z_image/train_turbo.sh index 93e0470..eedb285 100644 --- a/scripts/z_image/train_turbo.sh +++ b/scripts/z_image/train_turbo.sh @@ -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 \ diff --git a/scripts/z_image/train_turbo_lora.sh b/scripts/z_image/train_turbo_lora.sh index 2001a05..47c6e17 100644 --- a/scripts/z_image/train_turbo_lora.sh +++ b/scripts/z_image/train_turbo_lora.sh @@ -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 \ diff --git a/videox_fun/data/__init__.py b/videox_fun/data/__init__.py index babf155..a7d42ae 100644 --- a/videox_fun/data/__init__.py +++ b/videox_fun/data/__init__.py @@ -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, diff --git a/videox_fun/data/dataset_image_video.py b/videox_fun/data/dataset_image_video.py index 449a2f7..f95e7d4 100755 --- a/videox_fun/data/dataset_image_video.py +++ b/videox_fun/data/dataset_image_video.py @@ -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 diff --git a/videox_fun/data/dataset_video.py b/videox_fun/data/dataset_video.py index 3a49d05..9114760 100644 --- a/videox_fun/data/dataset_video.py +++ b/videox_fun/data/dataset_video.py @@ -8,7 +8,6 @@ import random from contextlib import contextmanager from threading import Thread -import albumentations import cv2 import librosa import numpy as np diff --git a/videox_fun/data/utils.py b/videox_fun/data/utils.py index 514a41b..87f4cad 100644 --- a/videox_fun/data/utils.py +++ b/videox_fun/data/utils.py @@ -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 diff --git a/videox_fun/pipeline/pipeline_qwenimage.py b/videox_fun/pipeline/pipeline_qwenimage.py index 038a2ef..93fa916 100644 --- a/videox_fun/pipeline/pipeline_qwenimage.py +++ b/videox_fun/pipeline/pipeline_qwenimage.py @@ -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, diff --git a/videox_fun/pipeline/pipeline_qwenimage_control.py b/videox_fun/pipeline/pipeline_qwenimage_control.py index 3d58a15..8e1a516 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_control.py +++ b/videox_fun/pipeline/pipeline_qwenimage_control.py @@ -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, diff --git a/videox_fun/pipeline/pipeline_qwenimage_edit.py b/videox_fun/pipeline/pipeline_qwenimage_edit.py index 53e882e..35bf4e6 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_edit.py +++ b/videox_fun/pipeline/pipeline_qwenimage_edit.py @@ -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, diff --git a/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py b/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py index ee36ce9..b7417d4 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py +++ b/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py @@ -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, diff --git a/videox_fun/pipeline/pipeline_qwenimage_instantx.py b/videox_fun/pipeline/pipeline_qwenimage_instantx.py index d2e97a5..e771fbb 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_instantx.py +++ b/videox_fun/pipeline/pipeline_qwenimage_instantx.py @@ -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, diff --git a/videox_fun/pipeline/pipeline_qwenimage_layered.py b/videox_fun/pipeline/pipeline_qwenimage_layered.py index b8fff49..eb8e2f3 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_layered.py +++ b/videox_fun/pipeline/pipeline_qwenimage_layered.py @@ -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, diff --git a/videox_fun/utils/group_offload.py b/videox_fun/utils/group_offload.py index cfc00b1..4929e8a 100644 --- a/videox_fun/utils/group_offload.py +++ b/videox_fun/utils/group_offload.py @@ -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): diff --git a/videox_fun/utils/lora_utils.py b/videox_fun/utils/lora_utils.py index 389b22a..95524b1 100755 --- a/videox_fun/utils/lora_utils.py +++ b/videox_fun/utils/lora_utils.py @@ -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