diff --git a/comfyui/cogvideox_fun/nodes.py b/comfyui/cogvideox_fun/nodes.py index 61b2bf6..06a9264 100755 --- a/comfyui/cogvideox_fun/nodes.py +++ b/comfyui/cogvideox_fun/nodes.py @@ -92,7 +92,7 @@ class LoadCogVideoXFunModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py index b648eb2..eaf231c 100755 --- a/comfyui/comfyui_nodes.py +++ b/comfyui/comfyui_nodes.py @@ -15,11 +15,20 @@ from .cogvideox_fun.nodes import (CogVideoXFunInpaintSampler, CogVideoXFunV2VSampler, LoadCogVideoXFunLora, LoadCogVideoXFunModel) from .comfyui_utils import script_directory -from .qwenimage.nodes import (CombineQwenImagePipeline, LoadQwenImageLora, - LoadQwenImageModel, LoadQwenImageProcessor, +from .flux2.nodes import (CombineFlux2Pipeline, Flux2ControlSampler, + Flux2T2ISampler, LoadFlux2ControlNetInModel, + LoadFlux2ControlNetInPipeline, LoadFlux2Lora, + LoadFlux2Model, LoadFlux2TextEncoderModel, + LoadFlux2TransformerModel, LoadFlux2VAEModel) +from .qwenimage.nodes import (CombineQwenImagePipeline, + LoadQwenImageControlNetInModel, + LoadQwenImageControlNetInPipeline, + LoadQwenImageLora, LoadQwenImageModel, + LoadQwenImageProcessor, LoadQwenImageTextEncoderModel, LoadQwenImageTransformerModel, - LoadQwenImageVAEModel, QwenImageEditSampler, + LoadQwenImageVAEModel, QwenImageControlSampler, + QwenImageEditPlusSampler, QwenImageEditSampler, QwenImageT2VSampler) from .wan2_1.nodes import (CombineWanPipeline, LoadWanClipEncoderModel, LoadWanLora, LoadWanModel, LoadWanTextEncoderModel, @@ -461,10 +470,26 @@ NODE_CLASS_MAPPINGS = { "LoadQwenImageVAEModel": LoadQwenImageVAEModel, "LoadQwenImageProcessor": LoadQwenImageProcessor, "CombineQwenImagePipeline": CombineQwenImagePipeline, + "LoadQwenImageControlNetInPipeline": LoadQwenImageControlNetInPipeline, + "LoadQwenImageControlNetInModel": LoadQwenImageControlNetInModel, "LoadQwenImageModel": LoadQwenImageModel, "QwenImageT2VSampler": QwenImageT2VSampler, "QwenImageEditSampler": QwenImageEditSampler, + "QwenImageEditPlusSampler": QwenImageEditPlusSampler, + "QwenImageControlSampler": QwenImageControlSampler, + + "LoadFlux2Lora": LoadFlux2Lora, + "LoadFlux2TransformerModel": LoadFlux2TransformerModel, + "LoadFlux2VAEModel": LoadFlux2VAEModel, + "LoadFlux2TextEncoderModel": LoadFlux2TextEncoderModel, + "CombineFlux2Pipeline": CombineFlux2Pipeline, + "LoadFlux2ControlNetInModel": LoadFlux2ControlNetInModel, + "LoadFlux2ControlNetInPipeline": LoadFlux2ControlNetInPipeline, + + "LoadFlux2Model": LoadFlux2Model, + "Flux2T2ISampler": Flux2T2ISampler, + "Flux2ControlSampler": Flux2ControlSampler, "LoadZImageLora": LoadZImageLora, "LoadZImageTextEncoderModel": LoadZImageTextEncoderModel, @@ -549,10 +574,26 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LoadQwenImageVAEModel": "Load QwenImage VAE Model", "LoadQwenImageProcessor": "Load QwenImage Processor", "CombineQwenImagePipeline": "Combine QwenImage Pipeline", + "LoadQwenImageControlNetInPipeline": "Load QwenImage ControlNet In Pipeline", + "LoadQwenImageControlNetInModel": "Load QwenImage ControlNet In Model", "LoadQwenImageModel": "Load QwenImage Model", "QwenImageT2VSampler": "QwenImage T2V Sampler", "QwenImageEditSampler": "QwenImage Edit Sampler", + "QwenImageEditPlusSampler": "QwenImage Edit Plus Sampler", + "QwenImageControlSampler": "QwenImage Control Sampler", + + "LoadFlux2Lora": "Load FLUX2 Lora", + "LoadFlux2TransformerModel": "Load FLUX2 Transformer Model", + "LoadFlux2VAEModel": "Load FLUX2 VAE Model", + "LoadFlux2TextEncoderModel": "Load FLUX2 Text Encoder Model", + "CombineFlux2Pipeline": "Combine FLUX2 Pipeline", + "LoadFlux2ControlNetInModel": "Load Flux2 ControlNet In Model", + "LoadFlux2ControlNetInPipeline": "Load Flux2 ControlNet In Pipeline", + + "LoadFlux2Model": "Load FLUX2 Model", + "Flux2T2ISampler": "FLUX2 Text to Image Sampler", + "Flux2ControlSampler": "FLUX2 Control Sampler", "LoadZImageLora": "Load ZImage Lora", "LoadZImageTextEncoderModel": "Load ZImage TextEncoder Model", diff --git a/comfyui/flux2/README.md b/comfyui/flux2/README.md new file mode 100644 index 0000000..b9f5cb5 --- /dev/null +++ b/comfyui/flux2/README.md @@ -0,0 +1,104 @@ +# FLUX.2-dev Model Setup Guide + +## a. Model Links and Storage Locations + +**Chunked loading is recommended** as it better aligns with ComfyUI's standard workflow. + +### 1. Chunked Loading Weights (Recommended) + +For chunked loading, it is recommended to directly download the FLUX.2-dev weights provided by ComfyUI official. Please organize the files according to the following directory structure: + +**Core Model Files:** + +| Component | File Name | +|-----------|-----------| +| Text Encoder | [`mistral_3_small_flux2_bf16.safetensors`](https://huggingface.co/Comfy-Org/flux2-dev/resolve/main/split_files/text_encoders/mistral_3_small_flux2_bf16.safetensors) | +| Diffusion Model | [`flux2_dev_fp8mixed.safetensors`](https://huggingface.co/Comfy-Org/flux2-dev/resolve/main/split_files/diffusion_models/flux2_dev_fp8mixed.safetensors) | +| VAE | [`flux2-vae.safetensors`](https://huggingface.co/Comfy-Org/flux2-dev/resolve/main/split_files/vae/flux2-vae.safetensors) | +| tokenizer | [`tokenizer`](https://huggingface.co/black-forest-labs/FLUX.2-dev/tree/main/tokenizer) | + +**ControlNet Model Files:** + +| Name | Storage | Hugging Face | Model Scope | Description | +|--|--|--|--|--| +| FLUX.2-dev-Fun-Controlnet-Union | - | [🤗Link](https://huggingface.co/alibaba-pai/FLUX.2-dev-Fun-Controlnet-Union) | [😄Link](https://modelscope.cn/models/PAI/FLUX.2-dev-Fun-Controlnet-Union) | ControlNet weights for FLUX.2-dev, supporting multiple control conditions such as Canny, Depth, Pose, MLSD, Scribble, etc. | + +**Storage Location:** + +``` +📂 ComfyUI/ +├── 📂 models/ +│ ├── 📂 text_encoders/ +│ │ └── mistral_3_small_flux2_bf16.safetensors +│ ├── 📂 diffusion_models/ +│ │ └── flux2_dev_fp8mixed.safetensors +│ ├── 📂 vae/ +│ │ └── flux2-vae.safetensors +│ ├── 📂 Fun_Models/ +│ │ └── flux2_tokenizer/ +│ └── 📂 model_patches/ +│ └── FLUX.2-dev-Fun-Controlnet-Union.safetensors +``` + +### 2. Preprocessing Weights (Optional) + +If you want to use the control preprocessing nodes, you can download the preprocessing weights to `ComfyUI/custom_nodes/Fun_Models/Third_Party/`. + +**Required Files:** + +| File Name | Download Link | Purpose | +|-----------|---------------|---------| +| `yolox_l.onnx` | [Download](https://huggingface.co/yzd-v/DWPose/resolve/main/yolox_l.onnx) | YOLO Detection Model | +| `dw-ll_ucoco_384.onnx` | [Download](https://huggingface.co/yzd-v/DWPose/resolve/main/dw-ll_ucoco_384.onnx) | DWPose Pose Estimation Model | +| `ZoeD_M12_N.pt` | [Download](https://huggingface.co/lllyasviel/Annotators/resolve/main/ZoeD_M12_N.pt) | ZoeDepth Depth Estimation Model | + +**Storage Location:** + +``` +📂 ComfyUI/ +├── 📂 models/ +│ └── 📂 Fun_Models/ +│ └── 📂 Third_Party +│ ├── yolox_l.onnx +│ ├── dw-ll_ucoco_384.onnx +│ └── ZoeD_M12_N.pt +``` + +### 3. Full Model Loading (Optional) + +If you prefer full model loading, you can directly download the diffusers weights. + +**Required Files:** + +| Name | Storage | Hugging Face | Model Scope | Description | +|--|--|--|--|--| +| FLUX.2-dev | [🤗Link](https://huggingface.co/black-forest-labs/FLUX.2-dev) | [😄Link](https://modelscope.cn/models/black-forest-labs/FLUX.2-dev) | Official FLUX.2-dev weights | + +For full model loading, use the diffusers version of FLUX.2-dev Turbo and place the model in `ComfyUI/models/Fun_Models/`. + +**Storage Location:** + +``` +📂 ComfyUI/ +├── 📂 models/ +│ └── 📂 Fun_Models/ +| └── 📂 FLUX.2-dev/ +``` + +## b. ComfyUI Json Workflows + +### 1. Chunked Loading (Recommended) + +[FLUX.2-dev Text to Image](v1/flux2_chunked_loading_workflow_t2i.json) + +[FLUX.2-dev Text to Image Control](v1/flux2_chunked_loading_workflow_t2i_control.json) + +[FLUX.2-dev Text to Image Inpaint](v1/flux2_chunked_loading_workflow_t2i_inpaint.json) + +### 2. Full Model Loading (Optional) + +[FLUX.2-dev Text to Image](v1/flux2_workflow_t2i.json) + +[FLUX.2-dev Text to Image Control](v1/flux2_workflow_t2i_control.json) + +[FLUX.2-dev Text to Image Inpaint](v1/flux2_workflow_t2i_inpaint.json) \ No newline at end of file diff --git a/comfyui/flux2/nodes.py b/comfyui/flux2/nodes.py new file mode 100644 index 0000000..bdfdaa8 --- /dev/null +++ b/comfyui/flux2/nodes.py @@ -0,0 +1,1357 @@ + +import copy +import gc +import inspect +import json +import os +from collections import OrderedDict + +import accelerate +import comfy.model_management as mm +import cv2 +import folder_paths +import numpy as np +import torch +from comfy.utils import ProgressBar, load_torch_file +from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers import __version__ as diffusers_version +from einops import rearrange +from omegaconf import OmegaConf +from safetensors.torch import load_file +from transformers import Mistral3Config + +if diffusers_version >= "0.33.0": + from diffusers.models.model_loading_utils import load_model_dict_into_meta +else: + from diffusers.models.modeling_utils import \ + load_model_dict_into_meta + +from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, + get_closest_ratio) +from ...videox_fun.models import (AutoencoderKLFlux2, + Flux2ControlTransformer2DModel, + Flux2Transformer2DModel, + Mistral3ForConditionalGeneration, + PixtralProcessor) +from ...videox_fun.pipeline import Flux2ControlPipeline, Flux2Pipeline +from ...videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload, + safe_remove_group_offloading) +from ...videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from ...videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler +from ...videox_fun.utils.fp8_optimization import ( + convert_model_weight_to_float8, convert_weight_dtype_wrapper, + replace_parameters_by_name, undo_convert_weight_dtype_wrapper) +from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from ...videox_fun.utils.utils import (filter_kwargs, get_autocast_dtype, + get_image, get_image_latent) +from ..comfyui_utils import (eas_cache_dir, script_directory, + search_model_in_possible_folders, + search_sub_dir_in_possible_folders, to_pil) + +transformer_cpu_cache = {} +lora_path_before = "" + +def get_flux2_scheduler(sampler_name, shift=1.0): + Chosen_Scheduler = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, + }[sampler_name] + + scheduler_kwargs = { + "_class_name": "FlowMatchEulerDiscreteScheduler", + "_diffusers_version": "0.36.0.dev0", + "base_image_seq_len": 256, + "base_shift": 0.5, + "invert_sigmas": False, + "max_image_seq_len": 4096, + "max_shift": 1.15, + "num_train_timesteps": 1000, + "shift": 3.0, + "shift_terminal": None, + "stochastic_sampling": False, + "time_shift_type": "exponential", + "use_beta_sigmas": False, + "use_dynamic_shifting": True, + "use_exponential_sigmas": False, + "use_karras_sigmas": False + } + scheduler_kwargs['shift'] = shift + scheduler = Chosen_Scheduler( + **filter_kwargs(Chosen_Scheduler, scheduler_kwargs) + ) + return scheduler + + +def convert_flux2_to_diffusers(original_state_dict): + weight_dtype = get_autocast_dtype() + converted = OrderedDict() + + def apply_scales(weight, prefix, layer_name): + weight_scale_key = f'{prefix}.{layer_name}.weight_scale' + input_scale_key = f'{prefix}.{layer_name}.input_scale' + + result = weight.to(weight_dtype) if weight.dtype in [torch.float8_e4m3fn, torch.float8_e5m2] else weight + + if weight_scale_key in original_state_dict: + weight_scale = original_state_dict[weight_scale_key] + weight_scale = weight_scale.to(weight_dtype) if weight_scale.dtype in [torch.float8_e4m3fn, torch.float8_e5m2] else weight_scale + result = result * weight_scale + + return result + + # Time and guidance embeddings + key = 'time_in.in_layer.weight' + if key in original_state_dict: + converted['time_guidance_embed.timestep_embedder.linear_1.weight'] = original_state_dict[key] + + key = 'time_in.out_layer.weight' + if key in original_state_dict: + converted['time_guidance_embed.timestep_embedder.linear_2.weight'] = original_state_dict[key] + + key = 'guidance_in.in_layer.weight' + if key in original_state_dict: + converted['time_guidance_embed.guidance_embedder.linear_1.weight'] = original_state_dict[key] + + key = 'guidance_in.out_layer.weight' + if key in original_state_dict: + converted['time_guidance_embed.guidance_embedder.linear_2.weight'] = original_state_dict[key] + + # Input projections + key = 'img_in.weight' + if key in original_state_dict: + converted['x_embedder.weight'] = original_state_dict[key] + + key = 'txt_in.weight' + if key in original_state_dict: + converted['context_embedder.weight'] = original_state_dict[key] + + # Modulations + key = 'double_stream_modulation_img.lin.weight' + if key in original_state_dict: + converted['double_stream_modulation_img.linear.weight'] = original_state_dict[key] + + key = 'double_stream_modulation_txt.lin.weight' + if key in original_state_dict: + converted['double_stream_modulation_txt.linear.weight'] = original_state_dict[key] + + key = 'single_stream_modulation.lin.weight' + if key in original_state_dict: + converted['single_stream_modulation.linear.weight'] = original_state_dict[key] + + # Double blocks (transformer_blocks) + for i in range(8): + prefix_old = f'double_blocks.{i}' + prefix_new = f'transformer_blocks.{i}' + + qkv_key = f'{prefix_old}.img_attn.qkv.weight' + if qkv_key in original_state_dict: + qkv_weight = original_state_dict[qkv_key] + qkv_weight = qkv_weight.to(weight_dtype) if qkv_weight.dtype in [torch.float8_e4m3fn, torch.float8_e5m2] else qkv_weight + total_dim = qkv_weight.shape[0] + single_dim = total_dim // 3 + converted[f'{prefix_new}.attn.to_q.weight'] = qkv_weight[:single_dim] + converted[f'{prefix_new}.attn.to_k.weight'] = qkv_weight[single_dim:2*single_dim] + converted[f'{prefix_new}.attn.to_v.weight'] = qkv_weight[2*single_dim:] + + # Norms + key = f'{prefix_old}.img_attn.norm.query_norm.scale' + if key in original_state_dict: + converted[f'{prefix_new}.attn.norm_q.weight'] = original_state_dict[key] + + key = f'{prefix_old}.img_attn.norm.key_norm.scale' + if key in original_state_dict: + converted[f'{prefix_new}.attn.norm_k.weight'] = original_state_dict[key] + + # Output projection + key = f'{prefix_old}.img_attn.proj.weight' + if key in original_state_dict: + converted[f'{prefix_new}.attn.to_out.0.weight'] = original_state_dict[key] + + # Text attention QKV (added) + qkv_added_key = f'{prefix_old}.txt_attn.qkv.weight' + if qkv_added_key in original_state_dict: + qkv_weight = original_state_dict[qkv_added_key] + qkv_weight = qkv_weight.to(weight_dtype) if qkv_weight.dtype in [torch.float8_e4m3fn, torch.float8_e5m2] else qkv_weight + total_dim = qkv_weight.shape[0] + single_dim = total_dim // 3 + converted[f'{prefix_new}.attn.add_q_proj.weight'] = qkv_weight[:single_dim] + converted[f'{prefix_new}.attn.add_k_proj.weight'] = qkv_weight[single_dim:2*single_dim] + converted[f'{prefix_new}.attn.add_v_proj.weight'] = qkv_weight[2*single_dim:] + + # Text norms + key = f'{prefix_old}.txt_attn.norm.query_norm.scale' + if key in original_state_dict: + converted[f'{prefix_new}.attn.norm_added_q.weight'] = original_state_dict[key] + + key = f'{prefix_old}.txt_attn.norm.key_norm.scale' + if key in original_state_dict: + converted[f'{prefix_new}.attn.norm_added_k.weight'] = original_state_dict[key] + + # Text output projection + key = f'{prefix_old}.txt_attn.proj.weight' + if key in original_state_dict: + converted[f'{prefix_new}.attn.to_add_out.weight'] = original_state_dict[key] + + # Image MLP with scales + key = f'{prefix_old}.img_mlp.0.weight' + if key in original_state_dict: + converted[f'{prefix_new}.ff.linear_in.weight'] = apply_scales( + original_state_dict[key], + prefix_old, 'img_mlp.0' + ) + asd = f'{prefix_new}.ff.linear_in.weight' + + + key = f'{prefix_old}.img_mlp.2.weight' + if key in original_state_dict: + converted[f'{prefix_new}.ff.linear_out.weight'] = apply_scales( + original_state_dict[key], + prefix_old, 'img_mlp.2' + ) + + # Text MLP with scales + key = f'{prefix_old}.txt_mlp.0.weight' + if key in original_state_dict: + converted[f'{prefix_new}.ff_context.linear_in.weight'] = apply_scales( + original_state_dict[key], + prefix_old, 'txt_mlp.0' + ) + + key = f'{prefix_old}.txt_mlp.2.weight' + if key in original_state_dict: + converted[f'{prefix_new}.ff_context.linear_out.weight'] = apply_scales( + original_state_dict[key], + prefix_old, 'txt_mlp.2' + ) + + # Single blocks (single_transformer_blocks) + for i in range(48): + prefix_old = f'single_blocks.{i}' + prefix_new = f'single_transformer_blocks.{i}' + + # QKV+MLP projection with scales + key = f'{prefix_old}.linear1.weight' + if key in original_state_dict: + converted[f'{prefix_new}.attn.to_qkv_mlp_proj.weight'] = apply_scales( + original_state_dict[key], + prefix_old, 'linear1' + ) + + # Norms + key = f'{prefix_old}.norm.query_norm.scale' + if key in original_state_dict: + converted[f'{prefix_new}.attn.norm_q.weight'] = original_state_dict[key] + + key = f'{prefix_old}.norm.key_norm.scale' + if key in original_state_dict: + converted[f'{prefix_new}.attn.norm_k.weight'] = original_state_dict[key] + + # Output projection with scales + key = f'{prefix_old}.linear2.weight' + if key in original_state_dict: + converted[f'{prefix_new}.attn.to_out.weight'] = apply_scales( + original_state_dict[key], + prefix_old, 'linear2' + ) + + # Final layer + key = 'final_layer.adaLN_modulation.1.weight' + if key in original_state_dict: + height = original_state_dict[key].size()[0] + height = int(height // 2) + converted['norm_out.linear.weight'] = torch.cat( + [original_state_dict[key][height:, :], original_state_dict[key][:height, :]], dim=0 + ) + + key = 'final_layer.linear.weight' + if key in original_state_dict: + converted['proj_out.weight'] = original_state_dict[key] + + return converted + + +class LoadFlux2TransformerModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": ( + folder_paths.get_filename_list("diffusion_models"), + {"default": "flux2_dev_fp8_e4m3fn.safetensors"}, + ), + "precision": ( + ["fp16", "bf16"], + {"default": "bf16"} + ), + }, + } + + RETURN_TYPES = ("TransformerModel", "STRING") + RETURN_NAMES = ("transformer", "model_name") + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, model_name, precision): + # Init weight_dtype and device + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + + mm.unload_all_models() + mm.cleanup_models_gc() + mm.soft_empty_cache() + transformer = None + + model_path = folder_paths.get_full_path("diffusion_models", model_name) + transformer_state_dict = load_torch_file(model_path, safe_load=True) + + model_name_in_pipeline = "FLUX.2-dev" + kwargs = { + "_class_name": "Flux2Transformer2DModel", + "_diffusers_version": "0.36.0.dev0", + "attention_head_dim": 128, + "axes_dims_rope": [ + 32, + 32, + 32, + 32 + ], + "eps": 1e-06, + "in_channels": 128, + "joint_attention_dim": 15360, + "mlp_ratio": 3.0, + "num_attention_heads": 48, + "num_layers": 8, + "num_single_layers": 48, + "out_channels": None, + "patch_size": 1, + "rope_theta": 2000, + "timestep_guidance_channels": 256 + } + + sig = inspect.signature(Flux2Transformer2DModel) + accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} + with accelerate.init_empty_weights(): + transformer = Flux2Transformer2DModel(**accepted) + + if 'time_in.in_layer.weight' in transformer_state_dict.keys(): + transformer_state_dict = convert_flux2_to_diffusers(transformer_state_dict) + + filtered_state_dict = {} + for key in transformer_state_dict: + if key in transformer.state_dict() and transformer.state_dict()[key].size() == transformer_state_dict[key].size(): + filtered_state_dict[key] = transformer_state_dict[key] + missing_keys = set(transformer.state_dict().keys()) - set(filtered_state_dict.keys()) + if missing_keys: + raise ValueError(f"Missing keys: {sorted(missing_keys)}") + + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + transformer, + transformer_state_dict, + dtype=weight_dtype, + model_name_or_path="", + ) + else: + transformer._convert_deprecated_attention_blocks(transformer_state_dict) + unexpected_keys = load_model_dict_into_meta( + transformer, + transformer_state_dict, + device=offload_device, + dtype=weight_dtype, + model_name_or_path="", + ) + + transformer = transformer.eval().to(weight_dtype) + return (transformer, model_name_in_pipeline) + +class LoadFlux2VAEModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": ( + folder_paths.get_filename_list("vae"), + {"default": "flux2_vae.safetensors"} + ), + "precision": ( + ["fp16", "bf16"], + {"default": "bf16"} + ), + }, + } + + RETURN_TYPES = ("VAEModel",) + RETURN_NAMES = ("vae",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, model_name, precision): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + model_path = folder_paths.get_full_path("vae", model_name) + vae_state_dict = load_torch_file(model_path, safe_load=True) + + kwargs = { + "_class_name": "AutoencoderKLFlux2", + "_diffusers_version": "0.36.0.dev0", + "act_fn": "silu", + "batch_norm_eps": 0.0001, + "batch_norm_momentum": 0.1, + "block_out_channels": [ + 128, + 256, + 512, + 512 + ], + "down_block_types": [ + "DownEncoderBlock2D", + "DownEncoderBlock2D", + "DownEncoderBlock2D", + "DownEncoderBlock2D" + ], + "force_upcast": True, + "in_channels": 3, + "latent_channels": 32, + "layers_per_block": 2, + "mid_block_add_attention": True, + "norm_num_groups": 32, + "out_channels": 3, + "patch_size": [ + 2, + 2 + ], + "sample_size": 1024, + "up_block_types": [ + "UpDecoderBlock2D", + "UpDecoderBlock2D", + "UpDecoderBlock2D", + "UpDecoderBlock2D" + ], + "use_post_quant_conv": True, + "use_quant_conv": True + } + + sig = inspect.signature(AutoencoderKLFlux2) + accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} + + vae = AutoencoderKLFlux2(**accepted) + vae.load_state_dict(vae_state_dict) + vae = vae.eval().to(device=offload_device, dtype=weight_dtype) + return (vae,) + + +class LoadFlux2TextEncoderModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": ( + folder_paths.get_filename_list("text_encoders"), + {"default": "mistral3_fp8_scaled.safetensors"} + ), + "precision": ( + ["fp16", "bf16"], + {"default": "bf16"} + ), + }, + } + + RETURN_TYPES = ("TextEncoderModel", "Tokenizer") + RETURN_NAMES = ("text_encoder", "tokenizer") + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, model_name, precision): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + model_path = folder_paths.get_full_path("text_encoders", model_name) + text_state_dict = load_torch_file(model_path, safe_load=True) + + config_kwargs = { + "architectures": [ + "Mistral3ForConditionalGeneration" + ], + "dtype": "bfloat16", + "image_token_index": 10, + "model_type": "mistral3", + "multimodal_projector_bias": False, + "projector_hidden_act": "gelu", + "spatial_merge_size": 2, + "text_config": { + "attention_dropout": 0.0, + "dtype": "bfloat16", + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 5120, + "initializer_range": 0.02, + "intermediate_size": 32768, + "max_position_embeddings": 131072, + "model_type": "mistral", + "num_attention_heads": 32, + "num_hidden_layers": 40, + "num_key_value_heads": 8, + "rms_norm_eps": 1e-05, + "rope_theta": 1000000000.0, + "sliding_window": None, + "use_cache": True, + "vocab_size": 131072 + }, + "transformers_version": "4.57.1", + "vision_config": { + "attention_dropout": 0.0, + "dtype": "bfloat16", + "head_dim": 64, + "hidden_act": "silu", + "hidden_size": 1024, + "image_size": 1540, + "initializer_range": 0.02, + "intermediate_size": 4096, + "model_type": "pixtral", + "num_attention_heads": 16, + "num_channels": 3, + "num_hidden_layers": 24, + "patch_size": 14, + "rope_theta": 10000.0 + }, + "vision_feature_layer": -1 + } + config = Mistral3Config(**config_kwargs) + text_encoder = Mistral3ForConditionalGeneration._from_config(config) + + if "tekken_model" in text_state_dict.keys(): + def convert_mistral3_to_diffusers(state_dict): + new_state_dict = {} + + for key, value in state_dict.items(): + if key == "tekken_model": + continue + if key.startswith("vision_tower."): + new_key = "model." + key + + elif key.startswith("multi_modal_projector."): + new_key = "model." + key + + elif key.startswith("model.layers."): + new_key = "model.language_" + key + + elif key.startswith("model.embed_tokens."): + new_key = "model.language_" + key + + elif key == "model.norm.weight": + new_key = "model.language_model.norm.weight" + + elif key.startswith("lm_head."): + new_key = key + + else: + print(f"Warning: Unmapped key: {key}") + new_key = key + + new_state_dict[new_key] = value + + return new_state_dict + else: + def convert_mistral3_to_diffusers(state_dict): + new_state_dict = {} + + for key, value in state_dict.items(): + if key.startswith('vision_tower.'): + new_key = 'model.' + key + + elif key.startswith('multi_modal_projector.'): + new_key = 'model.' + key + + elif key.startswith('language_model.model.embed_tokens.'): + new_key = key.replace('language_model.model.', 'model.language_model.') + + elif key.startswith('language_model.model.layers.'): + new_key = key.replace('language_model.model.', 'model.language_model.') + + elif key.startswith('language_model.model.norm.'): + new_key = key.replace('language_model.model.', 'model.language_model.') + + elif key.startswith('language_model.lm_head.'): + new_key = key.replace('language_model.', '') + + else: + new_key = key + + new_state_dict[new_key] = value + + return new_state_dict + + text_state_dict = convert_mistral3_to_diffusers(text_state_dict) + + text_encoder.load_state_dict(text_state_dict, strict=False) + text_encoder = text_encoder.eval().to(device=offload_device, dtype=weight_dtype) + + possible_folders = ["CogVideoX_Fun", "Fun_Models", "VideoX_Fun", "Wan-AI", "Qwen"] + \ + [os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "models/Diffusion_Transformer")] # Possible folder names to check + try: + tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="flux2_tokenizer") + except: + try: + tokenizer_path = os.path.join(search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="FLUX.2-dev"), "tokenizer") + except: + tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Mistral-Nemo-Instruct-2407") + + tokenizer = PixtralProcessor.from_pretrained(tokenizer_path) + return (text_encoder, tokenizer) + + +class CombineFlux2Pipeline: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "transformer": ("TransformerModel",), + "vae": ("VAEModel",), + "text_encoder": ("TextEncoderModel",), + "tokenizer": ("Tokenizer",), + "model_name": ("STRING",), + "GPU_memory_mode": ( + [ + "model_full_load", + "model_full_load_and_qfloat8", + "model_cpu_offload", + "model_cpu_offload_and_qfloat8", + "model_group_offload", + "sequential_cpu_offload" + ], + {"default": "model_cpu_offload"} + ), + }, + } + + RETURN_TYPES = ("FunModels",) + RETURN_NAMES = ("funmodels",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, model_name, GPU_memory_mode, transformer, vae, text_encoder, tokenizer): + # Get pipeline + weight_dtype = transformer.dtype if transformer.dtype not in [torch.float32, torch.float8_e4m3fn, torch.float8_e5m2] else get_autocast_dtype() + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + # Get pipeline + if hasattr(transformer, "control_transformer_blocks"): + model_type = "Control" + else: + model_type = "Inpaint" + + if model_type == "Inpaint": + pipeline = Flux2Pipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + ) + else: + pipeline = Flux2ControlPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + ) + + pipeline.remove_all_hooks() + safe_remove_group_offloading(pipeline) + undo_convert_weight_dtype_wrapper(transformer) + transformer = transformer.to(weight_dtype) + + if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) + elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device=offload_device, offload_type="leaf_level", use_stream=True) + elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) + else: + pipeline.to(device=device) + + funmodels = { + 'pipeline': pipeline, + 'GPU_memory_mode': GPU_memory_mode, + 'dtype': weight_dtype, + 'model_name': model_name, + 'model_type': model_type, + 'loras': [], + 'strength_model': [] + } + return (funmodels,) + + +class LoadFlux2Model: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ( + [ + 'FLUX.2-dev', + ], + {"default": 'FLUX.2-dev'} + ), + "GPU_memory_mode":( + [ + "model_full_load", "model_full_load_and_qfloat8", "model_cpu_offload", + "model_cpu_offload_and_qfloat8", "model_group_offload", "sequential_cpu_offload"], + { + "default": "model_cpu_offload", + } + ), + "precision": ( + ['fp16', 'bf16'], + { + "default": 'fp16' + } + ), + }, + } + + RETURN_TYPES = ("FunModels",) + RETURN_NAMES = ("funmodels",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, GPU_memory_mode, model, precision): + # Init weight_dtype and device + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + + mm.unload_all_models() + mm.cleanup_models_gc() + mm.soft_empty_cache() + + pbar = ProgressBar(5) + + # Detect model is existing or not + possible_folders = ["CogVideoX_Fun", "Fun_Models", "VideoX_Fun", "Wan-AI"] + \ + [os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "models/Diffusion_Transformer")] # Possible folder names to check + # Initialize model_name as None + model_name = search_model_in_possible_folders(possible_folders, model) + + print("Loading VAE...") + vae = AutoencoderKLFlux2.from_pretrained( + model_name, + subfolder="vae" + ).to(weight_dtype) + pbar.update(1) + + print("Loading Scheduler...") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model_name, + subfolder="scheduler" + ) + pbar.update(1) + + print("Loading Transformer...") + transformer = Flux2Transformer2DModel.from_pretrained( + model_name, + subfolder="transformer", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) + pbar.update(1) + + print("Loading Tokenizer...") + tokenizer = PixtralProcessor.from_pretrained( + model_name, + subfolder="tokenizer" + ) + pbar.update(1) + + print("Loading Text Encoder...") + text_encoder = Mistral3ForConditionalGeneration.from_pretrained( + model_name, + subfolder="text_encoder", + torch_dtype=weight_dtype, + low_cpu_mem_usage=True, + ) + pbar.update(1) + + model_type = "Inpaint" + pipeline = Flux2Pipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + ) + + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + + if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) + elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device=offload_device, offload_type="leaf_level", use_stream=True) + elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) + else: + pipeline.to(device=device) + + funmodels = { + 'pipeline': pipeline, + 'GPU_memory_mode': GPU_memory_mode, + 'dtype': weight_dtype, + 'model_name': model_name, + 'model_type': model_type, + 'loras': [], + 'strength_model': [] + } + return (funmodels,) + + +class LoadFlux2Lora: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "funmodels": ("FunModels",), + "lora_name": (folder_paths.get_filename_list("loras"), {"default": None,}), + "strength_model": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}), + "lora_cache": ([False, True], {"default": False,}), + } + } + RETURN_TYPES = ("FunModels",) + RETURN_NAMES = ("funmodels",) + FUNCTION = "load_lora" + CATEGORY = "CogVideoXFUNWrapper" + + def load_lora(self, funmodels, lora_name, strength_model, lora_cache): + new_funmodels = dict(funmodels) + if lora_name is not None: + loras = list(new_funmodels.get("loras", [])) + [folder_paths.get_full_path("loras", lora_name)] + strength_models = list(new_funmodels.get("strength_model", [])) + [strength_model] + new_funmodels['loras'] = loras + new_funmodels['strength_model'] = strength_models + new_funmodels['lora_cache'] = lora_cache + return (new_funmodels,) + + +class LoadFlux2ControlNetInPipeline: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "config": ( + [ + "flux2/flux2_control.yaml", + ], + { + "default": "flux2/flux2_control.yaml", + } + ), + "model_name": ( + folder_paths.get_filename_list("model_patches"), + {"default": "FLUX.2-dev-Fun-Controlnet-Union.safetensors", }, + ), + "funmodels": ("FunModels",), + }, + } + + RETURN_TYPES = ("FunModels",) + RETURN_NAMES = ("funmodels",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, config, model_name, funmodels): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + GPU_memory_mode = funmodels["GPU_memory_mode"] + weight_dtype = funmodels['dtype'] + + # Remove hooks + funmodels["pipeline"].remove_all_hooks() + safe_remove_group_offloading(funmodels["pipeline"]) + + # Get Transformer + transformer = funmodels["pipeline"].transformer + transformer = transformer.cpu() + + # Get state_dict + transformer_state_dict = transformer.state_dict() + del transformer + mm.soft_empty_cache() + gc.collect() + + # Load config + config_path = f"{script_directory}/config/{config}" + config = OmegaConf.load(config_path) + kwargs = { + "_class_name": "Flux2Transformer2DModel", + "_diffusers_version": "0.36.0.dev0", + "attention_head_dim": 128, + "axes_dims_rope": [ + 32, + 32, + 32, + 32 + ], + "eps": 1e-06, + "in_channels": 128, + "joint_attention_dim": 15360, + "mlp_ratio": 3.0, + "num_attention_heads": 48, + "num_layers": 8, + "num_single_layers": 48, + "out_channels": None, + "patch_size": 1, + "rope_theta": 2000, + "timestep_guidance_channels": 256 + } + kwargs.update(OmegaConf.to_container(config['transformer_additional_kwargs'])) + + # Get Model + sig = inspect.signature(Flux2ControlTransformer2DModel) + accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} + with accelerate.init_empty_weights(): + control_transformer = Flux2ControlTransformer2DModel.from_config(accepted).to(weight_dtype) + print(f"Load Flux Control Transformer") + + # Load Control state_dict + control_model_path = folder_paths.get_full_path("model_patches", model_name) + if control_model_path.endswith(".safetensors"): + control_state_dict = load_file(control_model_path) + else: + control_state_dict = torch.load(control_model_path) + + state_dict = {**transformer_state_dict, **control_state_dict} + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + control_transformer, + state_dict, + dtype=weight_dtype, + model_name_or_path="", + ) + else: + control_transformer._convert_deprecated_attention_blocks(state_dict) + load_model_dict_into_meta( + control_transformer, + state_dict, + device=offload_device, + dtype=weight_dtype, + model_name_or_path="", + ) + + # Create Pipeline + pipeline = Flux2ControlPipeline( + vae=funmodels["pipeline"].vae, + tokenizer=funmodels["pipeline"].tokenizer, + text_encoder=funmodels["pipeline"].text_encoder, + transformer=control_transformer, + scheduler=funmodels["pipeline"].scheduler, + ) + del funmodels["pipeline"] + mm.soft_empty_cache() + gc.collect() + + # Apply GPU memory mode + if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) + elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device=offload_device, offload_type="leaf_level", use_stream=True) + elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(control_transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(control_transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(control_transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(control_transformer, weight_dtype) + pipeline.to(device=device) + else: + pipeline.to(device=device) + funmodels["pipeline"] = pipeline + funmodels["model_type"] = "Control" + return (funmodels, ) + + +class LoadFlux2ControlNetInModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "config": ( + [ + "flux2/flux2_control.yaml", + ], + { + "default": "flux2/flux2_control.yaml", + } + ), + "model_name": ( + folder_paths.get_filename_list("model_patches"), + {"default": "FLUX.2-dev-Fun-Controlnet-Union.safetensors", }, + ), + "transformer": ("TransformerModel",), + }, + } + + RETURN_TYPES = ("TransformerModel",) + RETURN_NAMES = ("transformer",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, config, model_name, transformer): + offload_device = mm.unet_offload_device() + dtype = transformer.dtype + + # Get Transformer + transformer = transformer.cpu() + + # Get state_dict + transformer_state_dict = transformer.state_dict() + del transformer + mm.soft_empty_cache() + gc.collect() + + # Load config + config_path = f"{script_directory}/config/{config}" + config = OmegaConf.load(config_path) + kwargs = { + "_class_name": "Flux2Transformer2DModel", + "_diffusers_version": "0.36.0.dev0", + "attention_head_dim": 128, + "axes_dims_rope": [ + 32, + 32, + 32, + 32 + ], + "eps": 1e-06, + "in_channels": 128, + "joint_attention_dim": 15360, + "mlp_ratio": 3.0, + "num_attention_heads": 48, + "num_layers": 8, + "num_single_layers": 48, + "out_channels": None, + "patch_size": 1, + "rope_theta": 2000, + "timestep_guidance_channels": 256 + } + kwargs.update(OmegaConf.to_container(config['transformer_additional_kwargs'])) + + # Get Model + sig = inspect.signature(Flux2ControlTransformer2DModel) + accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} + with accelerate.init_empty_weights(): + control_transformer = Flux2ControlTransformer2DModel.from_config(accepted).to(dtype) + print(f"Load Flux Control Transformer") + + # Load Control state_dict + control_model_path = folder_paths.get_full_path("model_patches", model_name) + if control_model_path.endswith(".safetensors"): + control_state_dict = load_file(control_model_path) + else: + control_state_dict = torch.load(control_model_path) + + state_dict = {**transformer_state_dict, **control_state_dict} + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + control_transformer, + state_dict, + dtype=dtype, + model_name_or_path="", + ) + else: + control_transformer._convert_deprecated_attention_blocks(state_dict) + load_model_dict_into_meta( + control_transformer, + state_dict, + device=offload_device, + dtype=dtype, + model_name_or_path="", + ) + return (control_transformer, ) + + +class Flux2T2ISampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "funmodels": ("FunModels",), + "prompt": ("STRING_PROMPT",), + "width": ("INT", {"default": 1728, "min": 64, "max": 2048, "step": 16}), + "height": ("INT", {"default": 992, "min": 64, "max": 2048, "step": 16}), + "seed": ("INT", {"default": 43, "min": 0, "max": 0xffffffffffffffff}), + "steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}), + "cfg": ("FLOAT", {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.01}), + "scheduler": ( + ["Flow", "Flow_Unipc", "Flow_DPM++"], + {"default": 'Flow'} + ), + "shift": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "process" + CATEGORY = "CogVideoXFUNWrapper" + + def process( + self, + funmodels, + prompt, + width, + height, + seed, + steps, + cfg, + scheduler, + shift, + ): + global transformer_cpu_cache + global lora_path_before + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.soft_empty_cache() + gc.collect() + + # Get Pipeline + pipeline = funmodels['pipeline'] + model_name = funmodels['model_name'] + weight_dtype = funmodels['dtype'] + + # Load Sampler + pipeline.scheduler = get_flux2_scheduler(scheduler, shift) + + generator = torch.Generator(device).manual_seed(seed) + + with torch.no_grad(): + # Apply lora + if funmodels.get("lora_cache", False): + if len(funmodels.get("loras", [])) != 0: + # Save the original weights to cpu + if len(transformer_cpu_cache) == 0: + print('Save transformer state_dict to cpu memory') + transformer_state_dict = pipeline.transformer.state_dict() + for key in transformer_state_dict: + transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() + + lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) + if lora_path_now != lora_path_before: + print('Merge Lora with Cache') + lora_path_before = copy.deepcopy(lora_path_now) + pipeline.transformer.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + else: + print('Merge Lora') + # Clear lora when switch from lora_cache=True to lora_cache=False. + if len(transformer_cpu_cache) != 0: + pipeline.transformer.load_state_dict(transformer_cpu_cache) + transformer_cpu_cache = {} + lora_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + sample = pipeline( + prompt=prompt, + height=height, + width=width, + generator=generator, + guidance_scale=cfg, + num_inference_steps=steps, + comfyui_progressbar=True, + ).images + + image = torch.Tensor(np.array(sample[0])).unsqueeze(0) / 255 + + if not funmodels.get("lora_cache", False): + print('Unmerge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + return (image,) + + +class Flux2ControlSampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "funmodels": ( + "FunModels", + ), + "prompt": ( + "STRING_PROMPT", + ), + "width": ( + "INT", {"default": 992, "min": 64, "max": 20480, "step": 16} + ), + "height": ( + "INT", {"default": 1728, "min": 64, "max": 20480, "step": 16} + ), + "seed": ( + "INT", {"default": 43, "min": 0, "max": 0xffffffffffffffff} + ), + "steps": ( + "INT", {"default": 50, "min": 1, "max": 200, "step": 1} + ), + "cfg": ( + "FLOAT", {"default": 4.0, "min": 0.0, "max": 20.0, "step": 0.01} + ), + "scheduler": ( + ["Flow", "Flow_Unipc", "Flow_DPM++"], + { + "default": 'Flow' + } + ), + "shift": ( + "INT", {"default": 1, "min": 1, "max": 100, "step": 1} + ), + "control_context_scale": ( + "FLOAT", {"default": 0.75, "min": 0.0, "max": 2.0, "step": 0.01} + ), + }, + "optional":{ + "control_image": ("IMAGE",), + "inpaint_image": ("IMAGE",), + "mask_image": ("IMAGE",), + "image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES =("images",) + FUNCTION = "process" + CATEGORY = "CogVideoXFUNWrapper" + + def process(self, funmodels, prompt, width, height, seed, steps, cfg, scheduler, shift, control_context_scale, control_image=None, inpaint_image=None, mask_image=None, image=None): + global transformer_cpu_cache + global lora_path_before + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.soft_empty_cache() + gc.collect() + + # Get Pipeline + pipeline = funmodels['pipeline'] + model_name = funmodels['model_name'] + weight_dtype = funmodels['dtype'] + sample_size = [height, width] + + # Load Sampler + pipeline.scheduler = get_flux2_scheduler(scheduler, shift) + + generator = torch.Generator(device).manual_seed(seed) + + with torch.no_grad(): + # Apply lora + if funmodels.get("lora_cache", False): + if len(funmodels.get("loras", [])) != 0: + # Save the original weights to cpu + if len(transformer_cpu_cache) == 0: + print('Save transformer state_dict to cpu memory') + transformer_state_dict = pipeline.transformer.state_dict() + for key in transformer_state_dict: + transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() + + lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) + if lora_path_now != lora_path_before: + print('Merge Lora with Cache') + lora_path_before = copy.deepcopy(lora_path_now) + pipeline.transformer.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + else: + print('Merge Lora') + # Clear lora when switch from lora_cache=True to lora_cache=False. + if len(transformer_cpu_cache) != 0: + pipeline.transformer.load_state_dict(transformer_cpu_cache) + transformer_cpu_cache = {} + lora_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + # Process images + if inpaint_image is not None: + inpaint_image = [to_pil(inpaint_image) for inpaint_image in inpaint_image][0] + inpaint_image = get_image_latent(inpaint_image, sample_size=sample_size)[:, :, 0] + else: + inpaint_image = torch.zeros([1, 3, sample_size[0], sample_size[1]]) + + if mask_image is not None: + mask_image = [to_pil(mask_image) for mask_image in mask_image][0] + mask_image = get_image_latent(mask_image, sample_size=sample_size)[:, :1, 0] + else: + mask_image = torch.ones([1, 1, sample_size[0], sample_size[1]]) * 255 + + if control_image is not None: + control_image = [to_pil(control_image) for control_image in control_image][0] + control_image = get_image_latent(control_image, sample_size=sample_size)[:, :, 0] + + if image is not None: + image = [to_pil(image) for image in image] + + # Generate + sample = pipeline( + prompt = prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = cfg, + image = image, + inpaint_image = inpaint_image, + mask_image = mask_image, + control_image = control_image, + num_inference_steps = steps, + control_context_scale = control_context_scale, + comfyui_progressbar = True, + ).images + image = torch.Tensor(np.array(sample[0])).unsqueeze(0) / 255 + + # Unmerge lora + if not funmodels.get("lora_cache", False): + print('Unmerge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + return (image,) diff --git a/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i.json b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i.json new file mode 100644 index 0000000..4db8849 --- /dev/null +++ b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i.json @@ -0,0 +1,452 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 107, + "last_link_id": 112, + "nodes": [ + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 260.6739960937499, + -1.205332031249991 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 105 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 99, + "type": "LoadFlux2VAEModel", + "pos": [ + 766.8110625597004, + -472.8383839778354 + ], + "size": [ + 366.8770429687502, + 82 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 89 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2VAEModel" + }, + "widgets_values": [ + "flux2-vae.safetensors", + "bf16" + ] + }, + { + "id": 78, + "type": "Note", + "pos": [ + 24.634203125000003, + -2.051003906249974 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 98, + "type": "CombineFlux2Pipeline", + "pos": [ + 765.0797227159495, + -326.5219699153354 + ], + "size": [ + 370.22844140624966, + 142 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 92 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 89 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 107 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 108 + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 93 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 104 + ] + } + ], + "properties": { + "Node name for S&R": "CombineFlux2Pipeline" + }, + "widgets_values": [ + "", + "sequential_cpu_offload" + ] + }, + { + "id": 105, + "type": "Flux2T2ISampler", + "pos": [ + 726.9954453721992, + -55.18514179033538 + ], + "size": [ + 296.21252343749995, + 256.0413164062501 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 104 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 105 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 112 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2T2ISampler" + }, + "widgets_values": [ + 1728, + 992, + 599444093270696, + "randomize", + 25, + 4, + "Flow", + 3 + ] + }, + { + "id": 106, + "type": "LoadFlux2TextEncoderModel", + "pos": [ + 294.0149063096997, + -294.99910663408474 + ], + "size": [ + 436.84259765624995, + 102 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 107 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 108 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TextEncoderModel" + }, + "widgets_values": [ + "mistral_3_small_flux2_bf16.safetensors", + "bf16" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 112 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 100, + "type": "LoadFlux2TransformerModel", + "pos": [ + 313.73034380969943, + -465.4970871028352 + ], + "size": [ + 407.57818359374994, + 102 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 92 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 93 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TransformerModel" + }, + "widgets_values": [ + "flux2_dev_fp8mixed.safetensors", + "bf16" + ] + } + ], + "links": [ + [ + 89, + 99, + 0, + 98, + 1, + "VAEModel" + ], + [ + 92, + 100, + 0, + 98, + 0, + "TransformerModel" + ], + [ + 93, + 100, + 1, + 98, + 4, + "STRING" + ], + [ + 104, + 98, + 0, + 105, + 0, + "FunModels" + ], + [ + 105, + 75, + 0, + 105, + 1, + "STRING_PROMPT" + ], + [ + 107, + 106, + 0, + 98, + 2, + "TextEncoderModel" + ], + [ + 108, + 106, + 1, + 98, + 3, + "Tokenizer" + ], + [ + 112, + 105, + 0, + 88, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 985.5581665039062, + 393.7902526855469 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 228.67399609375008, + -78.20533203124995, + 449.1265312500001, + 257.2759179687498 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6806385776128845, + "offset": [ + 720.7131297880804, + 715.5354021239666 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75", + "comfy-core": "0.10.0" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_control.json b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_control.json new file mode 100644 index 0000000..6855398 --- /dev/null +++ b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_control.json @@ -0,0 +1,568 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 112, + "last_link_id": 126, + "nodes": [ + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 24.634203125000003, + -2.051003906249974 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 106, + "type": "LoadFlux2TextEncoderModel", + "pos": [ + 294.0149063096997, + -294.99910663408474 + ], + "size": [ + 436.84259765624995, + 102 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 107 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 108 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TextEncoderModel" + }, + "widgets_values": [ + "mistral_3_small_flux2_bf16.safetensors", + "bf16" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 126 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 100, + "type": "LoadFlux2TransformerModel", + "pos": [ + 313.73034380969943, + -465.4970871028352 + ], + "size": [ + 407.57818359374994, + 102 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 113 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 93 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TransformerModel" + }, + "widgets_values": [ + "flux2_dev_fp8mixed.safetensors", + "bf16" + ] + }, + { + "id": 108, + "type": "LoadFlux2ControlNetInModel", + "pos": [ + 778.396083325503, + -461.64928225980117 + ], + "size": [ + 545.4199394818678, + 82 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 113 + } + ], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 114 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2ControlNetInModel" + }, + "widgets_values": [ + "flux2/flux2_control.yaml", + "FLUX.2-dev-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 99, + "type": "LoadFlux2VAEModel", + "pos": [ + 1157.6607036173684, + -291.70678153123094 + ], + "size": [ + 378.0969752517299, + 82 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 89 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2VAEModel" + }, + "widgets_values": [ + "flux2-vae.safetensors", + "bf16" + ] + }, + { + "id": 98, + "type": "CombineFlux2Pipeline", + "pos": [ + 765.0797227159495, + -326.5219699153354 + ], + "size": [ + 370.22844140624966, + 142 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 114 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 89 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 107 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 108 + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 93 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 123 + ] + } + ], + "properties": { + "Node name for S&R": "CombineFlux2Pipeline" + }, + "widgets_values": [ + "", + "sequential_cpu_offload" + ] + }, + { + "id": 110, + "type": "LoadImage", + "pos": [ + 401.58427058709873, + 257.2700309330294 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 125 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "a7kXeQ5l9Dhspes7q3x3G (1).png", + "image" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 260.6739960937499, + -1.205332031249991 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 124 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 112, + "type": "Flux2ControlSampler", + "pos": [ + 726.6810910994532, + -44.26206223716924 + ], + "size": [ + 286.765625, + 330 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 123 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 124 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": 125 + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 126 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2ControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 43, + "randomize", + 25, + 4, + "Flow", + 3, + 0.75 + ] + } + ], + "links": [ + [ + 89, + 99, + 0, + 98, + 1, + "VAEModel" + ], + [ + 93, + 100, + 1, + 98, + 4, + "STRING" + ], + [ + 107, + 106, + 0, + 98, + 2, + "TextEncoderModel" + ], + [ + 108, + 106, + 1, + 98, + 3, + "Tokenizer" + ], + [ + 113, + 100, + 0, + 108, + 0, + "TransformerModel" + ], + [ + 114, + 108, + 0, + 98, + 0, + "TransformerModel" + ], + [ + 123, + 98, + 0, + 112, + 0, + "FunModels" + ], + [ + 124, + 75, + 0, + 112, + 1, + "STRING_PROMPT" + ], + [ + 125, + 110, + 0, + 112, + 2, + "IMAGE" + ], + [ + 126, + 112, + 0, + 88, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1324.985509221171, + 399.185002734652 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 228.67399609375008, + -78.20533203124995, + 449.1265312500001, + 257.2759179687498 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6806385776128846, + "offset": [ + 597.5416992665371, + 751.7438792427187 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75", + "comfy-core": "0.10.0" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_control_ref.json b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_control_ref.json new file mode 100644 index 0000000..fb06ea5 --- /dev/null +++ b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_control_ref.json @@ -0,0 +1,613 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 114, + "last_link_id": 131, + "nodes": [ + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 24.634203125000003, + -2.051003906249974 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 99, + "type": "LoadFlux2VAEModel", + "pos": [ + 1157.6607036173684, + -291.70678153123094 + ], + "size": [ + 378.0969752517299, + 82 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 89 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2VAEModel" + }, + "widgets_values": [ + "flux2-vae.safetensors", + "bf16" + ] + }, + { + "id": 98, + "type": "CombineFlux2Pipeline", + "pos": [ + 765.0797227159495, + -326.5219699153354 + ], + "size": [ + 370.22844140624966, + 142 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 114 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 89 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 107 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 108 + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 93 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 127 + ] + } + ], + "properties": { + "Node name for S&R": "CombineFlux2Pipeline" + }, + "widgets_values": [ + "", + "sequential_cpu_offload" + ] + }, + { + "id": 110, + "type": "LoadImage", + "pos": [ + 401.58427058709873, + 257.2700309330294 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 129 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "pose.jpg", + "image" + ] + }, + { + "id": 114, + "type": "LoadImage", + "pos": [ + 65.13447511647426, + 249.3438179015018 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 131 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "8.png", + "image" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 260.6739960937499, + -1.205332031249991 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 128 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "This is a panoramic portrait photo of a young woman. She has flowing long hair and a soft lavender like color. She is wearing a white sleeveless dress with a blue ribbon bow tied around the collar. She has a confident posture, with her left hand naturally hanging down and her right hand in her pocket, and her legs slightly apart. Look straight at the camera. The sea breeze gently brushed her long hair, and they stood on the sunny seaside path, surrounded by blooming purple seaside flowers and smooth pebbles, with the sparkling sea and blue sky behind them. The screen presents a bright summer atmosphere, with soft and natural lighting, realistic details, and 8K ultra high definition image quality, clearly presenting fine textures such as clothing and hair. " + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 130 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 108, + "type": "LoadFlux2ControlNetInModel", + "pos": [ + 778.396083325503, + -461.64928225980117 + ], + "size": [ + 545.4199394818678, + 82 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 113 + } + ], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 114 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2ControlNetInModel" + }, + "widgets_values": [ + "flux2/flux2_control.yaml", + "FLUX.2-dev-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 106, + "type": "LoadFlux2TextEncoderModel", + "pos": [ + 294.0149063096997, + -294.99910663408474 + ], + "size": [ + 436.84259765624995, + 102 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 107 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 108 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TextEncoderModel" + }, + "widgets_values": [ + "mistral_3_small_flux2_bf16.safetensors", + "bf16" + ] + }, + { + "id": 100, + "type": "LoadFlux2TransformerModel", + "pos": [ + 313.73034380969943, + -465.4970871028352 + ], + "size": [ + 407.57818359374994, + 102 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 113 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 93 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TransformerModel" + }, + "widgets_values": [ + "flux2_dev_fp8mixed.safetensors", + "bf16" + ] + }, + { + "id": 113, + "type": "Flux2ControlSampler", + "pos": [ + 726.6810910994532, + -44.26206223716924 + ], + "size": [ + 286.765625, + 350 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 127 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 128 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": 129 + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 131 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 130 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2ControlSampler" + }, + "widgets_values": [ + 992, + 1728, + 43, + "fixed", + 25, + 4, + "Flow", + 3, + 0.75 + ] + } + ], + "links": [ + [ + 89, + 99, + 0, + 98, + 1, + "VAEModel" + ], + [ + 93, + 100, + 1, + 98, + 4, + "STRING" + ], + [ + 107, + 106, + 0, + 98, + 2, + "TextEncoderModel" + ], + [ + 108, + 106, + 1, + 98, + 3, + "Tokenizer" + ], + [ + 113, + 100, + 0, + 108, + 0, + "TransformerModel" + ], + [ + 114, + 108, + 0, + 98, + 0, + "TransformerModel" + ], + [ + 127, + 98, + 0, + 113, + 0, + "FunModels" + ], + [ + 128, + 75, + 0, + 113, + 1, + "STRING_PROMPT" + ], + [ + 129, + 110, + 0, + 113, + 2, + "IMAGE" + ], + [ + 130, + 113, + 0, + 88, + 0, + "IMAGE" + ], + [ + 131, + 114, + 0, + 113, + 5, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1324.985509221171, + 399.185002734652 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 228.67399609375008, + -78.20533203124995, + 449.1265312500001, + 257.2759179687498 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7159450505734999, + "offset": [ + 362.4772087682636, + 717.3189141746142 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75", + "comfy-core": "0.10.0" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_inpaint.json b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_inpaint.json new file mode 100644 index 0000000..f0caded --- /dev/null +++ b/comfyui/flux2/v1/flux2_chunked_loading_workflow_t2i_inpaint.json @@ -0,0 +1,795 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 115, + "last_link_id": 130, + "nodes": [ + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 24.634203125000003, + -2.051003906249974 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 106, + "type": "LoadFlux2TextEncoderModel", + "pos": [ + 294.0149063096997, + -294.99910663408474 + ], + "size": [ + 436.84259765624995, + 102 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 107 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 108 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TextEncoderModel" + }, + "widgets_values": [ + "mistral_3_small_flux2_bf16.safetensors", + "bf16" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 126 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 100, + "type": "LoadFlux2TransformerModel", + "pos": [ + 313.73034380969943, + -465.4970871028352 + ], + "size": [ + 407.57818359374994, + 102 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 113 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 93 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2TransformerModel" + }, + "widgets_values": [ + "flux2_dev_fp8mixed.safetensors", + "bf16" + ] + }, + { + "id": 108, + "type": "LoadFlux2ControlNetInModel", + "pos": [ + 778.396083325503, + -461.64928225980117 + ], + "size": [ + 545.4199394818678, + 82 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 113 + } + ], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 114 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2ControlNetInModel" + }, + "widgets_values": [ + "flux2/flux2_control.yaml", + "FLUX.2-dev-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 99, + "type": "LoadFlux2VAEModel", + "pos": [ + 1157.6607036173684, + -291.70678153123094 + ], + "size": [ + 378.0969752517299, + 82 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 89 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2VAEModel" + }, + "widgets_values": [ + "flux2-vae.safetensors", + "bf16" + ] + }, + { + "id": 98, + "type": "CombineFlux2Pipeline", + "pos": [ + 765.0797227159495, + -326.5219699153354 + ], + "size": [ + 370.22844140624966, + 142 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 114 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 89 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 107 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 108 + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 93 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 123 + ] + } + ], + "properties": { + "Node name for S&R": "CombineFlux2Pipeline" + }, + "widgets_values": [ + "", + "sequential_cpu_offload" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 260.6739960937499, + -1.205332031249991 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 124 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 113, + "type": "PreviewImage", + "pos": [ + 604.9433567807472, + 356.8772558483988 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 127 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 114, + "type": "82cb7f31-667b-4bc2-aeef-5a46b2680f64", + "pos": [ + 422.53187751146385, + 359.0052922744091 + ], + "size": [ + 140, + 26 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 128 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 127, + 130 + ] + } + ], + "properties": { + "proxyWidgets": [] + }, + "widgets_values": [] + }, + { + "id": 115, + "type": "LoadImage", + "pos": [ + 117.62424527510734, + 358.76820232661834 + ], + "size": [ + 270, + 314.00000000000006 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 129 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 128 + ] + } + ], + "properties": { + "Node name for S&R": "LoadImage", + "image": "clipspace/clipspace-painted-masked-1766731857414.png [input]" + }, + "widgets_values": [ + "clipspace/clipspace-painted-masked-1766731857414.png [input]", + "image" + ] + }, + { + "id": 112, + "type": "Flux2ControlSampler", + "pos": [ + 726.6810910994532, + -44.26206223716924 + ], + "size": [ + 286.765625, + 330 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 123 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 124 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": 129 + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": 130 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 126 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2ControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 43, + "randomize", + 25, + 4, + "Flow", + 3, + 0.75 + ] + } + ], + "links": [ + [ + 89, + 99, + 0, + 98, + 1, + "VAEModel" + ], + [ + 93, + 100, + 1, + 98, + 4, + "STRING" + ], + [ + 107, + 106, + 0, + 98, + 2, + "TextEncoderModel" + ], + [ + 108, + 106, + 1, + 98, + 3, + "Tokenizer" + ], + [ + 113, + 100, + 0, + 108, + 0, + "TransformerModel" + ], + [ + 114, + 108, + 0, + 98, + 0, + "TransformerModel" + ], + [ + 123, + 98, + 0, + 112, + 0, + "FunModels" + ], + [ + 124, + 75, + 0, + 112, + 1, + "STRING_PROMPT" + ], + [ + 126, + 112, + 0, + 88, + 0, + "IMAGE" + ], + [ + 127, + 114, + 0, + 113, + 0, + "IMAGE" + ], + [ + 128, + 115, + 1, + 114, + 0, + "MASK" + ], + [ + 129, + 115, + 0, + 112, + 3, + "IMAGE" + ], + [ + 130, + 114, + 0, + 112, + 4, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1324.985509221171, + 399.185002734652 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 228.67399609375008, + -78.20533203124995, + 449.1265312500001, + 257.2759179687498 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "definitions": { + "subgraphs": [ + { + "id": "82cb7f31-667b-4bc2-aeef-5a46b2680f64", + "version": 1, + "state": { + "lastGroupId": 2, + "lastNodeId": 102, + "lastLinkId": 94, + "lastRerouteId": 0 + }, + "revision": 0, + "config": {}, + "name": "New Subgraph", + "inputNode": { + "id": -10, + "bounding": [ + 286.83779809080124, + 270.7433175266759, + 120, + 60 + ] + }, + "outputNode": { + "id": -20, + "bounding": [ + 666.8377980908012, + 270.7433175266759, + 120, + 60 + ] + }, + "inputs": [ + { + "id": "aca65ac3-68da-44e9-94f6-1444ee807b81", + "name": "mask", + "type": "MASK", + "linkIds": [ + 91 + ], + "localized_name": "遮罩", + "pos": [ + 55, + 20 + ] + } + ], + "outputs": [ + { + "id": "b96f1481-8afd-4089-8a3d-dcc523a2178f", + "name": "IMAGE", + "type": "IMAGE", + "linkIds": [ + 92, + 93 + ], + "localized_name": "图像", + "pos": [ + 20, + 20 + ] + } + ], + "widgets": [], + "nodes": [ + { + "id": 101, + "type": "MaskToImage", + "pos": [ + 466.83779809080124, + 302.7433175266759 + ], + "size": [ + 140, + 26 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [ + { + "localized_name": "遮罩", + "name": "mask", + "type": "MASK", + "link": 91 + } + ], + "outputs": [ + { + "localized_name": "图像", + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 92, + 93 + ] + } + ], + "properties": { + "Node name for S&R": "MaskToImage" + }, + "widgets_values": [] + } + ], + "groups": [], + "links": [ + { + "id": 91, + "origin_id": -10, + "origin_slot": 0, + "target_id": 101, + "target_slot": 0, + "type": "MASK" + }, + { + "id": 92, + "origin_id": 101, + "origin_slot": 0, + "target_id": -20, + "target_slot": 0, + "type": "IMAGE" + }, + { + "id": 93, + "origin_id": 101, + "origin_slot": 0, + "target_id": -20, + "target_slot": 0, + "type": "IMAGE" + } + ], + "extra": { + "workflowRendererVersion": "LG" + } + } + ] + }, + "config": {}, + "extra": { + "ds": { + "scale": 0.6806385776128846, + "offset": [ + 664.046341361251, + 735.4620644668554 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75", + "comfy-core": "0.10.0" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_workflow_t2i.json b/comfyui/flux2/v1/flux2_workflow_t2i.json new file mode 100644 index 0000000..003727c --- /dev/null +++ b/comfyui/flux2/v1/flux2_workflow_t2i.json @@ -0,0 +1,274 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 94, + "last_link_id": 77, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 76 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1068.092196899791, + -87.86407614285447 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 77 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 92, + "type": "LoadFlux2Model", + "pos": [ + 302.5105757682986, + -291.5089065020655 + ], + "size": [ + 340.5369411044164, + 114.70675475775249 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 75 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2Model" + }, + "widgets_values": [ + "FLUX.2-dev", + "sequential_cpu_offload", + "bf16" + ] + }, + { + "id": 94, + "type": "Flux2T2ISampler", + "pos": [ + 724.3815336609828, + -86.90266446701813 + ], + "size": [ + 296.597265625, + 342 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 75 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 76 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 77 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2T2ISampler" + }, + "widgets_values": [ + 1728, + 992, + 404922577542089, + "randomize", + 25, + 4, + "Flow", + 3 + ] + } + ], + "links": [ + [ + 75, + 92, + 0, + 94, + 0, + "FunModels" + ], + [ + 76, + 75, + 0, + 94, + 1, + "STRING_PROMPT" + ], + [ + 77, + 94, + 0, + 88, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 443.62835196237825, + 248.53632503063102 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7828871306993743, + "offset": [ + 615.0117311935076, + 663.8885747430894 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75", + "comfy-core": "0.10.0" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_workflow_t2i_control.json b/comfyui/flux2/v1/flux2_workflow_t2i_control.json new file mode 100644 index 0000000..1250ce0 --- /dev/null +++ b/comfyui/flux2/v1/flux2_workflow_t2i_control.json @@ -0,0 +1,390 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 99, + "last_link_id": 90, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 95, + "type": "LoadImage", + "pos": [ + 372.02490659065415, + 210.9970404718659 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 89 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "a7kXeQ5l9Dhspes7q3x3G (1).png", + "image" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1068.092196899791, + -87.86407614285447 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 90 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 92, + "type": "LoadFlux2Model", + "pos": [ + 302.5105757682986, + -291.5089065020655 + ], + "size": [ + 340.5369411044164, + 114.70675475775249 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2Model" + }, + "widgets_values": [ + "FLUX.2-dev", + "sequential_cpu_offload", + "bf16" + ] + }, + { + "id": 97, + "type": "LoadFlux2ControlNetInPipeline", + "pos": [ + 687.5636864647346, + -282.9468485282529 + ], + "size": [ + 450.0920822120196, + 85.47272281455344 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 81 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 87 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2ControlNetInPipeline" + }, + "widgets_values": [ + "flux2/flux2_control.yaml", + "FLUX.2-dev-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 88 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 99, + "type": "Flux2ControlSampler", + "pos": [ + 697.1217271030082, + -90.5350526983327 + ], + "size": [ + 286.765625, + 330 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 87 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 88 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": 89 + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 90 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2ControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 415216792885928, + "randomize", + 25, + 4, + "Flow", + 3, + 0.75 + ] + } + ], + "links": [ + [ + 81, + 92, + 0, + 97, + 0, + "FunModels" + ], + [ + 87, + 97, + 0, + 99, + 0, + "FunModels" + ], + [ + 88, + 75, + 0, + 99, + 1, + "STRING_PROMPT" + ], + [ + 89, + 95, + 0, + 99, + 2, + "IMAGE" + ], + [ + 90, + 99, + 0, + 88, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 973.0550255176311, + 230.3983562881154 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 443.62835196237825, + 248.53632503063102 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7828871306993743, + "offset": [ + 451.2000031409154, + 641.3907196125986 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "comfy-core": "0.10.0", + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_workflow_t2i_control_ref.json b/comfyui/flux2/v1/flux2_workflow_t2i_control_ref.json new file mode 100644 index 0000000..a9df604 --- /dev/null +++ b/comfyui/flux2/v1/flux2_workflow_t2i_control_ref.json @@ -0,0 +1,435 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 101, + "last_link_id": 92, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1068.092196899791, + -87.86407614285447 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 90 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 92, + "type": "LoadFlux2Model", + "pos": [ + 302.5105757682986, + -291.5089065020655 + ], + "size": [ + 340.5369411044164, + 114.70675475775249 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2Model" + }, + "widgets_values": [ + "FLUX.2-dev", + "sequential_cpu_offload", + "bf16" + ] + }, + { + "id": 97, + "type": "LoadFlux2ControlNetInPipeline", + "pos": [ + 687.5636864647346, + -282.9468485282529 + ], + "size": [ + 450.0920822120196, + 85.47272281455344 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 81 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 87 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2ControlNetInPipeline" + }, + "widgets_values": [ + "flux2/flux2_control.yaml", + "FLUX.2-dev-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 99, + "type": "Flux2ControlSampler", + "pos": [ + 697.1217271030082, + -90.5350526983327 + ], + "size": [ + 286.765625, + 350 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 87 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 88 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": 91 + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 92 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 90 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2ControlSampler" + }, + "widgets_values": [ + 992, + 1728, + 415216792885928, + "randomize", + 25, + 4, + "Flow", + 1, + 0.75 + ] + }, + { + "id": 101, + "type": "LoadImage", + "pos": [ + 76.43432472316154, + 195.04174265316996 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 92 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "8.png", + "image" + ] + }, + { + "id": 100, + "type": "LoadImage", + "pos": [ + 385.0208100997366, + 189.91725184250046 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 91 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "pose.jpg", + "image" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 88 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "This is a panoramic portrait photo of a young woman. She has flowing long hair and a soft lavender like color. She is wearing a white sleeveless dress with a blue ribbon bow tied around the collar. She has a confident posture, with her left hand naturally hanging down and her right hand in her pocket, and her legs slightly apart. Look straight at the camera. The sea breeze gently brushed her long hair, and they stood on the sunny seaside path, surrounded by blooming purple seaside flowers and smooth pebbles, with the sparkling sea and blue sky behind them. The screen presents a bright summer atmosphere, with soft and natural lighting, realistic details, and 8K ultra high definition image quality, clearly presenting fine textures such as clothing and hair. " + ] + } + ], + "links": [ + [ + 81, + 92, + 0, + 97, + 0, + "FunModels" + ], + [ + 87, + 97, + 0, + 99, + 0, + "FunModels" + ], + [ + 88, + 75, + 0, + 99, + 1, + "STRING_PROMPT" + ], + [ + 90, + 99, + 0, + 88, + 0, + "IMAGE" + ], + [ + 91, + 100, + 0, + 99, + 2, + "IMAGE" + ], + [ + 92, + 101, + 0, + 99, + 5, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 973.0550255176311, + 230.3983562881154 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 443.62835196237825, + 248.53632503063102 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7828871306993743, + "offset": [ + 579.865381328975, + 647.6525976761967 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "comfy-core": "0.10.0", + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/flux2/v1/flux2_workflow_t2i_inpaint.json b/comfyui/flux2/v1/flux2_workflow_t2i_inpaint.json new file mode 100644 index 0000000..1d138d0 --- /dev/null +++ b/comfyui/flux2/v1/flux2_workflow_t2i_inpaint.json @@ -0,0 +1,615 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 103, + "last_link_id": 97, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 92, + "type": "LoadFlux2Model", + "pos": [ + 302.5105757682986, + -291.5089065020655 + ], + "size": [ + 340.5369411044164, + 114.70675475775249 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2Model" + }, + "widgets_values": [ + "FLUX.2-dev", + "sequential_cpu_offload", + "bf16" + ] + }, + { + "id": 97, + "type": "LoadFlux2ControlNetInPipeline", + "pos": [ + 687.5636864647346, + -282.9468485282529 + ], + "size": [ + 450.0920822120196, + 85.47272281455344 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 81 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 87 + ] + } + ], + "properties": { + "Node name for S&R": "LoadFlux2ControlNetInPipeline" + }, + "widgets_values": [ + "flux2/flux2_control.yaml", + "FLUX.2-dev-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 88 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 99, + "type": "Flux2ControlSampler", + "pos": [ + 697.1217271030082, + -90.5350526983327 + ], + "size": [ + 286.765625, + 330 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 87 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 88 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": 94 + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": 97 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 90 + ] + } + ], + "properties": { + "Node name for S&R": "Flux2ControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 415216792885928, + "randomize", + 25, + 4, + "Flow", + 3, + 0.75 + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1068.092196899791, + -87.86407614285447 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 90 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 100, + "type": "LoadImage", + "pos": [ + 144.54659360458197, + 304.11286083504353 + ], + "size": [ + 270, + 314.00000000000006 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 94 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 95 + ] + } + ], + "properties": { + "Node name for S&R": "LoadImage", + "image": "clipspace/clipspace-painted-masked-1766731857414.png [input]" + }, + "widgets_values": [ + "clipspace/clipspace-painted-masked-1766731857414.png [input]", + "image" + ] + }, + { + "id": 103, + "type": "e19429a7-1aa4-43c9-8175-2536ab8b5433", + "pos": [ + 449.4542258409389, + 304.3499507828343 + ], + "size": [ + 140, + 26 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 95 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 96, + 97 + ] + } + ], + "properties": { + "proxyWidgets": [] + }, + "widgets_values": [] + }, + { + "id": 102, + "type": "PreviewImage", + "pos": [ + 631.8657051102216, + 302.221914356824 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 96 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + } + ], + "links": [ + [ + 81, + 92, + 0, + 97, + 0, + "FunModels" + ], + [ + 87, + 97, + 0, + 99, + 0, + "FunModels" + ], + [ + 88, + 75, + 0, + 99, + 1, + "STRING_PROMPT" + ], + [ + 90, + 99, + 0, + 88, + 0, + "IMAGE" + ], + [ + 94, + 100, + 0, + 99, + 3, + "IMAGE" + ], + [ + 95, + 100, + 1, + 103, + 0, + "MASK" + ], + [ + 96, + 103, + 0, + 102, + 0, + "IMAGE" + ], + [ + 97, + 103, + 0, + 99, + 4, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 973.0550255176311, + 230.3983562881154 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 443.62835196237825, + 248.53632503063102 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "definitions": { + "subgraphs": [ + { + "id": "e19429a7-1aa4-43c9-8175-2536ab8b5433", + "version": 1, + "state": { + "lastGroupId": 2, + "lastNodeId": 102, + "lastLinkId": 94, + "lastRerouteId": 0 + }, + "revision": 0, + "config": {}, + "name": "New Subgraph", + "inputNode": { + "id": -10, + "bounding": [ + 286.83779809080124, + 270.7433175266759, + 120, + 60 + ] + }, + "outputNode": { + "id": -20, + "bounding": [ + 666.8377980908012, + 270.7433175266759, + 120, + 60 + ] + }, + "inputs": [ + { + "id": "aca65ac3-68da-44e9-94f6-1444ee807b81", + "name": "mask", + "type": "MASK", + "linkIds": [ + 91 + ], + "localized_name": "遮罩", + "pos": [ + 55, + 20 + ] + } + ], + "outputs": [ + { + "id": "b96f1481-8afd-4089-8a3d-dcc523a2178f", + "name": "IMAGE", + "type": "IMAGE", + "linkIds": [ + 92, + 93 + ], + "localized_name": "图像", + "pos": [ + 20, + 20 + ] + } + ], + "widgets": [], + "nodes": [ + { + "id": 101, + "type": "MaskToImage", + "pos": [ + 466.83779809080124, + 302.7433175266759 + ], + "size": [ + 140, + 26 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [ + { + "localized_name": "遮罩", + "name": "mask", + "type": "MASK", + "link": 91 + } + ], + "outputs": [ + { + "localized_name": "图像", + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 92, + 93 + ] + } + ], + "properties": { + "Node name for S&R": "MaskToImage" + }, + "widgets_values": [] + } + ], + "groups": [], + "links": [ + { + "id": 91, + "origin_id": -10, + "origin_slot": 0, + "target_id": 101, + "target_slot": 0, + "type": "MASK" + }, + { + "id": 92, + "origin_id": 101, + "origin_slot": 0, + "target_id": -20, + "target_slot": 0, + "type": "IMAGE" + }, + { + "id": 93, + "origin_id": 101, + "origin_slot": 0, + "target_id": -20, + "target_slot": 0, + "type": "IMAGE" + } + ], + "extra": {} + } + ] + }, + "config": {}, + "extra": { + "ds": { + "scale": 0.7828871306993743, + "offset": [ + 474.4263317353785, + 530.1987254713293 + ] + }, + "frontendVersion": "1.37.11", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "VideoX-Fun": "7671af8b16701319cf043358941277b9a5a1cb75", + "comfy-core": "0.10.0" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/README.md b/comfyui/qwenimage/README.md index a0eb153..82f8789 100644 --- a/comfyui/qwenimage/README.md +++ b/comfyui/qwenimage/README.md @@ -2,13 +2,85 @@ ## a. Model Links and Storage Locations +**Chunked loading is recommended** as it better aligns with ComfyUI's standard workflow. + +### 1. Chunked Loading Weights (Recommended) + +For chunked loading, it is recommended to directly download the Qwen-Image weights provided by ComfyUI official. Please organize the files according to the following directory structure: + +**Core Model Files:** + +| Component | File Name | +|-----------|-----------| +| Text Encoder | [`qwen_2.5_vl_7b_fp8_scaled.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image_ComfyUI/resolve/main/split_files/text_encoders/qwen_2.5_vl_7b_fp8_scaled.safetensors) | +| Diffusion Model | [`qwen_image_fp8_e4m3fn.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image_ComfyUI/resolve/main/split_files/diffusion_models/qwen_image_fp8_e4m3fn.safetensors) | +| VAE | [`qwen_image_vae.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image_ComfyUI/resolve/main/split_files/vae/qwen_image_vae.safetensors) | +| tokenizer | [`tokenizer`](https://huggingface.co/Qwen/Qwen-Image-Edit/tree/main/tokenizer) | +| processor | [`processor`](https://huggingface.co/Qwen/Qwen-Image-Edit/tree/main/processor) | + +**ControlNet Model Files:** + +| Name | Storage | Hugging Face | Model Scope | Description | +|--|--|--|--|--| +| Qwen-Image-2512-Fun-Controlnet-Union | - | [🤗Link](https://huggingface.co/alibaba-pai/Qwen-Image-2512-Fun-Controlnet-Union) | [😄Link](https://modelscope.cn/models/PAI/Qwen-Image-2512-Fun-Controlnet-Union) | ControlNet weights for Qwen-Image-2512, supporting multiple control conditions such as Canny, Depth, Pose, MLSD, Scribble, etc. | + +**Storage Location:** + +``` +📂 ComfyUI/ +├── 📂 models/ +│ ├── 📂 text_encoders/ +│ │ └── qwen_2.5_vl_7b_fp8_scaled.safetensors +│ ├── 📂 diffusion_models/ +│ │ └── qwen_image_fp8_e4m3fn.safetensors` +│ ├── 📂 vae/ +│ │ └── qwen_image_vae.safetensors +│ ├── 📂 Fun_Models/ +│ │ ├── qwen2_tokenizer/ +│ │ └── qwen2_processor/ +│ └── 📂 model_patches/ +│ └── Qwen-Image-2512-Fun-Controlnet-Union.safetensors +``` + +### 2. Preprocessing Weights (Optional) + +If you want to use the control preprocessing nodes, you can download the preprocessing weights to `ComfyUI/custom_nodes/Fun_Models/Third_Party/`. + +**Required Files:** + +| File Name | Download Link | Purpose | +|-----------|---------------|---------| +| `yolox_l.onnx` | [Download](https://huggingface.co/yzd-v/DWPose/resolve/main/yolox_l.onnx) | YOLO Detection Model | +| `dw-ll_ucoco_384.onnx` | [Download](https://huggingface.co/yzd-v/DWPose/resolve/main/dw-ll_ucoco_384.onnx) | DWPose Pose Estimation Model | +| `ZoeD_M12_N.pt` | [Download](https://huggingface.co/lllyasviel/Annotators/resolve/main/ZoeD_M12_N.pt) | ZoeDepth Depth Estimation Model | + +**Storage Location:** + +``` +📂 ComfyUI/ +├── 📂 models/ +│ └── 📂 Fun_Models/ +│ └── 📂 Third_Party +│ ├── yolox_l.onnx +│ ├── dw-ll_ucoco_384.onnx +│ └── ZoeD_M12_N.pt +``` + +### 3. Full Model Loading (Optional) + +If you prefer full model loading, you can directly download the diffusers weights. + **Required Files:** | Name | Storage | Hugging Face | Model Scope | Description | |--|--|--|--|--| | Qwen-Image | [🤗Link](https://huggingface.co/Qwen/Qwen-Image) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image) | Official Qwen-Image weights | +| Qwen-Image-2512 | [🤗Link](https://huggingface.co/Qwen/Qwen-Image-2512) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image-2512) | Official Qwen-Image weights | | Qwen-Image-Edit | [🤗Link](https://huggingface.co/Qwen/Qwen-Image-Edit) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image-Edit) | Official Qwen-Image-Edit weights | | Qwen-Image-Edit-2509 | [🤗Link](https://huggingface.co/Qwen/Qwen-Image-Edit-2509) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image-Edit-2509) | Official Qwen-Image-Edit-2509 weights | +| Qwen-Image-Edit-2511 | [🤗Link](https://huggingface.co/Qwen/Qwen-Image-Edit-2511) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image-Edit-2511) | Official Qwen-Image-Edit-2511 weights | + +For full model loading, use the diffusers version of Qwen-Image Turbo and place the model in `ComfyUI/models/Fun_Models/`. **Storage Location:** @@ -26,10 +98,26 @@ [Qwen-Image Text to Image](v1/qwenimage_chunked_loading_workflow_t2i.json) +[Qwen-Image Text to Image Control](v1/qwenimage_chunked_loading_workflow_t2i_control.json) + +[Qwen-Image Text to Image Inpaint](v1/qwenimage_chunked_loading_workflow_t2i_inpaint.json) + [Qwen-Image Edit](v1/qwenimage_chunked_loading_workflow_edit.json) +[Qwen-Image Edit 2509](v1/qwenimage_chunked_loading_workflow_edit_2509.json) + +[Qwen-Image Edit 2511](v1/qwenimage_chunked_loading_workflow_edit_2511.json) + ### 2. Full Model Loading (Optional) [Qwen-Image Text to Image](v1/qwenimage_workflow_t2i.json) -[Qwen-Image Edit](v1/qwenimage_workflow_edit.json) \ No newline at end of file +[Qwen-Image Text to Image Control](v1/qwenimage_workflow_t2i_control.json) + +[Qwen-Image Text to Image Inpaint](v1/qwenimage_workflow_t2i_inpaint.json) + +[Qwen-Image Edit](v1/qwenimage_workflow_edit.json) + +[Qwen-Image Edit 2509](v1/qwenimage_workflow_edit_2509.json) + +[Qwen-Image Edit 2511](v1/qwenimage_workflow_edit_2511.json) \ No newline at end of file diff --git a/comfyui/qwenimage/nodes.py b/comfyui/qwenimage/nodes.py index 1189b68..c3e4523 100644 --- a/comfyui/qwenimage/nodes.py +++ b/comfyui/qwenimage/nodes.py @@ -6,6 +6,7 @@ import inspect import json import os +import accelerate import comfy.model_management as mm import cv2 import folder_paths @@ -13,18 +14,32 @@ import numpy as np import torch from comfy.utils import ProgressBar, load_torch_file from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers import __version__ as diffusers_version from einops import rearrange from omegaconf import OmegaConf -from PIL import Image +from safetensors.torch import load_file + +if diffusers_version >= "0.33.0": + from diffusers.models.model_loading_utils import load_model_dict_into_meta +else: + from diffusers.models.modeling_utils import \ + load_model_dict_into_meta from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, get_closest_ratio) from ...videox_fun.models import (AutoencoderKLQwenImage, Qwen2_5_VLConfig, Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer, Qwen2VLProcessor, + QwenImageControlTransformer2DModel, QwenImageTransformer2DModel) from ...videox_fun.models.cache_utils import get_teacache_coefficients -from ...videox_fun.pipeline import QwenImageEditPipeline, QwenImagePipeline +from ...videox_fun.pipeline import (QwenImageControlPipeline, + QwenImageEditPipeline, + QwenImageEditPlusPipeline, + QwenImagePipeline) +from ...videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload, + safe_remove_group_offloading) from ...videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from ...videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from ...videox_fun.utils.fp8_optimization import ( @@ -32,7 +47,7 @@ from ...videox_fun.utils.fp8_optimization import ( replace_parameters_by_name, undo_convert_weight_dtype_wrapper) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (filter_kwargs, get_autocast_dtype, - get_image) + get_image, get_image_latent) from ..comfyui_utils import (eas_cache_dir, script_directory, search_model_in_possible_folders, search_sub_dir_in_possible_folders, to_pil) @@ -77,7 +92,10 @@ class LoadQwenImageTransformerModel: "required": { "model_name": ( folder_paths.get_filename_list("diffusion_models"), - {"default": "Wan2_1-T2V-1_3B_bf16.safetensors,"}, + {"default": "qwen_image_fp8_e4m3fn.safetensors",}, + ), + "zero_cond_t":( + [False, True], {"default": False,} ), "precision": (["fp16", "bf16"], {"default": "bf16"} @@ -89,14 +107,14 @@ class LoadQwenImageTransformerModel: FUNCTION = "loadmodel" CATEGORY = "CogVideoXFUNWrapper" - def loadmodel(self, model_name, precision): + def loadmodel(self, model_name, zero_cond_t, precision): # Init weight_dtype and device device = mm.get_torch_device() offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() transformer = None @@ -118,14 +136,43 @@ class LoadQwenImageTransformerModel: "num_layers": 60, "out_channels": 16, "patch_size": 2, - "pooled_projection_dim": 768 + "zero_cond_t": zero_cond_t, } sig = inspect.signature(QwenImageTransformer2DModel) accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} - transformer = QwenImageTransformer2DModel(**accepted) - transformer.load_state_dict(transformer_state_dict) - transformer = transformer.eval().to(device=offload_device, dtype=weight_dtype) + with accelerate.init_empty_weights(): + transformer = QwenImageTransformer2DModel(**accepted) + + new_state_dict = {} + for key, value in transformer_state_dict.items(): + if key.startswith('model.diffusion_model.'): + new_key = key.replace('model.diffusion_model.', '') + new_state_dict[new_key] = value + else: + new_state_dict[key] = value + transformer_state_dict = new_state_dict + + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + transformer, + transformer_state_dict, + dtype=weight_dtype, + model_name_or_path="", + ) + else: + transformer._convert_deprecated_attention_blocks(transformer_state_dict) + unexpected_keys = load_model_dict_into_meta( + transformer, + transformer_state_dict, + device=offload_device, + dtype=weight_dtype, + model_name_or_path="", + ) + + transformer = transformer.eval().to(weight_dtype) return (transformer, model_name_in_pipeline) class LoadQwenImageVAEModel: @@ -135,7 +182,7 @@ class LoadQwenImageVAEModel: "required": { "model_name": ( folder_paths.get_filename_list("vae"), - {"default": "QwenImage2.1_VAE.pth"} + {"default": "qwen_image_vae.safetensors"} ), "precision": (["fp16", "bf16"], {"default": "bf16"} @@ -238,7 +285,7 @@ class LoadQwenImageTextEncoderModel: "required": { "model_name": ( folder_paths.get_filename_list("text_encoders"), - {"default": "models_t5_umt5-xxl-enc-bf16.pth"} + {"default": "qwen_2.5_vl_7b_fp8_scaled.safetensors", } ), "precision": (["fp16", "bf16"], {"default": "bf16"} @@ -259,9 +306,6 @@ class LoadQwenImageTextEncoderModel: model_path = folder_paths.get_full_path("text_encoders", model_name) text_state_dict = load_torch_file(model_path, safe_load=True) - if not any(k.startswith("model.") for k in text_state_dict.keys()): - text_state_dict = {f"model.{k}": v for k, v in text_state_dict.items()} - kwargs = { "attention_dropout": 0.0, "bos_token_id": 151643, @@ -396,11 +440,34 @@ class LoadQwenImageTextEncoderModel: } config = Qwen2_5_VLConfig(**kwargs) text_encoder = Qwen2_5_VLForConditionalGeneration._from_config(config) - def transform_key(key): - key = key.replace("model.", "model.language_model.") - key = key.replace("visual.", "model.visual.") - return key - text_state_dict = {transform_key(k): v for k, v in text_state_dict.items()} + + if not any(k.startswith("model.") for k in text_state_dict.keys()): + text_state_dict = {f"model.{k}": v for k, v in text_state_dict.items()} + + new_state_dict = {} + scale_dict = {} + for key, value in text_state_dict.items(): + if 'scale_input' in key or 'scale_weight' in key: + scale_dict[key] = value + + for key, value in text_state_dict.items(): + if 'scale_input' in key or 'scale_weight' in key or key == 'scaled_fp8': + continue + if key.startswith('visual.'): + new_key = 'model.' + key + elif key.startswith('model.layers.') or key.startswith('model.embed_tokens.') or key.startswith('model.norm.'): + new_key = 'model.language_' + key + else: + new_key = key + + if '.weight' in key and value.dtype == torch.float8_e4m3fn: + scale_key = key.replace('.weight', '.scale_weight') + if scale_key in scale_dict: + value = value.float() * scale_dict[scale_key].float() + + new_state_dict[new_key] = value + + text_state_dict = new_state_dict text_encoder.load_state_dict(text_state_dict) text_encoder = text_encoder.eval().to(device=offload_device, dtype=weight_dtype) @@ -461,7 +528,9 @@ class CombineQwenImagePipeline: "tokenizer": ("Tokenizer",), "model_name": ("STRING",), "GPU_memory_mode":( - ["model_full_load", "model_full_load_and_qfloat8","model_cpu_offload", "model_cpu_offload_and_qfloat8", "sequential_cpu_offload"], + [ + "model_full_load", "model_full_load_and_qfloat8", "model_cpu_offload", + "model_cpu_offload_and_qfloat8", "model_group_offload", "sequential_cpu_offload"], { "default": "model_cpu_offload", } @@ -484,17 +553,31 @@ class CombineQwenImagePipeline: offload_device = mm.unet_offload_device() # Get pipeline - model_type = "Inpaint" + if hasattr(transformer, "control_layers"): + model_type = "Control" + else: + model_type = "Inpaint" + if model_type == "Inpaint": if processor is not None: - pipeline = QwenImageEditPipeline( - vae=vae, - tokenizer=tokenizer, - text_encoder=text_encoder, - transformer=transformer, - scheduler=None, - processor=processor, - ) + if "2509" in model_name or "2511" in model_name: + pipeline = QwenImageEditPlusPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + processor=processor, + ) + else: + pipeline = QwenImageEditPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + processor=processor, + ) else: pipeline = QwenImagePipeline( vae=vae, @@ -504,15 +587,24 @@ class CombineQwenImagePipeline: scheduler=None, ) else: - raise ValueError("Not supported now.") + pipeline = QwenImageControlPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + ) 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": pipeline.enable_sequential_cpu_offload(device=device) + elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device=offload_device, offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) @@ -528,6 +620,7 @@ class CombineQwenImagePipeline: funmodels = { 'pipeline': pipeline, + 'GPU_memory_mode': GPU_memory_mode, 'dtype': weight_dtype, 'model_name': model_name, 'model_type': model_type, @@ -544,14 +637,19 @@ class LoadQwenImageModel: "model": ( [ 'Qwen-Image', + 'Qwen-Image-2512', 'Qwen-Image-Edit', + 'Qwen-Image-Edit-2509', + 'Qwen-Image-Edit-2511', ], { "default": 'Qwen-Image', } ), "GPU_memory_mode":( - ["model_full_load", "model_full_load_and_qfloat8","model_cpu_offload", "model_cpu_offload_and_qfloat8", "sequential_cpu_offload"], + [ + "model_full_load", "model_full_load_and_qfloat8", "model_cpu_offload", + "model_cpu_offload_and_qfloat8", "model_group_offload", "sequential_cpu_offload"], { "default": "model_cpu_offload", } @@ -577,7 +675,7 @@ class LoadQwenImageModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar @@ -640,14 +738,24 @@ class LoadQwenImageModel: model_type = "Inpaint" if model_type == "Inpaint": if need_processor: - pipeline = QwenImageEditPipeline( - vae=vae, - tokenizer=tokenizer, - text_encoder=text_encoder, - transformer=transformer, - scheduler=None, - processor=processor, - ) + if "2509" in model_name or "2511" in model_name: + pipeline = QwenImageEditPlusPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + processor=processor, + ) + else: + pipeline = QwenImageEditPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=None, + processor=processor, + ) else: pipeline = QwenImagePipeline( vae=vae, @@ -664,6 +772,9 @@ class LoadQwenImageModel: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) + elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device=offload_device, offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) @@ -679,6 +790,7 @@ class LoadQwenImageModel: funmodels = { 'pipeline': pipeline, + 'GPU_memory_mode': GPU_memory_mode, 'dtype': weight_dtype, 'model_name': model_name, 'model_type': model_type, @@ -713,6 +825,240 @@ class LoadQwenImageLora: new_funmodels['lora_cache'] = lora_cache return (new_funmodels,) +class LoadQwenImageControlNetInPipeline: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "config": ( + [ + "qwenimage/qwenimage_control.yaml", + ], + { + "default": "qwenimage/qwenimage_control.yaml", + } + ), + "model_name": ( + folder_paths.get_filename_list("model_patches"), + {"default": "Qwen-Image-2512-Fun-Controlnet-Union.safetensors", }, + ), + "sub_transformer_name":( + ["transformer", "transformer_2"], + { + "default": "transformer", + } + ), + "funmodels": ("FunModels",), + }, + } + + RETURN_TYPES = ("FunModels",) + RETURN_NAMES = ("funmodels",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, config, model_name, sub_transformer_name, funmodels): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + GPU_memory_mode = funmodels["GPU_memory_mode"] + weight_dtype = funmodels['dtype'] + + # Remove hooks + funmodels["pipeline"].remove_all_hooks() + safe_remove_group_offloading(funmodels["pipeline"]) + + # Get Transformer + transformer = getattr(funmodels["pipeline"], sub_transformer_name) + transformer = transformer.cpu() + + # Get state_dict + transformer_state_dict = transformer.state_dict() + del transformer + mm.soft_empty_cache() + gc.collect() + + # Load config + config_path = f"{script_directory}/config/{config}" + config = OmegaConf.load(config_path) + kwargs = { + "attention_head_dim": 128, + "axes_dims_rope": [ + 16, + 56, + 56 + ], + "guidance_embeds": False, + "in_channels": 64, + "joint_attention_dim": 3584, + "num_attention_heads": 24, + "num_layers": 60, + "out_channels": 16, + "patch_size": 2, + "pooled_projection_dim": 768 + } + kwargs.update(OmegaConf.to_container(config['transformer_additional_kwargs'])) + + # Get Model + sig = inspect.signature(QwenImageControlTransformer2DModel) + accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} + with accelerate.init_empty_weights(): + control_transformer = QwenImageControlTransformer2DModel(**accepted).to(weight_dtype) + print(f"Load Control Transformer") + + # Load Control state_dict + control_model_path = folder_paths.get_full_path("model_patches", model_name) + if control_model_path.endswith(".safetensors"): + control_state_dict = load_file(control_model_path) + else: + control_state_dict = torch.load(control_model_path) + + state_dict = {**transformer_state_dict, **control_state_dict} + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + control_transformer, + state_dict, + dtype=weight_dtype, + model_name_or_path="", + ) + else: + control_transformer._convert_deprecated_attention_blocks(state_dict) + load_model_dict_into_meta( + control_transformer, + state_dict, + device=offload_device, + dtype=weight_dtype, + model_name_or_path="", + ) + + pipeline = QwenImageControlPipeline( + vae=funmodels["pipeline"].vae, + tokenizer=funmodels["pipeline"].tokenizer, + text_encoder=funmodels["pipeline"].text_encoder, + transformer=control_transformer, + scheduler=funmodels["pipeline"].scheduler, + ) + del funmodels["pipeline"] + mm.soft_empty_cache() + gc.collect() + + if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) + elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device=offload_device, offload_type="leaf_level", use_stream=True) + elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(control_transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(control_transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) + elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(control_transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(control_transformer, weight_dtype) + pipeline.to(device=device) + else: + pipeline.to(device=device) + funmodels["pipeline"] = pipeline + funmodels["model_type"] = "Control" + return (funmodels, ) + +class LoadQwenImageControlNetInModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "config": ( + [ + "qwenimage/qwenimage_control.yaml", + ], + { + "default": "qwenimage/qwenimage_control.yaml", + } + ), + "model_name": ( + folder_paths.get_filename_list("model_patches"), + {"default": "Qwen-Image-2512-Fun-Controlnet-Union.safetensors", }, + ), + "transformer": ("TransformerModel",), + }, + } + + RETURN_TYPES = ("TransformerModel",) + RETURN_NAMES = ("transformer",) + FUNCTION = "loadmodel" + CATEGORY = "CogVideoXFUNWrapper" + + def loadmodel(self, config, model_name, transformer): + offload_device = mm.unet_offload_device() + dtype = transformer.dtype + + # Get Transformer + transformer = transformer.cpu() + + # Get state_dict + transformer_state_dict = transformer.state_dict() + del transformer + mm.soft_empty_cache() + gc.collect() + + # Load config + config_path = f"{script_directory}/config/{config}" + config = OmegaConf.load(config_path) + kwargs = { + "attention_head_dim": 128, + "axes_dims_rope": [ + 16, + 56, + 56 + ], + "guidance_embeds": False, + "in_channels": 64, + "joint_attention_dim": 3584, + "num_attention_heads": 24, + "num_layers": 60, + "out_channels": 16, + "patch_size": 2, + "pooled_projection_dim": 768 + } + kwargs.update(OmegaConf.to_container(config['transformer_additional_kwargs'])) + + # Get Model + sig = inspect.signature(QwenImageControlTransformer2DModel) + accepted = {k: v for k, v in kwargs.items() if k in sig.parameters} + with accelerate.init_empty_weights(): + control_transformer = QwenImageControlTransformer2DModel(**accepted).to(dtype) + print(f"Load Control Transformer") + + # Load Control state_dict + control_model_path = folder_paths.get_full_path("model_patches", model_name) + if control_model_path.endswith(".safetensors"): + control_state_dict = load_file(control_model_path) + else: + control_state_dict = torch.load(control_model_path) + + state_dict = {**transformer_state_dict, **control_state_dict} + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + control_transformer, + state_dict, + dtype=dtype, + model_name_or_path="", + ) + else: + control_transformer._convert_deprecated_attention_blocks(state_dict) + load_model_dict_into_meta( + control_transformer, + state_dict, + device=offload_device, + dtype=dtype, + model_name_or_path="", + ) + return (control_transformer, ) + class QwenImageT2VSampler: @classmethod def INPUT_TYPES(s): @@ -894,7 +1240,7 @@ class QwenImageEditSampler: "INT", {"default": 1, "min": 1, "max": 100, "step": 1} ), "teacache_threshold": ( - "FLOAT", {"default": 0.10, "min": 0.00, "max": 1.00, "step": 0.005} + "FLOAT", {"default": 0.250, "min": 0.00, "max": 1.00, "step": 0.005} ), "enable_teacache":( [False, True], {"default": True,} @@ -1003,3 +1349,334 @@ class QwenImageEditSampler: for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) return (image,) + +class QwenImageEditPlusSampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "funmodels": ( + "FunModels", + ), + "prompt": ( + "STRING_PROMPT", + ), + "negative_prompt": ( + "STRING_PROMPT", + ), + "width": ( + "INT", {"default": 1344, "min": 64, "max": 2048, "step": 16} + ), + "height": ( + "INT", {"default": 768, "min": 64, "max": 2048, "step": 16} + ), + "seed": ( + "INT", {"default": 43, "min": 0, "max": 0xffffffffffffffff} + ), + "steps": ( + "INT", {"default": 50, "min": 1, "max": 200, "step": 1} + ), + "cfg": ( + "FLOAT", {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.01} + ), + "scheduler": ( + ["Flow", "Flow_Unipc", "Flow_DPM++"], + { + "default": 'Flow' + } + ), + "shift": ( + "INT", {"default": 1, "min": 1, "max": 100, "step": 1} + ), + "teacache_threshold": ( + "FLOAT", {"default": 0.250, "min": 0.00, "max": 1.00, "step": 0.005} + ), + "enable_teacache":( + [False, True], {"default": True,} + ), + "num_skip_start_steps": ( + "INT", {"default": 5, "min": 0, "max": 50, "step": 1} + ), + "teacache_offload":( + [False, True], {"default": True,} + ), + "cfg_skip_ratio":( + "FLOAT", {"default": 0, "min": 0, "max": 1, "step": 0.01} + ), + }, + "optional":{ + "image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES =("images",) + FUNCTION = "process" + CATEGORY = "CogVideoXFUNWrapper" + + def process(self, funmodels, prompt, negative_prompt, width, height, seed, steps, cfg, scheduler, shift, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, cfg_skip_ratio, image=None): + global transformer_cpu_cache + global lora_path_before + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.soft_empty_cache() + gc.collect() + + # Get Pipeline + pipeline = funmodels['pipeline'] + model_name = funmodels['model_name'] + weight_dtype = funmodels['dtype'] + + # Change to QwenImageEditPlusPipeline + if not isinstance(pipeline, QwenImageEditPlusPipeline): + pipeline = QwenImageEditPlusPipeline( + vae=pipeline.vae, + tokenizer=pipeline.tokenizer, + text_encoder=pipeline.text_encoder, + transformer=pipeline.transformer, + processor=pipeline.processor, + scheduler=pipeline.scheduler, + ) + + # Load Sampler + pipeline.scheduler = get_qwen_scheduler(scheduler, shift) + + coefficients = get_teacache_coefficients(model_name) if enable_teacache else None + if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + else: + pipeline.transformer.disable_teacache() + + if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) + + generator= torch.Generator(device).manual_seed(seed) + + with torch.no_grad(): + # Apply lora + if funmodels.get("lora_cache", False): + if len(funmodels.get("loras", [])) != 0: + # Save the original weights to cpu + if len(transformer_cpu_cache) == 0: + print('Save transformer state_dict to cpu memory') + transformer_state_dict = pipeline.transformer.state_dict() + for key in transformer_state_dict: + transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() + + lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) + if lora_path_now != lora_path_before: + print('Merge Lora with Cache') + lora_path_before = copy.deepcopy(lora_path_now) + pipeline.transformer.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + else: + print('Merge Lora') + # Clear lora when switch from lora_cache=True to lora_cache=False. + if len(transformer_cpu_cache) != 0: + pipeline.transformer.load_state_dict(transformer_cpu_cache) + transformer_cpu_cache = {} + lora_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + image = [to_pil(image) for image in image] + image = get_image(image[0]) if image is not None else image + + sample = pipeline( + image = image, + prompt = prompt, + negative_prompt = negative_prompt, + height = height, + width = width, + generator = generator, + true_cfg_scale = cfg, + num_inference_steps = steps, + comfyui_progressbar = True, + ).images + image = torch.Tensor(np.array(sample[0])).unsqueeze(0) / 255 + + if not funmodels.get("lora_cache", False): + print('Unmerge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + return (image,) + +class QwenImageControlSampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "funmodels": ( + "FunModels", + ), + "prompt": ( + "STRING_PROMPT", + ), + "negative_prompt": ( + "STRING_PROMPT", + ), + "width": ( + "INT", {"default": 1568, "min": 64, "max": 20480, "step": 16} + ), + "height": ( + "INT", {"default": 1184, "min": 64, "max": 20480, "step": 16} + ), + "seed": ( + "INT", {"default": 43, "min": 0, "max": 0xffffffffffffffff} + ), + "steps": ( + "INT", {"default": 40, "min": 1, "max": 200, "step": 1} + ), + "cfg": ( + "FLOAT", {"default": 4.0, "min": 0.0, "max": 20.0, "step": 0.01} + ), + "scheduler": ( + ["Flow", "Flow_Unipc", "Flow_DPM++"], + { + "default": 'Flow' + } + ), + "shift": ( + "INT", {"default": 3, "min": 1, "max": 100, "step": 1} + ), + "teacache_threshold": ( + "FLOAT", {"default": 0.250, "min": 0.00, "max": 1.00, "step": 0.005} + ), + "enable_teacache":( + [False, True], {"default": True,} + ), + "num_skip_start_steps": ( + "INT", {"default": 5, "min": 0, "max": 50, "step": 1} + ), + "teacache_offload":( + [False, True], {"default": True,} + ), + "cfg_skip_ratio":( + "FLOAT", {"default": 0, "min": 0, "max": 1, "step": 0.01} + ), + "control_context_scale": ( + "FLOAT", {"default": 0.80, "min": 0.0, "max": 2.0, "step": 0.01} + ), + }, + "optional":{ + "control_image": ("IMAGE",), + "inpaint_image": ("IMAGE",), + "mask_image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES =("images",) + FUNCTION = "process" + CATEGORY = "CogVideoXFUNWrapper" + + def process(self, funmodels, prompt, negative_prompt, width, height, seed, steps, cfg, scheduler, shift, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, cfg_skip_ratio, control_context_scale, control_image=None, inpaint_image=None, mask_image=None): + global transformer_cpu_cache + global lora_path_before + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.soft_empty_cache() + gc.collect() + + # Get Pipeline + pipeline = funmodels['pipeline'] + model_name = funmodels['model_name'] + weight_dtype = funmodels['dtype'] + sample_size = [height, width] + + # Load Sampler + pipeline.scheduler = get_qwen_scheduler(scheduler, shift) + + coefficients = get_teacache_coefficients(model_name) if enable_teacache else None + if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + else: + pipeline.transformer.disable_teacache() + + if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) + + generator= torch.Generator(device).manual_seed(seed) + + with torch.no_grad(): + # Apply lora + if funmodels.get("lora_cache", False): + if len(funmodels.get("loras", [])) != 0: + # Save the original weights to cpu + if len(transformer_cpu_cache) == 0: + print('Save transformer state_dict to cpu memory') + transformer_state_dict = pipeline.transformer.state_dict() + for key in transformer_state_dict: + transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() + + lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) + if lora_path_now != lora_path_before: + print('Merge Lora with Cache') + lora_path_before = copy.deepcopy(lora_path_now) + pipeline.transformer.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + else: + print('Merge Lora') + # Clear lora when switch from lora_cache=True to lora_cache=False. + if len(transformer_cpu_cache) != 0: + pipeline.transformer.load_state_dict(transformer_cpu_cache) + transformer_cpu_cache = {} + lora_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + + if inpaint_image is not None: + inpaint_image = [to_pil(inpaint_image) for inpaint_image in inpaint_image][0] + inpaint_image = get_image_latent(inpaint_image, sample_size=sample_size)[:, :, 0] + else: + inpaint_image = torch.zeros([1, 3, sample_size[0], sample_size[1]]) + + if mask_image is not None: + mask_image = [to_pil(mask_image) for mask_image in mask_image][0] + mask_image = get_image_latent(mask_image, sample_size=sample_size)[:, :1, 0] + else: + mask_image = torch.ones([1, 1, sample_size[0], sample_size[1]]) * 255 + + if control_image is not None: + control_image = [to_pil(control_image) for control_image in control_image][0] + control_image = get_image_latent(control_image, sample_size=sample_size)[:, :, 0] + + sample = pipeline( + prompt, + negative_prompt = negative_prompt, + height = height, + width = width, + generator = generator, + guidance_scale = cfg, + num_inference_steps = steps, + image = inpaint_image, + mask_image = mask_image, + control_image = control_image, + control_context_scale = control_context_scale, + comfyui_progressbar = True, + ).images + image = torch.Tensor(np.array(sample[0])).unsqueeze(0) / 255 + + if not funmodels.get("lora_cache", False): + print('Unmerge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) + return (image,) diff --git a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit.json b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit.json index 1243fee..7198aaa 100644 --- a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit.json +++ b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit.json @@ -133,7 +133,7 @@ -275.1779479980469 ], "size": [ - 217.32675170898438, + 226.6099609375, 26 ], "flags": {}, @@ -287,6 +287,7 @@ }, "widgets_values": [ "Qwen-Image-Edit_bf16.safetensors", + false, "bf16" ] }, @@ -373,8 +374,8 @@ "Node name for S&R": "QwenImageEditSampler" }, "widgets_values": [ - 1344, - 768, + 1728, + 992, 373336117071181, "randomize", 40, @@ -452,7 +453,7 @@ }, "widgets_values": [ "", - "model_cpu_offload_and_qfloat8" + "model_group_offload" ] } ], @@ -579,18 +580,19 @@ "ds": { "scale": 0.6905497838871149, "offset": [ - 351.20689397709714, - 562.4271762478439 + 397.7447678565184, + 611.2107698221871 ] }, - "frontendVersion": "1.25.11", + "frontendVersion": "1.36.14", "workspace_info": { "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" }, "node_versions": { - "CogVideoX-Fun": "a97dd425909c3c3719fbbcb99e78061e2f0a237c", - "comfy-core": "0.3.57" - } + "CogVideoX-Fun": "ac114cc14285c8e0073a3e08e27525263d1264a7", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" }, "version": 0.4 } \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit_2509.json b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit_2509.json new file mode 100644 index 0000000..fa59ec7 --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit_2509.json @@ -0,0 +1,598 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 101, + "last_link_id": 89, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 84 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 98, + "type": "LoadImage", + "pos": [ + 312.6856384277344, + 418.9110107421875 + ], + "size": [ + 315, + 314.0000305175781 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 85 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "ref_1.png", + "image" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 83 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 94, + "type": "CombineQwenImagePipeline", + "pos": [ + 945.2576293945312, + -330.913330078125 + ], + "size": [ + 321.2720642089844, + 162 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 88 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 68 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 70 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 73 + }, + { + "name": "processor", + "shape": 7, + "type": "Processor", + "link": 74 + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 89 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "CombineQwenImagePipeline" + }, + "widgets_values": [ + "", + "model_group_offload" + ] + }, + { + "id": 93, + "type": "LoadQwenImageVAEModel", + "pos": [ + 775.0554809570312, + -470.7688293457031 + ], + "size": [ + 377.8583984375, + 84.69844055175781 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 68 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageVAEModel" + }, + "widgets_values": [ + "qwen_image_vae.safetensors", + "bf16" + ] + }, + { + "id": 91, + "type": "LoadQwenImageTextEncoderModel", + "pos": [ + 283.53765869140625, + -280.6837463378906 + ], + "size": [ + 407.4130859375, + 102 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 70 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 73 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTextEncoderModel" + }, + "widgets_values": [ + "qwen_2.5_vl_7b_fp8_scaled.safetensors", + "bf16" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 80 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "把相机转变成西瓜" + ] + }, + { + "id": 99, + "type": "QwenImageEditPlusSampler", + "pos": [ + 722.5752102270133, + -64.70894250180959 + ], + "size": [ + 298.1490234375, + 406 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 82 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 80 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 84 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 85 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 83 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageEditPlusSampler" + }, + "widgets_values": [ + 1728, + 992, + 275685855283225, + "randomize", + 50, + 4, + "Flow", + 1, + 0.25, + true, + 5, + true, + 0 + ] + }, + { + "id": 101, + "type": "LoadQwenImageTransformerModel", + "pos": [ + 268.4111302692554, + -465.72558354643127 + ], + "size": [ + 452.9282513511349, + 126 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 88 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 89 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTransformerModel" + }, + "widgets_values": [ + "qwen_image_edit_2509_fp8_e4m3fn.safetensors", + false, + "bf16" + ] + }, + { + "id": 96, + "type": "LoadQwenImageProcessor", + "pos": [ + 706.5747258956637, + -287.7528469446806 + ], + "size": [ + 226.6099609375, + 26 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "processor", + "type": "Processor", + "links": [ + 74 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageProcessor" + }, + "widgets_values": [] + } + ], + "links": [ + [ + 68, + 93, + 0, + 94, + 1, + "VAEModel" + ], + [ + 70, + 91, + 0, + 94, + 2, + "TextEncoderModel" + ], + [ + 73, + 91, + 1, + 94, + 3, + "Tokenizer" + ], + [ + 74, + 96, + 0, + 94, + 4, + "Processor" + ], + [ + 80, + 75, + 0, + 99, + 1, + "STRING_PROMPT" + ], + [ + 82, + 94, + 0, + 99, + 0, + "FunModels" + ], + [ + 83, + 99, + 0, + 88, + 0, + "IMAGE" + ], + [ + 84, + 73, + 0, + 99, + 2, + "STRING_PROMPT" + ], + [ + 85, + 98, + 0, + 99, + 3, + "IMAGE" + ], + [ + 88, + 101, + 0, + 94, + 0, + "TransformerModel" + ], + [ + 89, + 101, + 1, + 94, + 5, + "STRING" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1053.5875244140625, + 397.3387756347656 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6905497838871149, + "offset": [ + 553.0897408085864, + 691.4796929228297 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "13a802b574c3a4397e193a0b8ca4e90c480cc217", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit_2511.json b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit_2511.json new file mode 100644 index 0000000..e938ee0 --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_edit_2511.json @@ -0,0 +1,598 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 101, + "last_link_id": 89, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 84 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 98, + "type": "LoadImage", + "pos": [ + 312.6856384277344, + 418.9110107421875 + ], + "size": [ + 315, + 314.0000305175781 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 85 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "ref_1.png", + "image" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 83 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 94, + "type": "CombineQwenImagePipeline", + "pos": [ + 945.2576293945312, + -330.913330078125 + ], + "size": [ + 321.2720642089844, + 162 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 88 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 68 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 70 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 73 + }, + { + "name": "processor", + "shape": 7, + "type": "Processor", + "link": 74 + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 89 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "CombineQwenImagePipeline" + }, + "widgets_values": [ + "", + "model_group_offload" + ] + }, + { + "id": 93, + "type": "LoadQwenImageVAEModel", + "pos": [ + 775.0554809570312, + -470.7688293457031 + ], + "size": [ + 377.8583984375, + 84.69844055175781 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 68 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageVAEModel" + }, + "widgets_values": [ + "qwen_image_vae.safetensors", + "bf16" + ] + }, + { + "id": 91, + "type": "LoadQwenImageTextEncoderModel", + "pos": [ + 283.53765869140625, + -280.6837463378906 + ], + "size": [ + 407.4130859375, + 102 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 70 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 73 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTextEncoderModel" + }, + "widgets_values": [ + "qwen_2.5_vl_7b_fp8_scaled.safetensors", + "bf16" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 80 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "把相机转变成西瓜" + ] + }, + { + "id": 99, + "type": "QwenImageEditPlusSampler", + "pos": [ + 722.5752102270133, + -64.70894250180959 + ], + "size": [ + 298.1490234375, + 406 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 82 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 80 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 84 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 85 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 83 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageEditPlusSampler" + }, + "widgets_values": [ + 1728, + 992, + 275685855283225, + "randomize", + 50, + 4, + "Flow", + 1, + 0.25, + true, + 5, + true, + 0 + ] + }, + { + "id": 96, + "type": "LoadQwenImageProcessor", + "pos": [ + 706.5747258956637, + -287.7528469446806 + ], + "size": [ + 226.6099609375, + 26 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "processor", + "type": "Processor", + "links": [ + 74 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageProcessor" + }, + "widgets_values": [] + }, + { + "id": 101, + "type": "LoadQwenImageTransformerModel", + "pos": [ + 268.4111302692554, + -465.72558354643127 + ], + "size": [ + 452.9282513511349, + 126 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 88 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 89 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTransformerModel" + }, + "widgets_values": [ + "qwen_image_edit_2511_bf16.safetensors", + true, + "bf16" + ] + } + ], + "links": [ + [ + 68, + 93, + 0, + 94, + 1, + "VAEModel" + ], + [ + 70, + 91, + 0, + 94, + 2, + "TextEncoderModel" + ], + [ + 73, + 91, + 1, + 94, + 3, + "Tokenizer" + ], + [ + 74, + 96, + 0, + 94, + 4, + "Processor" + ], + [ + 80, + 75, + 0, + 99, + 1, + "STRING_PROMPT" + ], + [ + 82, + 94, + 0, + 99, + 0, + "FunModels" + ], + [ + 83, + 99, + 0, + 88, + 0, + "IMAGE" + ], + [ + 84, + 73, + 0, + 99, + 2, + "STRING_PROMPT" + ], + [ + 85, + 98, + 0, + 99, + 3, + "IMAGE" + ], + [ + 88, + 101, + 0, + 94, + 0, + "TransformerModel" + ], + [ + 89, + 101, + 1, + 94, + 5, + "STRING" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1053.5875244140625, + 397.3387756347656 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6905497838871149, + "offset": [ + 497.11079345101393, + 668.6095550725204 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "13a802b574c3a4397e193a0b8ca4e90c480cc217", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i.json b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i.json index 14be765..688f680 100644 --- a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i.json +++ b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i.json @@ -120,116 +120,6 @@ "color": "#432", "bgcolor": "#653" }, - { - "id": 92, - "type": "LoadQwenImageTransformerModel", - "pos": [ - 275.9798278808594, - -465.2391052246094 - ], - "size": [ - 416.3677673339844, - 106.13789367675781 - ], - "flags": {}, - "order": 4, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "transformer", - "type": "TransformerModel", - "links": [ - 78 - ] - }, - { - "name": "model_name", - "type": "STRING", - "links": [ - 82 - ] - } - ], - "properties": { - "Node name for S&R": "LoadQwenImageTransformerModel" - }, - "widgets_values": [ - "Qwen-Image_bf16.safetensors", - "bf16" - ] - }, - { - "id": 93, - "type": "LoadQwenImageVAEModel", - "pos": [ - 775.0554809570312, - -470.7688293457031 - ], - "size": [ - 377.8583984375, - 84.69844055175781 - ], - "flags": {}, - "order": 5, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "vae", - "type": "VAEModel", - "links": [ - 79 - ] - } - ], - "properties": { - "Node name for S&R": "LoadQwenImageVAEModel" - }, - "widgets_values": [ - "Qwen-Image-vae_bf16.safetensors", - "bf16" - ] - }, - { - "id": 91, - "type": "LoadQwenImageTextEncoderModel", - "pos": [ - 283.53765869140625, - -280.6837463378906 - ], - "size": [ - 407.4130859375, - 102 - ], - "flags": {}, - "order": 6, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "text_encoder", - "type": "TextEncoderModel", - "links": [ - 80 - ] - }, - { - "name": "tokenizer", - "type": "Tokenizer", - "links": [ - 81 - ] - } - ], - "properties": { - "Node name for S&R": "LoadQwenImageTextEncoderModel" - }, - "widgets_values": [ - "Qwen-Image-text_encoder_bf16.safetensors", - "bf16" - ] - }, { "id": 88, "type": "PreviewImage", @@ -301,8 +191,8 @@ "Node name for S&R": "QwenImageT2VSampler" }, "widgets_values": [ - 1344, - 768, + 1728, + 992, 867934328802019, "randomize", 40, @@ -380,7 +270,118 @@ }, "widgets_values": [ "", - "model_cpu_offload_and_qfloat8" + "model_group_offload" + ] + }, + { + "id": 92, + "type": "LoadQwenImageTransformerModel", + "pos": [ + 275.9798278808594, + -465.2391052246094 + ], + "size": [ + 416.3677673339844, + 126.29115625000003 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 78 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTransformerModel" + }, + "widgets_values": [ + "qwen_image_edit_fp8_e4m3fn.safetensors", + false, + "bf16" + ] + }, + { + "id": 93, + "type": "LoadQwenImageVAEModel", + "pos": [ + 775.0554809570312, + -470.7688293457031 + ], + "size": [ + 377.8583984375, + 84.69844055175781 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 79 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageVAEModel" + }, + "widgets_values": [ + "qwen_image_vae.safetensors", + "bf16" + ] + }, + { + "id": 91, + "type": "LoadQwenImageTextEncoderModel", + "pos": [ + 283.53765869140625, + -280.6837463378906 + ], + "size": [ + 407.4130859375, + 102 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 80 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTextEncoderModel" + }, + "widgets_values": [ + "qwen_2.5_vl_7b_fp8_scaled.safetensors", + "bf16" ] } ], @@ -495,14 +496,15 @@ 640.3416144465855 ] }, - "frontendVersion": "1.25.11", + "frontendVersion": "1.36.14", "workspace_info": { "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" }, "node_versions": { - "CogVideoX-Fun": "a97dd425909c3c3719fbbcb99e78061e2f0a237c", - "comfy-core": "0.3.57" - } + "CogVideoX-Fun": "13a802b574c3a4397e193a0b8ca4e90c480cc217", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" }, "version": 0.4 } \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i_control.json b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i_control.json new file mode 100644 index 0000000..948412b --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i_control.json @@ -0,0 +1,620 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 102, + "last_link_id": 103, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 101 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 91, + "type": "LoadQwenImageTextEncoderModel", + "pos": [ + 283.53765869140625, + -280.6837463378906 + ], + "size": [ + 407.4130859375, + 102 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 80 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTextEncoderModel" + }, + "widgets_values": [ + "qwen_2.5_vl_7b_fp8_scaled.safetensors", + "bf16" + ] + }, + { + "id": 92, + "type": "LoadQwenImageTransformerModel", + "pos": [ + 275.9798278808594, + -465.2391052246094 + ], + "size": [ + 416.3677673339844, + 106.13789367675781 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 87 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTransformerModel" + }, + "widgets_values": [ + "qwen_image_2512_fp8_e4m3fn.safetensors", + false, + "bf16" + ] + }, + { + "id": 98, + "type": "LoadQwenImageControlNetInModel", + "pos": [ + 753.1154601481974, + -466.61377999428464 + ], + "size": [ + 543.8293619798113, + 82 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 87 + } + ], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 88 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageControlNetInModel" + }, + "widgets_values": [ + "qwenimage/qwenimage_control.yaml", + "Qwen-Image-2512-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 103 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 100, + "type": "LoadImage", + "pos": [ + 408.7071040151891, + 403.5792288936819 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 102 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "a7kXeQ5l9Dhspes7q3x3G (1).png", + "image" + ] + }, + { + "id": 93, + "type": "LoadQwenImageVAEModel", + "pos": [ + 1133.7959245028942, + -257.8118537612475 + ], + "size": [ + 377.8583984375, + 84.69844055175781 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 79 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageVAEModel" + }, + "widgets_values": [ + "qwen_image_vae.safetensors", + "bf16" + ] + }, + { + "id": 96, + "type": "CombineQwenImagePipeline", + "pos": [ + 754.4458784421624, + -332.72602712957854 + ], + "size": [ + 342.5804748535156, + 162 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 88 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 79 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 80 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 81 + }, + { + "name": "processor", + "shape": 7, + "type": "Processor", + "link": null + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 82 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 99 + ] + } + ], + "properties": { + "Node name for S&R": "CombineQwenImagePipeline" + }, + "widgets_values": [ + "", + "model_group_offload" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 100 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 102, + "type": "QwenImageControlSampler", + "pos": [ + 743.3433452139876, + -72.05094395600588 + ], + "size": [ + 289.5494140625, + 470 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 99 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 100 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 101 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": 102 + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 103 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 563555207707640, + "randomize", + 40, + 4, + "Flow", + 3, + 0.25, + true, + 5, + true, + 0, + 0.8 + ] + } + ], + "links": [ + [ + 79, + 93, + 0, + 96, + 1, + "VAEModel" + ], + [ + 80, + 91, + 0, + 96, + 2, + "TextEncoderModel" + ], + [ + 81, + 91, + 1, + 96, + 3, + "Tokenizer" + ], + [ + 82, + 92, + 1, + 96, + 5, + "STRING" + ], + [ + 87, + 92, + 0, + 98, + 0, + "TransformerModel" + ], + [ + 88, + 98, + 0, + 96, + 0, + "TransformerModel" + ], + [ + 99, + 96, + 0, + 102, + 0, + "FunModels" + ], + [ + 100, + 75, + 0, + 102, + 1, + "STRING_PROMPT" + ], + [ + 101, + 73, + 0, + 102, + 2, + "STRING_PROMPT" + ], + [ + 102, + 100, + 0, + 102, + 3, + "IMAGE" + ], + [ + 103, + 102, + 0, + 88, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1350.7047455077704, + 391.2301145558972 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7034334668812453, + "offset": [ + 519.2334278821429, + 649.1136351768466 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "ac114cc14285c8e0073a3e08e27525263d1264a7", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i_inpaint.json b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i_inpaint.json new file mode 100644 index 0000000..ca98e78 --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_chunked_loading_workflow_t2i_inpaint.json @@ -0,0 +1,710 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 105, + "last_link_id": 108, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 101 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 91, + "type": "LoadQwenImageTextEncoderModel", + "pos": [ + 283.53765869140625, + -280.6837463378906 + ], + "size": [ + 407.4130859375, + 102 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "text_encoder", + "type": "TextEncoderModel", + "links": [ + 80 + ] + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTextEncoderModel" + }, + "widgets_values": [ + "qwen_2.5_vl_7b_fp8_scaled.safetensors", + "bf16" + ] + }, + { + "id": 92, + "type": "LoadQwenImageTransformerModel", + "pos": [ + 275.9798278808594, + -465.2391052246094 + ], + "size": [ + 416.3677673339844, + 106.13789367675781 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 87 + ] + }, + { + "name": "model_name", + "type": "STRING", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageTransformerModel" + }, + "widgets_values": [ + "qwen_image_2512_fp8_e4m3fn.safetensors", + false, + "bf16" + ] + }, + { + "id": 98, + "type": "LoadQwenImageControlNetInModel", + "pos": [ + 753.1154601481974, + -466.61377999428464 + ], + "size": [ + 543.8293619798113, + 82 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 87 + } + ], + "outputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "links": [ + 88 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageControlNetInModel" + }, + "widgets_values": [ + "qwenimage/qwenimage_control.yaml", + "Qwen-Image-2512-Fun-Controlnet-Union.safetensors" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 103 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 93, + "type": "LoadQwenImageVAEModel", + "pos": [ + 1133.7959245028942, + -257.8118537612475 + ], + "size": [ + 377.8583984375, + 84.69844055175781 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "VAEModel", + "links": [ + 79 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageVAEModel" + }, + "widgets_values": [ + "qwen_image_vae.safetensors", + "bf16" + ] + }, + { + "id": 96, + "type": "CombineQwenImagePipeline", + "pos": [ + 754.4458784421624, + -332.72602712957854 + ], + "size": [ + 342.5804748535156, + 162 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "transformer", + "type": "TransformerModel", + "link": 88 + }, + { + "name": "vae", + "type": "VAEModel", + "link": 79 + }, + { + "name": "text_encoder", + "type": "TextEncoderModel", + "link": 80 + }, + { + "name": "tokenizer", + "type": "Tokenizer", + "link": 81 + }, + { + "name": "processor", + "shape": 7, + "type": "Processor", + "link": null + }, + { + "name": "model_name", + "type": "STRING", + "widget": { + "name": "model_name" + }, + "link": 82 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 99 + ] + } + ], + "properties": { + "Node name for S&R": "CombineQwenImagePipeline" + }, + "widgets_values": [ + "", + "model_group_offload" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 100 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 104, + "type": "MaskToImage", + "pos": [ + 665.9127469276586, + 458.2047154396486 + ], + "size": [ + 140, + 26 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 105 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 104, + 106 + ] + } + ], + "properties": { + "Node name for S&R": "MaskToImage" + }, + "widgets_values": [] + }, + { + "id": 105, + "type": "LoadImage", + "pos": [ + 360.8304806417201, + 460.85158208210487 + ], + "size": [ + 270, + 314.00000000000006 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 108 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 105 + ] + } + ], + "properties": { + "Node name for S&R": "LoadImage", + "image": "clipspace/clipspace-painted-masked-1766731857414.png [input]" + }, + "widgets_values": [ + "clipspace/clipspace-painted-masked-1766731857414.png [input]", + "image" + ] + }, + { + "id": 103, + "type": "PreviewImage", + "pos": [ + 839.792105488772, + 453.4970846240923 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 104 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 102, + "type": "QwenImageControlSampler", + "pos": [ + 743.3433452139876, + -72.05094395600588 + ], + "size": [ + 289.5494140625, + 470 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 99 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 100 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 101 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": 108 + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": 106 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 103 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 284014972131965, + "randomize", + 40, + 4, + "Flow", + 3, + 0.25, + true, + 5, + true, + 0, + 0.8 + ] + } + ], + "links": [ + [ + 79, + 93, + 0, + 96, + 1, + "VAEModel" + ], + [ + 80, + 91, + 0, + 96, + 2, + "TextEncoderModel" + ], + [ + 81, + 91, + 1, + 96, + 3, + "Tokenizer" + ], + [ + 82, + 92, + 1, + 96, + 5, + "STRING" + ], + [ + 87, + 92, + 0, + 98, + 0, + "TransformerModel" + ], + [ + 88, + 98, + 0, + 96, + 0, + "TransformerModel" + ], + [ + 99, + 96, + 0, + 102, + 0, + "FunModels" + ], + [ + 100, + 75, + 0, + 102, + 1, + "STRING_PROMPT" + ], + [ + 101, + 73, + 0, + 102, + 2, + "STRING_PROMPT" + ], + [ + 103, + 102, + 0, + 88, + 0, + "IMAGE" + ], + [ + 104, + 104, + 0, + 103, + 0, + "IMAGE" + ], + [ + 105, + 105, + 1, + 104, + 0, + "MASK" + ], + [ + 106, + 104, + 0, + 102, + 5, + "IMAGE" + ], + [ + 108, + 105, + 0, + 102, + 4, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 227.96267700195312, + -546.4359741210938, + 1350.7047455077704, + 391.2301145558972 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7034334668812453, + "offset": [ + 333.2705669168081, + 543.6154736676693 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "ac114cc14285c8e0073a3e08e27525263d1264a7", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_workflow_edit.json b/comfyui/qwenimage/v1/qwenimage_workflow_edit.json index 34b151b..9507a5f 100644 --- a/comfyui/qwenimage/v1/qwenimage_workflow_edit.json +++ b/comfyui/qwenimage/v1/qwenimage_workflow_edit.json @@ -201,9 +201,9 @@ "Node name for S&R": "QwenImageEditSampler" }, "widgets_values": [ - 1344, - 768, - 686934831068040, + 1728, + 992, + 578345222670916, "randomize", 40, 4, @@ -278,7 +278,7 @@ }, "widgets_values": [ "Qwen-Image-Edit", - "model_cpu_offload_and_qfloat8", + "model_group_offload", "bf16" ] } @@ -358,18 +358,19 @@ "ds": { "scale": 0.6905497838871149, "offset": [ - 542.3419639311205, - 433.4877940663054 + 454.62313443699367, + 551.1306970771619 ] }, - "frontendVersion": "1.25.11", + "frontendVersion": "1.36.14", "workspace_info": { "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" }, "node_versions": { - "CogVideoX-Fun": "36287cdcab8d5b6972bb6a2d208539c6e4bd81e2", - "comfy-core": "0.3.57" - } + "CogVideoX-Fun": "ac114cc14285c8e0073a3e08e27525263d1264a7", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" }, "version": 0.4 } \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_workflow_edit_2509.json b/comfyui/qwenimage/v1/qwenimage_workflow_edit_2509.json new file mode 100644 index 0000000..e310ff5 --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_workflow_edit_2509.json @@ -0,0 +1,376 @@ +{ + "id": "0afeb9a9-c8d6-4e64-b303-61e02e44da9e", + "revision": 0, + "last_node_id": 100, + "last_link_id": 86, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 84 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 98, + "type": "LoadImage", + "pos": [ + 312.6856384277344, + 418.9110107421875 + ], + "size": [ + 315, + 314.0000305175781 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 83 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "ref_1.png", + "image" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 82 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 85 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "把相机转变成西瓜" + ] + }, + { + "id": 99, + "type": "LoadQwenImageModel", + "pos": [ + 294.5816650390625, + -309.4810485839844 + ], + "size": [ + 318.1009826660156, + 106 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 86 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageModel" + }, + "widgets_values": [ + "Qwen-Image-Edit-2509", + "model_group_offload", + "bf16" + ] + }, + { + "id": 100, + "type": "QwenImageEditPlusSampler", + "pos": [ + 732.0558807778726, + -67.107393762887 + ], + "size": [ + 298.1490234375, + 406 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 86 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 85 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 84 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 83 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageEditPlusSampler" + }, + "widgets_values": [ + 1728, + 992, + 437532779396225, + "randomize", + 50, + 4, + "Flow", + 1, + 0.25, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 82, + 100, + 0, + 88, + 0, + "IMAGE" + ], + [ + 83, + 98, + 0, + 100, + 3, + "IMAGE" + ], + [ + 84, + 73, + 0, + 100, + 2, + "STRING_PROMPT" + ], + [ + 85, + 75, + 0, + 100, + 1, + "STRING_PROMPT" + ], + [ + 86, + 99, + 0, + 100, + 0, + "FunModels" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 226.02244567871094, + -405.3177185058594, + 440.6474914550781, + 238.3169403076172 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6905497838871149, + "offset": [ + 456.52379392690324, + 557.1551088532143 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "13a802b574c3a4397e193a0b8ca4e90c480cc217", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_workflow_edit_2511.json b/comfyui/qwenimage/v1/qwenimage_workflow_edit_2511.json new file mode 100644 index 0000000..fa0cd2f --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_workflow_edit_2511.json @@ -0,0 +1,376 @@ +{ + "id": "9a274935-92d9-49d6-bcd9-f4dbdc9c9b60", + "revision": 0, + "last_node_id": 100, + "last_link_id": 86, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 84 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 98, + "type": "LoadImage", + "pos": [ + 312.6856384277344, + 418.9110107421875 + ], + "size": [ + 315, + 314.0000305175781 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 83 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "ref_1.png", + "image" + ] + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 82 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 85 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "把相机转变成西瓜" + ] + }, + { + "id": 100, + "type": "QwenImageEditPlusSampler", + "pos": [ + 732.0558807778726, + -67.107393762887 + ], + "size": [ + 298.1490234375, + 406 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 86 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 85 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 84 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 83 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 82 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageEditPlusSampler" + }, + "widgets_values": [ + 1728, + 992, + 496807707597563, + "randomize", + 50, + 4, + "Flow", + 1, + 0.25, + true, + 5, + true, + 0 + ] + }, + { + "id": 99, + "type": "LoadQwenImageModel", + "pos": [ + 294.5816650390625, + -309.4810485839844 + ], + "size": [ + 318.1009826660156, + 106 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 86 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageModel" + }, + "widgets_values": [ + "Qwen-Image-Edit-2511", + "model_group_offload", + "bf16" + ] + } + ], + "links": [ + [ + 82, + 100, + 0, + 88, + 0, + "IMAGE" + ], + [ + 83, + 98, + 0, + 100, + 3, + "IMAGE" + ], + [ + 84, + 73, + 0, + 100, + 2, + "STRING_PROMPT" + ], + [ + 85, + 75, + 0, + 100, + 1, + "STRING_PROMPT" + ], + [ + 86, + 99, + 0, + 100, + 0, + "FunModels" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 226.02244567871094, + -405.3177185058594, + 440.6474914550781, + 238.3169403076172 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6905497838871149, + "offset": [ + 462.22577239663184, + 589.4946038050374 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "13a802b574c3a4397e193a0b8ca4e90c480cc217", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_workflow_t2i.json b/comfyui/qwenimage/v1/qwenimage_workflow_t2i.json index 75cc7ce..db1620e 100644 --- a/comfyui/qwenimage/v1/qwenimage_workflow_t2i.json +++ b/comfyui/qwenimage/v1/qwenimage_workflow_t2i.json @@ -149,7 +149,7 @@ }, "widgets_values": [ "Qwen-Image", - "model_cpu_offload_and_qfloat8", + "model_group_offload", "bf16" ] }, @@ -224,8 +224,8 @@ "Node name for S&R": "QwenImageT2VSampler" }, "widgets_values": [ - 1344, - 768, + 1728, + 992, 43, "randomize", 40, diff --git a/comfyui/qwenimage/v1/qwenimage_workflow_t2i_control.json b/comfyui/qwenimage/v1/qwenimage_workflow_t2i_control.json new file mode 100644 index 0000000..b5ac99c --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_workflow_t2i_control.json @@ -0,0 +1,436 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 94, + "last_link_id": 77, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 75 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 77 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 86, + "type": "LoadQwenImageModel", + "pos": [ + 314.6495666503906, + -281.51666259765625 + ], + "size": [ + 276.705078125, + 106 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 71 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageModel" + }, + "widgets_values": [ + "Qwen-Image-2512", + "model_group_offload", + "bf16" + ] + }, + { + "id": 92, + "type": "LoadQwenImageControlNetInPipeline", + "pos": [ + 644.9080724681564, + -284.9326871492188 + ], + "size": [ + 513.0712528489155, + 106 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 71 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 73 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageControlNetInPipeline" + }, + "widgets_values": [ + "qwenimage/qwenimage_control.yaml", + "Qwen-Image-2512-Fun-Controlnet-Union.safetensors", + "transformer" + ] + }, + { + "id": 94, + "type": "LoadImage", + "pos": [ + 348.0102300009322, + 406.45848378043814 + ], + "size": [ + 270, + 314 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 76 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "a7kXeQ5l9Dhspes7q3x3G (1).png", + "image" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 74 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 93, + "type": "QwenImageControlSampler", + "pos": [ + 723.0842521497987, + -74.3739187955474 + ], + "size": [ + 289.5494140625, + 470 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 73 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 74 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 75 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": 76 + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 77 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 733079601689805, + "randomize", + 40, + 4, + "Flow", + 3, + 0.25, + true, + 5, + true, + 0, + 0.8 + ] + } + ], + "links": [ + [ + 71, + 86, + 0, + 92, + 0, + "FunModels" + ], + [ + 73, + 92, + 0, + 93, + 0, + "FunModels" + ], + [ + 74, + 75, + 0, + 93, + 1, + "STRING_PROMPT" + ], + [ + 75, + 73, + 0, + 93, + 2, + "STRING_PROMPT" + ], + [ + 76, + 94, + 0, + 93, + 3, + "IMAGE" + ], + [ + 77, + 93, + 0, + 88, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 954.3592031237638, + 226.9206439292882 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7828871306993743, + "offset": [ + 428.8268866172707, + 570.3944941409308 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "ac114cc14285c8e0073a3e08e27525263d1264a7", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/qwenimage/v1/qwenimage_workflow_t2i_inpaint.json b/comfyui/qwenimage/v1/qwenimage_workflow_t2i_inpaint.json new file mode 100644 index 0000000..ed8b843 --- /dev/null +++ b/comfyui/qwenimage/v1/qwenimage_workflow_t2i_inpaint.json @@ -0,0 +1,526 @@ +{ + "id": "dcf2fcac-6293-4a86-b30b-f63e420177f2", + "revision": 0, + "last_node_id": 97, + "last_link_id": 81, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 75 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -92, + -294 + ], + "size": [ + 351.1499938964844, + 130.12660217285156 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 20B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用20B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "PreviewImage", + "pos": [ + 1070.207763671875, + -73.63389587402344 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 77 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 86, + "type": "LoadQwenImageModel", + "pos": [ + 314.6495666503906, + -281.51666259765625 + ], + "size": [ + 276.705078125, + 106 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 71 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageModel" + }, + "widgets_values": [ + "Qwen-Image-2512", + "model_group_offload", + "bf16" + ] + }, + { + "id": 92, + "type": "LoadQwenImageControlNetInPipeline", + "pos": [ + 644.9080724681564, + -284.9326871492188 + ], + "size": [ + 513.0712528489155, + 106 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 71 + } + ], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 73 + ] + } + ], + "properties": { + "Node name for S&R": "LoadQwenImageControlNetInPipeline" + }, + "widgets_values": [ + "qwenimage/qwenimage_control.yaml", + "Qwen-Image-2512-Fun-Controlnet-Union.safetensors", + "transformer" + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 74 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\\n\\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\\n\\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock." + ] + }, + { + "id": 95, + "type": "PreviewImage", + "pos": [ + 726.8217221845972, + 452.93109914645987 + ], + "size": [ + 366.56134033203125, + 415.4429626464844 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 78 + } + ], + "outputs": [], + "properties": { + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 93, + "type": "QwenImageControlSampler", + "pos": [ + 723.0842521497987, + -74.3739187955474 + ], + "size": [ + 289.5494140625, + 470 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 73 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 74 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 75 + }, + { + "name": "control_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "inpaint_image", + "shape": 7, + "type": "IMAGE", + "link": 80 + }, + { + "name": "mask_image", + "shape": 7, + "type": "IMAGE", + "link": 81 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 77 + ] + } + ], + "properties": { + "Node name for S&R": "QwenImageControlSampler" + }, + "widgets_values": [ + 1184, + 1568, + 427877921479533, + "randomize", + 40, + 4, + "Flow", + 3, + 0.25, + true, + 5, + true, + 0, + 0.8 + ] + }, + { + "id": 96, + "type": "MaskToImage", + "pos": [ + 552.9423636234835, + 457.6387299620162 + ], + "size": [ + 140, + 26 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 79 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 78, + 81 + ] + } + ], + "properties": { + "Node name for S&R": "MaskToImage" + }, + "widgets_values": [] + }, + { + "id": 97, + "type": "LoadImage", + "pos": [ + 247.8600973375456, + 460.28559660447246 + ], + "size": [ + 270, + 314.00000000000006 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 80 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 79 + ] + } + ], + "properties": { + "Node name for S&R": "LoadImage", + "image": "clipspace/clipspace-painted-masked-1766731857414.png [input]" + }, + "widgets_values": [ + "clipspace/clipspace-painted-masked-1766731857414.png [input]", + "image" + ] + } + ], + "links": [ + [ + 71, + 86, + 0, + 92, + 0, + "FunModels" + ], + [ + 73, + 92, + 0, + 93, + 0, + "FunModels" + ], + [ + 74, + 75, + 0, + 93, + 1, + "STRING_PROMPT" + ], + [ + 75, + 73, + 0, + 93, + 2, + "STRING_PROMPT" + ], + [ + 77, + 93, + 0, + 88, + 0, + "IMAGE" + ], + [ + 78, + 96, + 0, + 95, + 0, + "IMAGE" + ], + [ + 79, + 97, + 1, + 96, + 0, + "MASK" + ], + [ + 80, + 97, + 0, + 93, + 4, + "IMAGE" + ], + [ + 81, + 96, + 0, + 93, + 5, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 954.3592031237638, + 226.9206439292882 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7117155733630676, + "offset": [ + 359.457252578857, + 482.3639645185654 + ] + }, + "frontendVersion": "1.36.14", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "ac114cc14285c8e0073a3e08e27525263d1264a7", + "comfy-core": "0.9.2" + }, + "workflowRendererVersion": "LG" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_1/nodes.py b/comfyui/wan2_1/nodes.py index 2be06c8..4dc7248 100755 --- a/comfyui/wan2_1/nodes.py +++ b/comfyui/wan2_1/nodes.py @@ -93,11 +93,7 @@ class LoadWanTransformerModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() - mm.soft_empty_cache() - - mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() transformer = None @@ -501,7 +497,7 @@ class LoadWanModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar diff --git a/comfyui/wan2_1_fun/nodes.py b/comfyui/wan2_1_fun/nodes.py index bf8a46d..8801924 100755 --- a/comfyui/wan2_1_fun/nodes.py +++ b/comfyui/wan2_1_fun/nodes.py @@ -105,7 +105,7 @@ class LoadWanFunModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar diff --git a/comfyui/wan2_2/nodes.py b/comfyui/wan2_2/nodes.py index c44a519..172da35 100755 --- a/comfyui/wan2_2/nodes.py +++ b/comfyui/wan2_2/nodes.py @@ -73,7 +73,7 @@ class LoadWan2_2TransformerModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() transformer = None @@ -318,7 +318,7 @@ class LoadWan2_2Model: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar diff --git a/comfyui/wan2_2_fun/nodes.py b/comfyui/wan2_2_fun/nodes.py index 0b1832f..385096a 100755 --- a/comfyui/wan2_2_fun/nodes.py +++ b/comfyui/wan2_2_fun/nodes.py @@ -106,7 +106,7 @@ class LoadWan2_2FunModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar diff --git a/comfyui/wan2_2_vace_fun/nodes.py b/comfyui/wan2_2_vace_fun/nodes.py index 9988314..43d404a 100644 --- a/comfyui/wan2_2_vace_fun/nodes.py +++ b/comfyui/wan2_2_vace_fun/nodes.py @@ -70,7 +70,7 @@ class LoadVaceWanTransformer3DModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() transformer = None @@ -267,7 +267,7 @@ class LoadWan2_2VaceFunModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar diff --git a/comfyui/z_image/nodes.py b/comfyui/z_image/nodes.py index 18d6c8f..9286453 100644 --- a/comfyui/z_image/nodes.py +++ b/comfyui/z_image/nodes.py @@ -70,7 +70,7 @@ class LoadZImageTransformerModel: "required": { "model_name": ( folder_paths.get_filename_list("diffusion_models"), - {"default": "Wan2_1-T2V-1_3B_bf16.safetensors,"}, + {"default": "z_image_turbo_bf16.safetensors", }, ), "precision": (["fp16", "bf16"], {"default": "bf16"} @@ -89,7 +89,7 @@ class LoadZImageTransformerModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() transformer = None @@ -196,7 +196,7 @@ class LoadZImageVAEModel: "required": { "model_name": ( folder_paths.get_filename_list("vae"), - {"default": "ZImage2.1_VAE.pth"} + {"default": "ae.safetensors", } ), "precision": (["fp16", "bf16"], {"default": "bf16"} @@ -371,7 +371,7 @@ class LoadZImageTextEncoderModel: "required": { "model_name": ( folder_paths.get_filename_list("text_encoders"), - {"default": "models_t5_umt5-xxl-enc-bf16.pth"} + {"default": "qwen_3_4b.safetensors", } ), "precision": (["fp16", "bf16"], {"default": "bf16"} @@ -569,7 +569,7 @@ class LoadZImageModel: weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.unload_all_models() - mm.cleanup_models() + mm.cleanup_models_gc() mm.soft_empty_cache() # Init processbar @@ -726,10 +726,14 @@ class LoadZImageControlNetInPipeline: def loadmodel(self, config, model_name, sub_transformer_name, funmodels): device = mm.get_torch_device() offload_device = mm.unet_offload_device() + # Get Transformer transformer = getattr(funmodels["pipeline"], sub_transformer_name) transformer = transformer.cpu() + # Remove hooks + funmodels["pipeline"].remove_all_hooks() + # Load config config_path = f"{script_directory}/config/{config}" config = OmegaConf.load(config_path) @@ -797,14 +801,14 @@ class LoadZImageControlNetInPipeline: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": - convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) - convert_weight_dtype_wrapper(transformer, weight_dtype) + convert_model_weight_to_float8(control_transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(control_transformer, weight_dtype) pipeline.enable_model_cpu_offload(device=device) elif GPU_memory_mode == "model_cpu_offload": pipeline.enable_model_cpu_offload(device=device) elif GPU_memory_mode == "model_full_load_and_qfloat8": - convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) - convert_weight_dtype_wrapper(transformer, weight_dtype) + convert_model_weight_to_float8(control_transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(control_transformer, weight_dtype) pipeline.to(device=device) else: pipeline.to(device=device) @@ -830,7 +834,7 @@ class LoadZImageControlNetInModel: ), "model_name": ( folder_paths.get_filename_list("model_patches"), - {"default": "Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors",}, + {"default": "Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors", }, ), "transformer": ("TransformerModel",), }, diff --git a/examples/flux2/predict_t2i.py b/examples/flux2/predict_t2i.py index b0ff689..9ba37e4 100644 --- a/examples/flux2/predict_t2i.py +++ b/examples/flux2/predict_t2i.py @@ -16,6 +16,8 @@ from videox_fun.models import (AutoencoderKLFlux2, PixtralProcessor, Flux2Transformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Flux2Pipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -33,6 +35,9 @@ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. GPU_memory_mode = "sequential_cpu_offload" @@ -161,6 +166,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/flux2_fun/predict_i2i_inpaint.py b/examples/flux2_fun/predict_i2i_inpaint.py index e5c1ec5..8b42e99 100644 --- a/examples/flux2_fun/predict_i2i_inpaint.py +++ b/examples/flux2_fun/predict_i2i_inpaint.py @@ -18,6 +18,8 @@ from videox_fun.models import (AutoencoderKLFlux2, PixtralProcessor, Flux2ControlTransformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Flux2ControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -39,6 +41,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent, # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. GPU_memory_mode = "model_cpu_offload" @@ -177,6 +182,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/flux2_fun/predict_t2i_control.py b/examples/flux2_fun/predict_t2i_control.py index 6bf4dd7..869f0b7 100644 --- a/examples/flux2_fun/predict_t2i_control.py +++ b/examples/flux2_fun/predict_t2i_control.py @@ -18,6 +18,8 @@ from videox_fun.models import (AutoencoderKLFlux2, PixtralProcessor, Flux2ControlTransformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Flux2ControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -39,6 +41,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent, # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. GPU_memory_mode = "model_cpu_offload" @@ -177,6 +182,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/flux2_fun/predict_t2i_control_ref.py b/examples/flux2_fun/predict_t2i_control_ref.py index a2ab1d8..5687b96 100644 --- a/examples/flux2_fun/predict_t2i_control_ref.py +++ b/examples/flux2_fun/predict_t2i_control_ref.py @@ -18,6 +18,8 @@ from videox_fun.models import (AutoencoderKLFlux2, PixtralProcessor, Flux2ControlTransformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Flux2ControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -39,6 +41,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent, # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. GPU_memory_mode = "model_cpu_offload" @@ -177,6 +182,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/qwenimage/predict_t2i.py b/examples/qwenimage/predict_t2i.py index 16da8c4..e77ec6f 100644 --- a/examples/qwenimage/predict_t2i.py +++ b/examples/qwenimage/predict_t2i.py @@ -16,6 +16,8 @@ from videox_fun.models import (AutoencoderKLQwenImage, Qwen2Tokenizer, QwenImageTransformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import QwenImagePipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -33,9 +35,12 @@ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. -GPU_memory_mode = "model_cpu_offload_and_qfloat8" +GPU_memory_mode = "model_group_offload" # Multi GPUs config # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. @@ -177,6 +182,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/qwenimage/predict_t2i_edit.py b/examples/qwenimage/predict_t2i_edit.py index 58b8faa..6b4c6d3 100644 --- a/examples/qwenimage/predict_t2i_edit.py +++ b/examples/qwenimage/predict_t2i_edit.py @@ -17,6 +17,8 @@ from videox_fun.models import (AutoencoderKLQwenImage, QwenImageTransformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import QwenImageEditPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -35,9 +37,12 @@ from videox_fun.utils.utils import get_image # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. -GPU_memory_mode = "model_cpu_offload_and_qfloat8" +GPU_memory_mode = "model_group_offload" # Multi GPUs config # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. @@ -188,6 +193,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/qwenimage/predict_t2i_edit_plus.py b/examples/qwenimage/predict_t2i_edit_plus.py index edded60..60a60ea 100644 --- a/examples/qwenimage/predict_t2i_edit_plus.py +++ b/examples/qwenimage/predict_t2i_edit_plus.py @@ -17,6 +17,8 @@ from videox_fun.models import (AutoencoderKLQwenImage, QwenImageTransformer2DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import QwenImageEditPlusPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, @@ -35,9 +37,12 @@ from videox_fun.utils.utils import get_image # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. -GPU_memory_mode = "model_cpu_offload_and_qfloat8" +GPU_memory_mode = "model_group_offload" # Multi GPUs config # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. @@ -188,6 +193,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) diff --git a/examples/qwenimage_fun/predict_i2i_inpaint.py b/examples/qwenimage_fun/predict_i2i_inpaint.py index c4bd2df..e7722b4 100644 --- a/examples/qwenimage_fun/predict_i2i_inpaint.py +++ b/examples/qwenimage_fun/predict_i2i_inpaint.py @@ -2,9 +2,8 @@ import os import sys import torch - +from diffusers import FlowMatchEulerDiscreteScheduler from omegaconf import OmegaConf -from diffusers import (FlowMatchEulerDiscreteScheduler) current_file_path = os.path.abspath(__file__) project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] @@ -14,16 +13,18 @@ for project_root in project_roots: from videox_fun.dist import set_multi_gpus_devices, shard_model from videox_fun.models import (AutoencoderKLQwenImage, Qwen2_5_VLForConditionalGeneration, - Qwen2Tokenizer, QwenImageControlTransformer2DModel) + Qwen2Tokenizer, + QwenImageControlTransformer2DModel) +from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import QwenImageControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, convert_weight_dtype_wrapper) from videox_fun.utils.lora_utils import merge_lora, unmerge_lora -from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, get_image, - get_video_to_video_latent, - save_videos_grid) +from videox_fun.utils.utils import get_image_latent, save_videos_grid # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. # model_full_load means that the entire model will be moved to the GPU. @@ -36,9 +37,12 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, ge # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. -GPU_memory_mode = "model_cpu_offload_and_qfloat8" +GPU_memory_mode = "model_group_offload" # Multi GPUs config # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. @@ -52,6 +56,21 @@ fsdp_text_encoder = False # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. compile_dit = False +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +teacache_threshold = 0.30 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference for acceleration +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + # Config path config_path = "config/qwenimage/qwenimage_control.yaml" # Model path @@ -163,6 +182,7 @@ if ulysses_degree > 1 or ring_degree > 1: print("Add FSDP DIT") if fsdp_text_encoder: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) text_encoder = shard_fn(text_encoder) @@ -175,6 +195,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) @@ -188,6 +211,17 @@ elif GPU_memory_mode == "model_full_load_and_qfloat8": else: pipeline.to(device=device) +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + +if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps) + generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: diff --git a/examples/qwenimage_fun/predict_t2i_control.py b/examples/qwenimage_fun/predict_t2i_control.py index 69fe07e..828748b 100644 --- a/examples/qwenimage_fun/predict_t2i_control.py +++ b/examples/qwenimage_fun/predict_t2i_control.py @@ -2,9 +2,8 @@ import os import sys import torch - +from diffusers import FlowMatchEulerDiscreteScheduler from omegaconf import OmegaConf -from diffusers import (FlowMatchEulerDiscreteScheduler) current_file_path = os.path.abspath(__file__) project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] @@ -14,16 +13,18 @@ for project_root in project_roots: from videox_fun.dist import set_multi_gpus_devices, shard_model from videox_fun.models import (AutoencoderKLQwenImage, Qwen2_5_VLForConditionalGeneration, - Qwen2Tokenizer, QwenImageControlTransformer2DModel) + Qwen2Tokenizer, + QwenImageControlTransformer2DModel) +from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import QwenImageControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, convert_weight_dtype_wrapper) from videox_fun.utils.lora_utils import merge_lora, unmerge_lora -from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, get_image, - get_video_to_video_latent, - save_videos_grid) +from videox_fun.utils.utils import get_image_latent, save_videos_grid # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. # model_full_load means that the entire model will be moved to the GPU. @@ -36,9 +37,12 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, ge # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, # and the transformer model has been quantized to float8, which can save more GPU memory. # +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, # resulting in slower speeds but saving a large amount of GPU memory. -GPU_memory_mode = "model_cpu_offload_and_qfloat8" +GPU_memory_mode = "model_group_offload" # Multi GPUs config # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. @@ -52,6 +56,21 @@ fsdp_text_encoder = False # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. compile_dit = False +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +teacache_threshold = 0.30 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference for acceleration +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + # Config path config_path = "config/qwenimage/qwenimage_control.yaml" # Model path @@ -163,6 +182,7 @@ if ulysses_degree > 1 or ring_degree > 1: print("Add FSDP DIT") if fsdp_text_encoder: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) text_encoder = shard_fn(text_encoder) @@ -175,6 +195,9 @@ if compile_dit: if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) @@ -188,6 +211,17 @@ elif GPU_memory_mode == "model_full_load_and_qfloat8": else: pipeline.to(device=device) +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + +if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps) + generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: diff --git a/examples/qwenimage_instantx/predict_t2i_control.py b/examples/qwenimage_instantx/predict_t2i_control.py new file mode 100644 index 0000000..d64155b --- /dev/null +++ b/examples/qwenimage_instantx/predict_t2i_control.py @@ -0,0 +1,281 @@ +import os +import sys + +import torch + +from omegaconf import OmegaConf +from diffusers import (FlowMatchEulerDiscreteScheduler) + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLQwenImage, QwenImageInstantXControlNetModel, + Qwen2_5_VLForConditionalGeneration, + Qwen2Tokenizer, QwenImageTransformer2DModel) +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import QwenImageControlNetPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image, + get_video_to_video_latent, + save_videos_grid) + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# model_full_load means that the entire model will be moved to the GPU. +# +# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory. +# +# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# +# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, +# resulting in slower speeds but saving a large amount of GPU memory. +GPU_memory_mode = "model_group_offload" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1. +ulysses_degree = 1 +ring_degree = 1 +# Use FSDP to save more GPU memory in multi gpus. +fsdp_dit = False +fsdp_text_encoder = False +# Compile will give a speedup in fixed resolution and need a little GPU memory. +# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. +compile_dit = False + +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +teacache_threshold = 0.30 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference for acceleration +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + +# Model path +model_name = "models/Diffusion_Transformer/Qwen-Image" +# Controlnet Model path +model_name_controlnet = "models/Diffusion_Transformer/Qwen-Image-ControlNet-Union" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow" + +# Load pretrained model if need +transformer_path = None +controlnet_path = None +vae_path = None +lora_path = None + +# Other params +sample_size = [1728, 992] + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_image = "asset/pose.jpg" +controlnet_conditioning_scale = 0.80 + +# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性 +# 在neg prompt中添加"安静,固定"等词语可以增加动态性。 +prompt = "画面中央是一位年轻女孩,她拥有一头令人印象深刻的亮紫色长发,发丝在海风中轻盈飘扬,营造出动感而唯美的效果。她的长发两侧各扎着黑色蝴蝶结发饰,增添了几分可爱与俏皮感。女孩身穿一袭纯白色无袖连衣裙,裙摆轻盈飘逸,与她清新的气质完美契合。她的妆容精致自然,淡粉色的唇妆和温柔的眼神流露出恬静优雅的气质。她单手叉腰,姿态自信从容,目光直视镜头,展现出既甜美又不失个性的魅力。背景是一片开阔的海景,湛蓝的海水在阳光照射下波光粼粼,闪烁着钻石般的光芒。天空呈现出清澈的蔚蓝色,点缀着几朵洁白的云朵,营造出晴朗明媚的夏日氛围。画面前景右下角可见粉紫色的小花丛和绿色植物,为整体构图增添了自然生机和色彩层次。整张照片色调明亮清新,紫色头发与白色裙装、蓝色海天形成鲜明而和谐的色彩对比。" +negative_prompt = " " +guidance_scale = 4.0 +seed = 43 +num_inference_steps = 50 +lora_weight = 0.55 +save_path = "samples/qwenimage-t2i-instantx-control" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) + +transformer = QwenImageTransformer2DModel.from_pretrained( + model_name, + subfolder="transformer", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +).to(weight_dtype) + +controlnet = QwenImageInstantXControlNetModel.from_pretrained( + model_name_controlnet, + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +).to(weight_dtype) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +if controlnet_path is not None: + print(f"From checkpoint: {controlnet_path}") + if controlnet_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(controlnet_path) + else: + state_dict = torch.load(controlnet_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = controlnet.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +vae = AutoencoderKLQwenImage.from_pretrained( + model_name, + subfolder="vae" +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get tokenizer and text_encoder +tokenizer = Qwen2Tokenizer.from_pretrained( + model_name, subfolder="tokenizer" +) +text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained( + model_name, subfolder="text_encoder", torch_dtype=weight_dtype +) + +# Get Scheduler +Chosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +scheduler = Chosen_Scheduler.from_pretrained( + model_name, + subfolder="scheduler" +) + +pipeline = QwenImageControlNetPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=scheduler, + controlnet=controlnet, +) + +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.transformer = shard_fn(pipeline.transformer) + print("Add FSDP DIT") + if fsdp_text_encoder: + from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) + text_encoder = shard_fn(text_encoder) + print("Add FSDP TEXT ENCODER") + +if compile_dit: + for i in range(len(pipeline.transformer.transformer_blocks)): + pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + +if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) + +with torch.no_grad(): + control_image_input = get_image(control_image) + + sample = pipeline( + prompt=prompt, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + true_cfg_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + control_image = control_image_input, + controlnet_conditioning_scale = controlnet_conditioning_scale + ).images + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + image_path = os.path.join(save_path, prefix + ".png") + image = sample[0] + image.save(image_path) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index c0adf3c..4d0558b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,13 +1,43 @@ [project] name = "videox-fun" +version = "1.0.1" description = "VideoX-Fun is a video generation pipeline that can be used to generate AI images and videos, as well as to train baseline and Lora models for Diffusion Transformer. We support direct prediction from pre-trained baseline models to generate videos with different resolutions, durations, and FPS. Additionally, we also support users in training their own baseline and Lora models to perform specific style transformations." -version = "1.0.0" -license = {file = "LICENSE"} -dependencies = ["Pillow", "einops", "safetensors", "timm", "tomesd", "torch>=2.1.2", "torchdiffeq", "torchsde", "decord", "datasets", "numpy", "scikit-image", "opencv-python", "omegaconf", "SentencePiece", "albumentations", "imageio[ffmpeg]", "imageio[pyav]", "tensorboard", "beautifulsoup4", "ftfy", "func_timeout", "accelerate>=0.25.0", "gradio>=3.41.2,<=3.48.0", "diffusers>=0.30.1,<=0.31.0", "transformers>=4.46.2"] +license = { file = "LICENSE" } +dependencies = [ + "Pillow", + "einops", + "safetensors", + "timm", + "tomesd", + "torch>=2.1.2", + "torchdiffeq", + "torchsde", + "decord", + "datasets", + "numpy", + "scikit-image", + "opencv-python", + "omegaconf", + "SentencePiece", + "albumentations", + "imageio[ffmpeg]", + "imageio[pyav]", + "tensorboard", + "beautifulsoup4", + "ftfy", + "func_timeout", + "accelerate>=0.25.0", + "gradio>=3.41.2", + "diffusers>=0.30.1", + "transformers>=4.46.2", +] [project.urls] Repository = "https://github.com/aigc-apps/VideoX-Fun" -# Used by Comfy Registry https://comfyregistry.org +# Used by Comfy Registry https://comfyregistry.org + +[tool.setuptools] +packages = ["videox_fun"] [tool.comfy] PublisherId = "bubbliiiing" diff --git a/scripts/flux/train.py b/scripts/flux/train.py index 1bb94d3..f6a5783 100644 --- a/scripts/flux/train.py +++ b/scripts/flux/train.py @@ -244,60 +244,69 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = FluxPipeline( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = FluxTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = FluxPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - text_encoder_2=accelerator.unwrap_model(text_encoder_2), - tokenizer=tokenizer, - tokenizer_2=tokenizer_2, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1632,28 +1641,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1661,28 +1668,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux/train_lora.py b/scripts/flux/train_lora.py index 5eb45e3..c6dae26 100644 --- a/scripts/flux/train_lora.py +++ b/scripts/flux/train_lora.py @@ -249,61 +249,69 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = FluxPipeline( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = FluxTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = FluxPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - text_encoder_2=accelerator.unwrap_model(text_encoder_2), - tokenizer=tokenizer, - tokenizer_2=tokenizer_2, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1700,21 +1708,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1722,21 +1729,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux2/train.py b/scripts/flux2/train.py index 23455e4..7c9a423 100644 --- a/scripts/flux2/train.py +++ b/scripts/flux2/train.py @@ -313,58 +313,67 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = Flux2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = Flux2Transformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1712,26 +1721,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1739,27 +1746,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux2/train_lora.py b/scripts/flux2/train_lora.py index 1b42963..2469bfe 100644 --- a/scripts/flux2/train_lora.py +++ b/scripts/flux2/train_lora.py @@ -318,59 +318,67 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = Flux2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = Flux2Transformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1690,19 +1698,18 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1710,20 +1717,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux2_fun/README_TRAIN.md b/scripts/flux2_fun/README_TRAIN.md new file mode 100644 index 0000000..a79d9a0 --- /dev/null +++ b/scripts/flux2_fun/README_TRAIN.md @@ -0,0 +1,159 @@ +## Training Code + +We can choose whether to use deepspeed or fsdp in flux2, which can save a lot of video memory +. +The metadata_control.json is a little different from normal json in flux2, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. + +```json +[ + { + "file_path": "train/00000002.jpg", + "control_file_path": "control/00000002.jpg", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "image" + }, + ..... +] +``` + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images at the center, but instead, it trains the entire images after grouping them into buckets based on resolution. +- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum. + - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` +- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + +Without deepspeed: + +Training flux2 without DeepSpeed may result in insufficient GPU memory. +```sh +export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/flux2_fun/train_control.py \ + --config_path="config/flux2/flux2_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_flux2_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --low_vram \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" \ + --resume_from_checkpoint="latest" +``` + +With Deepspeed Zero-2: + +```sh +export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/flux2_fun/train_control.py \ + --config_path="config/flux2/flux2_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_flux2_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --low_vram \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" \ + --resume_from_checkpoint="latest" +``` + +With FSDP: + +```sh +export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap Flux2SingleTransformerBlock,BaseFlux2TransformerBlock,Flux2ControlTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/flux2_fun/train_control.py \ + --config_path="config/flux2/flux2_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_flux2_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --low_vram \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" \ + --resume_from_checkpoint="latest" +``` \ No newline at end of file diff --git a/scripts/flux2_fun/train_control.py b/scripts/flux2_fun/train_control.py new file mode 100644 index 0000000..7ea5e49 --- /dev/null +++ b/scripts/flux2_fun/train_control.py @@ -0,0 +1,1910 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys +from typing import (Any, Callable, Dict, List, NamedTuple, Optional, Tuple, + Union) + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import torchvision.transforms.functional as TF +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (EMAModel, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers import AutoTokenizer +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from qwen_vl_utils import process_vision_info + +from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) +from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, + ImageVideoDataset, + ImageVideoSampler, + get_random_mask, + process_pose_file, + process_pose_params) +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLFlux2, AutoProcessor, + Flux2ControlTransformer2DModel, + Mistral3ForConditionalGeneration, + PixtralProcessor) +from videox_fun.pipeline import Flux2ControlPipeline +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) + + +if is_wandb_available(): + import wandb + + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def compute_empirical_mu(image_seq_len: int, num_steps: int) -> float: + a1, b1 = 8.73809524e-05, 1.89833333 + a2, b2 = 0.00016927, 0.45666666 + + if image_seq_len > 4300: + mu = a2 * image_seq_len + b2 + return float(mu) + + m_200 = a2 * image_seq_len + b2 + m_10 = a1 * image_seq_len + b1 + + a = (m_200 - m_10) / 190.0 + b = m_200 - 200.0 * a + mu = a * num_steps + b + + return float(mu) + +def calculate_shift( + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 4096, + base_shift: float = 0.5, + max_shift: float = 1.15, +): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu + +def _prepare_latent_ids( + latents: torch.Tensor, # (B, C, H, W) +): + r""" + Generates 4D position coordinates (T, H, W, L) for latent tensors. + + Args: + latents (torch.Tensor): + Latent tensor of shape (B, C, H, W) + + Returns: + torch.Tensor: + Position IDs tensor of shape (B, H*W, 4) All batches share the same coordinate structure: T=0, + H=[0..H-1], W=[0..W-1], L=0 + """ + + batch_size, _, height, width = latents.shape + + t = torch.arange(1) # [0] - time dimension + h = torch.arange(height) + w = torch.arange(width) + l = torch.arange(1) # [0] - layer dimension + + # Create position IDs: (H*W, 4) + latent_ids = torch.cartesian_prod(t, h, w, l) + + # Expand to batch: (B, H*W, 4) + latent_ids = latent_ids.unsqueeze(0).expand(batch_size, -1, -1) + + return latent_ids + +def _patchify_latents(latents): + batch_size, num_channels_latents, height, width = latents.shape + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 1, 3, 5, 2, 4) + latents = latents.reshape(batch_size, num_channels_latents * 4, height // 2, width // 2) + return latents + +def _pack_latents(latents): + """ + pack latents: (batch_size, num_channels, height, width) -> (batch_size, height * width, num_channels) + """ + + batch_size, num_channels, height, width = latents.shape + latents = latents.reshape(batch_size, num_channels, height * width).permute(0, 2, 1) + + return latents + +def format_text_input(prompts: List[str], system_message: str = None): + # Remove [IMG] tokens from prompts to avoid Pixtral validation issues + # when truncation is enabled. The processor counts [IMG] tokens and fails + # if the count changes after truncation. + cleaned_txt = [prompt.replace("[IMG]", "") for prompt in prompts] + + return [ + [ + { + "role": "system", + "content": [{"type": "text", "text": system_message}], + }, + {"role": "user", "content": [{"type": "text", "text": prompt}]}, + ] + for prompt in cleaned_txt + ] + +def _get_mistral_3_small_prompt_embeds( + text_encoder: Mistral3ForConditionalGeneration, + tokenizer: PixtralProcessor, + prompt: Union[str, List[str]], + dtype: Optional[torch.dtype] = None, + device: Optional[torch.device] = None, + max_sequence_length: int = 512, + # fmt: off + system_message: str = "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object attribution and actions without speculation.", + # fmt: on + hidden_states_layers: List[int] = (10, 20, 30), +): + dtype = text_encoder.dtype if dtype is None else dtype + device = text_encoder.device if device is None else device + + prompt = [prompt] if isinstance(prompt, str) else prompt + + # Format input messages + messages_batch = format_text_input(prompts=prompt, system_message=system_message) + + # Process all messages at once + inputs = tokenizer.apply_chat_template( + messages_batch, + add_generation_prompt=False, + tokenize=True, + return_dict=True, + return_tensors="pt", + padding="max_length", + truncation=True, + max_length=max_sequence_length, + ) + + # Move to device + input_ids = inputs["input_ids"].to(device) + attention_mask = inputs["attention_mask"].to(device) + + # Forward pass through the model + output = text_encoder( + input_ids=input_ids, + attention_mask=attention_mask, + output_hidden_states=True, + use_cache=False, + ) + + # Only use outputs from intermediate layers and stack them + out = torch.stack([output.hidden_states[k] for k in hidden_states_layers], dim=1) + out = out.to(dtype=dtype, device=device) + + batch_size, num_channels, seq_len, hidden_dim = out.shape + prompt_embeds = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, num_channels * hidden_dim) + + return prompt_embeds + +def _prepare_text_ids( + x: torch.Tensor, # (B, L, D) or (L, D) + t_coord: Optional[torch.Tensor] = None, +): + B, L, _ = x.shape + out_ids = [] + + for i in range(B): + t = torch.arange(1) if t_coord is None else t_coord[i] + h = torch.arange(1) + w = torch.arange(1) + l = torch.arange(L) + + coords = torch.cartesian_prod(t, h, w, l) + out_ids.append(coords) + + return torch.stack(out_ids) + +def encode_prompt( + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + text_encoder=None, + tokenizer=None, + num_images_per_prompt: int = 1, + prompt_embeds: Optional[torch.Tensor] = None, + max_sequence_length: int = 512, + text_encoder_out_layers: Tuple[int] = (10, 20, 30), + system_message = "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object attribution and actions without speculation." +): + if prompt is None: + prompt = "" + + prompt = [prompt] if isinstance(prompt, str) else prompt + + if prompt_embeds is None: + prompt_embeds = _get_mistral_3_small_prompt_embeds( + text_encoder=text_encoder, + tokenizer=tokenizer, + prompt=prompt, + device=device, + max_sequence_length=max_sequence_length, + system_message=system_message, + hidden_states_layers=text_encoder_out_layers, + ) + + batch_size, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) + + text_ids = _prepare_text_ids(prompt_embeds) + text_ids = text_ids.to(device) + return prompt_embeds, text_ids + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): + try: + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = Flux2ControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + control_image = Image.open(args.validation_paths[i]) + width, height = control_image.width, control_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + control_image = get_image_latent(control_image, sample_size=(height, width))[:, :, 0] + + sample = pipeline( + prompt = args.validation_prompts[i], + height = height, + width = width, + generator = generator, + num_inference_steps = 20, + control_context_scale = 0.90, + control_image = control_image, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--config_path", + type=str, + default=None, + help=( + "The config of the model in training." + ), + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=512, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--prompt_template_encode", + type=str, + default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n", + help=( + 'The prompt template for text encoder.' + ), + ) + parser.add_argument( + "--prompt_template_encode_start_idx", + type=int, + default=34, + help=( + 'The start idx for prompt template.' + ), + ) + parser.add_argument( + "--train_mode", + type=str, + default="normal", + help=( + 'The format of training data. Support `"normal"`' + ' (default), `"i2v"`.' + ), + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=3.5, + help="the FLUX.1 dev variant is a guidance distilled model", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + + # Get Tokenizer + tokenizer = PixtralProcessor.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer" + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + config = OmegaConf.load(args.config_path) + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So Mistral3ForConditionalGeneration and AutoencoderKLFlux2 will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = Mistral3ForConditionalGeneration.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype + ) + text_encoder = text_encoder.eval() + # Get Vae + vae = AutoencoderKLFlux2.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae" + ).to(weight_dtype) + vae.eval() + latents_bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(accelerator.device, weight_dtype) + latents_bn_std = torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) + vae.config.batch_norm_eps).to(accelerator.device, weight_dtype) + + # Get Transformer + transformer3d = Flux2ControlTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + low_cpu_mem_usage=True, + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + transformer3d.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # Create EMA for the transformer3d. + if args.use_ema: + if zero_stage == 3: + raise NotImplementedError("FSDP does not support EMA.") + + ema_transformer3d = Flux2ControlTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + ).to(weight_dtype) + + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=Flux2ControlTransformer2DModel, model_config=ema_transformer3d.config) + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = Flux2ControlTransformer2DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = Flux2ControlTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + ) + load_model = EMAModel(load_model.parameters(), model_cls=Flux2ControlTransformer2DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) + + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model + + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = Flux2ControlTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in transformer3d.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + if args.fix_sample_size is not None and args.enable_bucket: + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoControlDataset( + args.train_data_meta, args.train_data_dir, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, + enable_inpaint=True, + enable_camera_info=False, + enable_subject_info=False, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + return [first_element] + [other_elements_value] * (length - 1) + + MIN_TARGET = 1024 + + if sample_size < MIN_TARGET: + number_list = [1.0] + else: + max_allowed_ratio = sample_size / MIN_TARGET + base_ratios = [ + 1.0, + 1.1, 1.2, 1.25, 1.33, 1.5, + 1.75, 2.0, 2.25, 2.5, 2.75, + 3.0, 3.5, 4.0, 5.0, 6.0, 8.0 + ] + candidate_ratios = set(base_ratios + list(image_ratio)) + number_list = sorted([r for r in candidate_ratios if 1.0 <= r <= max_allowed_ratio]) + + if not number_list: + number_list = [1.0] + + if all_choices: + return number_list + + probs = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p=probs) + else: + return rng.choice(number_list, p=probs) + + # Create new output + new_examples = {} + new_examples["pixel_values"] = [] + new_examples["text"] = [] + + # Used in Control Mode + new_examples["control_pixel_values"] = [] + + # Used in Inpaint mode + new_examples["mask_pixel_values"] = [] + new_examples["mask"] = [] + + # Get downsample ratio in image + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + + if args.fix_sample_size is not None: + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + length = int(len(pixel_values) // 2) + new_examples["pixel_values"].append(transform(pixel_values)[length:length + 1]) + new_examples["control_pixel_values"].append(transform(control_pixel_values)[length:length + 1]) + + new_examples["text"].append(example["text"]) + + mask = get_random_mask(new_examples["pixel_values"][-1].size()) + mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) + + new_examples["mask_pixel_values"].append(mask_pixel_values[:1]) + new_examples["mask"].append(mask[:1]) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) + new_examples["control_pixel_values"] = torch.stack([example for example in new_examples["control_pixel_values"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) + + # Encode prompts when enable_text_encoder_in_dataloader=True + if args.enable_text_encoder_in_dataloader: + template = args.prompt_template_encode + drop_idx = args.prompt_template_encode_start_idx + + txt = [template.format(e) for e in batch['text']] + txt_tokens = tokenizer( + txt, max_length=args.tokenizer_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt" + ).to(accelerator.device) + encoder_hidden_states = text_encoder( + input_ids=txt_tokens.input_ids, + attention_mask=txt_tokens.attention_mask, + output_hidden_states=True, + ) + hidden_states = encoder_hidden_states.hidden_states[-1] + split_hidden_states = _extract_masked_hidden(hidden_states, txt_tokens.attention_mask) + split_hidden_states = [e[drop_idx:] for e in split_hidden_states] + attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] + max_seq_len = max([e.size(0) for e in split_hidden_states]) + prompt_embeds = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] + ) + encoder_attention_mask = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] + ) + + prompt_embeds = prompt_embeds.to(dtype=latents.dtype, device=accelerator.device) + + new_examples['encoder_attention_mask'] = encoder_attention_mask + new_examples['encoder_hidden_states'] = prompt_embeds + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + transformer3d.requires_grad_(True) + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + if fsdp_stage != 0 or zero_stage != 0: + from functools import partial + + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) + text_encoder = shard_fn(text_encoder) + + if args.use_ema: + ema_transformer3d.to(accelerator.device) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream and args.train_mode != "normal": + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + # Calculate the index we need】 + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step < 1: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + control_pixel_values = batch["control_pixel_values"].cpu() + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + control_pixel_values = rearrange(control_pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, control_pixel_value, text) in enumerate(zip(pixel_values, control_pixel_values, texts)): + pixel_value = pixel_value[None, ...] + control_pixel_value = control_pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + save_videos_grid(control_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_control.gif", rescale=True) + + mask_pixel_values, mask, texts = batch['mask_pixel_values'].cpu(), batch['mask'].cpu(), batch['text'] + mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w") + mask = torch.tile(rearrange(mask, "b f c h w -> b c f h w"), [1, 3, 1, 1, 1]) + for idx, (pixel_value, _mask, text) in enumerate(zip(mask_pixel_values, mask, texts)): + pixel_value = pixel_value[None, ...] + _mask = _mask[None, ...] + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/mask_pixel_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.gif", rescale=True) + save_videos_grid(_mask, f"{args.output_dir}/sanity_check/mask_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + control_pixel_values = batch["control_pixel_values"].to(weight_dtype) + mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) + mask = batch["mask"].to(weight_dtype) + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = pixel_values.squeeze(1) + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) + else: + latents = _batch_encode_vae(pixel_values) + + # Control Latents + control_latents = _batch_encode_vae(control_pixel_values) + control_latents = _patchify_latents(control_latents) + control_latents = ((control_latents - latents_bn_mean) / latents_bn_std).to(dtype=weight_dtype) + control_latents = _pack_latents(control_latents) + + for bs_index in range(control_latents.size()[0]): + if rng is None: + zero_init_control_conv_in = np.random.choice([0, 1], p = [0.90, 0.10]) + else: + zero_init_control_conv_in = rng.choice([0, 1], p = [0.90, 0.10]) + if zero_init_control_conv_in: + control_latents[bs_index] = control_latents[bs_index] * 0 + + mask = rearrange(mask, "b f c h w -> b c f h w").squeeze(2) + mask_conditions = F.interpolate(1 - mask, size=latents.size()[-2:], mode='nearest').to(accelerator.device, weight_dtype) + mask_conditions = _patchify_latents(mask_conditions) + mask_conditions = _pack_latents(mask_conditions) + + t2v_flag = [(_mask == 1).all() for _mask in mask] + new_t2v_flag = [] + for _mask in t2v_flag: + if _mask and np.random.rand() < 0.90: + new_t2v_flag.append(0) + else: + new_t2v_flag.append(1) + t2v_flag = torch.from_numpy(np.array(new_t2v_flag)).to(accelerator.device, dtype=weight_dtype) + + # Encode inpaint latents. + mask_latents = _batch_encode_vae(mask_pixel_values) + mask_latents = _patchify_latents(mask_latents) + mask_latents = ((mask_latents - latents_bn_mean) / latents_bn_std).to(dtype=weight_dtype) + mask_latents = _pack_latents(mask_latents) + mask_latents = t2v_flag[:, None, None] * mask_latents + + inpaint_latents = torch.concat([mask_conditions, mask_latents], dim=2) + control_context = torch.cat([control_latents, inpaint_latents], dim=2) + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + if args.enable_text_encoder_in_dataloader: + prompt_embeds = batch['prompt_embeds'].to(dtype=latents.dtype, device=accelerator.device) + text_ids = batch['text_ids'] + else: + with torch.no_grad(): + prompt_embeds, text_ids = encode_prompt( + batch['text'], device=accelerator.device, + text_encoder=text_encoder, + tokenizer=tokenizer, + ) + + if args.low_vram and not args.enable_text_encoder_in_dataloader: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, height, width = latents.size() + latents = _patchify_latents(latents) + latent_image_ids = _prepare_latent_ids(latents) + latents = ((latents - latents_bn_mean) / latents_bn_std).to(dtype=weight_dtype) + latents = _pack_latents(latents) + + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + # handle guidance + guidance = torch.tensor([args.guidance_scale], device=accelerator.device) + guidance = guidance.expand(latents.shape[0]) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + + sigmas = np.linspace(1.0, 1 / args.train_sampling_steps, args.train_sampling_steps) + image_seq_len = latents.shape[1] + mu = calculate_shift( + image_seq_len, + noise_scheduler.config.get("base_image_seq_len", 256), + noise_scheduler.config.get("max_image_seq_len", 4096), + noise_scheduler.config.get("base_shift", 0.5), + noise_scheduler.config.get("max_shift", 1.15), + ) + noise_scheduler.set_timesteps(sigmas=sigmas, device=latents.device, mu=mu) + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + noise_pred = transformer3d( + hidden_states=noisy_latents, + timestep=timesteps / 1000, + guidance=guidance, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + control_context=control_context, + return_dict=False, + )[0] + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + + if args.use_ema: + ema_transformer3d.step(transformer3d.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name not in name: + param.requires_grad = False + break + accelerator.save_state(save_path) + transformer3d.requires_grad_(True) + logger.info(f"Saved state to {save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/flux2_fun/train_control.sh b/scripts/flux2_fun/train_control.sh new file mode 100644 index 0000000..48d0c05 --- /dev/null +++ b/scripts/flux2_fun/train_control.sh @@ -0,0 +1,36 @@ +export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/flux2_fun/train_control.py \ + --config_path="config/flux2/flux2_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_flux2_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --low_vram \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" \ + --resume_from_checkpoint="latest" \ No newline at end of file diff --git a/scripts/qwenimage/train.py b/scripts/qwenimage/train.py index c943e5d..14e82f6 100644 --- a/scripts/qwenimage/train.py +++ b/scripts/qwenimage/train.py @@ -137,55 +137,67 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = QwenImagePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = QwenImageTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = QwenImagePipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + true_cfg_scale = 4.0, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1553,25 +1565,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1579,25 +1590,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/qwenimage/train_edit.py b/scripts/qwenimage/train_edit.py index 787eaae..1c39914 100644 --- a/scripts/qwenimage/train_edit.py +++ b/scripts/qwenimage/train_edit.py @@ -137,71 +137,87 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, processor, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = QwenImageTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - if args.train_mode == "qwen_image_edit": - pipeline = QwenImageEditPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" ) - else: - pipeline = QwenImageEditPlusPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + transformer3d.config = accelerator.unwrap_model(transformer3d).config + if args.train_mode == "qwen_image_edit": + pipeline = QwenImageEditPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + processor=processor, + scheduler=scheduler, + ) + else: + pipeline = QwenImageEditPlusPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + processor=processor, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode == "qwen_image_edit": - image = get_image(args.validation_image_paths[i]) - else: - image = [get_image(args.validation_image_paths[i])] - sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", - height = args.image_sample_size, - width = args.image_sample_size, - generator = generator, - image = image - ).images - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + for i in range(len(args.validation_prompts)): + with torch.no_grad(): + if args.train_mode == "qwen_image_edit": + image = get_image(args.validation_image_paths[i]) + else: + image = [get_image(args.validation_image_paths[i])] + sample = pipeline( + prompt = args.validation_prompts[i], + negative_prompt = "bad detailed", + height = args.image_sample_size, + width = args.image_sample_size, + generator = generator, + image = image, + true_cfg_scale = 4.0, + num_inference_steps = 20, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1729,25 +1745,25 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + processor, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1755,25 +1771,25 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + processor, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/qwenimage/train_edit_lora.py b/scripts/qwenimage/train_edit_lora.py index c20f3e5..54eb289 100644 --- a/scripts/qwenimage/train_edit_lora.py +++ b/scripts/qwenimage/train_edit_lora.py @@ -145,75 +145,88 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, processor, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = QwenImageTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - if args.train_mode == "qwen_image_edit": - pipeline = QwenImageEditPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" ) - else: - pipeline = QwenImageEditPlusPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.train_mode == "qwen_image_edit": + pipeline = QwenImageEditPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + processor=processor, + scheduler=scheduler, + ) + else: + pipeline = QwenImageEditPlusPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + processor=processor, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + for i in range(len(args.validation_prompts)): + with torch.no_grad(): + if args.train_mode == "qwen_image_edit": + image = get_image(args.validation_image_paths[i]) + else: + image = [get_image(args.validation_image_paths[i])] + sample = pipeline( + prompt = args.validation_prompts[i], + negative_prompt = "bad detailed", + height = args.image_sample_size, + width = args.image_sample_size, + generator = generator, + image = image, + true_cfg_scale = 4.0, + num_inference_steps = 20, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode == "qwen_image_edit": - image = get_image(args.validation_image_paths[i]) - else: - image = [get_image(args.validation_image_paths[i])] - sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", - height = args.image_sample_size, - width = args.image_sample_size, - generator = generator, - image = image - ).images - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1802,19 +1815,19 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + processor, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1822,19 +1835,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + processor, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/qwenimage/train_lora.py b/scripts/qwenimage/train_lora.py index 11cf808..545a8da 100644 --- a/scripts/qwenimage/train_lora.py +++ b/scripts/qwenimage/train_lora.py @@ -136,59 +136,69 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = QwenImagePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = QwenImageTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = QwenImagePipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + true_cfg_scale = 4.0, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1619,19 +1629,18 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1639,19 +1648,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/qwenimage_fun/README_TRAIN.md b/scripts/qwenimage_fun/README_TRAIN.md new file mode 100644 index 0000000..cabdebf --- /dev/null +++ b/scripts/qwenimage_fun/README_TRAIN.md @@ -0,0 +1,153 @@ +## Training Code + +We can choose whether to use deepspeed or fsdp in qwen_image, which can save a lot of video memory +. +The metadata_control.json is a little different from normal json in Qwen-Image, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. + +```json +[ + { + "file_path": "train/00000002.jpg", + "control_file_path": "control/00000002.jpg", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "image" + }, + ..... +] +``` + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images at the center, but instead, it trains the entire images after grouping them into buckets based on resolution. +- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum. + - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` +- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + +Without deepspeed: + +Training qwen_image without DeepSpeed may result in insufficient GPU memory. +```sh +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/qwenimage_fun/train_control.py \ + --config_path="config/qwenimage/qwenimage_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwenimage_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/Qwen-Image-2512-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" +``` + +With Deepspeed Zero-2: + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/qwenimage_fun/train_control.py \ + --config_path="config/qwenimage/qwenimage_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwenimage_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/Qwen-Image-2512-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" +``` + +With FSDP: + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap BaseQwenImageTransformerBlock,QwenImageControlTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/qwenimage_fun/train_control.py \ + --config_path="config/qwenimage/qwenimage_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwenimage_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/Qwen-Image-2512-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" +``` \ No newline at end of file diff --git a/scripts/qwenimage_fun/train_control.py b/scripts/qwenimage_fun/train_control.py new file mode 100644 index 0000000..55b2c5c --- /dev/null +++ b/scripts/qwenimage_fun/train_control.py @@ -0,0 +1,1731 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (EMAModel, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) +from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, + ImageVideoSampler, + get_random_mask) +from videox_fun.models import (AutoencoderKLQwenImage, + Qwen2_5_VLForConditionalGeneration, + Qwen2Tokenizer, QwenImageControlTransformer2DModel) +from videox_fun.pipeline import QwenImageControlPipeline +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) + +if is_wandb_available(): + pass + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def _pack_latents(latents, batch_size, num_channels_latents, height, width, num_frame=None): + if num_frame is None: + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 2, 4, 1, 3, 5) + latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4) + else: + latents = latents.view(batch_size, num_channels_latents, num_frame, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 2, 3, 5, 1, 4, 6) + latents = latents.reshape(batch_size, num_frame * (height // 2) * (width // 2), num_channels_latents * 4) + return latents + +def _extract_masked_hidden(hidden_states: torch.Tensor, mask: torch.Tensor): + bool_mask = mask.bool() + valid_lengths = bool_mask.sum(dim=1) + selected = hidden_states[bool_mask] + split_result = torch.split(selected, valid_lengths.tolist(), dim=0) + + return split_result + +def calculate_shift( + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 4096, + base_shift: float = 0.5, + max_shift: float = 1.15, +): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): + try: + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = QwenImageControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + control_image = Image.open(args.validation_paths[i]) + width, height = control_image.width, control_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + control_image = get_image_latent(control_image, sample_size=(height, width))[:, :, 0] + + sample = pipeline( + args.validation_prompts[i], + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + true_cfg_scale = 4.0, + num_inference_steps = 20, + control_image = control_image, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--token_sample_size", + type=int, + default=512, + help="Sample size of the token.", + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--config_path", + type=str, + default=None, + help=( + "The config of the model in training." + ), + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=1024, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--prompt_template_encode", + type=str, + default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n", + help=( + 'The prompt template for text encoder.' + ), + ) + parser.add_argument( + "--prompt_template_encode_start_idx", + type=int, + default=34, + help=( + 'The start idx for prompt template.' + ), + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=3.5, + help="the FLUX.1 dev variant is a guidance distilled model", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + + # Get Tokenizer + tokenizer = Qwen2Tokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer" + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + config = OmegaConf.load(args.config_path) + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype + ) + text_encoder = text_encoder.eval() + + # Get Vae + vae = AutoencoderKLQwenImage.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae" + ).to(weight_dtype) + vae.eval() + latents_mean = (torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1)).to(accelerator.device) + latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to(accelerator.device) + + # Get Transformer + transformer3d = QwenImageControlTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + low_cpu_mem_usage=True, + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + transformer3d.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # Create EMA for the transformer3d. + if args.use_ema: + if zero_stage == 3: + raise NotImplementedError("FSDP does not support EMA.") + + ema_transformer3d = QwenImageControlTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + ).to(weight_dtype) + + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=QwenImageControlTransformer2DModel, model_config=ema_transformer3d.config) + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = QwenImageControlTransformer2DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = QwenImageControlTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + ) + load_model = EMAModel(load_model.parameters(), model_cls=QwenImageControlTransformer2DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) + + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model + + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = QwenImageControlTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in transformer3d.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + if args.fix_sample_size is not None and args.enable_bucket: + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoControlDataset( + args.train_data_meta, args.train_data_dir, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, + enable_inpaint=True, + enable_camera_info=False, + enable_subject_info=False, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + return [first_element] + [other_elements_value] * (length - 1) + + MIN_TARGET = 1024 + + if sample_size < MIN_TARGET: + number_list = [1.0] + else: + max_allowed_ratio = sample_size / MIN_TARGET + base_ratios = [ + 1.0, + 1.1, 1.2, 1.25, 1.33, 1.5, + 1.75, 2.0, 2.25, 2.5, 2.75, + 3.0, 3.5, 4.0, 5.0, 6.0, 8.0 + ] + candidate_ratios = set(base_ratios + list(image_ratio)) + number_list = sorted([r for r in candidate_ratios if 1.0 <= r <= max_allowed_ratio]) + + if not number_list: + number_list = [1.0] + + if all_choices: + return number_list + + probs = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p=probs) + else: + return rng.choice(number_list, p=probs) + + # Create new output + new_examples = {} + new_examples["pixel_values"] = [] + new_examples["text"] = [] + + # Used in Control Mode + new_examples["control_pixel_values"] = [] + + # Used in Inpaint mode + new_examples["mask_pixel_values"] = [] + new_examples["mask"] = [] + + # Get downsample ratio in image + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + + if args.fix_sample_size is not None: + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + length = int(len(pixel_values) // 2) + new_examples["pixel_values"].append(transform(pixel_values)[length:length + 1]) + new_examples["control_pixel_values"].append(transform(control_pixel_values)[length:length + 1]) + + new_examples["text"].append(example["text"]) + + mask = get_random_mask(new_examples["pixel_values"][-1].size()) + mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) + + new_examples["mask_pixel_values"].append(mask_pixel_values[:1]) + new_examples["mask"].append(mask[:1]) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) + new_examples["control_pixel_values"] = torch.stack([example for example in new_examples["control_pixel_values"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) + + # Encode prompts when enable_text_encoder_in_dataloader=True + if args.enable_text_encoder_in_dataloader: + template = args.prompt_template_encode + drop_idx = args.prompt_template_encode_start_idx + + txt = [template.format(e) for e in batch['text']] + txt_tokens = tokenizer( + txt, max_length=args.tokenizer_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt" + ).to(accelerator.device) + encoder_hidden_states = text_encoder( + input_ids=txt_tokens.input_ids, + attention_mask=txt_tokens.attention_mask, + output_hidden_states=True, + ) + hidden_states = encoder_hidden_states.hidden_states[-1] + split_hidden_states = _extract_masked_hidden(hidden_states, txt_tokens.attention_mask) + split_hidden_states = [e[drop_idx:] for e in split_hidden_states] + attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] + max_seq_len = max([e.size(0) for e in split_hidden_states]) + prompt_embeds = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] + ) + encoder_attention_mask = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] + ) + + prompt_embeds = prompt_embeds.to(dtype=latents.dtype, device=accelerator.device) + + new_examples['encoder_attention_mask'] = encoder_attention_mask + new_examples['encoder_hidden_states'] = prompt_embeds + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + if fsdp_stage != 0: + from functools import partial + + from videox_fun.dist import shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) + text_encoder = shard_fn(text_encoder) + + if args.use_ema: + ema_transformer3d.to(accelerator.device) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream and args.train_mode != "normal": + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step < 1: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + control_pixel_values = batch["control_pixel_values"].cpu() + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + control_pixel_values = rearrange(control_pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, control_pixel_value, text) in enumerate(zip(pixel_values, control_pixel_values, texts)): + pixel_value = pixel_value[None, ...] + control_pixel_value = control_pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + save_videos_grid(control_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_control.gif", rescale=True) + + mask_pixel_values, mask, texts = batch['mask_pixel_values'].cpu(), batch['mask'].cpu(), batch['text'] + mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w") + mask = torch.tile(rearrange(mask, "b f c h w -> b c f h w"), [1, 3, 1, 1, 1]) + for idx, (pixel_value, _mask, text) in enumerate(zip(mask_pixel_values, mask, texts)): + pixel_value = pixel_value[None, ...] + _mask = _mask[None, ...] + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/mask_pixel_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.gif", rescale=True) + save_videos_grid(_mask, f"{args.output_dir}/sanity_check/mask_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + control_pixel_values = batch["control_pixel_values"].to(weight_dtype) + mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) + mask = batch["mask"].to(weight_dtype) + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) + else: + latents = _batch_encode_vae(pixel_values) + latents = ((latents - latents_mean) * latents_std).to(dtype=weight_dtype) + + control_latents = _batch_encode_vae(control_pixel_values) + control_latents = ((control_latents - latents_mean) * latents_std).to(dtype=weight_dtype) + + for bs_index in range(control_latents.size()[0]): + if rng is None: + zero_init_control_conv_in = np.random.choice([0, 1], p = [0.90, 0.10]) + else: + zero_init_control_conv_in = rng.choice([0, 1], p = [0.90, 0.10]) + if zero_init_control_conv_in: + control_latents[bs_index] = control_latents[bs_index] * 0 + + mask = mask.squeeze(1) + # mask = rearrange(mask, "b f c h w -> b c f h w") + mask_conditions = F.interpolate(1 - mask[:, :1], size=control_latents.size()[-2:], mode='nearest').to(accelerator.device, weight_dtype) + mask_conditions = mask_conditions.unsqueeze(2) + + # Encode inpaint latents. + t2v_flag = [(_mask == 1).all() for _mask in mask] + new_t2v_flag = [] + for _mask in t2v_flag: + if _mask and np.random.rand() < 0.90: + new_t2v_flag.append(0) + else: + new_t2v_flag.append(1) + t2v_flag = torch.from_numpy(np.array(new_t2v_flag)).to(accelerator.device, dtype=weight_dtype) + + mask_latents = _batch_encode_vae(mask_pixel_values) + mask_latents = ((mask_latents - latents_mean) * latents_std).to(dtype=weight_dtype) + mask_latents = t2v_flag[:, None, None] * mask_latents + + inpaint_latents = torch.concat([mask_conditions, mask_latents], dim=1) + control_context = torch.cat([control_latents, inpaint_latents], dim=1) + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + if args.enable_text_encoder_in_dataloader: + prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device) + encoder_attention_mask = batch['encoder_attention_mask'] + else: + with torch.no_grad(): + template = args.prompt_template_encode + drop_idx = args.prompt_template_encode_start_idx + + txt = [template.format(e) for e in batch['text']] + txt_tokens = tokenizer( + txt, max_length=args.tokenizer_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt" + ).to(accelerator.device) + encoder_hidden_states = text_encoder( + input_ids=txt_tokens.input_ids, + attention_mask=txt_tokens.attention_mask, + output_hidden_states=True, + ) + hidden_states = encoder_hidden_states.hidden_states[-1] + split_hidden_states = _extract_masked_hidden(hidden_states, txt_tokens.attention_mask) + split_hidden_states = [e[drop_idx:] for e in split_hidden_states] + attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] + max_seq_len = max([e.size(0) for e in split_hidden_states]) + prompt_embeds = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] + ) + encoder_attention_mask = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] + ) + + prompt_embeds = prompt_embeds.to(dtype=latents.dtype, device=accelerator.device) + + if args.low_vram and not args.enable_text_encoder_in_dataloader: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, num_frame, height, width = latents.size() + latents = _pack_latents(latents, bsz, channel, height, width, num_frame=num_frame) + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + control_context = _pack_latents(control_context, bsz, control_context.size(1), height, width, num_frame=num_frame) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + + sigmas = np.linspace(1.0, 1 / args.train_sampling_steps, args.train_sampling_steps) + image_seq_len = latents.shape[1] + mu = calculate_shift( + image_seq_len, + noise_scheduler.config.get("base_image_seq_len", 256), + noise_scheduler.config.get("max_image_seq_len", 4096), + noise_scheduler.config.get("base_shift", 0.5), + noise_scheduler.config.get("max_shift", 1.15), + ) + noise_scheduler.set_timesteps(sigmas=sigmas, device=latents.device, mu=mu) + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + img_shapes = [[(num_frame, height // 2, width // 2)]] * latents.size(0) + txt_seq_lens = encoder_attention_mask.sum(dim=1).tolist() if encoder_attention_mask is not None else None + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + noise_pred = transformer3d( + hidden_states=noisy_latents, + timestep=timesteps / 1000, + encoder_hidden_states_mask=encoder_attention_mask, + encoder_hidden_states=prompt_embeds, + img_shapes=img_shapes, + txt_seq_lens=txt_seq_lens, + control_context=control_context, + return_dict=False, + ) + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + + if args.use_ema: + ema_transformer3d.step(transformer3d.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/qwenimage_fun/train_control.sh b/scripts/qwenimage_fun/train_control.sh new file mode 100644 index 0000000..cfef283 --- /dev/null +++ b/scripts/qwenimage_fun/train_control.sh @@ -0,0 +1,34 @@ +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/qwenimage_fun/train_control.py \ + --config_path="config/qwenimage/qwenimage_control.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwen_image_fun_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/Qwen-Image-2512-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" \ No newline at end of file diff --git a/scripts/qwenimage_instantx/README_TRAIN.md b/scripts/qwenimage_instantx/README_TRAIN.md new file mode 100644 index 0000000..9c7ba5f --- /dev/null +++ b/scripts/qwenimage_instantx/README_TRAIN.md @@ -0,0 +1,153 @@ +## Training Code + +We can choose whether to use deepspeed or fsdp in qwen_image, which can save a lot of video memory +. +The metadata_control.json is a little different from normal json in Qwen-Image, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. + +```json +[ + { + "file_path": "train/00000002.jpg", + "control_file_path": "control/00000002.jpg", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "image" + }, + ..... +] +``` + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images at the center, but instead, it trains the entire images after grouping them into buckets based on resolution. +- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum. + - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` +- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + +Without deepspeed: + +Training qwen_image without DeepSpeed may result in insufficient GPU memory. +```sh +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export CN_MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-ControlNet-Union" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/qwenimage_instantx/train_control.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --cn_pretrained_model_name_or_path=$CN_MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwen_image_instantx_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +With Deepspeed Zero-2: + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export CN_MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-ControlNet-Union" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/qwenimage_instantx/train_control.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --cn_pretrained_model_name_or_path=$CN_MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwen_image_instantx_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +With FSDP: + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export CN_MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-ControlNet-Union" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap QwenImageTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/qwenimage_instantx/train_control.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --cn_pretrained_model_name_or_path=$CN_MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwen_image_instantx_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` \ No newline at end of file diff --git a/scripts/qwenimage_instantx/train_control.py b/scripts/qwenimage_instantx/train_control.py new file mode 100644 index 0000000..8fa883b --- /dev/null +++ b/scripts/qwenimage_instantx/train_control.py @@ -0,0 +1,1727 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from packaging import version +from PIL import Image +from qwen_vl_utils import process_vision_info +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) +from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, + ImageVideoSampler, + get_random_mask) +from videox_fun.models import (AutoencoderKLQwenImage, + QwenImageInstantXControlNetModel, + Qwen2_5_VLForConditionalGeneration, + Qwen2Tokenizer, QwenImageTransformer2DModel) +from videox_fun.pipeline import QwenImageControlNetPipeline +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) + +if is_wandb_available(): + pass + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def _pack_latents(latents, batch_size, num_channels_latents, height, width, num_frame=None): + if num_frame is None: + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 2, 4, 1, 3, 5) + latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4) + else: + latents = latents.view(batch_size, num_channels_latents, num_frame, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 2, 3, 5, 1, 4, 6) + latents = latents.reshape(batch_size, num_frame * (height // 2) * (width // 2), num_channels_latents * 4) + return latents + +def _extract_masked_hidden(hidden_states: torch.Tensor, mask: torch.Tensor): + bool_mask = mask.bool() + valid_lengths = bool_mask.sum(dim=1) + selected = hidden_states[bool_mask] + split_result = torch.split(selected, valid_lengths.tolist(), dim=0) + + return split_result + +def calculate_shift( + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 4096, + base_shift: float = 0.5, + max_shift: float = 1.15, +): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, cn_transformer, args, accelerator, weight_dtype, global_step): + try: + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = QwenImageControlNetPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + controlnet=cn_transformer, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + control_image = Image.open(args.validation_paths[i]) + width, height = control_image.width, control_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + sample = pipeline( + args.validation_prompts[i], + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + true_cfg_scale = 4.0, + num_inference_steps = 20, + controlnet_conditioning_scale = 0.90, + control_image = control_image, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--cn_pretrained_model_name_or_path", + type=str, + default="InstantX/Qwen-Image-ControlNet-Union", + required=False, + help="Path to controlnet pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--token_sample_size", + type=int, + default=512, + help="Sample size of the token.", + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--video_sample_size", + type=int, + default=512, + help="Sample size of the video.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--controlnet_path", + type=str, + default=None, + help=("If you want to load the weight from other controlnet, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=1024, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--prompt_template_encode", + type=str, + default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n", + help=( + 'The prompt template for text encoder.' + ), + ) + parser.add_argument( + "--prompt_template_encode_start_idx", + type=int, + default=34, + help=( + 'The start idx for prompt template.' + ), + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=3.5, + help="the FLUX.1 dev variant is a guidance distilled model", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + + # Get Tokenizer + tokenizer = Qwen2Tokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer" + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype + ) + text_encoder = text_encoder.eval() + + # Get Vae + vae = AutoencoderKLQwenImage.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae" + ).to(weight_dtype) + vae.eval() + latents_mean = (torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1)).to(accelerator.device) + latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to(accelerator.device) + + # Get Transformer + transformer3d = QwenImageTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + low_cpu_mem_usage=True, + ).to(weight_dtype) + cn_transformer = QwenImageInstantXControlNetModel.from_pretrained(args.cn_pretrained_model_name_or_path, torch_dtype=weight_dtype) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.controlnet_path is not None: + print(f"From checkpoint: {args.controlnet_path}") + if args.controlnet_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.controlnet_path) + else: + state_dict = torch.load(args.controlnet_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = cn_transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + cn_transformer.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in cn_transformer.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = QwenImageTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + cn_transformer.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, cn_transformer.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in cn_transformer.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + if args.fix_sample_size is not None and args.enable_bucket: + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoControlDataset( + args.train_data_meta, args.train_data_dir, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, + enable_inpaint=False, + enable_camera_info=False, + enable_subject_info=True, + video_sample_n_frames=81, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + return [first_element] + [other_elements_value] * (length - 1) + + MIN_TARGET = 1024 + + if sample_size < MIN_TARGET: + number_list = [1.0] + else: + max_allowed_ratio = sample_size / MIN_TARGET + base_ratios = [ + 1.0, + 1.1, 1.2, 1.25, 1.33, 1.5, + 1.75, 2.0, 2.25, 2.5, 2.75, + 3.0, 3.5, 4.0, 5.0, 6.0, 8.0 + ] + candidate_ratios = set(base_ratios + list(image_ratio)) + number_list = sorted([r for r in candidate_ratios if 1.0 <= r <= max_allowed_ratio]) + + if not number_list: + number_list = [1.0] + + if all_choices: + return number_list + + probs = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p=probs) + else: + return rng.choice(number_list, p=probs) + + # Create new output + new_examples = {} + new_examples["pixel_values"] = [] + new_examples["text"] = [] + + # Used in Control Mode + new_examples["control_pixel_values"] = [] + # Used in Control Ref Mode + new_examples["ref_pixel_values"] = [] + new_examples["clip_pixel_values"] = [] + new_examples["clip_idx"] = [] + + # Used in Inpaint mode + new_examples["mask_pixel_values"] = [] + new_examples["mask"] = [] + + new_examples["subject_images"] = [] + new_examples["subject_flags"] = [] + + # Get downsample ratio in image + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.001 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + number_list_prob = np.array(_create_special_list(len(pixel_values))) + clip_index = np.random.choice(list(range(len(pixel_values))), p = number_list_prob) + new_examples["clip_idx"].append(clip_index) + + ref_pixel_values = pixel_values[clip_index].permute(1, 2, 0).contiguous() + ref_pixel_values = Image.fromarray(np.uint8(ref_pixel_values * 255)) + ref_pixel_values = (torch.tensor(np.array(ref_pixel_values)).unsqueeze(0).permute(0, 3, 1, 2).contiguous() / 255 - 0.5) / 0.5 + new_examples["ref_pixel_values"].append(ref_pixel_values) + + clip_pixel_values = pixel_values[clip_index].permute(1, 2, 0).contiguous() + clip_pixel_values = clip_pixel_values * 255 + new_examples["clip_pixel_values"].append(clip_pixel_values) + + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + + _, channel, h, w = pixel_values.size() + new_subject_image = torch.zeros(4, channel, h, w) + num_subject = len(example["subject_image"]) + if num_subject != 0: + subject_image = torch.from_numpy(example["subject_image"]).permute(0, 3, 1, 2).contiguous() + new_subject_image[:num_subject] = subject_image + subject_image = new_subject_image / 255. + subject_flag = torch.from_numpy(np.array([1] * num_subject + [0] * (4 - num_subject))) + + if args.fix_sample_size is not None: + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + length = int(len(pixel_values) // 2) + new_examples["pixel_values"].append(transform(pixel_values)[length:length + 1]) + new_examples["control_pixel_values"].append(transform(control_pixel_values)[length:length + 1]) + + new_examples["text"].append(example["text"]) + + mask = get_random_mask(new_examples["pixel_values"][-1].size()) + mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) + + # Wan 2.1 use 0 for masked pixels + # + torch.ones_like(new_examples["pixel_values"][-1]) * -1 * mask + new_examples["mask_pixel_values"].append(mask_pixel_values[:1]) + new_examples["mask"].append(mask[:1]) + + new_examples["subject_images"].append(transform(subject_image)) + new_examples["subject_flags"].append(subject_flag) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) + new_examples["control_pixel_values"] = torch.stack([example for example in new_examples["control_pixel_values"]]) + new_examples["ref_pixel_values"] = torch.stack([example for example in new_examples["ref_pixel_values"]]) + new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) + new_examples["clip_idx"] = torch.tensor(new_examples["clip_idx"]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) + new_examples["subject_images"] = torch.stack([example for example in new_examples["subject_images"]]) + new_examples["subject_flags"] = torch.stack([example for example in new_examples["subject_flags"]]) + + # Encode prompts when enable_text_encoder_in_dataloader=True + if args.enable_text_encoder_in_dataloader: + template = args.prompt_template_encode + drop_idx = args.prompt_template_encode_start_idx + + txt = [template.format(e) for e in batch['text']] + txt_tokens = tokenizer( + txt, max_length=args.tokenizer_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt" + ).to(accelerator.device) + encoder_hidden_states = text_encoder( + input_ids=txt_tokens.input_ids, + attention_mask=txt_tokens.attention_mask, + output_hidden_states=True, + ) + hidden_states = encoder_hidden_states.hidden_states[-1] + split_hidden_states = _extract_masked_hidden(hidden_states, txt_tokens.attention_mask) + split_hidden_states = [e[drop_idx:] for e in split_hidden_states] + attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] + max_seq_len = max([e.size(0) for e in split_hidden_states]) + prompt_embeds = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] + ) + encoder_attention_mask = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] + ) + + prompt_embeds = prompt_embeds.to(dtype=latents.dtype, device=accelerator.device) + + new_examples['encoder_attention_mask'] = encoder_attention_mask + new_examples['encoder_hidden_states'] = prompt_embeds + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + cn_transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + cn_transformer, optimizer, train_dataloader, lr_scheduler + ) + + if fsdp_stage != 0: + from functools import partial + + from videox_fun.dist import shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) + text_encoder = shard_fn(text_encoder) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream: + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step < 1: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + control_pixel_values = batch["control_pixel_values"].cpu() + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + control_pixel_values = rearrange(control_pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, control_pixel_value, text) in enumerate(zip(pixel_values, control_pixel_values, texts)): + pixel_value = pixel_value[None, ...] + control_pixel_value = control_pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + save_videos_grid(control_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_control.gif", rescale=True) + + with accelerator.accumulate(cn_transformer): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + control_pixel_values = batch["control_pixel_values"].to(weight_dtype) + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) + else: + latents = _batch_encode_vae(pixel_values) + latents = ((latents - latents_mean) * latents_std).to(dtype=weight_dtype) + + control_latents = _batch_encode_vae(control_pixel_values) + control_latents = ((control_latents - latents_mean) * latents_std).to(dtype=weight_dtype) + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + if args.enable_text_encoder_in_dataloader: + prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device) + encoder_attention_mask = batch['encoder_attention_mask'] + else: + with torch.no_grad(): + template = args.prompt_template_encode + drop_idx = args.prompt_template_encode_start_idx + + txt = [template.format(e) for e in batch['text']] + txt_tokens = tokenizer( + txt, max_length=args.tokenizer_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt" + ).to(accelerator.device) + encoder_hidden_states = text_encoder( + input_ids=txt_tokens.input_ids, + attention_mask=txt_tokens.attention_mask, + output_hidden_states=True, + ) + hidden_states = encoder_hidden_states.hidden_states[-1] + split_hidden_states = _extract_masked_hidden(hidden_states, txt_tokens.attention_mask) + split_hidden_states = [e[drop_idx:] for e in split_hidden_states] + attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] + max_seq_len = max([e.size(0) for e in split_hidden_states]) + prompt_embeds = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] + ) + encoder_attention_mask = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] + ) + + prompt_embeds = prompt_embeds.to(dtype=latents.dtype, device=accelerator.device) + + if args.low_vram and not args.enable_text_encoder_in_dataloader: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, num_frame, height, width = latents.size() + latents = _pack_latents(latents, bsz, channel, height, width, num_frame=num_frame) + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + control_latents = _pack_latents(control_latents, bsz, control_latents.size(1), height, width, num_frame=num_frame) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + + sigmas = np.linspace(1.0, 1 / args.train_sampling_steps, args.train_sampling_steps) + image_seq_len = latents.shape[1] + mu = calculate_shift( + image_seq_len, + noise_scheduler.config.get("base_image_seq_len", 256), + noise_scheduler.config.get("max_image_seq_len", 4096), + noise_scheduler.config.get("base_shift", 0.5), + noise_scheduler.config.get("max_shift", 1.15), + ) + noise_scheduler.set_timesteps(sigmas=sigmas, device=latents.device, mu=mu) + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + img_shapes = [[(num_frame, height // 2, width // 2)]] * latents.size(0) + txt_seq_lens = encoder_attention_mask.sum(dim=1).tolist() if encoder_attention_mask is not None else None + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + controlnet_block_samples = cn_transformer( + hidden_states=noisy_latents, + timestep=timesteps / 1000, + controlnet_cond=control_latents, + encoder_hidden_states_mask=encoder_attention_mask, + encoder_hidden_states=prompt_embeds, + img_shapes=img_shapes, + txt_seq_lens=txt_seq_lens, + conditioning_scale=1, + return_dict=False, + ) + noise_pred = transformer3d( + hidden_states=noisy_latents, + timestep=timesteps / 1000, + encoder_hidden_states_mask=encoder_attention_mask, + encoder_hidden_states=prompt_embeds, + img_shapes=img_shapes, + txt_seq_lens=txt_seq_lens, + controlnet_block_samples=controlnet_block_samples, + return_dict=False, + ) + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + cn_transformer, + args, + accelerator, + weight_dtype, + global_step, + ) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + tokenizer_2, + transformer3d, + cn_transformer, + args, + accelerator, + weight_dtype, + global_step, + ) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/qwenimage_instantx/train_control.sh b/scripts/qwenimage_instantx/train_control.sh new file mode 100644 index 0000000..6c53203 --- /dev/null +++ b/scripts/qwenimage_instantx/train_control.sh @@ -0,0 +1,36 @@ +# This is an InstantX ControlNet architecture. +# Note that it differs from the Fun Control architecture. +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2512" +export CN_MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-ControlNet-Union" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/qwenimage_instantx/train_control.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --cn_pretrained_model_name_or_path=$CN_MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwen_image_instantx_control" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." \ No newline at end of file diff --git a/scripts/wan2.2_vace_fun/train.py b/scripts/wan2.2_vace_fun/train.py index 418e869..aeed55e 100644 --- a/scripts/wan2.2_vace_fun/train.py +++ b/scripts/wan2.2_vace_fun/train.py @@ -1477,10 +1477,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/z_image/train.py b/scripts/z_image/train.py index 7aeeb28..9f89dce 100644 --- a/scripts/z_image/train.py +++ b/scripts/z_image/train.py @@ -189,56 +189,67 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = ZImagePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = ZImageTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = ZImagePipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + guidance_scale = 0, + num_inference_steps = 8, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1551,25 +1562,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1577,25 +1587,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() @@ -1611,4 +1620,4 @@ def main(): if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/scripts/z_image/train_lora.py b/scripts/z_image/train_lora.py index 2242878..d74e509 100644 --- a/scripts/z_image/train_lora.py +++ b/scripts/z_image/train_lora.py @@ -192,59 +192,69 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = ZImagePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = ZImageTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = ZImagePipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + guidance_scale = 0, + num_inference_steps = 8, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1537,19 +1547,18 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1557,19 +1566,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/z_image_fun/train_control.py b/scripts/z_image_fun/train_control.py index cea7c6f..2bd7048 100644 --- a/scripts/z_image_fun/train_control.py +++ b/scripts/z_image_fun/train_control.py @@ -84,7 +84,9 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, ZImageControlTransformer2DModel) from videox_fun.pipeline import ZImageControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -191,56 +193,73 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = ZImageControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = ZImageControlTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = ZImageControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): + control_image = Image.open(args.validation_paths[i]) + width, height = control_image.width, control_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + control_image = get_image_latent(control_image, sample_size=(height, width))[:, :, 0] + sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", - height = args.image_sample_size, - width = args.image_sample_size, - generator = generator + height = height, + width = width, + generator = generator, + guidance_scale = 0, + num_inference_steps = 8, + control_image = control_image, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -299,6 +318,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1664,25 +1690,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1690,25 +1715,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() @@ -1724,4 +1748,4 @@ def main(): if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/scripts/z_image_fun/train_control_distill.py b/scripts/z_image_fun/train_control_distill.py index 8aac3d4..6998fe2 100644 --- a/scripts/z_image_fun/train_control_distill.py +++ b/scripts/z_image_fun/train_control_distill.py @@ -90,7 +90,9 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, ZImageControlTransformer2DModel) from videox_fun.pipeline import ZImageControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -197,56 +199,73 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = ZImageControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = ZImageControlTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = ZImageControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): + control_image = Image.open(args.validation_paths[i]) + width, height = control_image.width, control_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + control_image = get_image_latent(control_image, sample_size=(height, width))[:, :, 0] + sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", - height = args.image_sample_size, - width = args.image_sample_size, - generator = generator + height = height, + width = width, + generator = generator, + guidance_scale = 0, + num_inference_steps = 8, + control_image = control_image, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -305,6 +324,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -958,19 +984,6 @@ def main(): param.requires_grad = True break - # Create EMA for the transformer3d. - if args.use_ema: - if zero_stage == 3: - raise NotImplementedError("FSDP does not support EMA.") - - ema_transformer3d = ZImageControlTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="transformer", - torch_dtype=weight_dtype, - ).to(weight_dtype) - - ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=ZImageControlTransformer2DModel, model_config=ema_transformer3d.config) - # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format @@ -1998,25 +2011,17 @@ def main(): accelerator_fake_score_transformer3d.save_state(fake_score_save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + generator_transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"denoising_loss": denoising_loss.detach().item(), "dmd_loss": dmd_loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2024,25 +2029,17 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + generator_transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index b2d6a79..2be1234 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -34,6 +34,7 @@ from .longcatvideo_transformer3d import LongCatVideoTransformer3DModel from .longcatvideo_vae import AutoencoderKLLongCatVideo from .qwenimage_transformer2d import QwenImageTransformer2DModel from .qwenimage_transformer2d_control import QwenImageControlTransformer2DModel +from .qwenimage_transformer2d_instantx import QwenImageInstantXControlNetModel from .qwenimage_vae import AutoencoderKLQwenImage from .wan_audio_encoder import WanAudioEncoder from .wan_image_encoder import CLIPModel diff --git a/videox_fun/models/attention_utils.py b/videox_fun/models/attention_utils.py index 312a40d..0f2d10a 100644 --- a/videox_fun/models/attention_utils.py +++ b/videox_fun/models/attention_utils.py @@ -41,6 +41,26 @@ except: SAGE_ATTENTION_AVAILABLE = False +def convert_qkv_dtype(q, k, v): + try: + """Unify the dtype of q, k, v tensors""" + dtypes = {q.dtype, k.dtype, v.dtype} + + # If any tensor is float16/bfloat16 + if torch.float16 in dtypes or torch.bfloat16 in dtypes: + target_dtype = torch.bfloat16 if torch.bfloat16 in dtypes else torch.float16 + # If all tensors are float32 + elif dtypes == {torch.float32}: + target_dtype = torch.bfloat16 if (torch.cuda.is_available() and + torch.cuda.get_device_capability()[0] >= 8) else torch.float16 + else: + return q, k, v # No conversion for other cases + + return q.to(target_dtype), k.to(target_dtype), v.to(target_dtype) + except: + return q, k, v + + def flash_attention_naive( q, k, @@ -214,6 +234,7 @@ def attention( 'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.' ) + q, k, v = convert_qkv_dtype(q, k, v) out = sageattn( q, k, v, attn_mask=attn_mask, tensor_layout="NHD", is_causal=causal, dropout_p=dropout_p) diff --git a/videox_fun/models/flux2_transformer2d_control.py b/videox_fun/models/flux2_transformer2d_control.py index 2b35684..100abbb 100644 --- a/videox_fun/models/flux2_transformer2d_control.py +++ b/videox_fun/models/flux2_transformer2d_control.py @@ -237,6 +237,21 @@ class Flux2ControlTransformer2DModel(Flux2Transformer2DModel): torch.cat([text_rotary_emb[1], image_rotary_emb[1]], dim=0), ) + # Context Parallel + if self.sp_world_size > 1: + hidden_states = torch.chunk(hidden_states, self.sp_world_size, dim=1)[self.sp_world_rank] + if concat_rotary_emb is not None: + txt_rotary_emb = ( + concat_rotary_emb[0][:encoder_hidden_states.shape[1]], + concat_rotary_emb[1][:encoder_hidden_states.shape[1]] + ) + concat_rotary_emb = ( + torch.chunk(concat_rotary_emb[0][encoder_hidden_states.shape[1]:], self.sp_world_size, dim=0)[self.sp_world_rank], + torch.chunk(concat_rotary_emb[1][encoder_hidden_states.shape[1]:], self.sp_world_size, dim=0)[self.sp_world_rank], + ) + concat_rotary_emb = [torch.cat([_txt_rotary_emb, _image_rotary_emb], dim=0) \ + for _txt_rotary_emb, _image_rotary_emb in zip(txt_rotary_emb, concat_rotary_emb)] + # Arguments kwargs = dict( encoder_hidden_states=encoder_hidden_states, @@ -306,6 +321,9 @@ class Flux2ControlTransformer2DModel(Flux2Transformer2DModel): hidden_states = self.norm_out(hidden_states, temb) output = self.proj_out(hidden_states) + if self.sp_world_size > 1: + output = self.all_gather(output, dim=1) + if not return_dict: return (output,) diff --git a/videox_fun/models/qwenimage_transformer2d.py b/videox_fun/models/qwenimage_transformer2d.py index c2efb23..abbcafa 100644 --- a/videox_fun/models/qwenimage_transformer2d.py +++ b/videox_fun/models/qwenimage_transformer2d.py @@ -34,17 +34,11 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin from diffusers.loaders.single_file_model import FromOriginalModelMixin from diffusers.models.attention import Attention, FeedForward -from diffusers.models.attention_processor import ( - Attention, AttentionProcessor, CogVideoXAttnProcessor2_0, - FusedCogVideoXAttnProcessor2_0) -from diffusers.models.embeddings import (CogVideoXPatchEmbed, - TimestepEmbedding, Timesteps, - get_3d_sincos_pos_embed) +from diffusers.models.attention_processor import Attention, AttentionProcessor +from diffusers.models.embeddings import TimestepEmbedding, Timesteps from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_utils import ModelMixin -from diffusers.models.normalization import (AdaLayerNorm, - AdaLayerNormContinuous, - CogVideoXLayerNormZero, RMSNorm) +from diffusers.models.normalization import AdaLayerNormContinuous, RMSNorm from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers) from diffusers.utils.torch_utils import maybe_allow_in_graph @@ -859,60 +853,6 @@ class QwenImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fro self.teacache = None @cfg_skip() - def forward_bs(self, x, *args, **kwargs): - func = self.forward - sig = inspect.signature(func) - - bs = len(x) - bs_half = int(bs // 2) - - if bs >= 2: - # cond - x_i = x[bs_half:] - args_i = [ - arg[bs_half:] if - isinstance(arg, - (torch.Tensor, list, tuple, np.ndarray)) and - len(arg) == bs else arg for arg in args - ] - kwargs_i = { - k: (v[bs_half:] if - isinstance(v, - (torch.Tensor, list, tuple, - np.ndarray)) and len(v) == bs else v - ) for k, v in kwargs.items() - } - if 'cond_flag' in sig.parameters: - kwargs_i["cond_flag"] = True - - cond_out = func(x_i, *args_i, **kwargs_i) - - # uncond - uncond_x_i = x[:bs_half] - uncond_args_i = [ - arg[:bs_half] if - isinstance(arg, - (torch.Tensor, list, tuple, np.ndarray)) and - len(arg) == bs else arg for arg in args - ] - uncond_kwargs_i = { - k: (v[:bs_half] if - isinstance(v, - (torch.Tensor, list, tuple, - np.ndarray)) and len(v) == bs else v - ) for k, v in kwargs.items() - } - if 'cond_flag' in sig.parameters: - uncond_kwargs_i["cond_flag"] = False - uncond_out = func(uncond_x_i, *uncond_args_i, - **uncond_kwargs_i) - - x = torch.cat([uncond_out, cond_out], dim=0) - else: - x = func(x, *args, **kwargs) - - return x - def forward( self, hidden_states: torch.Tensor, @@ -923,6 +863,7 @@ class QwenImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fro txt_seq_lens: Optional[List[int]] = None, guidance: torch.Tensor = None, # TODO: this should probably be removed attention_kwargs: Optional[Dict[str, Any]] = None, + controlnet_block_samples=None, additional_t_cond=None, cond_flag: bool = True, return_dict: bool = True, @@ -1060,7 +1001,7 @@ class QwenImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fro ori_hidden_states = hidden_states.clone().cpu() if self.teacache.offload else hidden_states.clone() # 4. Transformer blocks - for i, block in enumerate(self.transformer_blocks): + for index_block, block in enumerate(self.transformer_blocks): if torch.is_grad_enabled() and self.gradient_checkpointing: def create_custom_forward(module): def custom_forward(*inputs): @@ -1091,11 +1032,17 @@ class QwenImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fro modulate_index=modulate_index, ) + if controlnet_block_samples is not None: + interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) + interval_control = int(np.ceil(interval_control)) + hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] + if cond_flag: self.teacache.previous_residual_cond = hidden_states.cpu() - ori_hidden_states if self.teacache.offload else hidden_states - ori_hidden_states else: self.teacache.previous_residual_uncond = hidden_states.cpu() - ori_hidden_states if self.teacache.offload else hidden_states - ori_hidden_states del ori_hidden_states + else: for index_block, block in enumerate(self.transformer_blocks): if torch.is_grad_enabled() and self.gradient_checkpointing: @@ -1128,6 +1075,12 @@ class QwenImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fro modulate_index=modulate_index, ) + # controlnet residual + if controlnet_block_samples is not None: + interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) + interval_control = int(np.ceil(interval_control)) + hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] + if self.zero_cond_t: temb = temb.chunk(2, dim=0)[0] # Use only the image part (hidden_states) from the dual-stream blocks diff --git a/videox_fun/models/qwenimage_transformer2d_control.py b/videox_fun/models/qwenimage_transformer2d_control.py index 6b3d293..69d80df 100644 --- a/videox_fun/models/qwenimage_transformer2d_control.py +++ b/videox_fun/models/qwenimage_transformer2d_control.py @@ -2,18 +2,22 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from math import prod from typing import Any, Dict, List, Optional, Tuple import torch import torch.nn as nn from diffusers.configuration_utils import register_to_config from diffusers.models.modeling_outputs import Transformer2DModelOutput -from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, +from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers) +from ..utils import cfg_skip from .qwenimage_transformer2d import (QwenImageTransformer2DModel, QwenImageTransformerBlock) +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + class QwenImageControlTransformerBlock(QwenImageTransformerBlock): def __init__( @@ -161,6 +165,7 @@ class QwenImageControlTransformer2DModel(QwenImageTransformer2DModel): hints = torch.unbind(c)[:-1] return hints + @cfg_skip() def forward( self, hidden_states: torch.Tensor, @@ -172,6 +177,7 @@ class QwenImageControlTransformer2DModel(QwenImageTransformer2DModel): guidance: torch.Tensor = None, # TODO: this should probably be removed attention_kwargs: Optional[Dict[str, Any]] = None, additional_t_cond=None, + cond_flag: bool=True, control_context=None, control_context_scale=1.0, return_dict: bool = True, @@ -231,20 +237,105 @@ class QwenImageControlTransformer2DModel(QwenImageTransformer2DModel): image_rotary_emb[1] ) - # Arguments - kwargs = dict( - encoder_hidden_states=encoder_hidden_states, - encoder_hidden_states_mask=encoder_hidden_states_mask, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=attention_kwargs, - modulate_index=modulate_index, - ) - hints = self.forward_control( - hidden_states, control_context, kwargs - ) + # TeaCache + if self.teacache is not None: + if cond_flag: + inp = hidden_states.clone() + temb_ = temb.clone() + encoder_hidden_states_ = encoder_hidden_states.clone() - for index_block, block in enumerate(self.transformer_blocks): + img_mod_params_ = self.transformer_blocks[0].img_mod(temb_) + img_mod1_, img_mod2_ = img_mod_params_.chunk(2, dim=-1) + img_normed_ = self.transformer_blocks[0].img_norm1(inp) + modulated_inp, img_gate1_ = self.transformer_blocks[0]._modulate(img_normed_, img_mod1_) + + skip_flag = self.teacache.cnt < self.teacache.num_skip_start_steps + if skip_flag: + self.should_calc = True + self.teacache.accumulated_rel_l1_distance = 0 + else: + if cond_flag: + rel_l1_distance = self.teacache.compute_rel_l1_distance(self.teacache.previous_modulated_input, modulated_inp) + self.teacache.accumulated_rel_l1_distance += self.teacache.rescale_func(rel_l1_distance) + + if torch.distributed.is_initialized(): + if not isinstance(self.teacache.accumulated_rel_l1_distance, torch.Tensor): + accumulated_distance_tensor = torch.tensor( + self.teacache.accumulated_rel_l1_distance, + device=hidden_states.device, + dtype=torch.float32 + ) + else: + accumulated_distance_tensor = self.teacache.accumulated_rel_l1_distance.clone() + + torch.distributed.broadcast(accumulated_distance_tensor, src=0) + self.teacache.accumulated_rel_l1_distance = accumulated_distance_tensor.item() + + if self.teacache.accumulated_rel_l1_distance < self.teacache.rel_l1_thresh: + self.should_calc = False + else: + self.should_calc = True + self.teacache.accumulated_rel_l1_distance = 0 + self.teacache.previous_modulated_input = modulated_inp + self.teacache.should_calc = self.should_calc + else: + self.should_calc = self.teacache.should_calc + + # TeaCache + if self.teacache is not None: + if not self.should_calc: + previous_residual = self.teacache.previous_residual_cond if cond_flag else self.teacache.previous_residual_uncond + hidden_states = hidden_states + previous_residual.to(hidden_states.device)[-hidden_states.size()[0]:,] + else: + ori_hidden_states = hidden_states.clone().cpu() if self.teacache.offload else hidden_states.clone() + + # Arguments + kwargs = dict( + encoder_hidden_states=encoder_hidden_states, + encoder_hidden_states_mask=encoder_hidden_states_mask, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=attention_kwargs, + modulate_index=modulate_index, + ) + hints = self.forward_control( + hidden_states, control_context, kwargs + ) + # 4. Transformer blocks + for index_block, block in enumerate(self.transformer_blocks): + # Arguments + kwargs = dict( + encoder_hidden_states=encoder_hidden_states, + encoder_hidden_states_mask=encoder_hidden_states_mask, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=attention_kwargs, + modulate_index=modulate_index, + hints=hints, + context_scale=control_context_scale + ) + if torch.is_grad_enabled() and self.gradient_checkpointing: + def create_custom_forward(module, **static_kwargs): + def custom_forward(*inputs): + return module(*inputs, **static_kwargs) + return custom_forward + + ckpt_kwargs = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block, **kwargs), + hidden_states, + **ckpt_kwargs, + ) + else: + encoder_hidden_states, hidden_states = block(hidden_states, **kwargs) + if cond_flag: + self.teacache.previous_residual_cond = hidden_states.cpu() - ori_hidden_states if self.teacache.offload else hidden_states - ori_hidden_states + else: + self.teacache.previous_residual_uncond = hidden_states.cpu() - ori_hidden_states if self.teacache.offload else hidden_states - ori_hidden_states + del ori_hidden_states + + else: # Arguments kwargs = dict( encoder_hidden_states=encoder_hidden_states, @@ -253,24 +344,37 @@ class QwenImageControlTransformer2DModel(QwenImageTransformer2DModel): image_rotary_emb=image_rotary_emb, joint_attention_kwargs=attention_kwargs, modulate_index=modulate_index, - hints=hints, - context_scale=control_context_scale ) - if torch.is_grad_enabled() and self.gradient_checkpointing: - def create_custom_forward(module, **static_kwargs): - def custom_forward(*inputs): - return module(*inputs, **static_kwargs) - return custom_forward - - ckpt_kwargs = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} - - encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( - create_custom_forward(block, **kwargs), - hidden_states, - **ckpt_kwargs, + hints = self.forward_control( + hidden_states, control_context, kwargs + ) + for index_block, block in enumerate(self.transformer_blocks): + # Arguments + kwargs = dict( + encoder_hidden_states=encoder_hidden_states, + encoder_hidden_states_mask=encoder_hidden_states_mask, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=attention_kwargs, + modulate_index=modulate_index, + hints=hints, + context_scale=control_context_scale ) - else: - encoder_hidden_states, hidden_states = block(hidden_states, **kwargs) + if torch.is_grad_enabled() and self.gradient_checkpointing: + def create_custom_forward(module, **static_kwargs): + def custom_forward(*inputs): + return module(*inputs, **static_kwargs) + return custom_forward + + ckpt_kwargs = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block, **kwargs), + hidden_states, + **ckpt_kwargs, + ) + else: + encoder_hidden_states, hidden_states = block(hidden_states, **kwargs) if self.zero_cond_t: temb = temb.chunk(2, dim=0)[0] @@ -285,4 +389,8 @@ class QwenImageControlTransformer2DModel(QwenImageTransformer2DModel): # remove `lora_scale` from each PEFT layer unscale_lora_layers(self, lora_scale) + if self.teacache is not None and cond_flag: + self.teacache.cnt += 1 + if self.teacache.cnt == self.teacache.num_steps: + self.teacache.reset() return output \ No newline at end of file diff --git a/videox_fun/models/qwenimage_transformer2d_instantx.py b/videox_fun/models/qwenimage_transformer2d_instantx.py new file mode 100644 index 0000000..cf7a5d7 --- /dev/null +++ b/videox_fun/models/qwenimage_transformer2d_instantx.py @@ -0,0 +1,243 @@ +# Modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/controlnets/controlnet_qwenimage.py +# Copyright 2025 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn as nn + +from .qwenimage_transformer2d import (USE_PEFT_BACKEND, ConfigMixin, + FromOriginalModelMixin, ModelMixin, + PeftAdapterMixin, QwenEmbedRope, + QwenImageTransformerBlock, + QwenTimestepProjEmbeddings, RMSNorm, + Transformer2DModelOutput, logging, + register_to_config, scale_lora_layers, + unscale_lora_layers) + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +def zero_module(module): + """Zero out the parameters of a module and return it.""" + for p in module.parameters(): + nn.init.zeros_(p) + return module + + +@dataclass +class QwenImageControlNetOutput(Transformer2DModelOutput): + controlnet_block_samples: Tuple[torch.Tensor] + + +class QwenImageInstantXControlNetModel( + ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin +): + _supports_gradient_checkpointing = True + + @register_to_config + def __init__( + self, + patch_size: int = 2, + in_channels: int = 64, + out_channels: Optional[int] = 16, + num_layers: int = 60, + attention_head_dim: int = 128, + num_attention_heads: int = 24, + joint_attention_dim: int = 3584, + axes_dims_rope: Tuple[int, int, int] = (16, 56, 56), + extra_condition_channels: int = 0, # for controlnet-inpainting + ): + super().__init__() + self.out_channels = out_channels or in_channels + self.inner_dim = num_attention_heads * attention_head_dim + + self.pos_embed = QwenEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True) + + self.time_text_embed = QwenTimestepProjEmbeddings(embedding_dim=self.inner_dim) + + self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6) + + self.img_in = nn.Linear(in_channels, self.inner_dim) + self.txt_in = nn.Linear(joint_attention_dim, self.inner_dim) + + self.transformer_blocks = nn.ModuleList( + [ + QwenImageTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + ) + for _ in range(num_layers) + ] + ) + + # controlnet_blocks + self.controlnet_blocks = nn.ModuleList([]) + for _ in range(len(self.transformer_blocks)): + self.controlnet_blocks.append(zero_module(nn.Linear(self.inner_dim, self.inner_dim))) + self.controlnet_x_embedder = zero_module( + torch.nn.Linear(in_channels + extra_condition_channels, self.inner_dim) + ) + + self.gradient_checkpointing = False + + @classmethod + def from_transformer( + cls, + transformer, + num_layers: int = 5, + attention_head_dim: int = 128, + num_attention_heads: int = 24, + load_weights_from_transformer=True, + extra_condition_channels: int = 0, + ): + config = dict(transformer.config) + config["num_layers"] = num_layers + config["attention_head_dim"] = attention_head_dim + config["num_attention_heads"] = num_attention_heads + config["extra_condition_channels"] = extra_condition_channels + + controlnet = cls.from_config(config) + + if load_weights_from_transformer: + controlnet.pos_embed.load_state_dict(transformer.pos_embed.state_dict()) + controlnet.time_text_embed.load_state_dict(transformer.time_text_embed.state_dict()) + controlnet.img_in.load_state_dict(transformer.img_in.state_dict()) + controlnet.txt_in.load_state_dict(transformer.txt_in.state_dict()) + controlnet.transformer_blocks.load_state_dict(transformer.transformer_blocks.state_dict(), strict=False) + controlnet.controlnet_x_embedder = zero_module(controlnet.controlnet_x_embedder) + + return controlnet + + def forward( + self, + hidden_states: torch.Tensor, + controlnet_cond: torch.Tensor, + conditioning_scale: float = 1.0, + encoder_hidden_states: torch.Tensor = None, + encoder_hidden_states_mask: torch.Tensor = None, + timestep: torch.LongTensor = None, + img_shapes: Optional[List[Tuple[int, int, int]]] = None, + txt_seq_lens: Optional[List[int]] = None, + joint_attention_kwargs: Optional[Dict[str, Any]] = None, + return_dict: bool = True, + ) -> Union[torch.FloatTensor, Transformer2DModelOutput]: + """ + The [`FluxTransformer2DModel`] forward method. + + Args: + hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): + Input `hidden_states`. + controlnet_cond (`torch.Tensor`): + The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. + conditioning_scale (`float`, defaults to `1.0`): + The scale factor for ControlNet outputs. + encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): + Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. + pooled_projections (`torch.FloatTensor` of shape `(batch_size, projection_dim)`): Embeddings projected + from the embeddings of input conditions. + timestep ( `torch.LongTensor`): + Used to indicate denoising step. + block_controlnet_hidden_states: (`list` of `torch.Tensor`): + A list of tensors that if specified are added to the residuals of transformer blocks. + joint_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain + tuple. + + Returns: + If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. + """ + if joint_attention_kwargs is not None: + joint_attention_kwargs = joint_attention_kwargs.copy() + lora_scale = joint_attention_kwargs.pop("scale", 1.0) + else: + lora_scale = 1.0 + + if USE_PEFT_BACKEND: + # weight the lora layers by setting `lora_scale` for each PEFT layer + scale_lora_layers(self, lora_scale) + else: + if joint_attention_kwargs is not None and joint_attention_kwargs.get("scale", None) is not None: + logger.warning( + "Passing `scale` via `joint_attention_kwargs` when not using the PEFT backend is ineffective." + ) + + if isinstance(encoder_hidden_states, list): + encoder_hidden_states = torch.stack(encoder_hidden_states) + encoder_hidden_states_mask = torch.stack(encoder_hidden_states_mask) + + hidden_states = self.img_in(hidden_states) + + # add + hidden_states = hidden_states + self.controlnet_x_embedder(controlnet_cond) + + temb = self.time_text_embed(timestep, hidden_states) + + image_rotary_emb = self.pos_embed(img_shapes, txt_seq_lens, device=hidden_states.device) + + timestep = timestep.to(hidden_states.dtype) + encoder_hidden_states = self.txt_norm(encoder_hidden_states) + encoder_hidden_states = self.txt_in(encoder_hidden_states) + + block_samples = () + for index_block, block in enumerate(self.transformer_blocks): + if torch.is_grad_enabled() and self.gradient_checkpointing: + encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( + block, + hidden_states, + encoder_hidden_states, + encoder_hidden_states_mask, + temb, + image_rotary_emb, + ) + + else: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + encoder_hidden_states_mask=encoder_hidden_states_mask, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=joint_attention_kwargs, + ) + block_samples = block_samples + (hidden_states,) + + # controlnet block + controlnet_block_samples = () + for block_sample, controlnet_block in zip(block_samples, self.controlnet_blocks): + block_sample = controlnet_block(block_sample) + controlnet_block_samples = controlnet_block_samples + (block_sample,) + + # scaling + controlnet_block_samples = [sample * conditioning_scale for sample in controlnet_block_samples] + controlnet_block_samples = None if len(controlnet_block_samples) == 0 else controlnet_block_samples + + if USE_PEFT_BACKEND: + # remove `lora_scale` from each PEFT layer + unscale_lora_layers(self, lora_scale) + + if not return_dict: + return controlnet_block_samples + + return QwenImageControlNetOutput( + controlnet_block_samples=controlnet_block_samples, + ) diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 36d906b..4d6b6ea 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -183,10 +183,10 @@ class WanRMSNorm(nn.Module): Args: x(Tensor): Shape [B, L, C] """ - return self._norm(x) * self.weight + return self._norm(x.float()).type_as(x) * self.weight def _norm(self, x): - return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps).to(x.dtype) + return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) class WanLayerNorm(nn.LayerNorm): @@ -199,7 +199,7 @@ class WanLayerNorm(nn.LayerNorm): Args: x(Tensor): Shape [B, L, C] """ - return super().forward(x) + return super().forward(x.float()).type_as(x) class WanSelfAttention(nn.Module): diff --git a/videox_fun/models/wan_vae.py b/videox_fun/models/wan_vae.py index 7815e7b..4d3fad1 100755 --- a/videox_fun/models/wan_vae.py +++ b/videox_fun/models/wan_vae.py @@ -733,6 +733,13 @@ class AutoencoderKLWanCompileQwenImage(ModelMixin, ConfigMixin, FromOriginalMode 4 ], dropout = 0.0, + num_res_blocks = 2, + temperal_downsample = [ + False, + True, + True + ], + z_dim = 16, latents_mean = [ -0.7571, -0.7089, @@ -769,13 +776,8 @@ class AutoencoderKLWanCompileQwenImage(ModelMixin, ConfigMixin, FromOriginalMode 2.8251, 1.916 ], - num_res_blocks = 2, - temperal_downsample = [ - False, - True, - True - ], - z_dim = 16 + temporal_compression_ratio=4, + spatial_compression_ratio=8 ): super().__init__() cfg = dict( @@ -797,6 +799,8 @@ class AutoencoderKLWanCompileQwenImage(ModelMixin, ConfigMixin, FromOriginalMode self.attn_scales = attn_scales self.temperal_downsample = temperal_downsample self.temperal_upsample = temperal_downsample[::-1] + self.temporal_compression_ratio = temporal_compression_ratio + self.spatial_compression_ratio = spatial_compression_ratio def _encode(self, x: torch.Tensor) -> torch.Tensor: x = [ diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index 8563e70..94018f6 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -10,6 +10,7 @@ from .pipeline_hunyuanvideo_i2v import HunyuanVideoI2VPipeline from .pipeline_longcatvideo import LongCatVideoPipeline from .pipeline_qwenimage import QwenImagePipeline from .pipeline_qwenimage_control import QwenImageControlPipeline +from .pipeline_qwenimage_instantx import QwenImageControlNetPipeline from .pipeline_qwenimage_edit import QwenImageEditPipeline from .pipeline_qwenimage_edit_plus import QwenImageEditPlusPipeline from .pipeline_wan import WanPipeline diff --git a/videox_fun/pipeline/pipeline_flux2.py b/videox_fun/pipeline/pipeline_flux2.py index 26a8741..06488f7 100644 --- a/videox_fun/pipeline/pipeline_flux2.py +++ b/videox_fun/pipeline/pipeline_flux2.py @@ -637,6 +637,7 @@ class Flux2Pipeline(DiffusionPipeline): callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: Tuple[int] = (10, 20, 30), + comfyui_progressbar: bool = False, ): r""" Function invoked when calling the pipeline for generation. @@ -733,6 +734,9 @@ class Flux2Pipeline(DiffusionPipeline): batch_size = prompt_embeds.shape[0] device = self._execution_device + if comfyui_progressbar: + from comfy.utils import ProgressBar + pbar = ProgressBar(num_inference_steps + 2) # 3. prepare text embeddings prompt_embeds, text_ids = self.encode_prompt( @@ -783,6 +787,8 @@ class Flux2Pipeline(DiffusionPipeline): generator=generator, latents=latents, ) + if comfyui_progressbar: + pbar.update(1) image_latents = None image_latent_ids = None @@ -815,6 +821,9 @@ class Flux2Pipeline(DiffusionPipeline): guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32) guidance = guidance.expand(latents.shape[0]) + if comfyui_progressbar: + pbar.update(1) + # 7. Denoising loop # We set the index here to remove DtoH sync, helpful especially during compilation. # Check out more details here: https://github.com/huggingface/diffusers/pull/11696 @@ -873,12 +882,14 @@ class Flux2Pipeline(DiffusionPipeline): if XLA_AVAILABLE: xm.mark_step() + if comfyui_progressbar: + pbar.update(1) + self._current_timestep = None if output_type == "latent": image = latents else: - torch.save({"pred": latents}, "pred_d.pt") latents = self._unpack_latents_with_ids(latents, latent_ids) latents_bn_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(latents.device, latents.dtype) diff --git a/videox_fun/pipeline/pipeline_flux2_control.py b/videox_fun/pipeline/pipeline_flux2_control.py index f2c5aee..c3b1613 100644 --- a/videox_fun/pipeline/pipeline_flux2_control.py +++ b/videox_fun/pipeline/pipeline_flux2_control.py @@ -649,6 +649,7 @@ class Flux2ControlPipeline(DiffusionPipeline): callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, text_encoder_out_layers: Tuple[int] = (10, 20, 30), + comfyui_progressbar: bool = False, ): r""" Function invoked when calling the pipeline for generation. @@ -746,6 +747,9 @@ class Flux2ControlPipeline(DiffusionPipeline): device = self._execution_device weight_dtype = self.text_encoder.dtype + if comfyui_progressbar: + from comfy.utils import ProgressBar + pbar = ProgressBar(num_inference_steps + 2) latents_bn_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(device, weight_dtype) latents_bn_std = torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) + self.vae.config.batch_norm_eps).to( @@ -758,14 +762,19 @@ class Flux2ControlPipeline(DiffusionPipeline): # Prepare mask latent variables if mask_image is not None: mask_condition = self.mask_processor.preprocess(mask_image, height=height, width=width) + mask_condition = torch.where(mask_condition >= 0.5, + torch.ones_like(mask_condition), + torch.zeros_like(mask_condition)) mask_condition = torch.tile(mask_condition, [1, 3, 1, 1]).to(dtype=weight_dtype, device=device) + else: + mask_condition = torch.ones([batch_size, 3, height, width]).to(dtype=weight_dtype, device=device) if inpaint_image is not None: init_image = self.diffusers_image_processor.preprocess(inpaint_image, height=height, width=width) init_image = init_image.to(dtype=weight_dtype, device=device) * (mask_condition < 0.5) inpaint_latent = self.vae.encode(init_image)[0].mode() else: - inpaint_latent = torch.zeros((batch_size, num_channels_latents * 4, height // 2 // self.vae_scale_factor, width // 2 // self.vae_scale_factor)).to(device, weight_dtype) + inpaint_latent = torch.zeros((batch_size, num_channels_latents, height // self.vae_scale_factor, width // self.vae_scale_factor)).to(device, weight_dtype) if control_image is not None: control_image = self.diffusers_image_processor.preprocess(control_image, height=height, width=width) @@ -840,6 +849,8 @@ class Flux2ControlPipeline(DiffusionPipeline): generator=generator, latents=latents, ) + if comfyui_progressbar: + pbar.update(1) image_latents = None image_latent_ids = None @@ -872,6 +883,9 @@ class Flux2ControlPipeline(DiffusionPipeline): guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32) guidance = guidance.expand(latents.shape[0]) + if comfyui_progressbar: + pbar.update(1) + # 7. Denoising loop # We set the index here to remove DtoH sync, helpful especially during compilation. # Check out more details here: https://github.com/huggingface/diffusers/pull/11696 @@ -947,6 +961,9 @@ class Flux2ControlPipeline(DiffusionPipeline): if XLA_AVAILABLE: xm.mark_step() + if comfyui_progressbar: + pbar.update(1) + self._current_timestep = None if output_type == "latent": diff --git a/videox_fun/pipeline/pipeline_qwenimage.py b/videox_fun/pipeline/pipeline_qwenimage.py index e61ac0e..038a2ef 100644 --- a/videox_fun/pipeline/pipeline_qwenimage.py +++ b/videox_fun/pipeline/pipeline_qwenimage.py @@ -692,8 +692,8 @@ class QwenImagePipeline(DiffusionPipeline): timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype) with torch.cuda.amp.autocast(dtype=latents.dtype), torch.cuda.device(device=latents.device): - noise_pred = self.transformer.forward_bs( - x=latent_model_input, + noise_pred = self.transformer( + hidden_states=latent_model_input, timestep=timestep / 1000, guidance=guidance, encoder_hidden_states_mask=prompt_embeds_mask_input, diff --git a/videox_fun/pipeline/pipeline_qwenimage_control.py b/videox_fun/pipeline/pipeline_qwenimage_control.py index 4564966..3d58a15 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_control.py +++ b/videox_fun/pipeline/pipeline_qwenimage_control.py @@ -490,7 +490,8 @@ class QwenImageControlPipeline(DiffusionPipeline): callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 512, - control_context_scale: float = 1.0 + control_context_scale: float = 1.0, + comfyui_progressbar: bool = False, ): r""" Function invoked when calling the pipeline for generation. @@ -603,6 +604,9 @@ class QwenImageControlPipeline(DiffusionPipeline): device = self._execution_device weight_dtype = self.text_encoder.dtype + if comfyui_progressbar: + from comfy.utils import ProgressBar + pbar = ProgressBar(num_inference_steps + 2) has_neg_prompt = negative_prompt is not None or ( negative_prompt_embeds is not None and negative_prompt_embeds_mask is not None @@ -625,6 +629,9 @@ class QwenImageControlPipeline(DiffusionPipeline): num_images_per_prompt=num_images_per_prompt, max_sequence_length=max_sequence_length, ) + if comfyui_progressbar: + pbar.update(1) + # 4. Prepare latent variables num_channels_latents = self.transformer.config.in_channels // 4 latents = self.prepare_latents( @@ -648,8 +655,8 @@ class QwenImageControlPipeline(DiffusionPipeline): torch.zeros_like(mask_condition)) mask_condition = torch.tile(mask_condition, [1, 3, 1, 1]).to(dtype=weight_dtype, device=device) else: - mask_condition = torch.zeros([batch_size, 3, height, width]).to(dtype=weight_dtype, device=device) - + mask_condition = torch.ones([batch_size, 3, height, width]).to(dtype=weight_dtype, device=device) + if image is not None: init_image = self.image_processor.preprocess(image, height=height, width=width) init_image = init_image.to(dtype=weight_dtype, device=device) * (mask_condition < 0.5) @@ -709,6 +716,8 @@ class QwenImageControlPipeline(DiffusionPipeline): negative_txt_seq_lens = ( negative_prompt_embeds_mask.sum(dim=1).tolist() if negative_prompt_embeds_mask is not None else None ) + if comfyui_progressbar: + pbar.update(1) # 6. Denoising loop self.scheduler.set_begin_index(0) @@ -745,10 +754,10 @@ class QwenImageControlPipeline(DiffusionPipeline): self._current_timestep = t # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype) - print(latent_model_input.size(), control_context_input.size()) + with torch.cuda.amp.autocast(dtype=latents.dtype), torch.cuda.device(device=latents.device): - noise_pred = self.transformer.forward_bs( - x=latent_model_input, + noise_pred = self.transformer( + hidden_states=latent_model_input, timestep=timestep / 1000, guidance=guidance, encoder_hidden_states_mask=prompt_embeds_mask_input, @@ -794,6 +803,9 @@ class QwenImageControlPipeline(DiffusionPipeline): if XLA_AVAILABLE: xm.mark_step() + if comfyui_progressbar: + pbar.update(1) + self._current_timestep = None if output_type == "latent": image = latents diff --git a/videox_fun/pipeline/pipeline_qwenimage_edit.py b/videox_fun/pipeline/pipeline_qwenimage_edit.py index fc0a3d3..53e882e 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_edit.py +++ b/videox_fun/pipeline/pipeline_qwenimage_edit.py @@ -876,8 +876,8 @@ class QwenImageEditPipeline(DiffusionPipeline): timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype) with torch.cuda.amp.autocast(dtype=latents.dtype), torch.cuda.device(device=latents.device): - noise_pred = self.transformer.forward_bs( - x=latent_model_input, + noise_pred = self.transformer( + hidden_states=latent_model_input, timestep=timestep / 1000, guidance=guidance, encoder_hidden_states_mask=prompt_embeds_mask_input, diff --git a/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py b/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py index 885550c..ee36ce9 100644 --- a/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py +++ b/videox_fun/pipeline/pipeline_qwenimage_edit_plus.py @@ -861,8 +861,8 @@ class QwenImageEditPlusPipeline(DiffusionPipeline): timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype) with torch.cuda.amp.autocast(dtype=latents.dtype), torch.cuda.device(device=latents.device): - noise_pred = self.transformer.forward_bs( - x=latent_model_input, + noise_pred = self.transformer( + hidden_states=latent_model_input, timestep=timestep / 1000, guidance=guidance, encoder_hidden_states_mask=prompt_embeds_mask_input, diff --git a/videox_fun/pipeline/pipeline_qwenimage_instantx.py b/videox_fun/pipeline/pipeline_qwenimage_instantx.py new file mode 100644 index 0000000..d2e97a5 --- /dev/null +++ b/videox_fun/pipeline/pipeline_qwenimage_instantx.py @@ -0,0 +1,939 @@ +# Modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py +# Copyright 2025 Qwen-Image Team, InstantX Team and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import inspect +import math +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Optional, Tuple, Union + +import numpy as np +import PIL.Image +import torch +import torch.nn.functional as F +from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers.image_processor import PipelineImageInput, VaeImageProcessor +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from diffusers.utils import (BaseOutput, deprecate, is_torch_xla_available, + logging, replace_example_docstring) +from diffusers.utils.torch_utils import randn_tensor + +from ..models import (AutoencoderKLQwenImage, + Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer, + QwenImageInstantXControlNetModel, + QwenImageTransformer2DModel, T5Tokenizer) + +if is_torch_xla_available(): + import torch_xla.core.xla_model as xm + + XLA_AVAILABLE = True +else: + XLA_AVAILABLE = False + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + +EXAMPLE_DOC_STRING = """ + Examples: + ``` + ``` +""" + + +# Coped from diffusers.pipelines.qwenimage.pipeline_qwenimage.calculate_shift +def calculate_shift( + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 4096, + base_shift: float = 0.5, + max_shift: float = 1.15, +): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents +def retrieve_latents( + encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample" +): + if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": + return encoder_output.latent_dist.sample(generator) + elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": + return encoder_output.latent_dist.mode() + elif hasattr(encoder_output, "latents"): + return encoder_output.latents + else: + raise AttributeError("Could not access latents of provided encoder_output") + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + sigmas: Optional[List[float]] = None, + **kwargs, +): + r""" + Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles + custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. + + Args: + scheduler (`SchedulerMixin`): + The scheduler to get timesteps from. + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` + must be `None`. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + timesteps (`List[int]`, *optional*): + Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, + `num_inference_steps` and `sigmas` must be `None`. + sigmas (`List[float]`, *optional*): + Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, + `num_inference_steps` and `timesteps` must be `None`. + + Returns: + `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the + second element is the number of inference steps. + """ + if timesteps is not None and sigmas is not None: + raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + elif sigmas is not None: + accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accept_sigmas: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" sigmas schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +@dataclass +class QwenImagePipelineOutput(BaseOutput): + """ + Output class for Stable Diffusion pipelines. + + Args: + images (`List[PIL.Image.Image]` or `np.ndarray`) + List of denoised PIL images of length `batch_size` or numpy array of shape `(batch_size, height, width, + num_channels)`. PIL images or numpy array present the denoised images of the diffusion pipeline. + """ + + images: Union[List[PIL.Image.Image], np.ndarray] + + +class QwenImageControlNetPipeline(DiffusionPipeline): + r""" + The QwenImage pipeline for text-to-image generation. + + Args: + transformer ([`QwenImageTransformer2DModel`]): + Conditional Transformer (MMDiT) architecture to denoise the encoded image latents. + scheduler ([`FlowMatchEulerDiscreteScheduler`]): + A scheduler to be used in combination with `transformer` to denoise the encoded image latents. + vae ([`AutoencoderKL`]): + Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations. + text_encoder ([`Qwen2.5-VL-7B-Instruct`]): + [Qwen2.5-VL-7B-Instruct](https://huggingface.co/Qwen/Qwen2.5-VL-7B-Instruct), specifically the + [Qwen2.5-VL-7B-Instruct](https://huggingface.co/Qwen/Qwen2.5-VL-7B-Instruct) variant. + tokenizer (`QwenTokenizer`): + Tokenizer of class + [CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer). + """ + + model_cpu_offload_seq = "text_encoder->transformer->vae" + _callback_tensor_inputs = ["latents", "prompt_embeds"] + + def __init__( + self, + scheduler: FlowMatchEulerDiscreteScheduler, + vae: AutoencoderKLQwenImage, + text_encoder: Qwen2_5_VLForConditionalGeneration, + tokenizer: Qwen2Tokenizer, + transformer: QwenImageTransformer2DModel, + controlnet: QwenImageInstantXControlNetModel, + ): + super().__init__() + + self.register_modules( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer, + scheduler=scheduler, + controlnet=controlnet, + ) + self.vae_scale_factor = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8 + # QwenImage latents are turned into 2x2 patches and packed. This means the latent width and height has to be divisible + # by the patch size. So the vae scale factor is multiplied by the patch size to account for this + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor * 2) + self.tokenizer_max_length = 1024 + self.prompt_template_encode = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + self.prompt_template_encode_start_idx = 34 + self.default_sample_size = 128 + + # Coped from diffusers.pipelines.qwenimage.pipeline_qwenimage.extract_masked_hidden + def _extract_masked_hidden(self, hidden_states: torch.Tensor, mask: torch.Tensor): + bool_mask = mask.bool() + valid_lengths = bool_mask.sum(dim=1) + selected = hidden_states[bool_mask] + split_result = torch.split(selected, valid_lengths.tolist(), dim=0) + + return split_result + + # Coped from diffusers.pipelines.qwenimage.pipeline_qwenimage.get_qwen_prompt_embeds + def _get_qwen_prompt_embeds( + self, + prompt: Union[str, List[str]] = None, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ): + device = device or self._execution_device + dtype = dtype or self.text_encoder.dtype + + prompt = [prompt] if isinstance(prompt, str) else prompt + + template = self.prompt_template_encode + drop_idx = self.prompt_template_encode_start_idx + txt = [template.format(e) for e in prompt] + txt_tokens = self.tokenizer( + txt, max_length=self.tokenizer_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt" + ).to(device) + encoder_hidden_states = self.text_encoder( + input_ids=txt_tokens.input_ids, + attention_mask=txt_tokens.attention_mask, + output_hidden_states=True, + ) + hidden_states = encoder_hidden_states.hidden_states[-1] + split_hidden_states = self._extract_masked_hidden(hidden_states, txt_tokens.attention_mask) + split_hidden_states = [e[drop_idx:] for e in split_hidden_states] + attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] + max_seq_len = max([e.size(0) for e in split_hidden_states]) + prompt_embeds = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] + ) + encoder_attention_mask = torch.stack( + [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] + ) + + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + return prompt_embeds, encoder_attention_mask + + # Coped from diffusers.pipelines.qwenimage.pipeline_qwenimage.encode_prompt + def encode_prompt( + self, + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + num_images_per_prompt: int = 1, + prompt_embeds: Optional[torch.Tensor] = None, + prompt_embeds_mask: Optional[torch.Tensor] = None, + max_sequence_length: int = 1024, + ): + r""" + + Args: + prompt (`str` or `List[str]`, *optional*): + prompt to be encoded + device: (`torch.device`): + torch device + num_images_per_prompt (`int`): + number of images that should be generated per prompt + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + """ + device = device or self._execution_device + + prompt = [prompt] if isinstance(prompt, str) else prompt + batch_size = len(prompt) if prompt_embeds is None else prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt_embeds, prompt_embeds_mask = self._get_qwen_prompt_embeds(prompt, device) + + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) + prompt_embeds_mask = prompt_embeds_mask.repeat(1, num_images_per_prompt, 1) + prompt_embeds_mask = prompt_embeds_mask.view(batch_size * num_images_per_prompt, seq_len) + + return prompt_embeds, prompt_embeds_mask + + def check_inputs( + self, + prompt, + height, + width, + negative_prompt=None, + prompt_embeds=None, + negative_prompt_embeds=None, + prompt_embeds_mask=None, + negative_prompt_embeds_mask=None, + callback_on_step_end_tensor_inputs=None, + max_sequence_length=None, + ): + if height % (self.vae_scale_factor * 2) != 0 or width % (self.vae_scale_factor * 2) != 0: + logger.warning( + f"`height` and `width` have to be divisible by {self.vae_scale_factor * 2} but are {height} and {width}. Dimensions will be resized accordingly" + ) + + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + + if prompt is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif prompt is None and prompt_embeds is None: + raise ValueError( + "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined." + ) + elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): + raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") + + if negative_prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + if prompt_embeds is not None and prompt_embeds_mask is None: + raise ValueError( + "If `prompt_embeds` are provided, `prompt_embeds_mask` also have to be passed. Make sure to generate `prompt_embeds_mask` from the same text encoder that was used to generate `prompt_embeds`." + ) + if negative_prompt_embeds is not None and negative_prompt_embeds_mask is None: + raise ValueError( + "If `negative_prompt_embeds` are provided, `negative_prompt_embeds_mask` also have to be passed. Make sure to generate `negative_prompt_embeds_mask` from the same text encoder that was used to generate `negative_prompt_embeds`." + ) + + if max_sequence_length is not None and max_sequence_length > 1024: + raise ValueError(f"`max_sequence_length` cannot be greater than 1024 but is {max_sequence_length}") + + @staticmethod + # Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage.QwenImagePipeline._pack_latents + def _pack_latents(latents, batch_size, num_channels_latents, height, width): + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 2, 4, 1, 3, 5) + latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4) + + return latents + + @staticmethod + # Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage.QwenImagePipeline._unpack_latents + def _unpack_latents(latents, height, width, vae_scale_factor): + batch_size, num_patches, channels = latents.shape + + # VAE applies 8x compression on images but we must also account for packing which requires + # latent height and width to be divisible by 2. + height = 2 * (int(height) // (vae_scale_factor * 2)) + width = 2 * (int(width) // (vae_scale_factor * 2)) + + latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2) + latents = latents.permute(0, 3, 1, 4, 2, 5) + + latents = latents.reshape(batch_size, channels // (2 * 2), 1, height, width) + + return latents + + def enable_vae_slicing(self): + r""" + Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to + compute decoding in several steps. This is useful to save some memory and allow larger batch sizes. + """ + depr_message = f"Calling `enable_vae_slicing()` on a `{self.__class__.__name__}` is deprecated and this method will be removed in a future version. Please use `pipe.vae.enable_slicing()`." + deprecate( + "enable_vae_slicing", + "0.40.0", + depr_message, + ) + self.vae.enable_slicing() + + def disable_vae_slicing(self): + r""" + Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to + computing decoding in one step. + """ + depr_message = f"Calling `disable_vae_slicing()` on a `{self.__class__.__name__}` is deprecated and this method will be removed in a future version. Please use `pipe.vae.disable_slicing()`." + deprecate( + "disable_vae_slicing", + "0.40.0", + depr_message, + ) + self.vae.disable_slicing() + + def enable_vae_tiling(self): + r""" + Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to + compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow + processing larger images. + """ + depr_message = f"Calling `enable_vae_tiling()` on a `{self.__class__.__name__}` is deprecated and this method will be removed in a future version. Please use `pipe.vae.enable_tiling()`." + deprecate( + "enable_vae_tiling", + "0.40.0", + depr_message, + ) + self.vae.enable_tiling() + + def disable_vae_tiling(self): + r""" + Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to + computing decoding in one step. + """ + depr_message = f"Calling `disable_vae_tiling()` on a `{self.__class__.__name__}` is deprecated and this method will be removed in a future version. Please use `pipe.vae.disable_tiling()`." + deprecate( + "disable_vae_tiling", + "0.40.0", + depr_message, + ) + self.vae.disable_tiling() + + # Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage.QwenImagePipeline.prepare_latents + def prepare_latents( + self, + batch_size, + num_channels_latents, + height, + width, + dtype, + device, + generator, + latents=None, + ): + # VAE applies 8x compression on images but we must also account for packing which requires + # latent height and width to be divisible by 2. + height = 2 * (int(height) // (self.vae_scale_factor * 2)) + width = 2 * (int(width) // (self.vae_scale_factor * 2)) + + shape = (batch_size, 1, num_channels_latents, height, width) + + if latents is not None: + return latents.to(device=device, dtype=dtype) + + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width) + + return latents + + # Copied from diffusers.pipelines.controlnet_sd3.pipeline_stable_diffusion_3_controlnet.StableDiffusion3ControlNetPipeline.prepare_image + def prepare_image( + self, + image, + width, + height, + batch_size, + num_images_per_prompt, + device, + dtype, + do_classifier_free_guidance=False, + guess_mode=False, + ): + if isinstance(image, torch.Tensor): + pass + else: + image = self.image_processor.preprocess(image, height=height, width=width) + + image_batch_size = image.shape[0] + + if image_batch_size == 1: + repeat_by = batch_size + else: + # image batch size is the same as prompt batch size + repeat_by = num_images_per_prompt + + image = image.repeat_interleave(repeat_by, dim=0) + + image = image.to(device=device, dtype=dtype) + + if do_classifier_free_guidance and not guess_mode: + image = torch.cat([image] * 2) + + return image + + @property + def guidance_scale(self): + return self._guidance_scale + + @property + def attention_kwargs(self): + return self._attention_kwargs + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def current_timestep(self): + return self._current_timestep + + @property + def interrupt(self): + return self._interrupt + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Union[str, List[str]] = None, + negative_prompt: Union[str, List[str]] = None, + true_cfg_scale: float = 4.0, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 50, + sigmas: Optional[List[float]] = None, + guidance_scale: Optional[float] = None, + control_guidance_start: Union[float, List[float]] = 0.0, + control_guidance_end: Union[float, List[float]] = 1.0, + control_image: PipelineImageInput = None, + controlnet_conditioning_scale: Union[float, List[float]] = 1.0, + num_images_per_prompt: int = 1, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.Tensor] = None, + prompt_embeds: Optional[torch.Tensor] = None, + prompt_embeds_mask: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds_mask: Optional[torch.Tensor] = None, + output_type: Optional[str] = "pil", + return_dict: bool = True, + attention_kwargs: Optional[Dict[str, Any]] = None, + callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 512, + ): + r""" + Function invoked when calling the pipeline for generation. + + Args: + prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. + instead. + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `true_cfg_scale` is + not greater than `1`). + true_cfg_scale (`float`, *optional*, defaults to 1.0): + Guidance scale as defined in [Classifier-Free Diffusion + Guidance](https://huggingface.co/papers/2207.12598). `true_cfg_scale` is defined as `w` of equation 2. + of [Imagen Paper](https://huggingface.co/papers/2205.11487). Classifier-free guidance is enabled by + setting `true_cfg_scale > 1` and a provided `negative_prompt`. Higher guidance scale encourages to + generate images that are closely linked to the text `prompt`, usually at the expense of lower image + quality. + height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The height in pixels of the generated image. This is set to 1024 by default for the best results. + width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The width in pixels of the generated image. This is set to 1024 by default for the best results. + num_inference_steps (`int`, *optional*, defaults to 50): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + sigmas (`List[float]`, *optional*): + Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in + their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed + will be used. + guidance_scale (`float`, *optional*, defaults to None): + A guidance scale value for guidance distilled models. Unlike the traditional classifier-free guidance + where the guidance scale is applied during inference through noise prediction rescaling, guidance + distilled models take the guidance scale directly as an input parameter during forward pass. Guidance + scale is enabled by setting `guidance_scale > 1`. Higher guidance scale encourages to generate images + that are closely linked to the text `prompt`, usually at the expense of lower image quality. This + parameter in the pipeline is there to support future guidance-distilled models when they come up. It is + ignored when not using guidance distilled models. To enable traditional classifier-free guidance, + please pass `true_cfg_scale > 1.0` and `negative_prompt` (even an empty negative prompt like " " should + enable classifier-free guidance computations). + num_images_per_prompt (`int`, *optional*, defaults to 1): + The number of images to generate per prompt. + generator (`torch.Generator` or `List[torch.Generator]`, *optional*): + One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) + to make generation deterministic. + latents (`torch.Tensor`, *optional*): + Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor will be generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generate image. Choose between + [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.qwenimage.QwenImagePipelineOutput`] instead of a plain tuple. + attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + callback_on_step_end (`Callable`, *optional*): + A function that calls at the end of each denoising steps during the inference. The function is called + with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, + callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by + `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`List`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeline class. + max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`. + + Examples: + + Returns: + [`~pipelines.qwenimage.QwenImagePipelineOutput`] or `tuple`: + [`~pipelines.qwenimage.QwenImagePipelineOutput`] if `return_dict` is True, otherwise a `tuple`. When + returning a tuple, the first element is a list with the generated images. + """ + + height = height or self.default_sample_size * self.vae_scale_factor + width = width or self.default_sample_size * self.vae_scale_factor + + if not isinstance(control_guidance_start, list) and isinstance(control_guidance_end, list): + control_guidance_start = len(control_guidance_end) * [control_guidance_start] + elif not isinstance(control_guidance_end, list) and isinstance(control_guidance_start, list): + control_guidance_end = len(control_guidance_start) * [control_guidance_end] + elif not isinstance(control_guidance_start, list) and not isinstance(control_guidance_end, list): + mult = 1 + control_guidance_start, control_guidance_end = ( + mult * [control_guidance_start], + mult * [control_guidance_end], + ) + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + height, + width, + negative_prompt=negative_prompt, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + prompt_embeds_mask=prompt_embeds_mask, + negative_prompt_embeds_mask=negative_prompt_embeds_mask, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + max_sequence_length=max_sequence_length, + ) + + self._guidance_scale = guidance_scale + self._attention_kwargs = attention_kwargs + self._current_timestep = None + self._interrupt = False + + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + + has_neg_prompt = negative_prompt is not None or ( + negative_prompt_embeds is not None and negative_prompt_embeds_mask is not None + ) + + if true_cfg_scale > 1 and not has_neg_prompt: + logger.warning( + f"true_cfg_scale is passed as {true_cfg_scale}, but classifier-free guidance is not enabled since no negative_prompt is provided." + ) + elif true_cfg_scale <= 1 and has_neg_prompt: + logger.warning( + " negative_prompt is passed but classifier-free guidance is not enabled since true_cfg_scale <= 1" + ) + + do_true_cfg = true_cfg_scale > 1 and has_neg_prompt + prompt_embeds, prompt_embeds_mask = self.encode_prompt( + prompt=prompt, + prompt_embeds=prompt_embeds, + prompt_embeds_mask=prompt_embeds_mask, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + ) + if do_true_cfg: + negative_prompt_embeds, negative_prompt_embeds_mask = self.encode_prompt( + prompt=negative_prompt, + prompt_embeds=negative_prompt_embeds, + prompt_embeds_mask=negative_prompt_embeds_mask, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + ) + + # 3. Prepare control image + num_channels_latents = self.transformer.config.in_channels // 4 + control_image = self.prepare_image( + image=control_image, + width=width, + height=height, + batch_size=batch_size * num_images_per_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + dtype=self.vae.dtype, + ) + height, width = control_image.shape[-2:] + + if control_image.ndim == 4: + control_image = control_image.unsqueeze(2) + + # vae encode + self.vae_scale_factor = 2 ** len(self.vae.temperal_downsample) + latents_mean = (torch.tensor(self.vae.config.latents_mean).view(1, self.vae.config.z_dim, 1, 1, 1)).to( + device + ) + latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to( + device + ) + + control_image = retrieve_latents(self.vae.encode(control_image), generator=generator) + control_image = (control_image - latents_mean) * latents_std + + control_image = control_image.permute(0, 2, 1, 3, 4) + + # pack + control_image = self._pack_latents( + control_image, + batch_size=control_image.shape[0], + num_channels_latents=num_channels_latents, + height=control_image.shape[3], + width=control_image.shape[4], + ).to(dtype=prompt_embeds.dtype, device=device) + + # 4. Prepare latent variables + num_channels_latents = self.transformer.config.in_channels // 4 + latents = self.prepare_latents( + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + img_shapes = [(1, height // self.vae_scale_factor // 2, width // self.vae_scale_factor // 2)] * batch_size + + # 5. Prepare timesteps + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas + image_seq_len = latents.shape[1] + mu = calculate_shift( + image_seq_len, + self.scheduler.config.get("base_image_seq_len", 256), + self.scheduler.config.get("max_image_seq_len", 4096), + self.scheduler.config.get("base_shift", 0.5), + self.scheduler.config.get("max_shift", 1.15), + ) + timesteps, num_inference_steps = retrieve_timesteps( + self.scheduler, + num_inference_steps, + device, + sigmas=sigmas, + mu=mu, + ) + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + self._num_timesteps = len(timesteps) + + controlnet_keep = [] + for i in range(len(timesteps)): + keeps = [ + 1.0 - float(i / len(timesteps) < s or (i + 1) / len(timesteps) > e) + for s, e in zip(control_guidance_start, control_guidance_end) + ] + controlnet_keep.append(keeps[0]) + + # handle guidance + if self.transformer.config.guidance_embeds and guidance_scale is None: + raise ValueError("guidance_scale is required for guidance-distilled model.") + elif self.transformer.config.guidance_embeds: + guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32) + guidance = guidance.expand(latents.shape[0]) + elif not self.transformer.config.guidance_embeds and guidance_scale is not None: + logger.warning( + f"guidance_scale is passed as {guidance_scale}, but ignored since the model is not guidance-distilled." + ) + guidance = None + elif not self.transformer.config.guidance_embeds and guidance_scale is None: + guidance = None + + if self.attention_kwargs is None: + self._attention_kwargs = {} + + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None + negative_txt_seq_lens = ( + negative_prompt_embeds_mask.sum(dim=1).tolist() if negative_prompt_embeds_mask is not None else None + ) + + # 6. Denoising loop + self.scheduler.set_begin_index(0) + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + self.controlnet.current_steps = i + self.transformer.current_steps = i + if self.interrupt: + continue + + # prepare inputs based on cfg mode + if do_true_cfg: + latent_model_input = torch.cat([latents] * 2) + control_image_input = torch.cat([control_image] * 2) + prompt_embeds_mask_input = [_negative_prompt_embeds_mask for _negative_prompt_embeds_mask in negative_prompt_embeds_mask] + [_prompt_embeds_mask for _prompt_embeds_mask in prompt_embeds_mask] + prompt_embeds_input = [_negative_prompt_embeds for _negative_prompt_embeds in negative_prompt_embeds] + [_prompt_embeds for _prompt_embeds in prompt_embeds] + img_shapes_input = img_shapes * 2 + txt_seq_lens_input = negative_txt_seq_lens + txt_seq_lens + else: + latent_model_input = latents + control_image_input = control_image + prompt_embeds_mask_input = prompt_embeds_mask + prompt_embeds_input = prompt_embeds + img_shapes_input = img_shapes + txt_seq_lens_input = txt_seq_lens + + if hasattr(self.scheduler, "scale_model_input"): + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + # handle guidance + if self.transformer.config.guidance_embeds: + guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32) + guidance = guidance.expand(latent_model_input.shape[0]) + else: + guidance = None + + self._current_timestep = t + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype) + + # prepare controlnet conditioning scale + if isinstance(controlnet_keep[i], list): + cond_scale = [c * s for c, s in zip(controlnet_conditioning_scale, controlnet_keep[i])] + else: + controlnet_cond_scale = controlnet_conditioning_scale + if isinstance(controlnet_cond_scale, list): + controlnet_cond_scale = controlnet_cond_scale[0] + cond_scale = controlnet_cond_scale * controlnet_keep[i] + + # controlnet + controlnet_block_samples = self.controlnet( + hidden_states=latents, + controlnet_cond=control_image, + conditioning_scale=cond_scale, + timestep=t.expand(latents.shape[0]).to(latents.dtype) / 1000, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, + img_shapes=img_shapes, + txt_seq_lens=txt_seq_lens, + return_dict=False, + ) + + with torch.cuda.amp.autocast(dtype=latents.dtype), torch.cuda.device(device=latents.device): + noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=guidance, + encoder_hidden_states=prompt_embeds_input, + encoder_hidden_states_mask=prompt_embeds_mask_input, + img_shapes=img_shapes_input, + controlnet_block_samples=controlnet_block_samples, + attention_kwargs=self.attention_kwargs, + txt_seq_lens=txt_seq_lens_input, + return_dict=False, + ) + + if do_true_cfg: + neg_noise_pred, noise_pred = noise_pred.chunk(2) + comb_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) + + cond_norm = torch.norm(noise_pred, dim=-1, keepdim=True) + noise_norm = torch.norm(comb_pred, dim=-1, keepdim=True) + noise_pred = comb_pred * (cond_norm / noise_norm) + + # compute the previous noisy sample x_t -> x_t-1 + latents_dtype = latents.dtype + latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] + + if latents.dtype != latents_dtype: + if torch.backends.mps.is_available(): + # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272 + latents = latents.to(latents_dtype) + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + + # call the callback, if provided + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + + if XLA_AVAILABLE: + xm.mark_step() + + self._current_timestep = None + if output_type == "latent": + image = latents + else: + latents = self._unpack_latents(latents, height, width, self.vae_scale_factor) + latents = latents.to(self.vae.dtype) + latents_mean = ( + torch.tensor(self.vae.config.latents_mean) + .view(1, self.vae.config.z_dim, 1, 1, 1) + .to(latents.device, latents.dtype) + ) + latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to( + latents.device, latents.dtype + ) + latents = latents / latents_std + latents_mean + image = self.vae.decode(latents, return_dict=False)[0][:, :, 0] + image = self.image_processor.postprocess(image, output_type=output_type) + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + return (image,) + + return QwenImagePipelineOutput(images=image) \ No newline at end of file diff --git a/videox_fun/pipeline/pipeline_z_image_control.py b/videox_fun/pipeline/pipeline_z_image_control.py index 814e30d..395bd42 100644 --- a/videox_fun/pipeline/pipeline_z_image_control.py +++ b/videox_fun/pipeline/pipeline_z_image_control.py @@ -461,7 +461,7 @@ class ZImageControlPipeline(DiffusionPipeline, FromSingleFileMixin): torch.zeros_like(mask_condition)) mask_condition = torch.tile(mask_condition, [1, 3, 1, 1]).to(dtype=weight_dtype, device=device) else: - mask_condition = torch.zeros([batch_size, 3, height, width]).to(dtype=weight_dtype, device=device) + mask_condition = torch.ones([batch_size, 3, height, width]).to(dtype=weight_dtype, device=device) if image is not None: init_image = self.image_processor.preprocess(image, height=height, width=width) diff --git a/videox_fun/utils/__init__.py b/videox_fun/utils/__init__.py index 009df37..7628f0c 100755 --- a/videox_fun/utils/__init__.py +++ b/videox_fun/utils/__init__.py @@ -1,17 +1,20 @@ import importlib.util +from .cfg_optimization import cfg_skip +from .discrete_sampler import DiscreteSampling from .fm_solvers import FlowDPMSolverMultistepScheduler from .fm_solvers_unipc import FlowUniPCMultistepScheduler from .fp8_optimization import (autocast_model_forward, convert_model_weight_to_float8, convert_weight_dtype_wrapper, replace_parameters_by_name) +from .group_offload import (register_auto_device_hook, + safe_enable_group_offload, + safe_remove_group_offloading) from .lora_utils import merge_lora, unmerge_lora -from .utils import (filter_kwargs, get_image_latent, get_image_to_video_latent, get_autocast_dtype, - get_video_to_video_latent, save_videos_grid) -from .cfg_optimization import cfg_skip -from .discrete_sampler import DiscreteSampling - +from .utils import (filter_kwargs, get_autocast_dtype, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, + save_videos_grid) # The pai_fuser is an internally developed acceleration package, which can be used on PAI. if importlib.util.find_spec("paifuser") is not None: @@ -19,7 +22,8 @@ if importlib.util.find_spec("paifuser") is not None: # FP8 Linear Kernel # --------------------------------------------------------------- # from paifuser.ops import (convert_model_weight_to_float8, - convert_weight_dtype_wrapper) + convert_weight_dtype_wrapper) + from . import fp8_optimization fp8_optimization.convert_model_weight_to_float8 = convert_model_weight_to_float8 fp8_optimization.convert_weight_dtype_wrapper = convert_weight_dtype_wrapper diff --git a/videox_fun/utils/cfg_optimization.py b/videox_fun/utils/cfg_optimization.py index 344a2ef..825febf 100755 --- a/videox_fun/utils/cfg_optimization.py +++ b/videox_fun/utils/cfg_optimization.py @@ -1,35 +1,94 @@ +import inspect + import numpy as np import torch def cfg_skip(): def decorator(func): - def wrapper(self, x, *args, **kwargs): - bs = len(x) + def wrapper(self, *args, **kwargs): + if torch.is_grad_enabled(): + return func(self, *args, **kwargs) + + if 'hidden_states' in kwargs and kwargs['hidden_states'] is not None: + main_input = kwargs['hidden_states'] + elif 'x' in kwargs and kwargs['x'] is not None: + main_input = kwargs['x'] + elif len(args) > 0: + main_input = args[0] + else: + raise ValueError("No input tensor found in args or kwargs") + + bs = len(main_input) if bs >= 2 and self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio): bs_half = int(bs // 2) - - new_x = x[bs_half:] - new_args = [] - for arg in args: - if isinstance(arg, (torch.Tensor, list, tuple, np.ndarray)): - new_args.append(arg[bs_half:]) - else: - new_args.append(arg) + new_x = main_input[bs_half:] + new_args = [ + arg[bs_half:] if + isinstance(arg, + (torch.Tensor, list, tuple, np.ndarray)) and + len(arg) == bs else arg for arg in args + ] - new_kwargs = {} - for key, content in kwargs.items(): - if isinstance(content, (torch.Tensor, list, tuple, np.ndarray)): - new_kwargs[key] = content[bs_half:] - else: - new_kwargs[key] = content + new_kwargs = { + k: (v[bs_half:] if + isinstance(v, + (torch.Tensor, list, tuple, + np.ndarray)) and len(v) == bs else v + ) for k, v in kwargs.items() + } else: - new_x = x + new_x = main_input new_args = args new_kwargs = kwargs - result = func(self, new_x, *new_args, **new_kwargs) + sig = inspect.signature(func) + + new_bs = len(new_x) + new_bs_half = int(new_bs // 2) + if new_bs >= 2: + # cond + args_i = [ + arg[new_bs_half:] if + isinstance(arg, + (torch.Tensor, list, tuple, np.ndarray)) and + len(arg) == new_bs else arg for arg in new_args + ] + kwargs_i = { + k: (v[new_bs_half:] if + isinstance(v, + (torch.Tensor, list, tuple, + np.ndarray)) and len(v) == new_bs else v + ) for k, v in new_kwargs.items() + } + if 'cond_flag' in sig.parameters: + kwargs_i["cond_flag"] = True + + cond_out = func(self, *args_i, **kwargs_i) + + # uncond + uncond_args_i = [ + arg[:new_bs_half] if + isinstance(arg, + (torch.Tensor, list, tuple, np.ndarray)) and + len(arg) == new_bs else arg for arg in new_args + ] + uncond_kwargs_i = { + k: (v[:new_bs_half] if + isinstance(v, + (torch.Tensor, list, tuple, + np.ndarray)) and len(v) == new_bs else v + ) for k, v in new_kwargs.items() + } + if 'cond_flag' in sig.parameters: + uncond_kwargs_i["cond_flag"] = False + uncond_out = func(self, *uncond_args_i, + **uncond_kwargs_i) + + result = torch.cat([uncond_out, cond_out], dim=0) + else: + result = func(self, *new_args, **new_kwargs) if bs >= 2 and self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio): result = torch.cat([result, result], dim=0) diff --git a/videox_fun/utils/group_offload.py b/videox_fun/utils/group_offload.py new file mode 100644 index 0000000..cfc00b1 --- /dev/null +++ b/videox_fun/utils/group_offload.py @@ -0,0 +1,1440 @@ +# Modified from https://github.com/huggingface/diffusers/blob/v0.36.0/src/diffusers/hooks/group_offloading.py +# Copyright 2025 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import functools +import hashlib +import os +import types +from contextlib import contextmanager, nullcontext +from dataclasses import dataclass, replace +from enum import Enum +from typing import Any, Dict, List, Optional, Set, Tuple, Union + +import safetensors.torch +import torch +from diffusers.utils import get_logger, is_accelerate_available + +if is_accelerate_available(): + from accelerate.hooks import AlignDevicesHook, CpuOffload + from accelerate.utils import send_to_device + +logger = get_logger(__name__) # pylint: disable=invalid-name + + +# fmt: off +_GROUP_OFFLOADING = "group_offloading" +_LAYER_EXECUTION_TRACKER = "layer_execution_tracker" +_LAZY_PREFETCH_GROUP_OFFLOADING = "lazy_prefetch_group_offloading" +_GROUP_ID_LAZY_LEAF = "lazy_leafs" +# fmt: on + +_GO_LC_SUPPORTED_PYTORCH_LAYERS = ( + torch.nn.Conv1d, + torch.nn.Conv2d, + torch.nn.Conv3d, + torch.nn.ConvTranspose1d, + torch.nn.ConvTranspose2d, + torch.nn.ConvTranspose3d, + torch.nn.Linear, + # TODO(aryan): look into torch.nn.LayerNorm, torch.nn.GroupNorm later, seems to be causing some issues with CogVideoX + # because of double invocation of the same norm layer in CogVideoXLayerNorm +) + + +class ModelHook: + r""" + A hook that contains callbacks to be executed just before and after the forward method of a model. + """ + + _is_stateful = False + + def __init__(self): + self.fn_ref: "HookFunctionReference" = None + + def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module: + r""" + Hook that is executed when a model is initialized. + + Args: + module (`torch.nn.Module`): + The module attached to this hook. + """ + return module + + def deinitalize_hook(self, module: torch.nn.Module) -> torch.nn.Module: + r""" + Hook that is executed when a model is deinitalized. + + Args: + module (`torch.nn.Module`): + The module attached to this hook. + """ + return module + + def pre_forward(self, module: torch.nn.Module, *args, **kwargs) -> Tuple[Tuple[Any], Dict[str, Any]]: + r""" + Hook that is executed just before the forward method of the model. + + Args: + module (`torch.nn.Module`): + The module whose forward pass will be executed just after this event. + args (`Tuple[Any]`): + The positional arguments passed to the module. + kwargs (`Dict[Str, Any]`): + The keyword arguments passed to the module. + Returns: + `Tuple[Tuple[Any], Dict[Str, Any]]`: + A tuple with the treated `args` and `kwargs`. + """ + return args, kwargs + + def post_forward(self, module: torch.nn.Module, output: Any) -> Any: + r""" + Hook that is executed just after the forward method of the model. + + Args: + module (`torch.nn.Module`): + The module whose forward pass been executed just before this event. + output (`Any`): + The output of the module. + Returns: + `Any`: The processed `output`. + """ + return output + + def detach_hook(self, module: torch.nn.Module) -> torch.nn.Module: + r""" + Hook that is executed when the hook is detached from a module. + + Args: + module (`torch.nn.Module`): + The module detached from this hook. + """ + return module + + def reset_state(self, module: torch.nn.Module): + if self._is_stateful: + raise NotImplementedError("This hook is stateful and needs to implement the `reset_state` method.") + return module + + +class HookFunctionReference: + def __init__(self) -> None: + """A container class that maintains mutable references to forward pass functions in a hook chain. + + Its mutable nature allows the hook system to modify the execution chain dynamically without rebuilding the + entire forward pass structure. + + Attributes: + pre_forward: A callable that processes inputs before the main forward pass. + post_forward: A callable that processes outputs after the main forward pass. + forward: The current forward function in the hook chain. + original_forward: The original forward function, stored when a hook provides a custom new_forward. + + The class enables hook removal by allowing updates to the forward chain through reference modification rather + than requiring reconstruction of the entire chain. When a hook is removed, only the relevant references need to + be updated, preserving the execution order of the remaining hooks. + """ + self.pre_forward = None + self.post_forward = None + self.forward = None + self.original_forward = None + + +class HookRegistry: + def __init__(self, module_ref: torch.nn.Module) -> None: + super().__init__() + + self.hooks: Dict[str, ModelHook] = {} + + self._module_ref = module_ref + self._hook_order = [] + self._fn_refs = [] + + def register_hook(self, hook: ModelHook, name: str) -> None: + if name in self.hooks.keys(): + raise ValueError( + f"Hook with name {name} already exists in the registry. Please use a different name or " + f"first remove the existing hook and then add a new one." + ) + + self._module_ref = hook.initialize_hook(self._module_ref) + + def create_new_forward(function_reference: HookFunctionReference): + def new_forward(module, *args, **kwargs): + args, kwargs = function_reference.pre_forward(module, *args, **kwargs) + output = function_reference.forward(*args, **kwargs) + return function_reference.post_forward(module, output) + + return new_forward + + forward = self._module_ref.forward + + fn_ref = HookFunctionReference() + fn_ref.pre_forward = hook.pre_forward + fn_ref.post_forward = hook.post_forward + fn_ref.forward = forward + + if hasattr(hook, "new_forward"): + fn_ref.original_forward = forward + fn_ref.forward = functools.update_wrapper( + functools.partial(hook.new_forward, self._module_ref), hook.new_forward + ) + + rewritten_forward = create_new_forward(fn_ref) + self._module_ref.forward = functools.update_wrapper( + functools.partial(rewritten_forward, self._module_ref), rewritten_forward + ) + + hook.fn_ref = fn_ref + self.hooks[name] = hook + self._hook_order.append(name) + self._fn_refs.append(fn_ref) + + def get_hook(self, name: str) -> Optional[ModelHook]: + return self.hooks.get(name, None) + + def remove_hook(self, name: str, recurse: bool = True) -> None: + if name in self.hooks.keys(): + num_hooks = len(self._hook_order) + hook = self.hooks[name] + index = self._hook_order.index(name) + fn_ref = self._fn_refs[index] + + old_forward = fn_ref.forward + if fn_ref.original_forward is not None: + old_forward = fn_ref.original_forward + + if index == num_hooks - 1: + self._module_ref.forward = old_forward + else: + self._fn_refs[index + 1].forward = old_forward + + self._module_ref = hook.deinitalize_hook(self._module_ref) + del self.hooks[name] + self._hook_order.pop(index) + self._fn_refs.pop(index) + + if recurse: + for module_name, module in self._module_ref.named_modules(): + if module_name == "": + continue + if hasattr(module, "_diffusers_hook"): + module._diffusers_hook.remove_hook(name, recurse=False) + + def reset_stateful_hooks(self, recurse: bool = True) -> None: + for hook_name in reversed(self._hook_order): + hook = self.hooks[hook_name] + if hook._is_stateful: + hook.reset_state(self._module_ref) + + if recurse: + for module_name, module in self._module_ref.named_modules(): + if module_name == "": + continue + if hasattr(module, "_diffusers_hook"): + module._diffusers_hook.reset_stateful_hooks(recurse=False) + + @classmethod + def check_if_exists_or_initialize(cls, module: torch.nn.Module) -> "HookRegistry": + if not hasattr(module, "_diffusers_hook"): + module._diffusers_hook = cls(module) + return module._diffusers_hook + + def __repr__(self) -> str: + registry_repr = "" + for i, hook_name in enumerate(self._hook_order): + if self.hooks[hook_name].__class__.__repr__ is not object.__repr__: + hook_repr = self.hooks[hook_name].__repr__() + else: + hook_repr = self.hooks[hook_name].__class__.__name__ + registry_repr += f" ({i}) {hook_name} - {hook_repr}" + if i < len(self._hook_order) - 1: + registry_repr += "\n" + return f"HookRegistry(\n{registry_repr}\n)" + + +class GroupOffloadingType(str, Enum): + BLOCK_LEVEL = "block_level" + LEAF_LEVEL = "leaf_level" + + +@dataclass +class GroupOffloadingConfig: + onload_device: torch.device + offload_device: torch.device + offload_type: GroupOffloadingType + non_blocking: bool + record_stream: bool + low_cpu_mem_usage: bool + num_blocks_per_group: Optional[int] = None + offload_to_disk_path: Optional[str] = None + stream: Optional[Union[torch.cuda.Stream, torch.Stream]] = None + block_modules: Optional[List[str]] = None + exclude_kwargs: Optional[List[str]] = None + module_prefix: Optional[str] = "" + + +class ModuleGroup: + def __init__( + self, + modules: List[torch.nn.Module], + offload_device: torch.device, + onload_device: torch.device, + offload_leader: torch.nn.Module, + onload_leader: Optional[torch.nn.Module] = None, + parameters: Optional[List[torch.nn.Parameter]] = None, + buffers: Optional[List[torch.Tensor]] = None, + non_blocking: bool = False, + stream: Union[torch.cuda.Stream, torch.Stream, None] = None, + record_stream: Optional[bool] = False, + low_cpu_mem_usage: bool = False, + onload_self: bool = True, + offload_to_disk_path: Optional[str] = None, + group_id: Optional[Union[int, str]] = None, + ) -> None: + self.modules = modules + self.offload_device = offload_device + self.onload_device = onload_device + self.offload_leader = offload_leader + self.onload_leader = onload_leader + self.parameters = parameters or [] + self.buffers = buffers or [] + self.non_blocking = non_blocking or stream is not None + self.stream = stream + self.record_stream = record_stream + self.onload_self = onload_self + self.low_cpu_mem_usage = low_cpu_mem_usage + + self.offload_to_disk_path = offload_to_disk_path + self._is_offloaded_to_disk = False + + if self.offload_to_disk_path is not None: + # Instead of `group_id or str(id(self))` we do this because `group_id` can be "" as well. + self.group_id = group_id if group_id is not None else str(id(self)) + short_hash = _compute_group_hash(self.group_id) + self.safetensors_file_path = os.path.join(self.offload_to_disk_path, f"group_{short_hash}.safetensors") + + all_tensors = [] + for module in self.modules: + all_tensors.extend(list(module.parameters())) + all_tensors.extend(list(module.buffers())) + all_tensors.extend(self.parameters) + all_tensors.extend(self.buffers) + all_tensors = list(dict.fromkeys(all_tensors)) # Remove duplicates + + self.tensor_to_key = {tensor: f"tensor_{i}" for i, tensor in enumerate(all_tensors)} + self.key_to_tensor = {v: k for k, v in self.tensor_to_key.items()} + self.cpu_param_dict = {} + else: + self.cpu_param_dict = self._init_cpu_param_dict() + + self._torch_accelerator_module = ( + getattr(torch, torch.accelerator.current_accelerator().type) + if hasattr(torch, "accelerator") + else torch.cuda + ) + + def _init_cpu_param_dict(self): + cpu_param_dict = {} + if self.stream is None: + return cpu_param_dict + + for module in self.modules: + for param in module.parameters(): + cpu_param_dict[param] = param.data.cpu() if self.low_cpu_mem_usage else param.data.cpu().pin_memory() + for buffer in module.buffers(): + cpu_param_dict[buffer] = ( + buffer.data.cpu() if self.low_cpu_mem_usage else buffer.data.cpu().pin_memory() + ) + + for param in self.parameters: + cpu_param_dict[param] = param.data.cpu() if self.low_cpu_mem_usage else param.data.cpu().pin_memory() + + for buffer in self.buffers: + cpu_param_dict[buffer] = buffer.data.cpu() if self.low_cpu_mem_usage else buffer.data.cpu().pin_memory() + + return cpu_param_dict + + @contextmanager + def _pinned_memory_tensors(self): + try: + pinned_dict = { + param: tensor.pin_memory() if not tensor.is_pinned() else tensor + for param, tensor in self.cpu_param_dict.items() + } + yield pinned_dict + finally: + pinned_dict = None + + def _transfer_tensor_to_device(self, tensor, source_tensor, default_stream): + tensor.data = source_tensor.to(self.onload_device, non_blocking=self.non_blocking) + if self.record_stream: + tensor.data.record_stream(default_stream) + + def _process_tensors_from_modules(self, pinned_memory=None, default_stream=None): + for group_module in self.modules: + for param in group_module.parameters(): + source = pinned_memory[param] if pinned_memory else param.data + self._transfer_tensor_to_device(param, source, default_stream) + for buffer in group_module.buffers(): + source = pinned_memory[buffer] if pinned_memory else buffer.data + self._transfer_tensor_to_device(buffer, source, default_stream) + + for param in self.parameters: + source = pinned_memory[param] if pinned_memory else param.data + self._transfer_tensor_to_device(param, source, default_stream) + + for buffer in self.buffers: + source = pinned_memory[buffer] if pinned_memory else buffer.data + self._transfer_tensor_to_device(buffer, source, default_stream) + + def _onload_from_disk(self): + if self.stream is not None: + # Wait for previous Host->Device transfer to complete + self.stream.synchronize() + + context = nullcontext() if self.stream is None else self._torch_accelerator_module.stream(self.stream) + current_stream = self._torch_accelerator_module.current_stream() if self.record_stream else None + + with context: + # Load to CPU (if using streams) or directly to target device, pin, and async copy to device + device = str(self.onload_device) if self.stream is None else "cpu" + loaded_tensors = safetensors.torch.load_file(self.safetensors_file_path, device=device) + + if self.stream is not None: + for key, tensor_obj in self.key_to_tensor.items(): + pinned_tensor = loaded_tensors[key].pin_memory() + tensor_obj.data = pinned_tensor.to(self.onload_device, non_blocking=self.non_blocking) + if self.record_stream: + tensor_obj.data.record_stream(current_stream) + else: + onload_device = ( + self.onload_device.type if isinstance(self.onload_device, torch.device) else self.onload_device + ) + loaded_tensors = safetensors.torch.load_file(self.safetensors_file_path, device=onload_device) + for key, tensor_obj in self.key_to_tensor.items(): + tensor_obj.data = loaded_tensors[key] + + def _onload_from_memory(self): + if self.stream is not None: + # Wait for previous Host->Device transfer to complete + self.stream.synchronize() + + context = nullcontext() if self.stream is None else self._torch_accelerator_module.stream(self.stream) + default_stream = self._torch_accelerator_module.current_stream() if self.stream is not None else None + + with context: + if self.stream is not None: + with self._pinned_memory_tensors() as pinned_memory: + self._process_tensors_from_modules(pinned_memory, default_stream=default_stream) + else: + self._process_tensors_from_modules(None) + + def _offload_to_disk(self): + # TODO: we can potentially optimize this code path by checking if the _all_ the desired + # safetensor files exist on the disk and if so, skip this step entirely, reducing IO + # overhead. Currently, we just check if the given `safetensors_file_path` exists and if not + # we perform a write. + # Check if the file has been saved in this session or if it already exists on disk. + if not self._is_offloaded_to_disk and not os.path.exists(self.safetensors_file_path): + os.makedirs(os.path.dirname(self.safetensors_file_path), exist_ok=True) + tensors_to_save = {key: tensor.data.to(self.offload_device) for tensor, key in self.tensor_to_key.items()} + safetensors.torch.save_file(tensors_to_save, self.safetensors_file_path) + + # The group is now considered offloaded to disk for the rest of the session. + self._is_offloaded_to_disk = True + + # We do this to free up the RAM which is still holding the up tensor data. + for tensor_obj in self.tensor_to_key.keys(): + tensor_obj.data = torch.empty_like(tensor_obj.data, device=self.offload_device) + + def _offload_to_memory(self): + if self.stream is not None: + if not self.record_stream: + self._torch_accelerator_module.current_stream().synchronize() + + for group_module in self.modules: + for param in group_module.parameters(): + param.data = self.cpu_param_dict[param] + for param in self.parameters: + param.data = self.cpu_param_dict[param] + for buffer in self.buffers: + buffer.data = self.cpu_param_dict[buffer] + else: + for group_module in self.modules: + group_module.to(self.offload_device, non_blocking=False) + for param in self.parameters: + param.data = param.data.to(self.offload_device, non_blocking=False) + for buffer in self.buffers: + buffer.data = buffer.data.to(self.offload_device, non_blocking=False) + + @torch.compiler.disable() + def onload_(self): + r"""Onloads the group of parameters to the onload_device.""" + if self.offload_to_disk_path is not None: + self._onload_from_disk() + else: + self._onload_from_memory() + + @torch.compiler.disable() + def offload_(self): + r"""Offloads the group of parameters to the offload_device.""" + if self.offload_to_disk_path: + self._offload_to_disk() + else: + self._offload_to_memory() + + +class GroupOffloadingHook(ModelHook): + r""" + A hook that offloads groups of torch.nn.Module to the CPU for storage and onloads to accelerator device for + computation. Each group has one "onload leader" module that is responsible for onloading, and an "offload leader" + module that is responsible for offloading. If prefetching is enabled, the onload leader of the previous module + group is responsible for onloading the current module group. + """ + + _is_stateful = False + + def __init__(self, group: ModuleGroup, *, config: GroupOffloadingConfig) -> None: + self.group = group + self.next_group: Optional[ModuleGroup] = None + self.config = config + + def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module: + if self.group.offload_leader == module: + self.group.offload_() + return module + + def pre_forward(self, module: torch.nn.Module, *args, **kwargs): + # If there wasn't an onload_leader assigned, we assume that the submodule that first called its forward + # method is the onload_leader of the group. + if self.group.onload_leader is None: + self.group.onload_leader = module + + # If the current module is the onload_leader of the group, we onload the group if it is supposed + # to onload itself. In the case of using prefetching with streams, we onload the next group if + # it is not supposed to onload itself. + if self.group.onload_leader == module: + if self.group.onload_self: + self.group.onload_() + + should_onload_next_group = self.next_group is not None and not self.next_group.onload_self + if should_onload_next_group: + self.next_group.onload_() + + should_synchronize = ( + not self.group.onload_self and self.group.stream is not None and not should_onload_next_group + ) + if should_synchronize: + # If this group didn't onload itself, it means it was asynchronously onloaded by the + # previous group. We need to synchronize the side stream to ensure parameters + # are completely loaded to proceed with forward pass. Without this, uninitialized + # weights will be used in the computation, leading to incorrect results + # Also, we should only do this synchronization if we don't already do it from the sync call in + # self.next_group.onload_, hence the `not should_onload_next_group` check. + self.group.stream.synchronize() + + args = send_to_device(args, self.group.onload_device, non_blocking=self.group.non_blocking) + + # Some Autoencoder models use a feature cache that is passed through submodules + # and modified in place. The `send_to_device` call returns a copy of this feature cache object + # which breaks the inplace updates. Use `exclude_kwargs` to mark these cache features + exclude_kwargs = self.config.exclude_kwargs or [] + if exclude_kwargs: + moved_kwargs = send_to_device( + {k: v for k, v in kwargs.items() if k not in exclude_kwargs}, + self.group.onload_device, + non_blocking=self.group.non_blocking, + ) + kwargs.update(moved_kwargs) + else: + kwargs = send_to_device(kwargs, self.group.onload_device, non_blocking=self.group.non_blocking) + + return args, kwargs + + def post_forward(self, module: torch.nn.Module, output): + if self.group.offload_leader == module: + self.group.offload_() + return output + + +class LazyPrefetchGroupOffloadingHook(ModelHook): + r""" + A hook, used in conjunction with GroupOffloadingHook, that applies lazy prefetching to groups of torch.nn.Module. + This hook is used to determine the order in which the layers are executed during the forward pass. Once the layer + invocation order is known, assignments of the next_group attribute for prefetching can be made, which allows + prefetching groups in the correct order. + """ + + _is_stateful = False + + def __init__(self): + self.execution_order: List[Tuple[str, torch.nn.Module]] = [] + self._layer_execution_tracker_module_names = set() + + def initialize_hook(self, module): + def make_execution_order_update_callback(current_name, current_submodule): + def callback(): + if not torch.compiler.is_compiling(): + logger.debug(f"Adding {current_name} to the execution order") + self.execution_order.append((current_name, current_submodule)) + + return callback + + # To every submodule that contains a group offloading hook (at this point, no prefetching is enabled for any + # of the groups), we add a layer execution tracker hook that will be used to determine the order in which the + # layers are executed during the forward pass. + for name, submodule in module.named_modules(): + if name == "" or not hasattr(submodule, "_diffusers_hook"): + continue + + registry = HookRegistry.check_if_exists_or_initialize(submodule) + group_offloading_hook = registry.get_hook(_GROUP_OFFLOADING) + + if group_offloading_hook is not None: + # For the first forward pass, we have to load in a blocking manner + group_offloading_hook.group.non_blocking = False + layer_tracker_hook = LayerExecutionTrackerHook(make_execution_order_update_callback(name, submodule)) + registry.register_hook(layer_tracker_hook, _LAYER_EXECUTION_TRACKER) + self._layer_execution_tracker_module_names.add(name) + + return module + + def post_forward(self, module, output): + # At this point, for the current modules' submodules, we know the execution order of the layers. We can now + # remove the layer execution tracker hooks and apply prefetching by setting the next_group attribute for each + # group offloading hook. + num_executed = len(self.execution_order) + execution_order_module_names = {name for name, _ in self.execution_order} + + # It may be possible that some layers were not executed during the forward pass. This can happen if the layer + # is not used in the forward pass, or if the layer is not executed due to some other reason. In such cases, we + # may not be able to apply prefetching in the correct order, which can lead to device-mismatch related errors + # if the missing layers end up being executed in the future. + if execution_order_module_names != self._layer_execution_tracker_module_names: + unexecuted_layers = list(self._layer_execution_tracker_module_names - execution_order_module_names) + if not torch.compiler.is_compiling(): + logger.warning( + "It seems like some layers were not executed during the forward pass. This may lead to problems when " + "applying lazy prefetching with automatic tracing and lead to device-mismatch related errors. Please " + "make sure that all layers are executed during the forward pass. The following layers were not executed:\n" + f"{unexecuted_layers=}" + ) + + # Remove the layer execution tracker hooks from the submodules + base_module_registry = module._diffusers_hook + registries = [submodule._diffusers_hook for _, submodule in self.execution_order] + group_offloading_hooks = [registry.get_hook(_GROUP_OFFLOADING) for registry in registries] + + for i in range(num_executed): + registries[i].remove_hook(_LAYER_EXECUTION_TRACKER, recurse=False) + + # Remove the current lazy prefetch group offloading hook so that it doesn't interfere with the next forward pass + base_module_registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=False) + + # LazyPrefetchGroupOffloadingHook is only used with streams, so we know that non_blocking should be True. + # We disable non_blocking for the first forward pass, but need to enable it for the subsequent passes to + # see the benefits of prefetching. + for hook in group_offloading_hooks: + hook.group.non_blocking = True + + # Set required attributes for prefetching + if num_executed > 0: + base_module_group_offloading_hook = base_module_registry.get_hook(_GROUP_OFFLOADING) + base_module_group_offloading_hook.next_group = group_offloading_hooks[0].group + base_module_group_offloading_hook.next_group.onload_self = False + + for i in range(num_executed - 1): + name1, _ = self.execution_order[i] + name2, _ = self.execution_order[i + 1] + if not torch.compiler.is_compiling(): + logger.debug(f"Applying lazy prefetch group offloading from {name1} to {name2}") + group_offloading_hooks[i].next_group = group_offloading_hooks[i + 1].group + group_offloading_hooks[i].next_group.onload_self = False + + return output + + +class LayerExecutionTrackerHook(ModelHook): + r""" + A hook that tracks the order in which the layers are executed during the forward pass by calling back to the + LazyPrefetchGroupOffloadingHook to update the execution order. + """ + + _is_stateful = False + + def __init__(self, execution_order_update_callback): + self.execution_order_update_callback = execution_order_update_callback + + def pre_forward(self, module, *args, **kwargs): + self.execution_order_update_callback() + return args, kwargs + + +def apply_group_offloading( + module: torch.nn.Module, + onload_device: Union[str, torch.device], + offload_device: Union[str, torch.device] = torch.device("cpu"), + offload_type: Union[str, GroupOffloadingType] = "block_level", + num_blocks_per_group: Optional[int] = None, + non_blocking: bool = False, + use_stream: bool = False, + record_stream: bool = False, + low_cpu_mem_usage: bool = False, + offload_to_disk_path: Optional[str] = None, + block_modules: Optional[List[str]] = None, + exclude_kwargs: Optional[List[str]] = None, +) -> None: + r""" + Applies group offloading to the internal layers of a torch.nn.Module. To understand what group offloading is, and + where it is beneficial, we need to first provide some context on how other supported offloading methods work. + + Typically, offloading is done at two levels: + - Module-level: In Diffusers, this can be enabled using the `ModelMixin::enable_model_cpu_offload()` method. It + works by offloading each component of a pipeline to the CPU for storage, and onloading to the accelerator device + when needed for computation. This method is more memory-efficient than keeping all components on the accelerator, + but the memory requirements are still quite high. For this method to work, one needs memory equivalent to size of + the model in runtime dtype + size of largest intermediate activation tensors to be able to complete the forward + pass. + - Leaf-level: In Diffusers, this can be enabled using the `ModelMixin::enable_sequential_cpu_offload()` method. It + works by offloading the lowest leaf-level parameters of the computation graph to the CPU for storage, and + onloading only the leafs to the accelerator device for computation. This uses the lowest amount of accelerator + memory, but can be slower due to the excessive number of device synchronizations. + + Group offloading is a middle ground between the two methods. It works by offloading groups of internal layers, + (either `torch.nn.ModuleList` or `torch.nn.Sequential`). This method uses lower memory than module-level + offloading. It is also faster than leaf-level/sequential offloading, as the number of device synchronizations is + reduced. + + Another supported feature (for CUDA devices with support for asynchronous data transfer streams) is the ability to + overlap data transfer and computation to reduce the overall execution time compared to sequential offloading. This + is enabled using layer prefetching with streams, i.e., the layer that is to be executed next starts onloading to + the accelerator device while the current layer is being executed - this increases the memory requirements slightly. + Note that this implementation also supports leaf-level offloading but can be made much faster when using streams. + + Args: + module (`torch.nn.Module`): + The module to which group offloading is applied. + onload_device (`torch.device`): + The device to which the group of modules are onloaded. + offload_device (`torch.device`, defaults to `torch.device("cpu")`): + The device to which the group of modules are offloaded. This should typically be the CPU. Default is CPU. + offload_type (`str` or `GroupOffloadingType`, defaults to "block_level"): + The type of offloading to be applied. Can be one of "block_level" or "leaf_level". Default is + "block_level". + offload_to_disk_path (`str`, *optional*, defaults to `None`): + The path to the directory where parameters will be offloaded. Setting this option can be useful in limited + RAM environment settings where a reasonable speed-memory trade-off is desired. + num_blocks_per_group (`int`, *optional*): + The number of blocks per group when using offload_type="block_level". This is required when using + offload_type="block_level". + non_blocking (`bool`, defaults to `False`): + If True, offloading and onloading is done with non-blocking data transfer. + use_stream (`bool`, defaults to `False`): + If True, offloading and onloading is done asynchronously using a CUDA stream. This can be useful for + overlapping computation and data transfer. + record_stream (`bool`, defaults to `False`): When enabled with `use_stream`, it marks the current tensor + as having been used by this stream. It is faster at the expense of slightly more memory usage. Refer to the + [PyTorch official docs](https://pytorch.org/docs/stable/generated/torch.Tensor.record_stream.html) more + details. + low_cpu_mem_usage (`bool`, defaults to `False`): + If True, the CPU memory usage is minimized by pinning tensors on-the-fly instead of pre-pinning them. This + option only matters when using streamed CPU offloading (i.e. `use_stream=True`). This can be useful when + the CPU memory is a bottleneck but may counteract the benefits of using streams. + block_modules (`List[str]`, *optional*): + List of module names that should be treated as blocks for offloading. If provided, only these modules will + be considered for block-level offloading. If not provided, the default block detection logic will be used. + exclude_kwargs (`List[str]`, *optional*): + List of kwarg keys that should not be processed by send_to_device. This is useful for mutable state like + caching lists that need to maintain their object identity across forward passes. If not provided, will be + inferred from the module's `_skip_keys` attribute if it exists. + + Example: + ```python + >>> from diffusers import CogVideoXTransformer3DModel + >>> from diffusers.hooks import apply_group_offloading + + >>> transformer = CogVideoXTransformer3DModel.from_pretrained( + ... "THUDM/CogVideoX-5b", subfolder="transformer", torch_dtype=torch.bfloat16 + ... ) + + >>> apply_group_offloading( + ... transformer, + ... onload_device=torch.device("cuda"), + ... offload_device=torch.device("cpu"), + ... offload_type="block_level", + ... num_blocks_per_group=2, + ... use_stream=True, + ... ) + ``` + """ + + onload_device = torch.device(onload_device) if isinstance(onload_device, str) else onload_device + offload_device = torch.device(offload_device) if isinstance(offload_device, str) else offload_device + offload_type = GroupOffloadingType(offload_type) + + stream = None + if use_stream: + if torch.cuda.is_available(): + stream = torch.cuda.Stream() + elif hasattr(torch, "xpu") and torch.xpu.is_available(): + stream = torch.Stream() + else: + raise ValueError("Using streams for data transfer requires a CUDA device, or an Intel XPU device.") + + if not use_stream and record_stream: + raise ValueError("`record_stream` cannot be True when `use_stream=False`.") + if offload_type == GroupOffloadingType.BLOCK_LEVEL and num_blocks_per_group is None: + raise ValueError("`num_blocks_per_group` must be provided when using `offload_type='block_level'.") + + _raise_error_if_accelerate_model_or_sequential_hook_present(module) + + if block_modules is None: + block_modules = getattr(module, "_group_offload_block_modules", None) + + if exclude_kwargs is None: + exclude_kwargs = getattr(module, "_skip_keys", None) + + config = GroupOffloadingConfig( + onload_device=onload_device, + offload_device=offload_device, + offload_type=offload_type, + num_blocks_per_group=num_blocks_per_group, + non_blocking=non_blocking, + stream=stream, + record_stream=record_stream, + low_cpu_mem_usage=low_cpu_mem_usage, + offload_to_disk_path=offload_to_disk_path, + block_modules=block_modules, + exclude_kwargs=exclude_kwargs, + ) + _apply_group_offloading(module, config) + + +def _apply_group_offloading(module: torch.nn.Module, config: GroupOffloadingConfig) -> None: + if config.offload_type == GroupOffloadingType.BLOCK_LEVEL: + _apply_group_offloading_block_level(module, config) + elif config.offload_type == GroupOffloadingType.LEAF_LEVEL: + _apply_group_offloading_leaf_level(module, config) + else: + assert False + + +def _apply_group_offloading_block_level(module: torch.nn.Module, config: GroupOffloadingConfig) -> None: + r""" + This function applies offloading to groups of torch.nn.ModuleList or torch.nn.Sequential blocks, and explicitly + defined block modules. In comparison to the "leaf_level" offloading, which is more fine-grained, this offloading is + done at the top-level blocks and modules specified in block_modules. + + When block_modules is provided, only those modules will be treated as blocks for offloading. For each specified + module, recursively apply block offloading to it. + """ + if config.stream is not None and config.num_blocks_per_group != 1: + logger.warning( + f"Using streams is only supported for num_blocks_per_group=1. Got {config.num_blocks_per_group=}. Setting it to 1." + ) + config.num_blocks_per_group = 1 + + block_modules = set(config.block_modules) if config.block_modules is not None else set() + + # Create module groups for ModuleList and Sequential blocks, and explicitly defined block modules + modules_with_group_offloading = set() + unmatched_modules = [] + matched_module_groups = [] + + for name, submodule in module.named_children(): + # Check if this is an explicitly defined block module + if name in block_modules: + # Track submodule using a prefix to avoid filename collisions during disk offload. + # Without this, submodules sharing the same model class would be assigned identical + # filenames (derived from the class name). + prefix = f"{config.module_prefix}{name}." if config.module_prefix else f"{name}." + submodule_config = replace(config, module_prefix=prefix) + + _apply_group_offloading_block_level(submodule, submodule_config) + modules_with_group_offloading.add(name) + + elif isinstance(submodule, (torch.nn.ModuleList, torch.nn.Sequential)): + # Handle ModuleList and Sequential blocks as before + for i in range(0, len(submodule), config.num_blocks_per_group): + current_modules = list(submodule[i : i + config.num_blocks_per_group]) + if len(current_modules) == 0: + continue + + group_id = f"{config.module_prefix}{name}_{i}_{i + len(current_modules) - 1}" + group = ModuleGroup( + modules=current_modules, + offload_device=config.offload_device, + onload_device=config.onload_device, + offload_to_disk_path=config.offload_to_disk_path, + offload_leader=current_modules[-1], + onload_leader=current_modules[0], + non_blocking=config.non_blocking, + stream=config.stream, + record_stream=config.record_stream, + low_cpu_mem_usage=config.low_cpu_mem_usage, + onload_self=True, + group_id=group_id, + ) + matched_module_groups.append(group) + for j in range(i, i + len(current_modules)): + modules_with_group_offloading.add(f"{name}.{j}") + else: + # This is an unmatched module + unmatched_modules.append((name, submodule)) + + # Apply group offloading hooks to the module groups + for i, group in enumerate(matched_module_groups): + for group_module in group.modules: + _apply_group_offloading_hook(group_module, group, config=config) + + # Parameters and Buffers of the top-level module need to be offloaded/onloaded separately + # when the forward pass of this module is called. This is because the top-level module is not + # part of any group (as doing so would lead to no VRAM savings). + parameters = _gather_parameters_with_no_group_offloading_parent(module, modules_with_group_offloading) + buffers = _gather_buffers_with_no_group_offloading_parent(module, modules_with_group_offloading) + parameters = [param for _, param in parameters] + buffers = [buffer for _, buffer in buffers] + + # Create a group for the remaining unmatched submodules of the top-level + # module so that they are on the correct device when the forward pass is called. + unmatched_modules = [unmatched_module for _, unmatched_module in unmatched_modules] + if len(unmatched_modules) > 0 or len(parameters) > 0 or len(buffers) > 0: + unmatched_group = ModuleGroup( + modules=unmatched_modules, + offload_device=config.offload_device, + onload_device=config.onload_device, + offload_to_disk_path=config.offload_to_disk_path, + offload_leader=module, + onload_leader=module, + parameters=parameters, + buffers=buffers, + non_blocking=False, + stream=None, + record_stream=False, + onload_self=True, + group_id=f"{config.module_prefix}{module.__class__.__name__}_unmatched_group", + ) + if config.stream is None: + _apply_group_offloading_hook(module, unmatched_group, config=config) + else: + _apply_lazy_group_offloading_hook(module, unmatched_group, config=config) + + +def _apply_group_offloading_leaf_level(module: torch.nn.Module, config: GroupOffloadingConfig) -> None: + r""" + This function applies offloading to groups of leaf modules in a torch.nn.Module. This method has minimal memory + requirements. However, it can be slower compared to other offloading methods due to the excessive number of device + synchronizations. When using devices that support streams to overlap data transfer and computation, this method can + reduce memory usage without any performance degradation. + """ + # Create module groups for leaf modules and apply group offloading hooks + modules_with_group_offloading = set() + for name, submodule in module.named_modules(): + if not isinstance(submodule, _GO_LC_SUPPORTED_PYTORCH_LAYERS): + continue + group = ModuleGroup( + modules=[submodule], + offload_device=config.offload_device, + onload_device=config.onload_device, + offload_to_disk_path=config.offload_to_disk_path, + offload_leader=submodule, + onload_leader=submodule, + non_blocking=config.non_blocking, + stream=config.stream, + record_stream=config.record_stream, + low_cpu_mem_usage=config.low_cpu_mem_usage, + onload_self=True, + group_id=name, + ) + _apply_group_offloading_hook(submodule, group, config=config) + modules_with_group_offloading.add(name) + + # Parameters and Buffers at all non-leaf levels need to be offloaded/onloaded separately when the forward pass + # of the module is called + module_dict = dict(module.named_modules()) + parameters = _gather_parameters_with_no_group_offloading_parent(module, modules_with_group_offloading) + buffers = _gather_buffers_with_no_group_offloading_parent(module, modules_with_group_offloading) + + # Find closest module parent for each parameter and buffer, and attach group hooks + parent_to_parameters = {} + for name, param in parameters: + parent_name = _find_parent_module_in_module_dict(name, module_dict) + if parent_name in parent_to_parameters: + parent_to_parameters[parent_name].append(param) + else: + parent_to_parameters[parent_name] = [param] + + parent_to_buffers = {} + for name, buffer in buffers: + parent_name = _find_parent_module_in_module_dict(name, module_dict) + if parent_name in parent_to_buffers: + parent_to_buffers[parent_name].append(buffer) + else: + parent_to_buffers[parent_name] = [buffer] + + parent_names = set(parent_to_parameters.keys()) | set(parent_to_buffers.keys()) + for name in parent_names: + parameters = parent_to_parameters.get(name, []) + buffers = parent_to_buffers.get(name, []) + parent_module = module_dict[name] + group = ModuleGroup( + modules=[], + offload_device=config.offload_device, + onload_device=config.onload_device, + offload_leader=parent_module, + onload_leader=parent_module, + offload_to_disk_path=config.offload_to_disk_path, + parameters=parameters, + buffers=buffers, + non_blocking=config.non_blocking, + stream=config.stream, + record_stream=config.record_stream, + low_cpu_mem_usage=config.low_cpu_mem_usage, + onload_self=True, + group_id=name, + ) + _apply_group_offloading_hook(parent_module, group, config=config) + + if config.stream is not None: + # When using streams, we need to know the layer execution order for applying prefetching (to overlap data transfer + # and computation). Since we don't know the order beforehand, we apply a lazy prefetching hook that will find the + # execution order and apply prefetching in the correct order. + unmatched_group = ModuleGroup( + modules=[], + offload_device=config.offload_device, + onload_device=config.onload_device, + offload_to_disk_path=config.offload_to_disk_path, + offload_leader=module, + onload_leader=module, + parameters=None, + buffers=None, + non_blocking=False, + stream=None, + record_stream=False, + low_cpu_mem_usage=config.low_cpu_mem_usage, + onload_self=True, + group_id=_GROUP_ID_LAZY_LEAF, + ) + _apply_lazy_group_offloading_hook(module, unmatched_group, config=config) + + +def _apply_group_offloading_hook( + module: torch.nn.Module, + group: ModuleGroup, + *, + config: GroupOffloadingConfig, +) -> None: + registry = HookRegistry.check_if_exists_or_initialize(module) + + # We may have already registered a group offloading hook if the module had a torch.nn.Parameter whose parent + # is the current module. In such cases, we don't want to overwrite the existing group offloading hook. + if registry.get_hook(_GROUP_OFFLOADING) is None: + hook = GroupOffloadingHook(group, config=config) + registry.register_hook(hook, _GROUP_OFFLOADING) + + +def _apply_lazy_group_offloading_hook( + module: torch.nn.Module, + group: ModuleGroup, + *, + config: GroupOffloadingConfig, +) -> None: + registry = HookRegistry.check_if_exists_or_initialize(module) + + # We may have already registered a group offloading hook if the module had a torch.nn.Parameter whose parent + # is the current module. In such cases, we don't want to overwrite the existing group offloading hook. + if registry.get_hook(_GROUP_OFFLOADING) is None: + hook = GroupOffloadingHook(group, config=config) + registry.register_hook(hook, _GROUP_OFFLOADING) + + lazy_prefetch_hook = LazyPrefetchGroupOffloadingHook() + registry.register_hook(lazy_prefetch_hook, _LAZY_PREFETCH_GROUP_OFFLOADING) + + +def _gather_parameters_with_no_group_offloading_parent( + module: torch.nn.Module, modules_with_group_offloading: Set[str] +) -> List[torch.nn.Parameter]: + parameters = [] + for name, parameter in module.named_parameters(): + has_parent_with_group_offloading = False + atoms = name.split(".") + while len(atoms) > 0: + parent_name = ".".join(atoms) + if parent_name in modules_with_group_offloading: + has_parent_with_group_offloading = True + break + atoms.pop() + if not has_parent_with_group_offloading: + parameters.append((name, parameter)) + return parameters + + +def _gather_buffers_with_no_group_offloading_parent( + module: torch.nn.Module, modules_with_group_offloading: Set[str] +) -> List[torch.Tensor]: + buffers = [] + for name, buffer in module.named_buffers(): + has_parent_with_group_offloading = False + atoms = name.split(".") + while len(atoms) > 0: + parent_name = ".".join(atoms) + if parent_name in modules_with_group_offloading: + has_parent_with_group_offloading = True + break + atoms.pop() + if not has_parent_with_group_offloading: + buffers.append((name, buffer)) + return buffers + + +def _find_parent_module_in_module_dict(name: str, module_dict: Dict[str, torch.nn.Module]) -> str: + atoms = name.split(".") + while len(atoms) > 0: + parent_name = ".".join(atoms) + if parent_name in module_dict: + return parent_name + atoms.pop() + return "" + + +def _raise_error_if_accelerate_model_or_sequential_hook_present(module: torch.nn.Module) -> None: + if not is_accelerate_available(): + return + for name, submodule in module.named_modules(): + if not hasattr(submodule, "_hf_hook"): + continue + if isinstance(submodule._hf_hook, (AlignDevicesHook, CpuOffload)): + raise ValueError( + f"Cannot apply group offloading to a module that is already applying an alternative " + f"offloading strategy from Accelerate. If you want to apply group offloading, please " + f"disable the existing offloading strategy first. Offending module: {name} ({type(submodule)})" + ) + + +def _get_top_level_group_offload_hook(module: torch.nn.Module) -> Optional[GroupOffloadingHook]: + for submodule in module.modules(): + if hasattr(submodule, "_diffusers_hook"): + group_offloading_hook = submodule._diffusers_hook.get_hook(_GROUP_OFFLOADING) + if group_offloading_hook is not None: + return group_offloading_hook + return None + + +def _is_group_offload_enabled(module: torch.nn.Module) -> bool: + top_level_group_offload_hook = _get_top_level_group_offload_hook(module) + return top_level_group_offload_hook is not None + + +def _get_group_onload_device(module: torch.nn.Module) -> torch.device: + top_level_group_offload_hook = _get_top_level_group_offload_hook(module) + if top_level_group_offload_hook is not None: + return top_level_group_offload_hook.config.onload_device + raise ValueError("Group offloading is not enabled for the provided module.") + + +def _compute_group_hash(group_id): + hashed_id = hashlib.sha256(group_id.encode("utf-8")).hexdigest() + # first 16 characters for a reasonably short but unique name + return hashed_id[:16] + + +def _maybe_remove_and_reapply_group_offloading(module: torch.nn.Module) -> None: + r""" + Removes the group offloading hook from the module and re-applies it. This is useful when the module has been + modified in-place and the group offloading hook references-to-tensors needs to be updated. The in-place + modification can happen in a number of ways, for example, fusing QKV or unloading/loading LoRAs on-the-fly. + + In this implementation, we make an assumption that group offloading has only been applied at the top-level module, + and therefore all submodules have the same onload and offload devices. If this assumption is not true, say in the + case where user has applied group offloading at multiple levels, this function will not work as expected. + + There is some performance penalty associated with doing this when non-default streams are used, because we need to + retrace the execution order of the layers with `LazyPrefetchGroupOffloadingHook`. + """ + top_level_group_offload_hook = _get_top_level_group_offload_hook(module) + + if top_level_group_offload_hook is None: + return + + registry = HookRegistry.check_if_exists_or_initialize(module) + registry.remove_hook(_GROUP_OFFLOADING, recurse=True) + registry.remove_hook(_LAYER_EXECUTION_TRACKER, recurse=True) + registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=True) + + _apply_group_offloading(module, top_level_group_offload_hook.config) + + +def remove_group_offloading( + module: torch.nn.Module, + exclude_modules: Optional[Union[str, List[str]]] = None, +) -> None: + """ + Removes group offloading hooks from a module and its submodules. + + Args: + module (`torch.nn.Module`): + The module from which to remove group offloading hooks. + exclude_modules (`Union[str, List[str]]`, *optional*, defaults to `None`): + List of modules to exclude from hook removal. + """ + if isinstance(exclude_modules, str): + exclude_modules = [exclude_modules] + elif exclude_modules is None: + exclude_modules = [] + + # Check if this is a pipeline with components + if hasattr(module, 'components'): + unknown = set(exclude_modules) - module.components.keys() + if unknown: + logger.info( + f"The following modules are not present in pipeline: {', '.join(unknown)}. Ignore if this is expected." + ) + + # Remove hooks from each component + for name, component in module.components.items(): + if name not in exclude_modules and isinstance(component, torch.nn.Module): + registry = HookRegistry.check_if_exists_or_initialize(component) + registry.remove_hook(_GROUP_OFFLOADING, recurse=True) + registry.remove_hook(_LAYER_EXECUTION_TRACKER, recurse=True) + registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=True) + else: + # Original behavior for single modules + registry = HookRegistry.check_if_exists_or_initialize(module) + registry.remove_hook(_GROUP_OFFLOADING, recurse=True) + registry.remove_hook(_LAYER_EXECUTION_TRACKER, recurse=True) + registry.remove_hook(_LAZY_PREFETCH_GROUP_OFFLOADING, recurse=True) + + +def safe_remove_group_offloading(obj, *args, **kwargs): + """Safely call remove_group_offloading""" + return remove_group_offloading(obj, *args, **kwargs) + + +def enable_group_offload( + self, + onload_device: torch.device, + offload_device: torch.device = torch.device("cpu"), + offload_type: str = "block_level", + num_blocks_per_group: Optional[int] = None, + non_blocking: bool = False, + use_stream: bool = False, + record_stream: bool = False, + low_cpu_mem_usage=False, + offload_to_disk_path: Optional[str] = None, + exclude_modules: Optional[Union[str, List[str]]] = None, +) -> None: + r""" + Applies group offloading to the internal layers of a torch.nn.Module. To understand what group offloading is, + and where it is beneficial, we need to first provide some context on how other supported offloading methods + work. + + Typically, offloading is done at two levels: + - Module-level: In Diffusers, this can be enabled using the `ModelMixin::enable_model_cpu_offload()` method. It + works by offloading each component of a pipeline to the CPU for storage, and onloading to the accelerator + device when needed for computation. This method is more memory-efficient than keeping all components on the + accelerator, but the memory requirements are still quite high. For this method to work, one needs memory + equivalent to size of the model in runtime dtype + size of largest intermediate activation tensors to be able + to complete the forward pass. + - Leaf-level: In Diffusers, this can be enabled using the `ModelMixin::enable_sequential_cpu_offload()` method. + It + works by offloading the lowest leaf-level parameters of the computation graph to the CPU for storage, and + onloading only the leafs to the accelerator device for computation. This uses the lowest amount of accelerator + memory, but can be slower due to the excessive number of device synchronizations. + + Group offloading is a middle ground between the two methods. It works by offloading groups of internal layers, + (either `torch.nn.ModuleList` or `torch.nn.Sequential`). This method uses lower memory than module-level + offloading. It is also faster than leaf-level/sequential offloading, as the number of device synchronizations + is reduced. + + Another supported feature (for CUDA devices with support for asynchronous data transfer streams) is the ability + to overlap data transfer and computation to reduce the overall execution time compared to sequential + offloading. This is enabled using layer prefetching with streams, i.e., the layer that is to be executed next + starts onloading to the accelerator device while the current layer is being executed - this increases the + memory requirements slightly. Note that this implementation also supports leaf-level offloading but can be made + much faster when using streams. + + Args: + onload_device (`torch.device`): + The device to which the group of modules are onloaded. + offload_device (`torch.device`, defaults to `torch.device("cpu")`): + The device to which the group of modules are offloaded. This should typically be the CPU. Default is + CPU. + offload_type (`str` or `GroupOffloadingType`, defaults to "block_level"): + The type of offloading to be applied. Can be one of "block_level" or "leaf_level". Default is + "block_level". + offload_to_disk_path (`str`, *optional*, defaults to `None`): + The path to the directory where parameters will be offloaded. Setting this option can be useful in + limited RAM environment settings where a reasonable speed-memory trade-off is desired. + num_blocks_per_group (`int`, *optional*): + The number of blocks per group when using offload_type="block_level". This is required when using + offload_type="block_level". + non_blocking (`bool`, defaults to `False`): + If True, offloading and onloading is done with non-blocking data transfer. + use_stream (`bool`, defaults to `False`): + If True, offloading and onloading is done asynchronously using a CUDA stream. This can be useful for + overlapping computation and data transfer. + record_stream (`bool`, defaults to `False`): When enabled with `use_stream`, it marks the current tensor + as having been used by this stream. It is faster at the expense of slightly more memory usage. Refer to + the [PyTorch official docs](https://pytorch.org/docs/stable/generated/torch.Tensor.record_stream.html) + more details. + low_cpu_mem_usage (`bool`, defaults to `False`): + If True, the CPU memory usage is minimized by pinning tensors on-the-fly instead of pre-pinning them. + This option only matters when using streamed CPU offloading (i.e. `use_stream=True`). This can be + useful when the CPU memory is a bottleneck but may counteract the benefits of using streams. + exclude_modules (`Union[str, List[str]]`, defaults to `None`): List of modules to exclude from offloading. + + Example: + ```python + >>> from diffusers import DiffusionPipeline + >>> import torch + + >>> pipe = DiffusionPipeline.from_pretrained("Qwen/Qwen-Image", torch_dtype=torch.bfloat16) + + >>> pipe.enable_group_offload( + ... onload_device=torch.device("cuda"), + ... offload_device=torch.device("cpu"), + ... offload_type="leaf_level", + ... use_stream=True, + ... ) + >>> image = pipe("a beautiful sunset").images[0] + ``` + """ + if isinstance(exclude_modules, str): + exclude_modules = [exclude_modules] + elif exclude_modules is None: + exclude_modules = [] + + unknown = set(exclude_modules) - self.components.keys() + if unknown: + logger.info( + f"The following modules are not present in pipeline: {', '.join(unknown)}. Ignore if this is expected." + ) + + group_offload_kwargs = { + "onload_device": onload_device, + "offload_device": offload_device, + "offload_type": offload_type, + "num_blocks_per_group": num_blocks_per_group, + "non_blocking": non_blocking, + "use_stream": use_stream, + "record_stream": record_stream, + "low_cpu_mem_usage": low_cpu_mem_usage, + "offload_to_disk_path": offload_to_disk_path, + } + for name, component in self.components.items(): + if name not in exclude_modules and isinstance(component, torch.nn.Module): + apply_group_offloading(module=component, **group_offload_kwargs) + + if exclude_modules: + for module_name in exclude_modules: + module = getattr(self, module_name, None) + if module is not None and isinstance(module, torch.nn.Module): + module.to(onload_device) + logger.debug(f"Placed `{module_name}` on {onload_device} device as it was in `exclude_modules`.") + + +def safe_enable_group_offload(obj, *args, **kwargs): + """Safely call enable_group_offload, register default implementation if not exists""" + + if not hasattr(obj, 'enable_group_offload'): + obj.enable_group_offload = types.MethodType(enable_group_offload, obj) + + return obj.enable_group_offload(*args, **kwargs) + + +def register_auto_device_hook(model): + """ + Register forward pre-hooks for all modules to automatically transfer device + + Args: + model: The model to process + + Returns: + model: The model with registered hooks + """ + + def auto_device_hook(module, input: Tuple[Any, ...]): + """ + Forward pre-hook function to automatically transfer device before forward + + Args: + module: Current module + input: Forward input arguments (in tuple form) + """ + # Get the device of input tensor + input_device = None + + # Traverse input tuple to find the first tensor + for item in input: + if isinstance(item, torch.Tensor): + input_device = item.device + break + # Handle nested cases (like list, tuple, etc.) + elif isinstance(item, (list, tuple)): + for sub_item in item: + if isinstance(sub_item, torch.Tensor): + input_device = sub_item.device + break + if input_device is not None: + break + + # If no tensor input found, return directly + if input_device is None: + return + + # Get current device of the module + module_device = None + try: + # Try to get device from parameters + module_device = next(module.parameters()).device + except StopIteration: + # If no parameters, try to get from buffers + try: + module_device = next(module.buffers()).device + except StopIteration: + # No parameters or buffers, no need to transfer + return + + # Check if device transfer is needed + # Condition: module_device is not 'meta' and different from input_device + if module_device.type != 'meta' and module_device != input_device: + # print(f"Moving {module.__class__.__name__} from {module_device} to {input_device}") + module.to(input_device) + + # Register hooks for all submodules + hooks = [] + for module in model.modules(): + hook = module.register_forward_pre_hook(auto_device_hook) + hooks.append(hook) + + # Save hooks to model for later removal + model._auto_device_hooks = hooks + + return model + + +def remove_auto_device_hook(model): + """ + Remove previously registered auto device hooks + + Args: + model: The model to process + """ + if hasattr(model, '_auto_device_hooks'): + for hook in model._auto_device_hooks: + hook.remove() + delattr(model, '_auto_device_hooks') + print("Auto device hooks removed") \ No newline at end of file diff --git a/videox_fun/utils/utils.py b/videox_fun/utils/utils.py index 730d5af..d208458 100755 --- a/videox_fun/utils/utils.py +++ b/videox_fun/utils/utils.py @@ -1,5 +1,6 @@ import gc import inspect +import math import os import shutil import subprocess @@ -142,6 +143,15 @@ def merge_video_audio(video_path: str, audio_path: str): os.remove(temp_output) print(f"merge_video_audio failed with error: {e}") +def calculate_dimensions(target_area, ratio): + width = math.sqrt(target_area * ratio) + height = width / ratio + + width = round(width / 32) * 32 + height = round(height / 32) * 32 + + return width, height + def get_image_to_video_latent(validation_image_start, validation_image_end, video_length, sample_size): if validation_image_start is not None and validation_image_end is not None: if type(validation_image_start) is str and os.path.isfile(validation_image_start):