diff --git a/examples/qwenimage21_fun/predict_i2i_inpaint.py b/examples/qwenimage21_fun/predict_i2i_inpaint.py new file mode 100644 index 0000000..d9eb3cf --- /dev/null +++ b/examples/qwenimage21_fun/predict_i2i_inpaint.py @@ -0,0 +1,253 @@ +import os +import sys + +import torch + +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf +from PIL import Image + +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 (AutoencoderKLQwenImage21, + Qwen3VLForConditionalGeneration, + Qwen3VLProcessor, + QwenImage21ControlTransformer2DModel) +from videox_fun.pipeline import QwenImage21ControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) +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 + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, 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 +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 + +# Config path (control_layers / control_in_dim live here and must match the trained adapter) +config_path = "config/qwenimage21/qwenimage21_control.yaml" +# model path +model_name = "models/Diffusion_Transformer/Qwen-Image-2.1" + +# Choose the sampler. Qwen-Image 2.1 is a flow-matching model sampled with the Euler discrete scheduler. +sampler_name = "Flow" + +# Load pretrained model if need +transformer_path = "models/Personalized_Model/Qwen-Image-2.1-Fun-Controlnet-Union.safetensors" +vae_path = None +lora_path = None + +# Other params +sample_size = [1728, 992] +# Cache the text and condition-image keys/values after the first denoising step. Valid because the +# transformer modulates those tokens from t = 0, making their activations step-independent. +use_kv_cache = True + +# Use torch.float16 if GPU does not support torch.bfloat16 +# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +# Inpaint source image, and its mask. In the mask, WHITE (>= 0.5) marks the region to REGENERATE and BLACK +# marks the region to KEEP -- matching the mask convention used during training and the *_mask assets in asset/. +image_path = "asset/pose.jpg" +mask_path = "asset/mask.png" +# Optional edge/pose control map layered on top of the inpaint branch. None -> its 64 channels are zeroed, +# which is exactly the trained "inpaint only" regime. Set a path to run control + inpaint together. +control_image_path = None +# Strength of the control/union branch. 1.0 is the value the adapter is trained to consume. +control_context_scale = 1.0 + +# Describe what should appear inside the masked (white) region. +prompt = "A young woman with long straight black hair in an elegant three-quarter pose, wearing a white off-shoulder top with delicate lace trim, soft studio lighting against a dark blue-grey gradient background, high-fashion portrait photography, shallow depth of field." +negative_prompt = "低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。" +guidance_scale = 1.0 +seed = 43 +num_inference_steps = 40 +lora_weight = 1.0 +save_path = "samples/qwenimage21-inpaint-images" + +assert ring_degree == 1, ( + "Qwen-Image 2.1 only supports Ulysses (head-parallel) sequence parallelism; ring_degree must be 1, " + "because ring attention cannot express the block-causal mask or the prefix KV cache." +) +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) + +# Transformer +transformer = QwenImage21ControlTransformer2DModel.from_pretrained( + model_name, + subfolder="transformer", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), +).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, safe_open + 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)}") + +# Get Vae +vae = AutoencoderKLQwenImage21.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, safe_open + 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 processor and text_encoder. Qwen-Image 2.1 encodes the prompt (and any condition images) with a +# Qwen3-VL model, so a processor replaces the plain tokenizer used by the earlier Qwen-Image families. +processor = Qwen3VLProcessor.from_pretrained( + model_name, subfolder="processor" +) +text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( + model_name, subfolder="text_encoder", torch_dtype=weight_dtype +) + +# Get Scheduler +Chosen_Scheduler = { + "Flow": FlowMatchEulerDiscreteScheduler, +}[sampler_name] +scheduler = Chosen_Scheduler.from_pretrained( + model_name, + subfolder="scheduler" +) + +pipeline = QwenImage21ControlPipeline( + vae=vae, + text_encoder=text_encoder, + processor=processor, + transformer=transformer, + scheduler=scheduler, +) + +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, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.control_blocks)) + 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.model.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", "time_text_embed", "modulation"], 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", "time_text_embed", "modulation"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +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) + +# Load the source image and mask as PIL images; the pipeline's image_processor / mask_processor resize them to +# (height, width) and build control_context = [control_latents(64) | mask(1) | masked-image latents(64)] = 129 ch. +inpaint_image = Image.open(image_path).convert("RGB") +mask_image = Image.open(mask_path) +if control_image_path is not None: + control_image = Image.open(control_image_path).convert("RGB") +else: + control_image = None + +with torch.no_grad(): + sample = pipeline( + 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, + image = inpaint_image, + mask_image = mask_image, + control_image = control_image, + control_context_scale = control_context_scale, + use_kv_cache = use_kv_cache, + ).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) + # 2.1's VAE decodes to RGBA; JPEG cannot store an alpha channel, so every preview is saved as PNG. + 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() diff --git a/examples/qwenimage21_fun/predict_t2i_control.py b/examples/qwenimage21_fun/predict_t2i_control.py new file mode 100644 index 0000000..f272290 --- /dev/null +++ b/examples/qwenimage21_fun/predict_t2i_control.py @@ -0,0 +1,239 @@ +import os +import sys + +import torch + +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf + +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 (AutoencoderKLQwenImage21, + Qwen3VLForConditionalGeneration, + Qwen3VLProcessor, + QwenImage21ControlTransformer2DModel) +from videox_fun.pipeline import QwenImage21ControlPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) +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 get_image_latent + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, 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 +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 + +# Config path (control_layers / control_in_dim live here and must match the trained adapter) +config_path = "config/qwenimage21/qwenimage21_control.yaml" +# model path +model_name = "models/Diffusion_Transformer/Qwen-Image-2.1" + +# Choose the sampler. Qwen-Image 2.1 is a flow-matching model sampled with the Euler discrete scheduler. +sampler_name = "Flow" + +# Load pretrained model if need +transformer_path = "models/Personalized_Model/Qwen-Image-2.1-Fun-Controlnet-Union.safetensors" +vae_path = None +lora_path = None + +# Other params +sample_size = [1728, 992] +# Cache the text and condition-image keys/values after the first denoising step. Valid because the +# transformer modulates those tokens from t = 0, making their activations step-independent. +use_kv_cache = True + +# Use torch.float16 if GPU does not support torch.bfloat16 +# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_image = "asset/pose.jpg" +control_context_scale = 1.0 + +# Please use as detailed a prompt as possible to describe the object that needs to be generated. +prompt = "画面中央是一位年轻女孩,她拥有一头令人印象深刻的亮紫色长发,发丝在海风中轻盈飘扬,营造出动感而唯美的效果。她的长发两侧各扎着黑色蝴蝶结发饰,增添了几分可爱与俏皮感。女孩身穿一袭纯白色无袖连衣裙,裙摆轻盈飘逸,与她清新的气质完美契合。她的妆容精致自然,淡粉色的唇妆和温柔的眼神流露出恬静优雅的气质。她单手叉腰,姿态自信从容,目光直视镜头,展现出既甜美又不失个性的魅力。背景是一片开阔的海景,湛蓝的海水在阳光照射下波光粼粼,闪烁着钻石般的光芒。天空呈现出清澈的蔚蓝色,点缀着几朵洁白的云朵,营造出晴朗明媚的夏日氛围。画面前景右下角可见粉紫色的小花丛和绿色植物,为整体构图增添了自然生机和色彩层次。整张照片色调明亮清新,紫色头发与白色裙装、蓝色海天形成鲜明而和谐的色彩对比。" +negative_prompt = " " +guidance_scale = 1.0 +seed = 43 +num_inference_steps = 40 +lora_weight = 0.55 +save_path = "samples/qwenimage21-control-images" + +assert ring_degree == 1, ( + "Qwen-Image 2.1 only supports Ulysses (head-parallel) sequence parallelism; ring_degree must be 1, " + "because ring attention cannot express the block-causal mask or the prefix KV cache." +) +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) + +# Transformer +transformer = QwenImage21ControlTransformer2DModel.from_pretrained( + model_name, + subfolder="transformer", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), +).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, safe_open + 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)}") + +# Get Vae +vae = AutoencoderKLQwenImage21.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, safe_open + 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 processor and text_encoder. Qwen-Image 2.1 encodes the prompt (and any condition images) with a +# Qwen3-VL model, so a processor replaces the plain tokenizer used by the earlier Qwen-Image families. +processor = Qwen3VLProcessor.from_pretrained( + model_name, subfolder="processor" +) +text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( + model_name, subfolder="text_encoder", torch_dtype=weight_dtype +) + +# Get Scheduler +Chosen_Scheduler = { + "Flow": FlowMatchEulerDiscreteScheduler, +}[sampler_name] +scheduler = Chosen_Scheduler.from_pretrained( + model_name, + subfolder="scheduler" +) + +pipeline = QwenImage21ControlPipeline( + vae=vae, + text_encoder=text_encoder, + processor=processor, + transformer=transformer, + scheduler=scheduler, +) + +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, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.control_blocks)) + 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.model.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", "time_text_embed", "modulation"], 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", "time_text_embed", "modulation"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +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) + +# Load the control image as a single-frame (1, 3, h, w) tensor, matching scripts/qwenimage21_fun/train_control.py +# validation (get_image_latent(... )[:, :, 0]) so inference preprocessing is identical to training. +control_image = get_image_latent(control_image, sample_size=(sample_size[0], sample_size[1]))[:, :, 0] + +with torch.no_grad(): + sample = pipeline( + 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, + control_context_scale = control_context_scale, + use_kv_cache = use_kv_cache, + ).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) + # 2.1's VAE decodes to RGBA; JPEG cannot store an alpha channel, so every preview is saved as PNG. + image_path = os.path.join(save_path, prefix + f"-{control_context_scale}.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() diff --git a/examples/wan2.1_flex_forcing/predict_t2v.py b/examples/wan2.1_flex_forcing/predict_t2v.py index 50862f0..ea86636 100644 --- a/examples/wan2.1_flex_forcing/predict_t2v.py +++ b/examples/wan2.1_flex_forcing/predict_t2v.py @@ -78,7 +78,7 @@ shift = 5 # Any Wan2.1 / CausVid / Self-Forcing checkpoint loads as-is: the Flex-Forcing # backbone inherits every parameter name and only the new `flex_kproj.*` tensors # are reported missing (they are identity-initialised, so step 0 is unchanged). -transformer_path = "output_dir_wan2.1_flex_forcing_distill/checkpoint-1000/diffusion_pytorch_model.safetensors" +transformer_path = "output_dir_wan2.1_flex_forcing_distill/checkpoint-3000/diffusion_pytorch_model.safetensors" vae_path = None lora_path = None @@ -99,6 +99,12 @@ fps = 16 # so there is no second number to keep in sync. Levels only ever # *add* boundaries, so a KV cache written at a coarse level stays # valid at a finer one. +# "full_then_blocks" -> first denoising step runs the whole clip as one +# bidirectional ("full") chunk, every later step is the block-major +# Self-Forcing schedule over `num_frame_per_block`. A fixed 2-level +# ladder - coarser than the binary pyramid, no `min_num_frame_per_ +# block` involvement; needs `num_inference_steps >= 2` for the +# block-major steps to actually run. # An int instead pins a truncated pyramid of exactly that many levels; for 21 # latent frames (= 81 pixel frames) that ladder is # 2 -> [[21], [11, 10]] 3 -> [[21], [11, 10], [6, 5, 5, 5]] @@ -109,12 +115,6 @@ denoise_mode = "pyramid" # 3 -> leaves stay 3-frame blocks (classic Self-Forcing granularity); the # ladder then converges early and later steps reuse its finest level. min_num_frame_per_block = 1 -# Advanced: the 3.1 partition itself can also be pinned on the pipeline call -# (`chunk_spec = "18-3" / "ar" / "uniform:3"`); the pyramid above does not need -# it, since it derives every level from the whole-clip level 0. -# 3.3's K-Projection (the noise-level aligned Pi_{t<-0} of the cached clean keys) -# is deliberately not configurable here: the model builds `diag_rank1` and applies -# it on every call, so there is nothing left to set. # --- Causal backbone (inherited from Self-Forcing) ------------------------- # `num_frame_per_block` only takes effect once the pyramid is off; the rollout # derives the block size from the partition itself otherwise. `context_noise` diff --git a/scripts/qwenimage21/train.py b/scripts/qwenimage21/train.py index c5616dc..37ea929 100644 --- a/scripts/qwenimage21/train.py +++ b/scripts/qwenimage21/train.py @@ -1112,7 +1112,7 @@ def main(): 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] + fix_sample_size = [int(x / 32) * 32 for x in args.fix_sample_size] # 32 = vae_scale_factor(16)*2: keeps the latent grid even so h*w % 4 == 0 (joint stream expands image slots 4x) elif args.random_ratio_crop: if rng is None: random_sample_size = aspect_ratio_random_crop_sample_size[ @@ -1122,10 +1122,10 @@ def main(): 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] + random_sample_size = [int(x / 32) * 32 for x in random_sample_size] # 32 = vae_scale_factor(16)*2: keep latent dims even 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] + closest_size = [int(x / 32) * 32 for x in closest_size] # 32 = vae_scale_factor(16)*2: keep latent dims even for example in examples: if args.fix_sample_size is not None: diff --git a/scripts/qwenimage21_fun/README_TRAIN.md b/scripts/qwenimage21_fun/README_TRAIN.md new file mode 100644 index 0000000..edadaa4 --- /dev/null +++ b/scripts/qwenimage21_fun/README_TRAIN.md @@ -0,0 +1,548 @@ +# Qwen-Image 2.1 Control (ControlNet-Union) Training Guide + +This document provides a complete workflow for training a **ControlNet-Union** adapter on top of the frozen +**Qwen-Image 2.1** base transformer: environment setup, data preparation, distributed training, CFG +distillation, and inference testing. + +A parallel chain of zero-initialized `control_blocks` produces a per-layer skip (`hints`) that is added back into +the frozen base blocks, so the adapter starts as an identity skip and learns control gradually. Only the control +modules are trained (`--trainable_modules "control"`). + +The adapter is a **union** of control + inpaint: the conditioning tensor `control_context` packs +`[control_latents (64) | mask (1) | masked-image latents (64)] = 129` channels, so one adapter handles both +spatial control (depth / canny / pose / …) and inpainting. + +--- + +## Table of Contents +- [1. Environment Setup](#1-environment-setup) +- [2. Data Preparation](#2-data-preparation) + - [2.1 Quick Test Dataset](#21-quick-test-dataset) + - [2.2 Dataset Structure](#22-dataset-structure) + - [2.3 metadata.json Format](#23-metadatajson-format) + - [2.4 Relative vs Absolute Path Usage](#24-relative-vs-absolute-path-usage) +- [3. Control Training](#3-control-training) + - [3.1 Download Pre-trained Model](#31-download-pre-trained-model) + - [3.2 Quick Start (DeepSpeed-Zero-2)](#32-quick-start-deepspeed-zero-2) + - [3.3 Common Training Parameters](#33-common-training-parameters) + - [3.4 Training Validation](#34-training-validation) + - [3.5 Training with FSDP](#35-training-with-fsdp) + - [3.6 Other Backends](#36-other-backends) + - [3.7 Multi-machine Distributed Training](#37-multi-machine-distributed-training) + - [3.8 CFG Distillation](#38-cfg-distillation) +- [4. Inference Testing](#4-inference-testing) + - [4.1 Inference Parameters](#41-inference-parameters) + - [4.2 Single GPU Inference](#42-single-gpu-inference) + - [4.3 Multi-GPU Parallel Inference](#43-multi-gpu-parallel-inference) +- [5. Additional Resources](#5-additional-resources) + +--- + +## 1. Environment Setup + +**Option 1: Using requirements.txt** + +```bash +pip install -r requirements.txt +``` + +**Option 2: Manual Installation** + +```bash +pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image +pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime +pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2" +pip install yunchang xfuser modelscope openpyxl +pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y +pip install opencv-python-headless +pip install deepspeed==0.17.0 numpy==1.26.4 +``` + +**Option 3: Using Docker** + +When using Docker, please ensure that the GPU drivers and CUDA environment are correctly installed, then execute the following commands: + +``` +# pull image +docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun + +# enter image +docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun +``` + +> **Qwen-Image 2.1 specific**: the text encoder is a **Qwen3-VL** model, so the environment needs a `transformers` +> build that ships the `qwen3_vl` architecture (newer than the base pin in `requirements.txt`). If +> `Qwen3VLForConditionalGeneration` / `Qwen3VLProcessor` import as `None`, your `transformers` is too old. +> +> The YOLO object-mask feature (see [2.3](#23-metadatajson-format)) needs `ultralytics`; `yolov8x-seg.pt` +> downloads automatically on first use. + +--- + +## 2. Data Preparation + +Control training uses `ImageVideoControlDataset`. + +### 2.1 Quick Test Dataset + +We provide a test dataset containing several training samples with corresponding control files. + +```bash +# Download official example dataset +modelscope download --dataset PAI/X-Fun-Images-Controls-Demo --local_dir ./datasets/X-Fun-Images-Controls-Demo +``` + +### 2.2 Dataset Structure + +``` +📦 datasets/ +├── 📂 my_dataset/ +│ ├── 📂 train/ # target images (what the model should generate) +│ │ ├── 📄 image001.jpg +│ │ └── 📄 ... +│ ├── 📂 control/ # paired control / condition images (pose, canny, depth, ...) +│ │ ├── 📄 image001.png +│ │ └── 📄 ... +│ └── 📄 metadata.json +``` + +### 2.3 metadata.json Format + +The manifest is the standard image metadata JSON plus one extra `control_file_path` field that pairs each +**target** image with its **control** image. + +```json +[ + { + "file_path": "train/image001.jpg", + "control_file_path": "control/image001.png", + "text": "A young woman, studio lighting, high quality.", + "width": 1024, + "height": 1024, + "type": "image" + } +] +``` + +**Key field descriptions**: +- `file_path`: the **target** image (relative or absolute). +- `control_file_path`: the **control / condition** image (pose map, edge map, depth, gray sketch, …). It is loaded + as RGB and resized/cropped with the **same** transform as the target, so they stay pixel-aligned. +- `text`: caption. +- `width` / `height`: recommended for bucket training; use `scripts/process_json_add_width_and_height.py` to add + them to a JSON that lacks them. +- `type`: `"image"`. + +> **You only supply the target + control images. You do NOT supply masks.** The inpaint mask is generated on the +> fly: +> - A random rectangular hole via `get_random_mask` in the collate. +> - Then, on a random ~70% subset of frames, an **irregular object-shaped mask** produced by a **YOLO-seg** +> detector (`ObjectInstanceDetector`), with its edges randomly dilated / eroded / Gaussian-blurred. This mirrors +> `scripts/qwenimage_fun/train_control.py` and gives the model realistic, object-shaped inpaint holes instead of +> only rectangles. +> - The masked image fed to the union branch is always `target * (1 - mask)`. + +> **RGBA note**: the 2.1 VAE reads RGBA. Training images are loaded as RGB and automatically composited over an +> opaque alpha channel before encoding, so you do not need to provide RGBA data. + +### 2.4 Relative vs Absolute Path Usage + +**Relative paths** (small, local dataset): +```bash +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.json" +``` + +**Absolute paths** (NAS / OSS / multi-machine shared data): +```bash +export DATASET_NAME="" +export DATASET_META_NAME="/mnt/data/metadata.json" +``` + +> If the dataset is stored on external storage or shared across machines, prefer absolute paths. + +--- + +## 3. Control Training + +### 3.1 Download Pre-trained Model + +Point `MODEL_NAME` at a local **Qwen-Image 2.1** checkpoint directory. Its `transformer/` subfolder supplies the +frozen base weights; the control modules are zero-initialized on load. + +```bash +mkdir -p models/Diffusion_Transformer +# Place your Qwen-Image 2.1 weights here, e.g. +# models/Diffusion_Transformer/Qwen-Image-2.1/{transformer,vae,text_encoder,...} +``` + +> **No released 2.1 ControlNet-Union checkpoint.** Unlike Qwen-Image 2512, there is currently no published +> `...-Fun-Controlnet-Union.safetensors` for 2.1, so training starts from scratch with the **zero-initialized** +> control branch. Consequently the launcher leaves `--transformer_path` out; only add it to resume or fine-tune a +> control checkpoint you have already trained (or produced with `scripts/*/extract_control_weights.py`). + +### 3.2 Quick Start (DeepSpeed-Zero-2) + +It is recommended to use DeepSpeed-Zero-2 or FSDP for training, which can save a significant amount of GPU memory. + +After following **2.1 Quick Test Dataset** and **3.1 Download Pre-trained Model**, you can directly copy and run the following command: + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.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/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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=1024 \ + --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_21_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 "control" +``` + +### 3.3 Common Training Parameters + +| Parameter | Description | Example Value | +|-----------|-------------|---------------| +| `--config_path` | **Required.** Builds the control transformer with `control_layers` / `control_in_dim` | `config/qwenimage21/qwenimage21_control.yaml` | +| `--pretrained_model_name_or_path` | Base Qwen-Image 2.1 model (frozen weights) | `models/Diffusion_Transformer/Qwen-Image-2.1` | +| `--train_data_dir` / `--train_data_meta` | Dataset root / manifest JSON | `""` / `/path/metadata.json` | +| `--trainable_modules` | `"control"` trains only `control_blocks.*` + `control_img_in.*`; base stays frozen | `"control"` | +| `--transformer_path` | **Omit for from-scratch** training; add only to resume/finetune a trained control branch | *(none)* | +| `--image_sample_size` | Max training resolution, auto bucketing | `1024` | +| `--train_batch_size` / `--gradient_accumulation_steps` | Per-device batch / accumulation | `1` / `1` | +| `--learning_rate` | Initial learning rate | `2e-05` | +| `--lr_scheduler` / `--lr_warmup_steps` | Scheduler / warmup | `constant_with_warmup` / `100` | +| `--checkpointing_steps` | Save a checkpoint every N steps | `50` | +| `--gradient_checkpointing` | Activation recomputation | flag | +| `--vae_mini_batch` | Mini-batch size for VAE encoding (control encodes 3 streams) | `1` | +| `--max_grad_norm` | Gradient clipping | `0.05` | +| `--enable_bucket` | Bucket training by resolution without cropping | flag | +| `--uniform_sampling` | Uniform timestep sampling | flag | +| `--low_vram` | Offload VAE / text encoder when idle to save memory | flag (optional) | + +> **Memory**: each control step encodes **three** latent streams (target, control, masked image), so it is heavier +> than base training. Keep `--vae_mini_batch=1`; add `--low_vram` if you are tight on memory. + +### 3.4 Training Validation + +Configure validation during training to periodically render control previews: + +```bash + --validation_paths "asset/pose.jpg" \ + --validation_steps=50 \ + --validation_epochs=500 \ + --validation_prompts="1girl, black_hair, brown_eyes, ... solo, upper_body" +``` + +- `--validation_prompts` and `--validation_paths` must have **matching counts**; entry `i` of each pair is used + together. The output resolution is derived from each control image's aspect ratio via + `calculate_dimensions(image_sample_size^2, w/h)`. +- Validation triggers on either `--validation_steps` or `--validation_epochs`. +- Previews are written to `{output_dir}/sample/`. Because the 2.1 VAE decodes to **RGBA**, previews are saved as + **`.png`** (JPEG cannot store alpha). +- `log_validation` is wrapped in `try/except`: a bad control path only logs `Eval error on rank N` and never + crashes training. To confirm validation actually produced images, check `output_dir/sample/` **and** grep the log + for `Eval error`. + +### 3.5 Training with FSDP + +If DeepSpeed-Zero-2 runs out of GPU memory, you can switch to FSDP for training. The launcher `scripts/qwenimage21_fun/train_control.sh` runs +exactly the command below; edit the paths at the top (`MODEL_NAME`, `DATASET_META_NAME`, …) and run it with `bash` if you prefer +(the wrap classes must match the control model's blocks): + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.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=BaseQwenImage21TransformerBlock,QwenImage21ControlTransformerBlock \ + --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \ + --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \ + scripts/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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=1024 \ + --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_21_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 "control" +``` + +### 3.6 Other Backends + +#### 3.6.1 Training without DeepSpeed and FSDP + +Using neither DeepSpeed nor FSDP may result in insufficient GPU memory; only recommended when GPU memory is +sufficient. Plain DDP also has to replicate the 2.1 base transformer on every GPU, so it is generally not +recommended: + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.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/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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=1024 \ + --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_21_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 "control" +``` + +### 3.7 Multi-machine Distributed Training + +**Suitable for**: Ultra-large-scale datasets, faster training speed + +#### 3.7.1 Environment Configuration + +When using multi-machine training, please set the following environment variables: + +```bash +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/qwenimage21_fun/train_control.py \ + [other training parameters...] +``` + +#### 3.7.2 Multi-machine Training Considerations + +- **Network Requirements**: + - Recommended: RDMA/InfiniBand (high performance) + - Without RDMA, add environment variables: + ```bash + export NCCL_IB_DISABLE=1 + export NCCL_P2P_DISABLE=1 + ``` + +- **Data Synchronization**: All machines must be able to access the same data paths (NFS/shared storage) + +### 3.8 CFG Distillation + +`train_control_distill.py` / `train_control_distill.sh` are an **optional second stage**. They distill +classifier-free guidance (CFG) into the trained control branch, so that **inference needs no guidance scale**. + +Algorithm (identical in spirit to `scripts/minimax_h3_fun` / `scripts/flux2_fun` control distillation): +- A **frozen teacher** (a second copy of the same control model, loaded with the **same** `--transformer_path` — + your stage-1 trained control branch) runs two forward passes per step, on the prompt and on an **empty** + negative prompt, both **with the control condition**. The two velocities combine into the CFG target: + `target = uncond + (cond - uncond) * real_guidance_scale`. +- The trainable **student** (control branch) runs a single conditional forward and regresses onto that target + (MSE in velocity space). Only `--trainable_modules "control"` is trained. + +Run it after you have a trained control checkpoint: + +```bash +# Set CONTROL_TRANSFORMER_PATH to the stage-1 checkpoint's +# output_dir_qwen_image_21_control//checkpoint-/diffusion_pytorch_model.safetensors +bash scripts/qwenimage21_fun/train_control_distill.sh +``` + +Distillation-specific parameters: + +| Parameter | Description | Example Value | +|-----------|-------------|---------------| +| `--transformer_path` | **Required**: the trained control branch both student and teacher load | `/root/diffusion_pytorch_model.safetensors` | +| `--real_guidance_scale` | CFG scale applied to the teacher to build the target | `3.5` | +| `--learning_rate` | Lower LR for distillation | `2e-06` | +| `--output_dir` | Separate output dir for the distilled adapter | `output_dir_qwen_image_21_control_distill` | + +> The teacher is a separate, unsharded bf16 copy on each GPU (only the student is FSDP-sharded), which is +> memory-heavy. Add `--low_vram` to stream the teacher only for its two forward passes. A distilled student is run +> at inference with `guidance_scale = 1.0` (CFG is already baked into the weights). + +--- + +## 4. Inference Testing + +### 4.1 Inference Parameters + +| Parameter | Description | Example Value | +|-----------|-------------|---------------| +| `config_path` | Must match the trained adapter's config | `config/qwenimage21/qwenimage21_control.yaml` | +| `model_name` | Base Qwen-Image 2.1 path | `models/Diffusion_Transformer/Qwen-Image-2.1` | +| `transformer_path` | Trained control weights (`control_*` keys load with `strict=False`), or `None` for the base | `output_dir_qwen_image_21_control/.../diffusion_pytorch_model.safetensors` | +| `sampler_name` | Flow-matching sampler | `Flow` | +| `sample_size` | Output canvas `[height, width]` | `[1728, 992]` | +| `control_image` | Control condition image (`predict_t2i_control.py`) | `asset/pose.jpg` | +| `control_image_path` | Optional control condition image (`predict_i2i_inpaint.py`, defaults to `None`) | `asset/pose.jpg` | +| `control_context_scale` | Control-branch strength (the value the adapter was trained to consume) | `1.0` | +| `image_path` / `mask_path` | Inpaint source image / mask (only `predict_i2i_inpaint.py`, see 4.2) | `asset/pose.jpg` / `asset/mask.png` | +| `guidance_scale` | CFG strength. `1.0` for a CFG-distilled checkpoint | `1.0` | +| `weight_dtype` | Use `torch.float16` on GPUs without bf16 (v100, 2080Ti, …) | `torch.bfloat16` | +| `GPU_memory_mode` | GPU memory management mode, see table below | `model_group_offload` | +| `ulysses_degree` / `ring_degree` | Multi-GPU parallelism (see 4.3). `ring_degree` must stay `1` | `1` / `1` | +| `num_inference_steps` / `seed` | Sampling steps / seed | `40` / `43` | +| `save_path` | Output directory | `samples/qwenimage21-control-images` | + +**GPU Memory Management Modes**: + +| Mode | Description | Memory Usage | +|------|------|---------| +| `model_full_load` | Load entire model to GPU | Highest | +| `model_full_load_and_qfloat8` | Full load + FP8 quantization | High | +| `model_cpu_offload` | Offload model to CPU after use | Medium | +| `model_cpu_offload_and_qfloat8` | CPU offload + FP8 quantization | Medium-Low | +| `model_group_offload` | Switch layer groups between CPU/CUDA | Low | +| `sequential_cpu_offload` | Layer-by-layer offload (slowest) | Lowest | + +### 4.2 Single GPU Inference + +#### Quick Start + +```bash +python examples/qwenimage21_fun/predict_t2i_control.py +``` + +Edit the top-of-file constants to match your setup. The pipeline preprocesses and VAE-encodes `control_image` +(accepts a PIL image / a path), builds the 129-channel `control_context`, and injects the control skips: + +```python +GPU_memory_mode = "model_group_offload" +model_name = "models/Diffusion_Transformer/Qwen-Image-2.1" +transformer_path = "models/Personalized_Model/Qwen-Image-2.1-Fun-Controlnet-Union.safetensors" # or your trained checkpoint's diffusion_pytorch_model.safetensors +control_image = "asset/pose.jpg" +control_context_scale = 1.0 +prompt = "A young woman with long straight black hair ..." +sample_size = [1728, 992] +num_inference_steps = 40 +``` + +Results are saved to `samples/qwenimage21-control-images/*.png`. + +> The KV cache is **disabled automatically** when `control_context` is present: the control skip depends on the +> per-step base joint stream, so a prefix cache would be invalid. + +**Image Inpainting Inference**: + +The union adapter also does inpainting. `predict_i2i_inpaint.py` feeds `image_path` + `mask_path` (and leaves +`control_image_path` as `None`), so only the inpaint half of the 129-channel context is used: + +```bash +python examples/qwenimage21_fun/predict_i2i_inpaint.py +``` + +`mask_path` semantics: **white (`>= 0.5`) = repaint**, black = keep. You can supply a control image *and* the +inpaint pair together to use the full union. + +### 4.3 Multi-GPU Parallel Inference + +**Suitable for**: High-resolution generation, faster inference + +Qwen-Image 2.1 supports **Ulysses (head-parallel) sequence parallelism only**. + +#### Install Parallel Inference Dependencies + +```bash +pip install xfuser==0.4.2 yunchang==0.6.2 +``` + +#### Configure Parallel Strategy + +Edit `examples/qwenimage21_fun/predict_t2i_control.py`: + +```python +# Ensure ulysses_degree × ring_degree = number of GPUs +# For example, using 4 GPUs: +ulysses_degree = 4 # Head dimension parallelism +ring_degree = 1 # Sequence dimension parallelism, must stay 1 +``` + +**Configuration Principles**: +- `ulysses_degree` must divide `num_attention_heads` (32): one of `1/2/4/8/16/32`. +- `ring_degree` **must stay 1** — ring attention rotates KV chunks and cannot express 2.1's block-causal mask or + its prefix KV cache. + +**Example Configurations**: + +| GPU count | ulysses_degree | ring_degree | +|-----------|----------------|-------------| +| 1 | 1 | 1 | +| 4 | 4 | 1 | +| 8 | 8 | 1 | + +#### Run Multi-GPU Inference + +```bash +# Set ulysses_degree > 1, keep ring_degree = 1, and GPU count = ulysses_degree * ring_degree. +torchrun --nproc-per-node=4 examples/qwenimage21_fun/predict_t2i_control.py +``` + +--- + +## 5. Additional Resources + +- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun +- **Qwen-Image Official Repository**: https://github.com/QwenLM/Qwen-Image diff --git a/scripts/qwenimage21_fun/README_TRAIN_zh-CN.md b/scripts/qwenimage21_fun/README_TRAIN_zh-CN.md new file mode 100644 index 0000000..22a2e5f --- /dev/null +++ b/scripts/qwenimage21_fun/README_TRAIN_zh-CN.md @@ -0,0 +1,531 @@ +# Qwen-Image 2.1 Control(ControlNet-Union)训练指南 + +本文档提供在冻结的 **Qwen-Image 2.1** 基座 transformer 之上训练 **ControlNet-Union** 适配器的完整流程,包括环境配置、 +数据准备、分布式训练、CFG 蒸馏与推理测试。 + +一整套零初始化的 `control_blocks` 并行链会产出逐层残差(`hints`),再加回冻结的基座 block,因此适配器初始时等价于恒等 +跳连,并逐步学习控制信号。训练时只更新控制模块(`--trainable_modules "control"`)。 + +该适配器是 control + inpaint 的 **union**:条件张量 `control_context` 打包了 +`[control_latents(64) | mask(1) | masked-image latents(64)] = 129` 个通道,因此单个适配器同时处理空间控制 +(depth / canny / pose / …)与图像修补(inpainting)。 + +--- + +## 目录 +- [一、环境配置](#一环境配置) +- [二、数据准备](#二数据准备) + - [2.1 快速测试数据集](#21-快速测试数据集) + - [2.2 数据集结构](#22-数据集结构) + - [2.3 metadata.json 格式](#23-metadatajson-格式) + - [2.4 相对路径与绝对路径使用方案](#24-相对路径与绝对路径使用方案) +- [三、Control 训练](#三control-训练) + - [3.1 下载预训练模型](#31-下载预训练模型) + - [3.2 快速开始(DeepSpeed-Zero-2)](#32-快速开始deepspeed-zero-2) + - [3.3 训练常用参数解析](#33-训练常用参数解析) + - [3.4 训练验证](#34-训练验证) + - [3.5 使用 FSDP 训练](#35-使用-fsdp-训练) + - [3.6 其他后端](#36-其他后端) + - [3.7 多机分布式训练](#37-多机分布式训练) + - [3.8 CFG 蒸馏](#38-cfg-蒸馏) +- [四、推理测试](#四推理测试) + - [4.1 推理参数解析](#41-推理参数解析) + - [4.2 单卡推理](#42-单卡推理) + - [4.3 多卡并行推理](#43-多卡并行推理) +- [五、更多资源](#五更多资源) + +--- + +## 一、环境配置 + +**方式 1:使用requirements.txt** + +```bash +pip install -r requirements.txt +``` + +**方式 2:手动安装依赖** + +```bash +pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image +pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime +pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2" +pip install yunchang xfuser modelscope openpyxl +pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y +pip install opencv-python-headless +pip install deepspeed==0.17.0 numpy==1.26.4 +``` + +**方式 3:使用docker** + +使用docker的情况下,请保证机器中已经正确安装显卡驱动与CUDA环境,然后依次执行以下命令: + +``` +# pull image +docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun + +# enter image +docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun +``` + +> **Qwen-Image 2.1 特有**:文本编码器是 **Qwen3-VL** 模型,因此环境需要一个包含 `qwen3_vl` 结构的 `transformers` +> 版本(比 `requirements.txt` 里的基线更新)。如果 `Qwen3VLForConditionalGeneration` / `Qwen3VLProcessor` 导入为 +> `None`,说明你的 `transformers` 太旧。 +> +> YOLO 目标掩膜功能(见 [2.3](#23-metadatajson-格式))需要 `ultralytics`;`yolov8x-seg.pt` 会在首次使用时自动下载。 + +--- + +## 二、数据准备 + +Control 训练使用 `ImageVideoControlDataset`。 + +### 2.1 快速测试数据集 + +我们提供了一个测试的数据集,其中包含若干训练数据以及对应的控制文件。 + +```bash +# 下载官方示例数据集 +modelscope download --dataset PAI/X-Fun-Images-Controls-Demo --local_dir ./datasets/X-Fun-Images-Controls-Demo +``` + +### 2.2 数据集结构 + +``` +📦 datasets/ +├── 📂 my_dataset/ +│ ├── 📂 train/ # 目标图(模型应生成的内容) +│ │ ├── 📄 image001.jpg +│ │ └── 📄 ... +│ ├── 📂 control/ # 配对的 control / 条件图(pose、canny、depth…) +│ │ ├── 📄 image001.png +│ │ └── 📄 ... +│ └── 📄 metadata.json +``` + +### 2.3 metadata.json 格式 + +清单文件是标准图像 metadata JSON,外加一个 `control_file_path` 字段,把每张**目标图**与其 **control 图**配对。 + +```json +[ + { + "file_path": "train/image001.jpg", + "control_file_path": "control/image001.png", + "text": "A young woman, studio lighting, high quality.", + "width": 1024, + "height": 1024, + "type": "image" + } +] +``` + +**关键字段说明**: +- `file_path`:**目标图**(相对或绝对路径)。 +- `control_file_path`:**control / 条件图**(姿态图、边缘图、深度图、线稿…)。它以 RGB 载入,并与目标图使用**完全相同**的 + 变换做 resize / crop,从而保证像素对齐。 +- `text`:描述(caption)。 +- `width` / `height`:推荐提供,用于 bucket 训练;可用 `scripts/process_json_add_width_and_height.py` 为缺失字段的 + JSON 补上。 +- `type`:图像数据为 `"image"`。 + +> **你只需提供目标图 + control 图,不需要提供掩膜。** inpaint 掩膜是即时生成的: +> - 先在 collate 中用 `get_random_mask` 生成随机矩形遮挡。 +> - 然后在随机约 70% 的帧上,再由 **YOLO-seg** 检测器(`ObjectInstanceDetector`)生成**不规则的目标形状掩膜**, +> 其边缘随机做膨胀 / 腐蚀 / 高斯模糊。这与 `scripts/qwenimage_fun/train_control.py` 一致,能让模型见到真实、 +> 目标形状的修补空洞,而不只是矩形。 +> - 送进 union 分支的被遮罩图始终是 `target * (1 - mask)`。 + +> **RGBA 说明**:2.1 VAE 读取 RGBA。训练图以 RGB 载入,在编码前会自动合成到不透明 alpha 通道,因此你不需要提供 RGBA 数据。 + +### 2.4 相对路径与绝对路径使用方案 + +**相对路径**(小型本地数据集): +```bash +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.json" +``` + +**绝对路径**(NAS / OSS / 多机共享数据): +```bash +export DATASET_NAME="" +export DATASET_META_NAME="/mnt/data/metadata.json" +``` + +> 如果数据集存放在外部存储或被多台机器共享,推荐使用绝对路径。 + +--- + +## 三、Control 训练 + +### 3.1 下载预训练模型 + +将 `MODEL_NAME` 指向本地的 **Qwen-Image 2.1** checkpoint 目录。其 `transformer/` 子目录提供冻结的基座权重;控制模块在 +载入时零初始化。 + +```bash +mkdir -p models/Diffusion_Transformer +# 将你的 Qwen-Image 2.1 权重放到这里,例如 +# models/Diffusion_Transformer/Qwen-Image-2.1/{transformer,vae,text_encoder,...} +``` + +> **没有公开的 2.1 ControlNet-Union checkpoint。** 与 Qwen-Image 2512 不同,2.1 目前并没有发布的 +> `...-Fun-Controlnet-Union.safetensors`,因此训练从零开始,control 分支**零初始化**。所以启动脚本不写 `--transformer_path`; +> 只有在你需要 resume / fine-tune 一个已训练好的 control checkpoint(或你自己用 `scripts/*/extract_control_weights.py` 得到的) +> 时才加上它。 + +### 3.2 快速开始(DeepSpeed-Zero-2) + +推荐使用 DeepSpeed-Zero-2 或 FSDP 方案进行训练,可以节省大量显存。 + +如果按照 **2.1 快速测试数据集**下载数据与 **3.1 下载预训练模型**放置权重后,直接复制以下启动指令进行启动。 + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.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/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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=1024 \ + --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_21_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 "control" +``` + +### 3.3 训练常用参数解析 + +| 参数 | 说明 | 示例值 | +|------|------|--------| +| `--config_path` | **必填。** 用 `control_layers` / `control_in_dim` 构建 control transformer | `config/qwenimage21/qwenimage21_control.yaml` | +| `--pretrained_model_name_or_path` | 基座 Qwen-Image 2.1 模型(冻结权重) | `models/Diffusion_Transformer/Qwen-Image-2.1` | +| `--train_data_dir` / `--train_data_meta` | 数据根目录 / 清单 JSON | `""` / `/path/metadata.json` | +| `--trainable_modules` | `"control"` 只训练 `control_blocks.*` + `control_img_in.*`,基座冻结 | `"control"` | +| `--transformer_path` | **从零训练时省略**;仅在 resume / fine-tune 已训练的 control 分支时加上 | *(无)* | +| `--image_sample_size` | 最大训练分辨率,自动 bucket | `1024` | +| `--train_batch_size` / `--gradient_accumulation_steps` | 单卡 batch / 梯度累积 | `1` / `1` | +| `--learning_rate` | 初始学习率 | `2e-05` | +| `--lr_scheduler` / `--lr_warmup_steps` | 学习率调度 / 预热 | `constant_with_warmup` / `100` | +| `--checkpointing_steps` | 每 N 步保存 checkpoint | `50` | +| `--gradient_checkpointing` | 激活重计算 | flag | +| `--vae_mini_batch` | VAE 编码 mini-batch(control 要编码 3 路 latent) | `1` | +| `--max_grad_norm` | 梯度裁剪 | `0.05` | +| `--enable_bucket` | 按分辨率分组、不裁剪的 bucket 训练 | flag | +| `--uniform_sampling` | 均匀 timestep 采样 | flag | +| `--low_vram` | 空闲时卸载 VAE / 文本编码器以省显存 | flag(可选) | + +> **显存**:每个 control step 要编码**三路** latent(目标图、control 图、被遮罩图),比普通训练更重。请保持 +> `--vae_mini_batch=1`;显存紧张时再加 `--low_vram`。 + +### 3.4 训练验证 + +在训练时配置验证参数,定期渲染 control 预览: + +```bash + --validation_paths "asset/pose.jpg" \ + --validation_steps=50 \ + --validation_epochs=500 \ + --validation_prompts="1girl, black_hair, brown_eyes, ... solo, upper_body" +``` + +- `--validation_prompts` 与 `--validation_paths` 的数量必须**匹配**,每对的第 i 项一起使用。输出分辨率由每张 control 图的 + 宽高比经 `calculate_dimensions(image_sample_size^2, w/h)` 推出。 +- 验证在 `--validation_steps` 或 `--validation_epochs` 任一满足时触发。 +- 预览写入 `{output_dir}/sample/`。由于 2.1 VAE 解码为 **RGBA**,预览保存为 **`.png`**(JPEG 无法存 alpha)。 +- `log_validation` 包在 `try/except` 里:错误的 control 路径只会打印 `Eval error on rank N`,不会让训练崩溃。要确认验证 + 是否真的产出了图,请查看 `output_dir/sample/` **并**在日志里 grep `Eval error`。 + +### 3.5 使用 FSDP 训练 + +如果 DeepSpeed-Zero-2 显存不足,可以切换使用 FSDP 进行训练。配套的启动脚本 `scripts/qwenimage21_fun/train_control.sh` 运行的就是 +下面这条命令,修改其顶部路径(`MODEL_NAME`、`DATASET_META_NAME` 等)后可直接 `bash` 运行(wrap 类必须与 control 模型的 block 同名): + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.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=BaseQwenImage21TransformerBlock,QwenImage21ControlTransformerBlock \ + --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \ + --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \ + scripts/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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=1024 \ + --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_21_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 "control" +``` + +### 3.6 其他后端 + +#### 3.6.1 不使用 DeepSpeed 与 FSDP 训练 + +不使用 DeepSpeed 或 FSDP 可能会导致显存不足,仅建议在显存充足的情况下使用;普通 DDP 还需要在每张卡上完整复制 2.1 基座 +transformer,通常并不推荐: + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +export DATASET_NAME="datasets/X-Fun-Images-Controls-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Controls-Demo/metadata_add_width_height.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/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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=1024 \ + --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_21_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 "control" +``` + +### 3.7 多机分布式训练 + +**适合场景**:超大规模数据集、需要更快的训练速度 + +#### 3.7.1 环境配置 + +当使用多机训练时,请设置以下环境变量: + +```bash +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/qwenimage21_fun/train_control.py \ + [其他训练参数...] +``` + +#### 3.7.2 多机训练注意事项 + +- **网络要求**: + - 推荐 RDMA/InfiniBand(高性能) + - 无 RDMA 时添加环境变量: + ```bash + export NCCL_IB_DISABLE=1 + export NCCL_P2P_DISABLE=1 + ``` + +- **数据同步**:所有机器必须能够访问相同的数据路径(NFS/共享存储) + +### 3.8 CFG 蒸馏 + +`train_control_distill.py` / `train_control_distill.sh` 是一个**可选的第二阶段**。它把 classifier-free guidance(CFG) +蒸馏进已训练好的 control 分支,使**推理阶段无需 guidance scale**。 + +算法(思路与 `scripts/minimax_h3_fun` / `scripts/flux2_fun` 的 control 蒸馏完全一致): +- **冻结的 teacher**(同一 control 模型的另一份拷贝,用**相同**的 `--transformer_path` 载入,即你第一阶段训练好的 control 分支) + 每步做两次前向,分别在 prompt 与**空**负 prompt 上,两者都**带 control 条件**。两个速度合成 CFG 目标: + `target = uncond + (cond - uncond) * real_guidance_scale`。 +- 可训练的 **student**(control 分支)只跑一次带条件的前向,向该目标回归(速度空间 MSE)。同样只训练 + `--trainable_modules "control"`。 + +在你已有一个训练好的 control checkpoint 后运行: + +```bash +# 将 CONTROL_TRANSFORMER_PATH 指向第一阶段 checkpoint 的 +# output_dir_qwen_image_21_control//checkpoint-/diffusion_pytorch_model.safetensors +bash scripts/qwenimage21_fun/train_control_distill.sh +``` + +蒸馏专有参数: + +| 参数 | 说明 | 示例值 | +|------|------|--------| +| `--transformer_path` | **必填**:student 与 teacher 都载入的、已训练好的 control 分支 | `/root/diffusion_pytorch_model.safetensors` | +| `--real_guidance_scale` | 作用于 teacher 以合成目标的 CFG scale | `3.5` | +| `--learning_rate` | 蒸馏使用更低的学习率 | `2e-06` | +| `--output_dir` | 蒸馏适配器单独的输出目录 | `output_dir_qwen_image_21_control_distill` | + +> teacher 是每卡上一份独立的、不分片的 bf16 拷贝(只有 student 被 FSDP 分片),因此很吃显存。可加 `--low_vram` 让 teacher +> 只在其两次前向时上卡。蒸馏后的 student 推理时用 `guidance_scale = 1.0`(CFG 已烘焙进权重)。 + +--- + +## 四、推理测试 + +### 4.1 推理参数解析 + +| 参数 | 说明 | 示例值 | +|------|------|--------| +| `config_path` | 必须与训练好的适配器 config 一致 | `config/qwenimage21/qwenimage21_control.yaml` | +| `model_name` | 基座 Qwen-Image 2.1 路径 | `models/Diffusion_Transformer/Qwen-Image-2.1` | +| `transformer_path` | 训练好的 control 权重(`control_*` 键以 `strict=False` 载入),基线可为 `None` | `output_dir_qwen_image_21_control/.../diffusion_pytorch_model.safetensors` | +| `sampler_name` | flow-matching 采样器 | `Flow` | +| `sample_size` | 输出画布 `[height, width]` | `[1728, 992]` | +| `control_image` | control 条件图(`predict_t2i_control.py`) | `asset/pose.jpg` | +| `control_image_path` | 可选 control 条件图(`predict_i2i_inpaint.py`,默认 `None`) | `asset/pose.jpg` | +| `control_context_scale` | control 分支强度(适配器训练时消费的取值) | `1.0` | +| `image_path` / `mask_path` | inpaint 输入图 / 掩码(仅 `predict_i2i_inpaint.py`,见 4.2) | `asset/pose.jpg` / `asset/mask.png` | +| `guidance_scale` | CFG 强度。CFG 蒸馏后的 checkpoint 用 `1.0` | `1.0` | +| `weight_dtype` | 不支持 bf16 的卡(v100、2080Ti…)用 `torch.float16` | `torch.bfloat16` | +| `GPU_memory_mode` | 显存管理模式,可选值见下表 | `model_group_offload` | +| `ulysses_degree` / `ring_degree` | 多卡并行(见 4.3)。`ring_degree` 必须为 `1` | `1` / `1` | +| `num_inference_steps` / `seed` | 采样步数 / 随机种子 | `40` / `43` | +| `save_path` | 输出目录 | `samples/qwenimage21-control-images` | + +**显存管理模式说明**: + +| 模式 | 说明 | 显存占用 | +|------|------|---------| +| `model_full_load` | 整个模型加载到 GPU | 最高 | +| `model_full_load_and_qfloat8` | 全量加载 + FP8 量化 | 高 | +| `model_cpu_offload` | 使用后将模型卸载到 CPU | 中等 | +| `model_cpu_offload_and_qfloat8` | CPU 卸载 + FP8 量化 | 中低 | +| `model_group_offload` | 层组在 CPU/CUDA 间切换 | 低 | +| `sequential_cpu_offload` | 逐层卸载(速度最慢) | 最低 | + +### 4.2 单卡推理 + +#### 快速开始 + +```bash +python examples/qwenimage21_fun/predict_t2i_control.py +``` + +修改文件顶部常量以匹配你的环境。pipeline 会预处理并 VAE 编码 `control_image`(接受 PIL 图 / 路径),构建 129 通道的 +`control_context`,并注入 control 残差: + +```python +GPU_memory_mode = "model_group_offload" +model_name = "models/Diffusion_Transformer/Qwen-Image-2.1" +transformer_path = "models/Personalized_Model/Qwen-Image-2.1-Fun-Controlnet-Union.safetensors" # 或训练输出的 diffusion_pytorch_model.safetensors +control_image = "asset/pose.jpg" +control_context_scale = 1.0 +prompt = "A young woman with long straight black hair ..." +sample_size = [1728, 992] +num_inference_steps = 40 +``` + +结果保存到 `samples/qwenimage21-control-images/*.png`。 + +> 当存在 `control_context` 时,KV cache 会**自动禁用**:control 残差依赖每步的基座 joint 流,前缀缓存会失效。 + +**图像修补推理**: + +union 适配器也能做 inpainting。`predict_i2i_inpaint.py` 传入 `image_path` + `mask_path`(并把 `control_image_path` 留为 `None`), +于是只用 129 通道条件的 inpaint 半边: + +```bash +python examples/qwenimage21_fun/predict_i2i_inpaint.py +``` + +`mask_path` 语义:**白色(`>= 0.5`)= 重绘**,黑色 = 保留。你可以同时提供 control 图与 inpaint 对,以使用完整 union。 + +### 4.3 多卡并行推理 + +**适合场景**:高分辨率生成、加速推理 + +Qwen-Image 2.1 **仅支持 Ulysses(head 并行)序列并行**。 + +#### 安装并行推理依赖 + +```bash +pip install xfuser==0.4.2 yunchang==0.6.2 +``` + +#### 配置并行策略 + +编辑 `examples/qwenimage21_fun/predict_t2i_control.py`: + +```python +# 确保 ulysses_degree × ring_degree = GPU 数量 +# 例如使用 4 张 GPU: +ulysses_degree = 4 # Head 维度并行 +ring_degree = 1 # Sequence 维度并行,必须保持为 1 +``` + +**配置原则**: +- `ulysses_degree` 必须整除 `num_attention_heads`(32):取 `1/2/4/8/16/32`。 +- `ring_degree` **必须为 1** —— ring attention 会轮转 KV chunk,无法表达 2.1 的 block-causal mask 或其前缀 KV cache。 + +**示例配置**: + +| GPU 数 | ulysses_degree | ring_degree | +|--------|----------------|-------------| +| 1 | 1 | 1 | +| 4 | 4 | 1 | +| 8 | 8 | 1 | + +#### 运行多卡推理 + +```bash +# 设 ulysses_degree > 1,保持 ring_degree = 1,且 GPU 数 = ulysses_degree * ring_degree。 +torchrun --nproc-per-node=4 examples/qwenimage21_fun/predict_t2i_control.py +``` + +--- + +## 五、更多资源 + +- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun +- **Qwen-Image 官方仓库**:https://github.com/QwenLM/Qwen-Image diff --git a/scripts/qwenimage21_fun/extract_control_weights.py b/scripts/qwenimage21_fun/extract_control_weights.py new file mode 100644 index 0000000..88adaeb --- /dev/null +++ b/scripts/qwenimage21_fun/extract_control_weights.py @@ -0,0 +1,119 @@ +# Extract the control branch of a trained QwenImage21ControlTransformer2DModel checkpoint. +# +# `train_control.py` / `train_control_distill.py` save the whole transformer (frozen base branch + trainable control +# branch) in the diffusers layout; with FSDP the gathered state dict is written as +# `/diffusion_pytorch_model.safetensors` (no `transformer/` subdir, no `_control` suffix). The control +# branch is everything that `QwenImage21ControlTransformer2DModel` adds on top of the base model: the +# `control_blocks.*` list (one block per `control_layers` entry) and the `control_img_in.*` input projection. This +# script writes just those tensors to a standalone safetensors file, which can be re-applied onto a fresh base model +# built from `config/qwenimage21/qwenimage21_control.yaml` with `transformer.load_state_dict(..., strict=False)` +# (only the `control_*` keys are consumed; the base keys already match the freshly-loaded base weights). +# +# Usage: +# python scripts/qwenimage21_fun/extract_control_weights.py \ +# --model_path /path/to/train_control/checkpoint-xxx/diffusion_pytorch_model.safetensors \ +# --output_path /path/to/control_weights.safetensors +import argparse +import json +import os + +import torch +from safetensors.torch import load_file, save_file + +CONTROL_PREFIXES = ("control_blocks.", "control_img_in.") +# FSDP / DeepSpeed unwrap may leave wrapper prefixes on the keys; strip them to the bare model namespace. +WRAPPER_PREFIXES = ("_fsdp_wrapped_module.", "_fsdp_wrapped_module_", "module.", "_orig_mod.") + + +def parse_args(): + parser = argparse.ArgumentParser( + description="Extract the control-branch weights (control_blocks / control_img_in) of a trained " + "Qwen-Image 2.1 control transformer into a standalone safetensors file." + ) + parser.add_argument( + "--model_path", type=str, default="output_dir_qwen_image_21_control/checkpoint-10000/diffusion_pytorch_model.safetensors", + help="Path to the saved transformer: a directory containing diffusion_pytorch_model*.safetensors " + "(e.g. the FSDP `` dir), or a single .safetensors file.", + ) + parser.add_argument( + "--output_path", type=str, default="output_dir_qwen_image_21_control/checkpoint-10000/diffusion_pytorch_model_control.safetensors", + help="Where to write the extracted control weights.", + ) + return parser.parse_args() + + +def resolve_safetensor_files(model_path): + if os.path.isdir(model_path): + shards = sorted( + os.path.join(model_path, name) + for name in os.listdir(model_path) + if name.endswith(".safetensors") + ) + if not shards: + raise FileNotFoundError(f"No .safetensors files found under {model_path}.") + return shards + if os.path.isfile(model_path) and model_path.endswith(".safetensors"): + return [model_path] + raise FileNotFoundError(f"--model_path must be a safetensors file or a directory of them, got {model_path}.") + + +def unwrap_key(key): + changed = True + while changed: + changed = False + for prefix in WRAPPER_PREFIXES: + if key.startswith(prefix): + key = key[len(prefix):] + changed = True + return key + + +def main(): + args = parse_args() + + state_dict = {} + for shard in resolve_safetensor_files(args.model_path): + state_dict.update(load_file(shard, device="cpu")) + + control_state_dict = {} + for key, value in state_dict.items(): + bare_key = unwrap_key(key) + if bare_key.startswith(CONTROL_PREFIXES): + control_state_dict[bare_key] = value.contiguous() + + if not control_state_dict: + raise ValueError( + f"No control-branch keys (control_blocks.* / control_img_in.*) found in {args.model_path}; " + "this checkpoint does not look like a Qwen-Image 2.1 control training output." + ) + + # Carry the branch layout next to the weights so a loader can rebuild the same control model without + # inspecting the full training config. `config.json` of the saved transformer records both fields; keep + # the safetensors metadata strings-only. + metadata = {"format": "pt"} + config_path = os.path.join(args.model_path, "config.json") if os.path.isdir(args.model_path) else None + if config_path is not None and os.path.isfile(config_path): + with open(config_path, "r") as file: + config = json.load(file) + for field in ("control_layers", "control_in_dim"): + if field in config: + metadata[field] = json.dumps(config[field]) + + os.makedirs(os.path.dirname(os.path.abspath(args.output_path)), exist_ok=True) + save_file(control_state_dict, args.output_path, metadata=metadata) + + num_params = sum(value.numel() for value in control_state_dict.values()) + block_ids = sorted({ + int(key.split(".")[1]) for key in control_state_dict if key.startswith("control_blocks.") + }) + print(f"Extracted {len(control_state_dict)} control tensors ({num_params / 1e9:.3f}B params) -> {args.output_path}") + print(f" control_blocks indices: {block_ids}") + if "control_in_dim" in metadata: + print(f" control_in_dim: {metadata['control_in_dim']}, control_layers: {metadata['control_layers']}") + for key in sorted(control_state_dict): + if key.endswith(".weight") and control_state_dict[key].dim() >= 2: + print(f" {key}: {list(control_state_dict[key].shape)} {control_state_dict[key].dtype}") + + +if __name__ == "__main__": + main() diff --git a/scripts/qwenimage21_fun/train_control.py b/scripts/qwenimage21_fun/train_control.py new file mode 100644 index 0000000..9bf02e5 --- /dev/null +++ b/scripts/qwenimage21_fun/train_control.py @@ -0,0 +1,1739 @@ +"""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 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 import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + ImageVideoControlDataset, ImageVideoSampler, + RandomSampler, get_closest_ratio, 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.fsdp_ema import FSDPEMA +from videox_fun.utils.tqdm_bar import PauseAwareTqdm +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=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else 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("DeepSpeed Zero-3 does not support EMA.") + + ema_module = QwenImageControlTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + low_cpu_mem_usage=True, + ) + if args.use_fsdp: + # The EMA copy gets the same FSDP wrap as the live model so that + # every local shard of the copy pairs 1:1 with the live shard. + ema_transformer3d = FSDPEMA(ema_module, source=transformer3d, accelerator=accelerator, fsdp_plugin=fsdp_plugin) + else: + ema_module = ema_module.to(weight_dtype) + ema_transformer3d = EMAModel(ema_module.parameters(), model_cls=QwenImageControlTransformer2DModel, model_config=ema_module.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 or zero_stage == 3: + 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) + if args.use_ema: + # Every rank joins the FULL_STATE_DICT all-gather inside. + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_transformer3d.load_pretrained(os.path.join(input_dir, "transformer_ema")) + 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", + low_cpu_mem_usage=True, + ) + 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", + low_cpu_mem_usage=True, + ) + 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 Exception: + 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 / 32) * 32 for x in args.fix_sample_size] # 32 = vae_scale_factor(16)*2: keeps the latent grid even so h*w % 4 == 0 (joint stream expands image slots 4x) + 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 / 32) * 32 for x in random_sample_size] # 32 = vae_scale_factor(16)*2: keep latent dims even + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 32) * 32 for x in closest_size] # 32 = vae_scale_factor(16)*2: keep latent dims even + + 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 or zero_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 = PauseAwareTqdm( + 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}") + # Keep the checkpoint out of the progress bar rate: a minute-long save would + # otherwise land in the next step's interval and be shown as a slow step. The + # save also stages the whole state in host RAM (safetensors materializes every + # tensor as bytes) and leaves the freed blocks in the allocator caches, so the + # cache flushes run inside the same window. + with progress_bar.paused(): + accelerator.save_state(save_path) + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + logger.info(f"Saved state to {save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + with progress_bar.paused(): + 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: + with progress_bar.paused(): + 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()) + + # Close the bar before the end-of-run checkpoint: tqdm keeps redrawing a live bar whenever + # something else writes to the console. PauseAwareTqdm.close() rebases the closing line onto + # the smoothed rate, so the worker warm-up and the first dataloader fetch do not dilute it. + progress_bar.close() + + # 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/qwenimage21_fun/train_control.sh b/scripts/qwenimage21_fun/train_control.sh new file mode 100644 index 0000000..02a8dd3 --- /dev/null +++ b/scripts/qwenimage21_fun/train_control.sh @@ -0,0 +1,39 @@ +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +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=BaseQwenImage21TransformerBlock,QwenImage21ControlTransformerBlock \ + --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \ + --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/qwenimage21_fun/train_control.py \ + --config_path="config/qwenimage21/qwenimage21_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_21_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 \ + --random_hw_adapt \ + --low_vram \ + --uniform_sampling \ + --trainable_modules "control" \ + --resume_from_checkpoint="latest" \ No newline at end of file diff --git a/scripts/qwenimage21_fun/train_control_distill.py b/scripts/qwenimage21_fun/train_control_distill.py new file mode 100644 index 0000000..fc10197 --- /dev/null +++ b/scripts/qwenimage21_fun/train_control_distill.py @@ -0,0 +1,1851 @@ +"""CFG distillation of the Qwen-Image 2.1 control branch -- the image counterpart of +scripts/minimax_h3_fun/train_control_distill.py (which itself mirrors scripts/flux2_fun/train_control_distill.py). +A frozen teacher copy of QwenImage21ControlTransformer2DModel (the real score) runs two forward passes per step -- +on the prompt and on an empty negative prompt -- and the two predicted velocities combine with +--real_guidance_scale into the classifier-free-guided target that the trainable student (the control branch, via +--trainable_modules control) regresses onto with an MSE loss. Both copies load the same --transformer_path (a +trained control branch produced by scripts/qwenimage21_fun/train_control.sh). The student takes no guidance input, so +the guidance is distilled into its control branch and inference needs no CFG. + +Original training code 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 torchvision.transforms.functional as TF +import transformers +from accelerate import Accelerator, FullyShardedDataParallelPlugin +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.distributed.fsdp.fully_sharded_data_parallel import ( + FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, + ShardedStateDictConfig) +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +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 import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + ImageVideoControlDataset, + ImageVideoSampler, RandomSampler, + get_closest_ratio, get_random_mask) +from videox_fun.models import (AutoencoderKLQwenImage21, + Qwen3VLForConditionalGeneration, + Qwen3VLProcessor, + QwenImage21ControlTransformer2DModel) +from videox_fun.pipeline import QwenImage21ControlPipeline +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.fsdp_ema import FSDPEMA +from videox_fun.utils.tqdm_bar import PauseAwareTqdm +from videox_fun.utils.utils import (calculate_dimensions, get_image_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 _pack_latents(latents, batch_size, num_channels_latents, height, width): + # 2.1 consumes latents unpatched (patch_size = 1), so packing is a plain spatial flatten. + return latents.view(batch_size, num_channels_latents, height * width).transpose(1, 2) + +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") + +SYS_PROMPT = "Comprehend and analyze the provided prompt." +# The prompt is built as a raw template string and passed straight to the processor, rather than going +# through apply_chat_template: the two tokenize differently and the checkpoint expects this one. +PROMPT_TEMPLATE_T2I = ( + f"<|im_start|>system\n{SYS_PROMPT}<|im_end|>\n" + f"<|im_start|>user\n{{}}<|im_end|>\n" + f"<|im_start|>assistant\n" +) + +def get_prompt_drop_idx(processor): + # Number of leading system-role tokens to drop from the hidden states. Derived from the tokenized + # system message rather than hardcoded, so it tracks the processor's chat template. + sys_message = [{"role": "system", "content": [{"type": "text", "text": SYS_PROMPT}]}] + sys_tokens = processor.apply_chat_template(sys_message, tokenize=True, return_dict=False) + return len(sys_tokens[0]) + +def get_qwen_prompt_embeds(text_encoder, processor, prompt, template, drop_idx, device, weight_dtype): + # Mirrors QwenImage21Pipeline._get_qwen_prompt_embeds for the text-only (t2i) case. + prompt = [" " if not p else p for p in prompt] + prompts = [template.format(t) for t in prompt] + + model_inputs = processor( + text=prompts, padding=True, padding_side="left", return_tensors="pt" + ).to(device) + + # The hidden states have to be read before the vision-language model's final RMSNorm: that is what + # the transformer was trained on. A forward hook returning the module's input neutralizes the norm. + text_model = getattr(text_encoder.model, "language_model", text_encoder.model) + handle = text_model.norm.register_forward_hook(lambda module, args, output: args[0]) + try: + outputs = text_encoder( + input_ids=model_inputs.input_ids, + attention_mask=model_inputs.attention_mask, + output_hidden_states=True, + ) + finally: + handle.remove() + hidden_states = outputs.hidden_states[-1] + + split_hidden_states = list(_extract_masked_hidden(hidden_states, model_inputs.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] + ) + return prompt_embeds.to(dtype=weight_dtype), encoder_attention_mask + +def log_validation(vae, text_encoder, processor, 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 = QwenImage21ControlPipeline( + vae=vae, + text_encoder=text_encoder, + processor=processor, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else 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 validation drives generation from a control image: derive the output resolution from that + # image's aspect ratio (as in scripts/qwenimage_fun/train_control.py) and load it as a single-frame + # (1, 3, h, w) tensor via get_image_latent. + 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], + height = height, + width = width, + generator = generator, + num_inference_steps = 20, + control_image = control_image, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + # 2.1's VAE decodes to RGBA; JPEG cannot store an alpha channel (it raises "cannot write mode RGBA + # as JPEG"), so save the validation preview as PNG -- matching examples/qwenimage21/predict_t2i.py. + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.png" + ) + ) + + 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 images 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( + "--real_guidance_scale", + type=float, + default=3.5, + help="Classifier-free guidance scale applied to the frozen teacher (the real score) to build the " + "distillation target. The student takes no guidance input; guidance enters only via the teacher's " + "cond/uncond forward pair.", + ) + 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=1024, + 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=PROMPT_TEMPLATE_T2I, + help=( + 'The prompt template for text encoder.' + ), + ) + 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`.", + ) + + 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 + processor = Qwen3VLProcessor.from_pretrained( + args.pretrained_model_name_or_path, subfolder="processor" + ) + # 2.1 drops a fixed number of leading system-role tokens from the hidden states; derive it from + # the processor chat template instead of hardcoding, so it tracks template changes. + prompt_drop_idx = get_prompt_drop_idx(processor) + + 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 = Qwen3VLForConditionalGeneration.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype + ) + text_encoder = text_encoder.eval() + # Get Vae + vae = AutoencoderKLQwenImage21.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) + + # Load the training config so the control transformer is built with the extra control kwargs + # (control_layers / control_in_dim) declared in config/qwenimage21/qwenimage21_control.yaml. + config = OmegaConf.load(args.config_path) + + # Get Transformer + transformer3d = QwenImage21ControlTransformer2DModel.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) + + # Distillation teacher (the "real score"): a frozen second copy of the same control model. Both it and the + # trainable student load the same --transformer_path (a trained control branch), so the teacher reproduces the + # model's own conditional score and the CFG of its cond/uncond predictions is the target the student's control + # branch regresses onto. + real_score_transformer = QwenImage21ControlTransformer2DModel.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) + real_score_transformer.requires_grad_(False) + real_score_transformer.eval() + + 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 + + # Load the SAME trained control branch onto the frozen teacher so its score matches the student's init. + m_t, u_t = real_score_transformer.load_state_dict(state_dict, strict=False) + print(f"teacher missing keys: {len(m_t)}, unexpected keys: {len(u_t)}") + assert len(u_t) == 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. + # Full finetune: trainable_modules = ['transformer_blocks', 'img_in', 'txt_in', 'time_text_embed', 'modulation', 'norm_out', 'proj_out', 'pos_embed'] + # Partial finetune: trainable_modules = ['transformer_blocks'] + 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("DeepSpeed Zero-3 does not support EMA.") + + ema_module = QwenImage21ControlTransformer2DModel.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']), + ) + if args.use_fsdp: + # The EMA copy gets the same FSDP wrap as the live model so that + # every local shard of the copy pairs 1:1 with the live shard. + ema_transformer3d = FSDPEMA(ema_module, source=transformer3d, accelerator=accelerator, fsdp_plugin=fsdp_plugin) + else: + ema_module = ema_module.to(weight_dtype) + ema_transformer3d = EMAModel(ema_module.parameters(), model_cls=QwenImage21ControlTransformer2DModel, model_config=ema_module.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 or zero_stage == 3: + 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) + if args.use_ema: + # Every rank joins the FULL_STATE_DICT all-gather inside. + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_transformer3d.load_pretrained(os.path.join(input_dir, "transformer_ema")) + 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 = QwenImage21ControlTransformer2DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = QwenImage21ControlTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + low_cpu_mem_usage=True, + ) + load_model = EMAModel(load_model.parameters(), model_cls=QwenImage21ControlTransformer2DModel, 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 = QwenImage21ControlTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer", + low_cpu_mem_usage=True, + ) + 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 Exception: + 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] + if length >= 2: + first_element = 0.90 + 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 + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p = number_list_prob) + else: + return rng.choice(number_list, p = number_list_prob) + + # 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 / 32) * 32 for x in args.fix_sample_size] # 32 = vae_scale_factor(16)*2: keeps the latent grid even so h*w % 4 == 0 (joint stream expands image slots 4x) + 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 / 32) * 32 for x in random_sample_size] # 32 = vae_scale_factor(16)*2: keep latent dims even + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 32) * 32 for x in closest_size] # 32 = vae_scale_factor(16)*2: keep latent dims even + + for example in examples: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # 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: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # 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: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # 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), + ]) + new_examples["pixel_values"].append(transform(pixel_values)) + + # Control mode: run the SAME spatial transform on the control image, then build the inpaint + # mask / masked-image pair from the transformed target. control_pixel_values is (f, c, h, w); + # get_random_mask returns a (f, 1, h, w) uint8 mask that broadcasts over channels. This mirrors + # scripts/qwenimage_fun/train_control.py's bucket collate. + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + new_examples["control_pixel_values"].append(transform(control_pixel_values)) + + 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) + new_examples["mask"].append(mask) + + new_examples["text"].append(example["text"]) + + # 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: + prompt_embeds, encoder_attention_mask = get_qwen_prompt_embeds( + text_encoder, processor, batch["text"], args.prompt_template_encode, + prompt_drop_idx, accelerator.device, weight_dtype, + ) + + 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 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.model.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) + + # The frozen teacher is not accelerator.prepare()-wrapped (only the trainable student is sharded); keep it a + # plain bf16 module. Under --low_vram park it on CPU and stream it to the GPU only for its forward passes. + real_score_transformer.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + + # Precompute the teacher's unconditional (empty-prompt) embedding once; it is constant across steps and is the + # uncond branch of the CFG distillation target. Temporarily place the text encoder on the GPU to encode it, then + # restore its previous device so the --low_vram / in-dataloader parking below is unchanged. + _te_prev_device = next(text_encoder.parameters()).device + text_encoder.to(accelerator.device) + with torch.no_grad(): + neg_prompt_embeds, neg_encoder_attention_mask = get_qwen_prompt_embeds( + text_encoder, processor, [" "], args.prompt_template_encode, + prompt_drop_idx, accelerator.device, weight_dtype, + ) + neg_prompt_embeds = neg_prompt_embeds.to(dtype=weight_dtype) + neg_encoder_attention_mask = neg_encoder_attention_mask.to(device=accelerator.device) + text_encoder.to(_te_prev_device) + + # 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 = PauseAwareTqdm( + 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 == 0: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + pixel_values = rearrange(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, text) in enumerate(zip(pixel_values, texts)): + pixel_value = 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) + + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + # Convert images to latent space + pixel_values = batch["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(): + # 2.1's VAE reads RGBA, so composite the RGB batch over an opaque alpha channel. + alpha = torch.ones_like(pixel_values[:, :, :1]) + pixel_values = torch.cat([pixel_values, alpha], dim=2) + # 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) + + # 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) + + # --- Control conditioning (control_context) --- + # Encode the control image and the inpaint pair (masked image + latent-resolution mask) and concat them + # into the (b, 129, 1, h', w') tensor the transformer scatters into the joint stream: 64 control-latent + # channels + 1 mask channel + 64 masked-image latent channels. Mirrors + # scripts/qwenimage_fun/train_control.py, adapted to 2.1: the VAE reads RGBA (composite an opaque alpha) + # and t2v_flag scales with 5D-safe broadcasting. Runs after the base latents encode has synced and + # before the VAE is offloaded under --low_vram, so the VAE is never double-booked across CUDA streams. + with torch.no_grad(): + 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) + + control_pixel_values = torch.cat([control_pixel_values, torch.ones_like(control_pixel_values[:, :, :1])], dim=2) + control_latents = _batch_encode_vae(control_pixel_values) + control_latents = ((control_latents - latents_mean) * latents_std).to(dtype=weight_dtype) + + # Drop the control latents 10% of the time so the adapter also sees the unconditional (t2i) path. + for bs_index in range(control_latents.size()[0]): + zero_init_control_conv_in = (np.random.choice([0, 1], p=[0.90, 0.10]) if rng is None + else rng.choice([0, 1], p=[0.90, 0.10])) + if zero_init_control_conv_in: + control_latents[bs_index] = control_latents[bs_index] * 0 + + # Downsample the mask to latent resolution. mask is (b, f=1, 1, h, w); squeeze the frame dim so + # interpolate sees a 4D (b, 1, h, w) tensor, then restore a unit frame dim -> (b, 1, 1, h', w'). + mask = mask.squeeze(1) + mask_conditions = F.interpolate(1 - mask[:, :1], size=control_latents.size()[-2:], mode='nearest').to(latents.device, weight_dtype) + mask_conditions = mask_conditions.unsqueeze(2) + + # A full-frame mask carries no inpaint signal, so gate the masked latents off 90% of the time in + # that case (t2v_flag) to stop the model leaning on them. + 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(latents.device, dtype=weight_dtype) + + mask_pixel_values = torch.cat([mask_pixel_values, torch.ones_like(mask_pixel_values[:, :, :1])], dim=2) + 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, None, None] * mask_latents + + inpaint_latents = torch.cat([mask_conditions, mask_latents], dim=1) + control_context = torch.cat([control_latents, inpaint_latents], dim=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(): + prompt_embeds, encoder_attention_mask = get_qwen_prompt_embeds( + text_encoder, processor, batch["text"], args.prompt_template_encode, + prompt_drop_idx, accelerator.device, weight_dtype, + ) + 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() + + with accelerator.accumulate(transformer3d): + bsz, channel, num_frame, height, width = latents.size() + latents = _pack_latents(latents, bsz, channel, height, width) + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + # Pack the control conditioning into the same (seq, dim) layout the transformer scatters into the joint + # image positions: (b, 129, 1, h', w') -> (b, h'*w', 129). 2.1 keeps latents unpatched (no num_frame). + control_context = _pack_latents(control_context, bsz, control_context.size(1), height, width) + + 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 + + # NOTE: the flow-matching velocity target (noise - latents) is intentionally not used here -- this + # is CFG distillation, so the student regresses onto the teacher's CFG velocity built below. + + # 2.1 keeps latents unpatched, so one img_shapes entry spans the full latent grid, while each + # vision-language image slot in img_mask stands for a 2x2 group of those latent tokens. + img_shapes = [[(1, height, width)]] * latents.size(0) + img_mask = torch.cat( + [ + encoder_attention_mask.new_zeros(bsz, prompt_embeds.size(1), dtype=torch.bool), + encoder_attention_mask.new_ones(bsz, latents.size(1) // 4, dtype=torch.bool), + ], + dim=1, + ) + + # 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, + img_mask=img_mask, + control_context=control_context, + return_dict=False, + )[0][:, -noisy_latents.size(1):] + + # Distillation target: the frozen teacher scores the SAME noised rows twice (prompt, then empty + # prompt); the two velocities combine with --real_guidance_scale into the classifier-free-guided + # target the student regresses onto. Runs under no_grad; the teacher is never prepare()-wrapped. + if args.low_vram: + real_score_transformer.to(accelerator.device) + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + teacher_cond = real_score_transformer( + hidden_states=noisy_latents, + timestep=timesteps / 1000, + encoder_hidden_states_mask=encoder_attention_mask, + encoder_hidden_states=prompt_embeds, + img_shapes=img_shapes, + img_mask=img_mask, + control_context=control_context, + return_dict=False, + )[0][:, -noisy_latents.size(1):] + # Uncond: same image/control rows, empty prompt. The joint layout depends on the text length, so + # rebuild img_mask / attention mask for the short negative embedding (expanded to this batch). + neg_embeds_b = neg_prompt_embeds.expand(bsz, -1, -1) + neg_attn_b = neg_encoder_attention_mask.expand(bsz, -1).to(device=latents.device) + neg_img_mask = torch.cat( + [ + neg_attn_b.new_zeros(bsz, neg_embeds_b.size(1), dtype=torch.bool), + neg_attn_b.new_ones(bsz, latents.size(1) // 4, dtype=torch.bool), + ], + dim=1, + ) + teacher_uncond = real_score_transformer( + hidden_states=noisy_latents, + timestep=timesteps / 1000, + encoder_hidden_states_mask=neg_attn_b, + encoder_hidden_states=neg_embeds_b.to(dtype=weight_dtype), + img_shapes=img_shapes, + img_mask=neg_img_mask, + control_context=control_context, + return_dict=False, + )[0][:, -noisy_latents.size(1):] + if args.low_vram: + real_score_transformer.to("cpu") + torch.cuda.empty_cache() + + teacher_cfg_target = teacher_uncond + (teacher_cond - teacher_uncond) * args.real_guidance_scale + + 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) + # CFG-distillation loss: the student's conditional velocity regresses onto the teacher's CFG velocity. + loss = custom_mse_loss(noise_pred.float(), teacher_cfg_target.detach().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}") + # Keep the checkpoint out of the progress bar rate: a minute-long save would + # otherwise land in the next step's interval and be shown as a slow step. The + # save also stages the whole state in host RAM (safetensors materializes every + # tensor as bytes) and leaves the freed blocks in the allocator caches, so the + # cache flushes run inside the same window. + with progress_bar.paused(): + accelerator.save_state(save_path) + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + logger.info(f"Saved state to {save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + with progress_bar.paused(): + 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, + 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) + + if global_step >= args.max_train_steps: + break + + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + with progress_bar.paused(): + 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, + processor, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Close the bar before the end-of-run checkpoint: tqdm keeps redrawing a live bar whenever + # something else writes to the console. PauseAwareTqdm.close() rebases the closing line onto + # the smoothed rate, so the worker warm-up and the first dataloader fetch do not dilute it. + progress_bar.close() + + # 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/qwenimage21_fun/train_control_distill.sh b/scripts/qwenimage21_fun/train_control_distill.sh new file mode 100644 index 0000000..437487a --- /dev/null +++ b/scripts/qwenimage21_fun/train_control_distill.sh @@ -0,0 +1,40 @@ +export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image-2.1" +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=BaseQwenImage21TransformerBlock,QwenImage21ControlTransformerBlock \ + --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \ + --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/qwenimage21_fun/train_control_distill.py \ + --config_path="config/qwenimage21/qwenimage21_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-06 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_qwen_image_21_control_distill" \ + --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 \ + --random_hw_adapt \ + --low_vram \ + --uniform_sampling \ + --transformer_path="models/Personalized_Model/Qwen-Image-2.1-Fun-Controlnet-Union.safetensors" \ + --trainable_modules "control" \ + --resume_from_checkpoint="latest \ No newline at end of file diff --git a/scripts/wan2.1_flex_forcing/train_distill.py b/scripts/wan2.1_flex_forcing/train_distill.py index 9c95a48..23a51d4 100644 --- a/scripts/wan2.1_flex_forcing/train_distill.py +++ b/scripts/wan2.1_flex_forcing/train_distill.py @@ -81,8 +81,8 @@ from videox_fun.pipeline import (WanI2VPipeline, WanPipeline, WanFlexForcingPipeline, WanSelfForcingPipeline) from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.flex_chunking import (UNIFORM_BLOCK_PROB, - broadcast_chunk_sizes, +from videox_fun.utils.flex_chunking import (broadcast_chunk_sizes, + build_full_then_blocks_partitions, build_pyramid_partitions, chunk_boundaries, sample_flexible_chunks, @@ -317,7 +317,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer # would validate out of distribution too. The uniform # block layout is the one band that is both fixed # across checkpoints and actually trained: - # UNIFORM_BLOCK_PROB of iterations use it, against + # FLEX_ARM_PROB of iterations use it, against # ~1.3% for the most common random partition. flex_kwargs = dict( chunk_spec=args.num_frame_per_block, @@ -379,14 +379,16 @@ def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=Non return torch.clip(t.to(torch.int32), low, high - 1) -# Fraction of pyramid iterations whose level 0 is the whole clip as a single -# chunk, i.e. fully bidirectional. 3.2's coarse planning step is exactly that at -# inference (`denoise_mode="pyramid"` with no `chunk_spec` builds `[[F], ...]`), -# so it has to appear in training; the remaining iterations keep the random -# partition of 3.1 so the causal end and every layout in between stay covered. -# Deliberately not a CLI flag - it is a property of the paper's schedule, not a -# knob the launcher should have to keep in sync. -COARSE_GLOBAL_PROB = 0.5 +# Each training iteration draws its level-0 layout from four equally likely arms +# (FLEX_ARM_PROB = 25% of iterations each): the whole-clip planning chunk (3.2's +# coarse step, `[[F], ...]` at inference), the launcher's own uniform +# `num_frame_per_block`, the fixed "first step full, then block-major" ladder +# (`denoise_mode="full_then_blocks"`), and a random 3.1 partition. Equal shares +# keep the coarse, causal and every in-between layout covered while giving the +# full_then_blocks ladder enough iterations to train the block-major-at-high-noise +# trajectory the binary pyramid never reaches. Deliberately not a CLI flag - it is +# a property of the schedule mixture, not a knob the launcher keeps in sync. +FLEX_ARM_PROB = 0.25 # How many steps of a launch report the partition they drew. Counted from this # process rather than from `global_step`, so a resumed run still gets its own @@ -406,59 +408,110 @@ def sample_flex_partitions(args, num_frames, num_denoising_steps, torch_rng, per denoising step. Returns ``None`` when Flex-Forcing is off, which leaves the inherited uniform ``num_frame_per_block`` masks untouched. - The level-0 layout is a three-way mixture decided by a single uniform draw, - so the constants below are the actual iteration shares: ``COARSE_GLOBAL_PROB`` - of iterations use one chunk over the whole clip (what inference's coarse - planning step uses), ``UNIFORM_BLOCK_PROB`` pin the launcher's own uniform - ``num_frame_per_block``, and the rest draw a random 2..10 partition. All of - them are then refined into the same ladder, so a single set of weights covers - `denoise_mode="pyramid"` whether or not the caller pins `chunk_spec`. With no - pyramid the coarse band is empty and the split is 10% uniform / 90% random. + The level-0 layout is a four-way mixture decided by a single uniform draw, + each arm taking ``FLEX_ARM_PROB`` (25%) of iterations: one chunk over the + whole clip (inference's coarse planning step), the launcher's own uniform + ``num_frame_per_block``, the fixed "first step full, then block-major" ladder + (`denoise_mode="full_then_blocks"`), and a random 2..10 partition. The first, + second and last are then refined into the same binary pyramid, so a single set + of weights covers `denoise_mode="pyramid"` whether or not the caller pins + `chunk_spec`; the full_then_blocks arm bypasses that refinement and emits its + two-level ladder whole, drawing the block-major width between fully causal (1 + frame / block) and `num_frame_per_block`. With no pyramid the coarse arm is + empty (a `[F]` level 0 would be a degenerate single-chunk rollout), so the + split is 25% uniform / 25% full_then_blocks / 50% random. """ if not args.flex_forcing: return None - u = torch.rand((), generator=torch_rng, device=device).item() - coarse_prob = COARSE_GLOBAL_PROB if args.flex_pyramid_levels > 1 else 0.0 - if u < coarse_prob: + # The arm selector must be IDENTICAL on every rank, not just its broadcast + # count. Each arm now post-processes the received partition its own way (the + # coarse/uniform/random trio funnel through the shared refine tail below, but + # full_then_blocks sets the ladder directly from a different tensor), so a + # per-rank `u` would leave ranks in different arms building different ladders + # -> different walk forward counts -> FSDP all-gather desync (NCCL hang). + # Broadcast rank 0's draw so all ranks take the same arm; the arm's existing + # single broadcast then reconciles its within-arm randomness (the random base, + # the block width) to rank 0, making every ladder byte-identical. + u = torch.rand((), generator=torch_rng, device=device) + if dist.is_available() and dist.is_initialized(): + dist.broadcast(u.reshape(1), src=0) + u = u.item() + pyramid = args.flex_pyramid_levels > 1 + # Four equally likely arms; the coarse (whole-clip planning) one only when a + # pyramid is on, else its 25% folds into the random arm. + coarse_hi = FLEX_ARM_PROB if pyramid else 0.0 + uniform_hi = coarse_hi + FLEX_ARM_PROB + ftb_hi = uniform_hi + FLEX_ARM_PROB + ladder = None + if u < coarse_hi: # Coarse end of 3.1: reuse the ladder builder's own level-0 rule so the # `independent_first_frame` handling cannot drift from inference's. base = build_pyramid_partitions( num_frames, num_levels=1, base_chunks=None, independent_first_frame=args.independent_first_frame)[0] arm = "whole clip, the 3.2 coarse planning layout" - elif u < coarse_prob + UNIFORM_BLOCK_PROB: + elif u < uniform_hi: # `uniform_chunks`, not `normalize_chunk_spec`: the latter takes no # `independent_first_frame` argument (it encodes that as a leading 1), so - # it would silently disagree with the two bands around it. This is the - # band `log_validation` renders when there is no pyramid. + # it would silently disagree with the bands around it. This is the band + # `log_validation` renders when there is no pyramid. base = uniform_chunks( num_frames, args.num_frame_per_block, independent_first_frame=args.independent_first_frame) arm = f"uniform {args.num_frame_per_block}-frame blocks, the launcher's own layout" + elif u < ftb_hi: + # "First step full, every later step block-major": a fixed two-level + # ladder that does NOT go through the binary pyramid refinement, so the + # block-major level is reached at high noise (the 2nd step) instead of as + # a deep refinement - the one trajectory the pyramid arms never sample. + # The block-major width is drawn between fully causal (1 frame / block) + # and the launcher's `num_frame_per_block`, so both the tight-AR and the + # coarse-block refinement of the whole-clip plan get trained. The block + # width is a per-rank draw, so it has to be reconciled with EXACTLY ONE + # broadcast - the same collective count every other arm issues (they each + # broadcast their single base) - or the ranks desync and NCCL deadlocks. + # Level 0 is just [F], identical on every rank and needing no sync, so we + # broadcast only the block-major level; rank 0's draw wins and every rank + # rebuilds the same two-level ladder locally. + block_choices = sorted({1, max(1, int(args.num_frame_per_block))}) + pick = int(torch.rand((), generator=torch_rng, device=device).item() + * len(block_choices)) + block = block_choices[min(pick, len(block_choices) - 1)] + block = min(block, int(num_frames)) + ftb = build_full_then_blocks_partitions( + num_frames, block, + independent_first_frame=args.independent_first_frame) + ftb[1] = broadcast_chunk_sizes(ftb[1], device=device) + ladder = ftb[:num_denoising_steps] + arm = (f"full clip -> uniform {max(ftb[1])}-frame blocks, " + "the full_then_blocks ladder") else: base = sample_flexible_chunks( num_frames, min_chunk=args.flex_chunk_min, max_chunk=args.flex_chunk_max, generator=torch_rng, device=device, independent_first_frame=args.independent_first_frame) arm = f"random {args.flex_chunk_min}..{args.flex_chunk_max}-frame blocks, the 3.1 spectrum" - # Every rank has to train the same layout: the FlexAttention mask, and the - # `num_frame_per_block` derived from it, must agree across the SP/FSDP group. - base = broadcast_chunk_sizes(base, device=device) - ladder = [base] - if args.flex_pyramid_levels > 1: - ladder = build_pyramid_partitions( - num_frames, num_levels=args.flex_pyramid_levels, - min_num_frame_per_block=args.flex_min_num_frame_per_block, base_chunks=base, - independent_first_frame=args.independent_first_frame) - if len(ladder) > num_denoising_steps: - # Same short-circuit as at inference, where the rollout stops refining at - # the last step: deeper levels would never be reached, so drop them - # instead of reporting a pyramid that was not actually trained. - if verbose: - print(f"--flex_pyramid_levels={args.flex_pyramid_levels} builds " - f"{len(ladder)} levels but only {num_denoising_steps} denoising " - f"steps are trained; keeping the first {num_denoising_steps}.") - ladder = ladder[:num_denoising_steps] + if ladder is None: + # Every rank has to train the same layout: the FlexAttention mask, and the + # `num_frame_per_block` derived from it, must agree across the SP/FSDP + # group. (The full_then_blocks arm above is deterministic given `args` and + # already broadcast each level, so it skips this path.) + base = broadcast_chunk_sizes(base, device=device) + ladder = [base] + if pyramid: + ladder = build_pyramid_partitions( + num_frames, num_levels=args.flex_pyramid_levels, + min_num_frame_per_block=args.flex_min_num_frame_per_block, base_chunks=base, + independent_first_frame=args.independent_first_frame) + if len(ladder) > num_denoising_steps: + # Same short-circuit as at inference, where the rollout stops + # refining at the last step: deeper levels would never be reached, + # so drop them instead of reporting a pyramid not actually trained. + if verbose: + print(f"--flex_pyramid_levels={args.flex_pyramid_levels} builds " + f"{len(ladder)} levels but only {num_denoising_steps} denoising " + f"steps are trained; keeping the first {num_denoising_steps}.") + ladder = ladder[:num_denoising_steps] if verbose: # Every level, not just the drawn one: level 0 is what the mixture above # picked, the rest are derived from it, and they are the sub-spans the diff --git a/videox_fun/models/qwenimage21_transformer2d_control.py b/videox_fun/models/qwenimage21_transformer2d_control.py new file mode 100644 index 0000000..f33c654 --- /dev/null +++ b/videox_fun/models/qwenimage21_transformer2d_control.py @@ -0,0 +1,437 @@ +# Modified from videox_fun/models/qwenimage_transformer2d_control.py (the Qwen-Image 2.0 Fun control model). +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +# +# NOTE: This is the Qwen-Image 2.1 counterpart of `qwenimage_transformer2d_control.py`. It keeps the same VACE-style +# ControlNet-Union design (a parallel chain of zero-init `control_blocks` produces per-layer skip `hints` that are +# added back into the frozen base blocks), but adapts it to 2.1's single-stream transformer: +# * 2.1 concatenates text and image tokens into one `joint_hidden_states` sequence and every block returns a single +# tensor (not the 2.0 `(encoder_hidden_states, hidden_states)` tuple), so the control stream is a joint stream too +# and the skips are added to the whole joint sequence. +# * 2.1 computes one shared `modulation` for all blocks and threads `rotary_emb` / block-causal `segments` / +# `key_valid` / prefix-KV-cache args through every block; the control blocks reuse the exact same block kwargs. +# * The control conditioning (`control_context`) is image-space and is scattered into the joint sequence at the image +# token positions, mirroring how `img_in(hidden_states)` is scattered in the base forward. + +import math +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.utils.checkpoint +from diffusers.configuration_utils import register_to_config +from diffusers.models.modeling_outputs import Transformer2DModelOutput +from diffusers.utils import logging + +from ..dist import sequence_parallel_all_gather +from .qwenimage21_transformer2d import (QwenImage21KVCache, + QwenImage21Transformer2DModel, + QwenImage21TransformerBlock, + _IMG_TOKENS_PER_SLOT, + _qwenimage21_prefix_segments) + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class QwenImage21ControlTransformerBlock(QwenImage21TransformerBlock): + """A control block in the parallel VACE chain. + + Mirrors `QwenImageControlTransformerBlock`: the first control block (``block_id == 0``) merges the control + conditioning into the base joint stream through a zero-init ``before_proj``; every block emits a zero-init + ``after_proj`` skip. The running control stream ``c`` is carried between blocks as a stack + ``[skip_0, ..., skip_{n-1}, c_n]`` exactly like the 2.0 model, so ``forward_control`` can unbind the skips. + """ + + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + mlp_ratio: int = 3, + eps: float = 1e-6, + block_id: int = 0, + ): + super().__init__(dim, num_attention_heads, attention_head_dim, mlp_ratio, eps) + self.block_id = block_id + if block_id == 0: + self.before_proj = nn.Linear(dim, dim) + nn.init.zeros_(self.before_proj.weight) + nn.init.zeros_(self.before_proj.bias) + self.after_proj = nn.Linear(dim, dim) + nn.init.zeros_(self.after_proj.weight) + nn.init.zeros_(self.after_proj.bias) + + def forward(self, c, x, **kwargs): + if self.block_id == 0: + # `before_proj` is zero-init, so at step 0 the control stream starts as an exact copy of the base joint + # stream `x` and the whole adapter is a no-op until training moves `before_proj` / `after_proj`. + c = self.before_proj(c) + x + all_c = [] + else: + all_c = list(torch.unbind(c)) + c = all_c.pop(-1) + + c = super().forward(c, **kwargs) + c_skip = self.after_proj(c) + all_c += [c_skip, c] + c = torch.stack(all_c) + return c + + +class BaseQwenImage21TransformerBlock(QwenImage21TransformerBlock): + """A frozen base block that optionally adds one control skip (`hints[block_id]`) to its joint output.""" + + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + mlp_ratio: int = 3, + eps: float = 1e-6, + block_id: Optional[int] = None, + ): + super().__init__(dim, num_attention_heads, attention_head_dim, mlp_ratio, eps) + self.block_id = block_id + + def forward(self, hidden_states, hints=None, context_scale: float = 1.0, **kwargs): + hidden_states = super().forward(hidden_states, **kwargs) + if self.block_id is not None and hints is not None: + hidden_states = hidden_states + hints[self.block_id] * context_scale + return hidden_states + + +class QwenImage21ControlTransformer2DModel(QwenImage21Transformer2DModel): + """Qwen-Image 2.1 transformer with a VACE-style ControlNet-Union adapter. + + The base blocks are rebuilt as `BaseQwenImage21TransformerBlock` (so they can receive skips) and a parallel chain of + `QwenImage21ControlTransformerBlock` plus a `control_img_in` projection is added. Only the control parameters are + meant to be trained (`--trainable_modules "control"`); the base weights stay frozen. + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["BaseQwenImage21TransformerBlock", "QwenImage21ControlTransformerBlock"] + _skip_layerwise_casting_patterns = ["pos_embed", "norm"] + _repeated_blocks = ["BaseQwenImage21TransformerBlock", "QwenImage21ControlTransformerBlock"] + + @register_to_config + def __init__( + self, + control_layers=None, + control_in_dim=None, + patch_size: int = 1, + in_channels: int = 64, + out_channels: Optional[int] = 64, + num_layers: int = 32, + attention_head_dim: int = 128, + num_attention_heads: int = 32, + context_in_dim: int = 4096, + mlp_ratio: int = 3, + axes_dims_rope: Tuple[int, int, int] = (16, 56, 56), + eps: float = 1e-6, + causal_condition: bool = True, + ): + super().__init__( + patch_size, in_channels, out_channels, num_layers, attention_head_dim, num_attention_heads, + context_in_dim, mlp_ratio, axes_dims_rope, eps, causal_condition, + ) + + self.control_layers = [i for i in range(0, num_layers, 2)] if control_layers is None else list(control_layers) + self.control_in_dim = in_channels if control_in_dim is None else control_in_dim + + assert 0 in self.control_layers, "control_layers must contain 0 (the first control block merges the conditioning)." + self.control_layers_mapping = {i: n for n, i in enumerate(self.control_layers)} + + # Rebuild the base blocks so each knows whether it consumes a control skip (and which one). + self.transformer_blocks = nn.ModuleList( + [ + BaseQwenImage21TransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + mlp_ratio=mlp_ratio, + eps=eps, + block_id=self.control_layers_mapping[i] if i in self.control_layers else None, + ) + for i in range(num_layers) + ] + ) + + # Parallel control chain. `block_id` here is the *layer index* (only layer 0 gets `before_proj`), matching 2.0. + self.control_blocks = nn.ModuleList( + [ + QwenImage21ControlTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + mlp_ratio=mlp_ratio, + eps=eps, + block_id=i, + ) + for i in self.control_layers + ] + ) + + # Projects the packed control conditioning (control latents + inpaint mask/latents) into the joint stream dim. + self.control_img_in = nn.Linear(self.control_in_dim, self.inner_dim) + + @classmethod + def from_pretrained( + cls, pretrained_model_path, subfolder=None, transformer_additional_kwargs=None, + low_cpu_mem_usage=False, torch_dtype=torch.bfloat16, + ): + model = super().from_pretrained( + pretrained_model_path, subfolder=subfolder, transformer_additional_kwargs=transformer_additional_kwargs, + low_cpu_mem_usage=low_cpu_mem_usage, torch_dtype=torch_dtype, + ) + # There are no pretrained 2.1 control weights, so the control parameters are "missing keys" that the base + # `from_pretrained` auto-initializes (xavier for 2D weights under low_cpu_mem_usage). That would destroy the + # VACE no-op start, so re-zero `before_proj` / `after_proj` and reset `control_img_in` to the default Linear + # init. A subsequently loaded `--transformer_path` control checkpoint overrides these again. + with torch.no_grad(): + for block in model.control_blocks: + if hasattr(block, "before_proj"): + nn.init.zeros_(block.before_proj.weight) + nn.init.zeros_(block.before_proj.bias) + nn.init.zeros_(block.after_proj.weight) + nn.init.zeros_(block.after_proj.bias) + model.control_img_in.reset_parameters() + return model + + def forward_control(self, control_joint, x_joint, kwargs): + """Run the parallel control chain and return the per-layer skips (`hints`). + + `control_joint` is the control conditioning scattered into a joint-shaped stream (image positions filled, text + positions zero); `x_joint` is the base joint stream merged in at the first control block. Both are already + sliced for kv-cache decode / sequence parallel exactly like the base stream. + """ + c = control_joint + for block in self.control_blocks: + 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 + + c = torch.utils.checkpoint.checkpoint( + create_custom_forward(block, x=x_joint, **kwargs), + c, + use_reentrant=False, + ) + else: + c = block(c, x_joint, **kwargs) + + # Drop the last entry (the running stream); keep only the `len(control_layers)` skips. + hints = torch.unbind(c)[:-1] + return hints + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + timestep: torch.Tensor, + img_shapes: List[List[Tuple[int, int, int]]], + img_mask: torch.Tensor, + encoder_hidden_states_mask: Optional[torch.Tensor] = None, + attention_kwargs: Optional[Dict[str, Any]] = None, + kv_cache: Optional[QwenImage21KVCache] = None, + kv_cache_mode: Optional[str] = None, + control_context: Optional[torch.Tensor] = None, + control_context_scale: float = 1.0, + return_dict: bool = True, + ) -> Union[torch.Tensor, Transformer2DModelOutput]: + r""" + Same as `QwenImage21Transformer2DModel.forward` plus the control conditioning: + + control_context (`torch.Tensor`, *optional*): Packed control conditioning of shape + `(batch_size, image_sequence_length, control_in_dim)` -- the VAE-encoded control image concatenated with the + inpaint condition (mask + masked latents). When `None`, the model behaves exactly like the base transformer. + control_context_scale (`float`, defaults to `1.0`): Weight of the injected control skips. + """ + batch_size = hidden_states.shape[0] + if kv_cache is not None and not self.config.causal_condition: + raise ValueError( + "kv_cache requires `causal_condition=True`. The cache is only valid because text and condition-image " + "tokens modulate from t=0, which makes their activations independent of the denoising step." + ) + if kv_cache is not None and kv_cache_mode not in ("extract", "cached"): + raise ValueError( + f"kv_cache_mode must be 'extract' or 'cached' when kv_cache is provided, got {kv_cache_mode!r}." + ) + if kv_cache is None and kv_cache_mode is not None: + raise ValueError(f"kv_cache_mode is {kv_cache_mode!r} but no kv_cache was passed to hold the prefix.") + + hidden_states = self.img_in(hidden_states) + encoder_hidden_states = self.txt_in(encoder_hidden_states) + + # Each vision-language image slot stands for 2x2 latent tokens, so expand those positions four-fold and drop + # the actual latents into them. Samples share a layout, hence the single row. + repeats = torch.where(img_mask, _IMG_TOKENS_PER_SLOT, 1)[0] + image_pad_mask = torch.repeat_interleave(img_mask[0], repeats) + + target_tokens = math.prod(img_shapes[0][-1]) + joint_hidden_states = torch.cat( + [ + encoder_hidden_states, + encoder_hidden_states.new_zeros(batch_size, target_tokens // 4, encoder_hidden_states.shape[2]), + ], + dim=1, + ) + joint_hidden_states = joint_hidden_states.repeat_interleave(repeats, dim=1) + joint_hidden_states[:, image_pad_mask] = hidden_states + + # Build the control stream in lockstep with the base joint stream: project the packed control conditioning and + # scatter it into the image token positions (text positions stay zero; `before_proj` merges `x` in at block 0). + control_joint = None + if control_context is not None: + control_features = self.control_img_in(control_context) + control_joint = torch.zeros_like(joint_hidden_states) + control_joint[:, image_pad_mask] = control_features + + rotary_emb = self.pos_embed(img_shapes[0], image_pad_mask, device=hidden_states.device) + image_ids, target_token_mask = self.build_token_metadata(image_pad_mask, img_shapes[0]) + + timestep = timestep.to(hidden_states.dtype) + if self.config.causal_condition: + timestep = torch.cat([timestep, timestep.new_zeros(1)], dim=0) + modulation_mask = target_token_mask + else: + modulation_mask = None + temb = self.time_text_embed(timestep, hidden_states) + modulation = self.modulation(temb) + + joint_key_valid = None + if encoder_hidden_states_mask is not None: + joint_key_valid = torch.ones( + batch_size, image_pad_mask.shape[0], dtype=torch.bool, device=hidden_states.device + ) + text_positions = (~image_pad_mask).nonzero(as_tuple=True)[0] + vlm_text_positions = ~img_mask[0][: encoder_hidden_states_mask.shape[1]] + joint_key_valid[:, text_positions] = encoder_hidden_states_mask.bool()[:, vlm_text_positions] + + prefix_len = int((~target_token_mask).sum()) + + if kv_cache_mode == "cached": + joint_hidden_states = joint_hidden_states[:, prefix_len:] + if control_joint is not None: + control_joint = control_joint[:, prefix_len:] + rotary_emb = rotary_emb[prefix_len:] + modulation_mask = modulation_mask[prefix_len:] + attention_mask = None if joint_key_valid is None else joint_key_valid[:, None, None, :] + cache_write_slice = None + block_segments, block_key_valid = None, None + else: + attention_mask = None + block_segments = _qwenimage21_prefix_segments(image_ids, prefix_len) + cache_write_slice = slice(0, prefix_len) if kv_cache_mode == "extract" else None + block_key_valid = joint_key_valid + + # Ulysses sequence parallel: pad / chunk the control stream identically to the base stream so the skips line up. + sp_size = self.sp_world_size + sp_pad_len = 0 + sp_active_len = joint_hidden_states.shape[1] + if sp_size > 1: + sp_padded_len = math.ceil(sp_active_len / sp_size) * sp_size + sp_pad_len = sp_padded_len - sp_active_len + if sp_pad_len > 0: + joint_hidden_states = F.pad(joint_hidden_states, (0, 0, 0, sp_pad_len)) + if control_joint is not None: + control_joint = F.pad(control_joint, (0, 0, 0, sp_pad_len)) + rotary_emb = torch.cat( + [rotary_emb, rotary_emb.new_zeros(sp_pad_len, rotary_emb.shape[-1])], dim=0 + ) + if modulation_mask is not None: + modulation_mask = torch.cat( + [modulation_mask, modulation_mask.new_zeros(sp_pad_len, dtype=torch.bool)], dim=0 + ) + if kv_cache_mode == "cached": + if joint_key_valid is None: + joint_key_valid = torch.ones( + batch_size, prefix_len + sp_active_len, dtype=torch.bool, + device=joint_hidden_states.device, + ) + joint_key_valid = torch.cat( + [joint_key_valid, joint_key_valid.new_zeros(batch_size, sp_pad_len, dtype=torch.bool)], + dim=1, + ) + attention_mask = joint_key_valid[:, None, None, :] + else: + if block_key_valid is None: + block_key_valid = torch.ones( + batch_size, sp_active_len, dtype=torch.bool, device=joint_hidden_states.device + ) + block_key_valid = torch.cat( + [block_key_valid, block_key_valid.new_zeros(batch_size, sp_pad_len, dtype=torch.bool)], + dim=1, + ) + sp_local = sp_padded_len // sp_size + sp_lo = self.sp_world_rank * sp_local + joint_hidden_states = joint_hidden_states[:, sp_lo:sp_lo + sp_local] + if control_joint is not None: + control_joint = control_joint[:, sp_lo:sp_lo + sp_local] + rotary_emb = rotary_emb[sp_lo:sp_lo + sp_local] + if modulation_mask is not None: + modulation_mask = modulation_mask[sp_lo:sp_lo + sp_local] + + # Control blocks never read/write the prefix KV cache: they recompute the (static) conditioning every step, so + # they get the base block kwargs with the cache args cleared. + hints = None + if control_joint is not None: + control_kwargs = dict( + modulation=modulation, + rotary_emb=rotary_emb, + attention_mask=attention_mask, + target_token_mask=modulation_mask, + layer_cache=None, + kv_cache_mode=None, + cache_write_slice=None, + segments=block_segments, + key_valid=block_key_valid, + ) + hints = self.forward_control(control_joint, joint_hidden_states, control_kwargs) + + for index_block, block in enumerate(self.transformer_blocks): + layer_cache = kv_cache.get_layer(index_block) if kv_cache is not None else None + kwargs = dict( + modulation=modulation, + rotary_emb=rotary_emb, + attention_mask=attention_mask, + target_token_mask=modulation_mask, + layer_cache=layer_cache, + kv_cache_mode=kv_cache_mode, + cache_write_slice=cache_write_slice, + segments=block_segments, + key_valid=block_key_valid, + hints=hints, + context_scale=control_context_scale, + ) + if torch.is_grad_enabled() and self.gradient_checkpointing: + # diffusers 0.32.2 has no `ModelMixin._gradient_checkpointing_func`; call `torch.utils.checkpoint` + # directly. `hints` (a tuple of tensors) and the other non-tensor args ride in the closure so only the + # joint stream is passed positionally; `use_reentrant=False` tracks them correctly. + def create_custom_forward(module, **static_kwargs): + def custom_forward(*inputs): + return module(*inputs, **static_kwargs) + + return custom_forward + + joint_hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block, **kwargs), + joint_hidden_states, + use_reentrant=False, + ) + else: + joint_hidden_states = block(joint_hidden_states, **kwargs) + + joint_hidden_states = self.norm_out(joint_hidden_states, temb, modulation_mask) + output = self.proj_out(joint_hidden_states) + if sp_size > 1: + output = sequence_parallel_all_gather(output, dim=1) + if sp_pad_len > 0: + output = output[:, :sp_active_len] + + if not return_dict: + return (output,) + + return Transformer2DModelOutput(sample=output) diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index 364ef33..2e53c18 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -19,18 +19,22 @@ from .pipeline_ltx2 import LTX2Pipeline from .pipeline_ltx2_i2v import LTX2I2VPipeline from .pipeline_ltx2_latent_upsample import LTX2LatentUpsamplePipeline from .pipeline_minimax_h3 import (MiniMaxH3AudioReference, - MiniMaxH3ImageReference, - MiniMaxH3Pipeline, + MiniMaxH3ImageReference, MiniMaxH3Pipeline, MiniMaxH3VideoReference) from .pipeline_minimax_h3_control import MiniMaxH3ControlPipeline from .pipeline_mova import MOVAPipeline from .pipeline_qwenimage import QwenImagePipeline from .pipeline_qwenimage21 import QwenImage21Pipeline +from .pipeline_qwenimage21_control import QwenImage21ControlPipeline from .pipeline_qwenimage_control import QwenImageControlPipeline from .pipeline_qwenimage_edit import QwenImageEditPipeline from .pipeline_qwenimage_edit_plus import QwenImageEditPlusPipeline from .pipeline_qwenimage_instantx import QwenImageControlNetPipeline from .pipeline_qwenimage_layered import QwenImageLayeredPipeline +from .pipeline_taomate_h3 import (MiniMaxH3StreamingPipeline, + TaomateH3StreamPhase, TaomateH3StreamPlan, + TaomateH3TeacherArtifact, + TaomateH3TeacherError) from .pipeline_wan import WanPipeline from .pipeline_wan2_2 import Wan2_2Pipeline from .pipeline_wan2_2_animate import Wan2_2AnimatePipeline @@ -38,11 +42,6 @@ from .pipeline_wan2_2_fun_control import Wan2_2FunControlPipeline from .pipeline_wan2_2_fun_inpaint import Wan2_2FunInpaintPipeline from .pipeline_wan2_2_s2v import Wan2_2S2VPipeline from .pipeline_wan2_2_ti2v import Wan2_2TI2VPipeline -from .pipeline_taomate_h3 import (MiniMaxH3StreamingPipeline, - TaomateH3StreamPhase, - TaomateH3StreamPlan, - TaomateH3TeacherArtifact, - TaomateH3TeacherError) from .pipeline_wan2_2_vace_fun import Wan2_2VaceFunPipeline from .pipeline_wan_flex_forcing import WanFlexForcingPipeline from .pipeline_wan_fun_control import WanFunControlPipeline diff --git a/videox_fun/pipeline/pipeline_qwenimage21_control.py b/videox_fun/pipeline/pipeline_qwenimage21_control.py new file mode 100644 index 0000000..a0c14ef --- /dev/null +++ b/videox_fun/pipeline/pipeline_qwenimage21_control.py @@ -0,0 +1,844 @@ +# Modified from https://github.com/huggingface/diffusers/blob/cp-support-qwenimage2.1/src/diffusers/pipelines/qwenimage/pipeline_qwenimage21.py +# Copyright 2026 Qwen-Image Team, 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 typing import Any, Callable, List, Optional, Union + +import numpy as np +import torch +import torch.nn.functional as F +from diffusers.image_processor import PipelineImageInput, VaeImageProcessor +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from diffusers.utils import (is_torch_xla_available, logging, + replace_example_docstring) +from diffusers.utils.torch_utils import randn_tensor +from PIL import Image as PILImage + +from ..models import (AutoencoderKLQwenImage21, + Qwen3VLForConditionalGeneration, Qwen3VLProcessor, + QwenImage21ControlTransformer2DModel, QwenImage21KVCache, + QwenImage21Transformer2DModel) +from .pipeline_qwenimage import QwenImagePipelineOutput + +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: + ```py + >>> import torch + >>> from videox_fun.pipeline import QwenImage21Pipeline + + >>> pipe = QwenImage21Pipeline.from_pretrained("Qwen/Qwen-Image-2.1", torch_dtype=torch.bfloat16) + >>> pipe.to("cuda") + >>> prompt = "A capybara wearing a wizard hat, reading a book by candlelight, oil painting" + >>> image = pipe(prompt).images[0] + >>> image.save("qwenimage21.png") + ``` +""" + + +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 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`. + """ + 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 + + +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") + + +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, None + + +class QwenImage21ControlPipeline(DiffusionPipeline): + r""" + Text-to-image and image-conditioned generation with Qwen-Image 2.1. + + Prompt and condition images are encoded together by a Qwen3-VL model, so a condition image occupies the vision + slots the encoder reserved for it and the transformer sees one interleaved text/image sequence. + + Args: + scheduler ([`FlowMatchEulerDiscreteScheduler`]): + Scheduler used to denoise the encoded image latents. + vae ([`AutoencoderKLQwenImage21`]): + Variational auto-encoder mapping images to and from the 64-channel latent space. + text_encoder ([`Qwen3VLForConditionalGeneration`]): + Qwen3-VL model producing the joint text/image embeddings. + processor ([`Qwen3VLProcessor`]): + Processor that builds the chat template and tokenizes prompt and condition images. + transformer ([`QwenImage21Transformer2DModel`]): + The single-stream block-causal transformer that denoises the latents. + """ + + model_cpu_offload_seq = "text_encoder->transformer->vae" + _callback_tensor_inputs = ["latents", "prompt_embeds"] + + def __init__( + self, + scheduler: FlowMatchEulerDiscreteScheduler, + vae: AutoencoderKLQwenImage21, + text_encoder: Qwen3VLForConditionalGeneration, + processor: Qwen3VLProcessor, + transformer: QwenImage21ControlTransformer2DModel, + ): + super().__init__() + + self.register_modules( + vae=vae, + text_encoder=text_encoder, + processor=processor, + transformer=transformer, + scheduler=scheduler, + ) + # The VAE compresses 16x spatially and the transformer consumes latents unpatched, so one token covers a 16x16 + # pixel tile. + self.vae_scale_factor = 16 + self.latent_channels = self.vae.config.z_dim if getattr(self, "vae", None) else 64 + self.image_processor = VaeImageProcessor( + vae_scale_factor=self.vae_scale_factor, vae_latent_channels=self.latent_channels + ) + self.mask_processor = VaeImageProcessor( + vae_scale_factor=self.vae_scale_factor, do_normalize=False + ) + self.sys_prompt = "Comprehend and analyze the provided prompt." + # The prompt is built as a raw template string and passed straight to + # `self.processor(text=..., images=...)`, rather than going through `apply_chat_template`: + # the two tokenize differently and the checkpoint expects this one. The "Picture 1: ..." + # vision prefix only appears in the image-conditioned template. + self.prompt_template_t2i = ( + f"<|im_start|>system\n{self.sys_prompt}<|im_end|>\n" + f"<|im_start|>user\n{{}}<|im_end|>\n" + f"<|im_start|>assistant\n" + ) + self.prompt_template_ti2i = ( + f"<|im_start|>system\n{self.sys_prompt}<|im_end|>\n" + f"<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{{}}<|im_end|>\n" + f"<|im_start|>assistant\n" + ) + # Number of leading system-role tokens to drop from the hidden states. Derived from the + # tokenized system message rather than hardcoded, so it tracks the processor's template. + sys_message = [{"role": "system", "content": [{"type": "text", "text": self.sys_prompt}]}] + sys_tokens = self.processor.apply_chat_template(sys_message, tokenize=True, return_dict=False) + self._drop_idx = len(sys_tokens[0]) + self._img_token_id = self.processor.tokenizer.encode("<|image_pad|>")[0] + + 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] + return torch.split(selected, valid_lengths.tolist(), dim=0) + + def _get_qwen_prompt_embeds( + self, + prompt: Union[str, List[str]] = None, + image: Optional[list] = None, + device: Optional[torch.device] = None, + ): + device = device or self._execution_device + prompt = [prompt] if isinstance(prompt, str) else prompt + # Qwen has no bos token, so an empty string leaves the encoder with nothing to read. + prompt = [" " if not p else p for p in prompt] + is_t2i = image is None + + if is_t2i: + prompts = [self.prompt_template_t2i.format(t) for t in prompt] + else: + prompts = [] + condition_pil_list = [] + for t in prompt: + n_imgs = len(image) + replace = "<|vision_start|><|image_pad|><|vision_end|>" + for i in range(2, n_imgs + 1): + replace += f" <|vision_start|><|image_pad|><|vision_end|>" + template = self.prompt_template_ti2i.replace( + "<|vision_start|><|image_pad|><|vision_end|>", replace + ) + prompts.append(template.format(t)) + # Each prompt's template repeats the `<|image_pad|>` placeholders, so hand the processor one set of + # images per prompt, in the order the placeholders appear. + for _ in prompt: + for img in image: + if not isinstance(img, PILImage.Image): + img = PILImage.fromarray(img) + if img.mode == "RGBA": + # The checkpoint was trained with the alpha composited over white for the vision encoder. + # Only this copy is flattened; the VAE still reads all four channels. + white = PILImage.new("RGB", img.size, (255, 255, 255)) + white.paste(img, mask=img.getchannel("A")) + img = white + condition_pil_list.append(img) + + # Left padding, as the checkpoint was trained with. `_extract_masked_hidden` drops the padding either way, + # but the side decides the positions the encoder sees for a batch of prompts of different lengths. + processor_kwargs = { + "text": prompts, + "padding": True, + "padding_side": "left", + "return_tensors": "pt", + } + if not is_t2i: + processor_kwargs["images"] = condition_pil_list + + model_inputs = self.processor(**processor_kwargs).to(device) + + forward_kwargs = { + "input_ids": model_inputs.input_ids, + "attention_mask": model_inputs.attention_mask, + "output_hidden_states": True, + } + if not is_t2i and hasattr(model_inputs, "pixel_values"): + forward_kwargs.update(pixel_values=model_inputs.pixel_values, image_grid_thw=model_inputs.image_grid_thw) + if hasattr(model_inputs, "mm_token_type_ids"): + forward_kwargs["mm_token_type_ids"] = model_inputs.mm_token_type_ids + + # `hidden_states[-1]` has to be the last decoder layer's output, before the text encoder's final RMSNorm: + # that is what the transformer was trained on. A forward hook returning the module's input replaces its + # output, which neutralizes the norm for this call on either transformers version. + text_model = getattr(self.text_encoder.model, "language_model", self.text_encoder.model) + handle = text_model.norm.register_forward_hook(lambda module, args, output: args[0]) + try: + outputs = self.text_encoder(**forward_kwargs) + finally: + handle.remove() + hidden_states = outputs.hidden_states[-1] + + split_hidden_states = list(self._extract_masked_hidden(hidden_states, model_inputs.attention_mask)) + split_hidden_states = [e[self._drop_idx :] for e in split_hidden_states] + + image_pad_mask = [ + (sample_ids[sample_mask.bool()] == self._img_token_id) + for sample_ids, sample_mask in zip(model_inputs.input_ids, model_inputs.attention_mask) + ] + image_pad_mask = [e[self._drop_idx :] for e in image_pad_mask] + + 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] + ) + image_pad_mask = torch.stack([torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in image_pad_mask]) + + return prompt_embeds, encoder_attention_mask, image_pad_mask + + def encode_prompt( + self, + prompt: Union[str, List[str]], + image: Optional[List[PipelineImageInput]] = None, + device: Optional[torch.device] = None, + num_images_per_prompt: int = 1, + prompt_embeds: Optional[torch.Tensor] = None, + prompt_embeds_mask: Optional[torch.Tensor] = None, + image_pad_mask: Optional[torch.Tensor] = None, + ): + r""" + Encode the prompt (and optional condition images) into joint text/image embeddings. + """ + 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, image_pad_mask = self._get_qwen_prompt_embeds(prompt, image, device) + elif image_pad_mask is None: + if image is not None: + raise ValueError( + "Pass `image_pad_mask` alongside `prompt_embeds` when the embeddings cover condition images, so " + "the transformer knows which positions hold image tokens." + ) + # Embeddings supplied without a mask can only be text, so no position holds an image token. + image_pad_mask = prompt_embeds.new_zeros(prompt_embeds.shape[:2], dtype=torch.bool) + + _, 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) + # `repeat(1, n)` on the 2D mask, so its rows interleave the same way the 3D embeddings' do. + if prompt_embeds_mask is not None: + prompt_embeds_mask = prompt_embeds_mask.repeat(1, num_images_per_prompt) + prompt_embeds_mask = prompt_embeds_mask.view(batch_size * num_images_per_prompt, seq_len) + + # Without padding there is nothing to mask, and a mask that carries no information costs the attention + # backends that reject one outright. + if prompt_embeds_mask is not None and prompt_embeds_mask.all(): + prompt_embeds_mask = None + + return prompt_embeds, prompt_embeds_mask, image_pad_mask + + def check_inputs(self, prompt, height, width, prompt_embeds, callback_on_step_end_tensor_inputs): + 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 " + f"{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 " + f"{[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("Pass either `prompt` or `prompt_embeds`, not both.") + if prompt is None and prompt_embeds is None: + raise ValueError("Pass one of `prompt` or `prompt_embeds`.") + + @staticmethod + def _pack_latents(latents, batch_size, num_channels_latents, height, width): + # 2.1 consumes latents unpatched, so packing is a plain spatial flatten. + return latents.view(batch_size, num_channels_latents, height * width).transpose(1, 2) + + @staticmethod + def _unpack_latents(latents, height, width, vae_scale_factor): + batch_size, _, channels = latents.shape + height = 2 * (int(height) // (vae_scale_factor * 2)) + width = 2 * (int(width) // (vae_scale_factor * 2)) + latents = latents.transpose(1, 2).reshape(batch_size, channels, 1, height, width) + return latents + + def _encode_vae_image(self, image: torch.Tensor, generator: torch.Generator): + if isinstance(generator, list): + image_latents = [ + retrieve_latents(self.vae.encode(image[i : i + 1]), generator=generator[i], sample_mode="argmax") + for i in range(image.shape[0]) + ] + image_latents = torch.cat(image_latents, dim=0) + else: + image_latents = retrieve_latents(self.vae.encode(image), generator=generator, sample_mode="argmax") + latents_mean = ( + torch.tensor(self.vae.config.latents_mean) + .view(1, self.latent_channels, 1, 1, 1) + .to(image_latents.device, image_latents.dtype) + ) + latents_std = ( + torch.tensor(self.vae.config.latents_std) + .view(1, self.latent_channels, 1, 1, 1) + .to(image_latents.device, image_latents.dtype) + ) + image_latents = (image_latents - latents_mean) / latents_std + + return image_latents + + def prepare_latents( + self, images, batch_size, num_channels_latents, height, width, dtype, device, generator, latents=None + ): + height = 2 * (int(height) // (self.vae_scale_factor * 2)) + width = 2 * (int(width) // (self.vae_scale_factor * 2)) + + 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." + ) + + image_latents = None + if images is not None: + all_image_latents = [] + for image in images: + image = image.to(device=device, dtype=dtype) + encoded = self._encode_vae_image(image, generator) + if batch_size > encoded.shape[0]: + if batch_size % encoded.shape[0] != 0: + raise ValueError( + f"Cannot duplicate `image` of batch size {encoded.shape[0]} to {batch_size} text prompts." + ) + encoded = torch.cat([encoded] * (batch_size // encoded.shape[0]), dim=0) + image_latent_height, image_latent_width = encoded.shape[3:] + all_image_latents.append( + self._pack_latents( + encoded, batch_size, num_channels_latents, image_latent_height, image_latent_width + ) + ) + image_latents = torch.cat(all_image_latents, dim=1) + + if latents is None: + shape = (batch_size, 1, num_channels_latents, height, width) + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width) + else: + latents = latents.to(device=device, dtype=dtype) + + return latents, image_latents + + @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, + image: Optional[PipelineImageInput] = None, + control_image: Optional[PipelineImageInput] = None, + mask_image: Optional[PipelineImageInput] = None, + negative_prompt: Optional[Union[str, List[str]]] = None, + true_cfg_scale: float = 1.0, + control_context_scale: float = 1.0, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 40, + sigmas: Optional[List[float]] = None, + 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] = None, + callback_on_step_end: Optional[Callable[[int, int, dict], None]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + output_resolution: int = 1024, + use_kv_cache: bool = True, + ): + r""" + Function invoked when calling the pipeline for generation. + + Args: + prompt (`str` or `list[str]`, *optional*): + The prompt to guide image generation. Pass `prompt_embeds` instead to supply embeddings directly. + image (`PipelineImageInput`, *optional*): + One or more condition images, as a PIL image or a numpy array. A list is one set of images shared by + every prompt in the batch, not one entry per prompt. + negative_prompt (`str` or `list[str]`, *optional*): + The prompt not to guide image generation. Ignored when `true_cfg_scale` is not greater than 1. + true_cfg_scale (`float`, *optional*, defaults to 1.0): + Classifier-free guidance scale. Enabled by `true_cfg_scale > 1` together with a negative prompt. + height (`int`, *optional*): + Height in pixels of the generated image. Derived from the condition image's aspect ratio if omitted. + width (`int`, *optional*): + Width in pixels of the generated image. Derived from the condition image's aspect ratio if omitted. + num_inference_steps (`int`, *optional*, defaults to 40): + Number of denoising steps. + sigmas (`list[float]`, *optional*): + Custom sigmas for the denoising schedule. + num_images_per_prompt (`int`, *optional*, defaults to 1): + Number of images generated per prompt. + generator (`torch.Generator` or `list[torch.Generator]`, *optional*): + Generator(s) to make generation deterministic. + latents (`torch.Tensor`, *optional*): + Pre-generated noisy latents. + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings, which skip prompt encoding. + prompt_embeds_mask (`torch.Tensor`, *optional*): + Bool mask marking the valid positions of `prompt_embeds`. + negative_prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated negative text embeddings, used in place of `negative_prompt`. + negative_prompt_embeds_mask (`torch.Tensor`, *optional*): + Bool mask marking the valid positions of `negative_prompt_embeds`. + output_type (`str`, *optional*, defaults to `"pil"`): + Output format, `"pil"`, `"np"`, `"pt"` or `"latent"`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether to return a [`QwenImagePipelineOutput`] instead of a plain tuple. + attention_kwargs (`dict`, *optional*): + Passed through to the attention processor. + callback_on_step_end (`Callable`, *optional*): + Called at the end of each denoising step. + callback_on_step_end_tensor_inputs (`list[str]`, *optional*, defaults to `["latents"]`): + Tensors from the denoising loop to hand to `callback_on_step_end`. + output_resolution (`int`, *optional*, defaults to 1024): + Reference side length, not the output size once `height`/`width` are provided. It is + the fallback used to derive `height`/`width` when they are omitted, and the target + used to resize condition images (`image`) before they are encoded. + use_kv_cache (`bool`, *optional*, defaults to `True`): + Cache the text and condition-image keys and values after the first step. + + Examples: + + Returns: + [`QwenImagePipelineOutput`] or `tuple`: + [`QwenImagePipelineOutput`] if `return_dict` is True, otherwise a `tuple` whose first element is a list + with the generated images. + """ + if image is not None: + # The text encoder reads each condition image as vision context, so the pixels have to be there. Normalize + # to PIL up front, and everything downstream — the aspect ratio below, the resize, the VAE — sees one type. + image = image if isinstance(image, list) else [image] + condition_images = [] + for img in image: + if isinstance(img, PILImage.Image): + condition_images.append(img) + elif isinstance(img, np.ndarray): + condition_images.append(PILImage.fromarray(img)) + elif isinstance(img, (list, tuple)): + raise ValueError( + "`image` is one flat set of condition images that applies to every prompt in the batch, so it " + "cannot be nested per prompt. Call the pipeline once per prompt when they need different " + "condition images." + ) + else: + raise ValueError( + f"`image` accepts a PIL image or a numpy array, or a list of either, but got " + f"{type(img).__name__}. Latents cannot stand in for a condition image here, because the text " + f"encoder has to see the image itself." + ) + image = condition_images + calculated_width, calculated_height, _ = calculate_dimensions( + output_resolution * output_resolution, image[-1].size[0] / image[-1].size[1] + ) + height = height or calculated_height + width = width or calculated_width + height = height or output_resolution + width = width or output_resolution + + self.check_inputs(prompt, height, width, prompt_embeds, callback_on_step_end_tensor_inputs) + + multiple_of = self.vae_scale_factor * 2 + width = width // multiple_of * multiple_of + height = height // multiple_of * multiple_of + + self._attention_kwargs = attention_kwargs or {} + self._current_timestep = None + self._interrupt = False + + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + + # 1. The control pipeline does not use the base ti2i (text+image) conditioning path. `image` is instead the + # inpaint source folded into `control_context` (step 3b), so the joint stream carries only the target latents + # and `control_context` lines up with them one-to-one. + input_image_sizes, input_images, vae_images = [], None, None + + # 2. Encode prompt + has_neg_prompt = negative_prompt is not None or negative_prompt_embeds is not None + do_true_cfg = true_cfg_scale > 1 and has_neg_prompt + 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 " + f"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" + ) + + prompt_embeds, prompt_embeds_mask, image_pad_mask = self.encode_prompt( + image=input_images, + prompt=prompt, + prompt_embeds=prompt_embeds, + prompt_embeds_mask=prompt_embeds_mask, + device=device, + num_images_per_prompt=num_images_per_prompt, + ) + if do_true_cfg: + negative_prompt_embeds, negative_prompt_embeds_mask, negative_image_pad_mask = self.encode_prompt( + image=input_images, + 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, + ) + + # 3. Prepare latents + num_channels_latents = self.transformer.config.in_channels + latents, input_images_latents = self.prepare_latents( + vae_images, + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + + # 3b. VACE-style control conditioning: control-image latents (64) + inpaint mask (1) + masked-image latents + # (64) = 129 channels, packed onto the target latent sequence so it lines up one-to-one with `image_pad_mask`. + # Mirrors the 2.0 Fun-Controlnet-Union pipeline, adapted to 2.1's z_dim=64 / patch_size=1 VAE (reads RGBA). + weight_dtype = prompt_embeds.dtype + ctrl_batch = batch_size * num_images_per_prompt + lat_h = 2 * (int(height) // (self.vae_scale_factor * 2)) + lat_w = 2 * (int(width) // (self.vae_scale_factor * 2)) + + def _to_ctrl_batch(t): + return t.repeat(ctrl_batch, *[1] * (t.dim() - 1)) if (t.shape[0] == 1 and ctrl_batch > 1) else t + + def _rgba(t): + # 2.1's VAE reads RGBA (in_channels=4); composite an opaque alpha over 3-channel batches. + return torch.cat([t, torch.ones_like(t[:, :1])], dim=1) if t.shape[1] == 3 else t + + if mask_image is not None: + mask_condition = self.mask_processor.preprocess(mask_image, height=height, width=width) + mask_condition = (mask_condition >= 0.5).to(weight_dtype)[:, :1] + else: + mask_condition = torch.ones(ctrl_batch, 1, height, width, dtype=weight_dtype) + mask_condition = _to_ctrl_batch(mask_condition).to(device) + + if image is not None: + inpaint_image = self.image_processor.preprocess(image, height=height, width=width) + inpaint_image = inpaint_image * (mask_condition < 0.5) # zero out the region to regenerate + inpaint_latent = self._encode_vae_image( + _rgba(inpaint_image).unsqueeze(2).to(device=device, dtype=weight_dtype), generator + ) + else: + inpaint_latent = torch.zeros( + ctrl_batch, self.latent_channels, 1, lat_h, lat_w, device=device, dtype=weight_dtype + ) + inpaint_latent = _to_ctrl_batch(inpaint_latent) + + if control_image is not None: + ctrl_image = self.image_processor.preprocess(control_image, height=height, width=width) + control_latents = self._encode_vae_image( + _rgba(ctrl_image).unsqueeze(2).to(device=device, dtype=weight_dtype), generator + ) + else: + control_latents = torch.zeros_like(inpaint_latent) + control_latents = _to_ctrl_batch(control_latents) + + # 1-channel inpaint condition at latent resolution (1 = keep, 0 = regenerate), matching the 2.0 convention. + mask_latent = F.interpolate(1 - mask_condition, size=(lat_h, lat_w), mode="nearest").unsqueeze(2) + + control_context = torch.cat([control_latents, mask_latent, inpaint_latent], dim=1) # [B, 129, 1, lat_h, lat_w] + control_context = self._pack_latents(control_context, ctrl_batch, control_context.shape[1], lat_h, lat_w) + control_context = control_context.to(device=device, dtype=weight_dtype) + + img_shapes = [ + [ + *[ + (1, vae_height // self.vae_scale_factor, vae_width // self.vae_scale_factor) + for vae_width, vae_height in input_image_sizes + ], + (1, height // self.vae_scale_factor, width // self.vae_scale_factor), + ] + ] * batch_size + + # 4. Prepare timesteps + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas + mu = calculate_shift( + latents.shape[1], + 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) + + # The transformer's `img_mask` spans the joint sequence, so append one slot per 2x2 group of target latents. + def append_target_slots(mask): + return torch.cat([mask, mask.new_ones(mask.shape[0], latents.shape[1] // 4)], dim=1) + + image_pad_mask = append_target_slots(image_pad_mask) + if do_true_cfg: + negative_image_pad_mask = append_target_slots(negative_image_pad_mask) + + # Text and condition-image keys and values are step-independent under `causal_condition`, so the first step + # prefills them and later steps only recompute the target image's tokens. + num_blocks = len(self.transformer.transformer_blocks) + cache_enabled = use_kv_cache and self.transformer.config.causal_condition and control_context is None + cond_cache = QwenImage21KVCache(num_blocks) if cache_enabled else None + neg_cache = QwenImage21KVCache(num_blocks) if cache_enabled and do_true_cfg else None + + # 5. 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): + if self.interrupt: + # `continue` would skip the step that prefills the cache and leave the next one decoding from an + # empty one, so stop the loop instead. + break + + self._current_timestep = t + kv_mode = "extract" if (cache_enabled and i == 0) else ("cached" if cache_enabled else None) + + latent_model_input = latents + if input_images_latents is not None: + latent_model_input = torch.cat([input_images_latents, latents], dim=1) + + timestep = t.expand(latents.shape[0]).to(latents.dtype) + noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, + img_shapes=img_shapes, + img_mask=image_pad_mask, + attention_kwargs=self.attention_kwargs, + kv_cache=cond_cache, + kv_cache_mode=kv_mode, + control_context=control_context, + control_context_scale=control_context_scale, + return_dict=False, + )[0] + noise_pred = noise_pred[:, -latents.size(1) :] + + if do_true_cfg: + neg_noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + encoder_hidden_states=negative_prompt_embeds, + encoder_hidden_states_mask=negative_prompt_embeds_mask, + img_shapes=img_shapes, + img_mask=negative_image_pad_mask, + attention_kwargs=self.attention_kwargs, + kv_cache=neg_cache, + kv_cache_mode=kv_mode, + control_context=control_context, + control_context_scale=control_context_scale, + return_dict=False, + )[0] + neg_noise_pred = neg_noise_pred[:, -latents.size(1) :] + noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) + + latents_dtype = latents.dtype + latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] + if latents.dtype != latents_dtype and 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) + + 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 = ( + 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) + + self.maybe_free_model_hooks() + + if not return_dict: + return (image,) + + return QwenImagePipelineOutput(images=image) diff --git a/videox_fun/pipeline/pipeline_wan_flex_forcing.py b/videox_fun/pipeline/pipeline_wan_flex_forcing.py index cf2c272..00414a6 100644 --- a/videox_fun/pipeline/pipeline_wan_flex_forcing.py +++ b/videox_fun/pipeline/pipeline_wan_flex_forcing.py @@ -16,8 +16,10 @@ from ..models import (AutoencoderKLWan, AutoTokenizer, WanT5EncoderModel, from ..utils.fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas) from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler -from ..utils.flex_chunking import (build_pyramid_partitions, chunk_boundaries, - normalize_chunk_spec, uniform_chunks, +from ..utils.flex_chunking import (build_full_then_blocks_partitions, + build_pyramid_partitions, + chunk_boundaries, normalize_chunk_spec, + uniform_chunks, validate_nested_partitions) from .pipeline_wan_self_forcing import (WanSelfForcingPipeline, WanSelfForcingPipelineOutput, @@ -26,6 +28,12 @@ from .pipeline_wan_self_forcing import (WanSelfForcingPipeline, logger = logging.get_logger(__name__) # pylint: disable=invalid-name +# `denoise_mode` strings that select the "first step full, every later step +# block-major" ladder built by `build_full_then_blocks_partitions`, i.e. one +# bidirectional planning chunk followed by the uniform `num_frame_per_block` +# partition. Kept as a tuple so the aliases stay in one place. +FULL_THEN_BLOCKS_MODES = ("full_then_blocks", "full_first", "one_shot") + # The partitions evaluated in the paper (5 s / 21 latent frames and the long # video regimes), kept as a reference for choosing `chunk_spec`. Any list of # positive ints summing to the latent frame count is accepted. @@ -168,6 +176,10 @@ class WanFlexForcingPipeline(WanSelfForcingPipeline): from ``num_inference_steps`` so the two never need syncing. An int pins a truncated pyramid of exactly that many levels, e.g. ``2`` for the paper's ``[21] -> [11, 10]`` ladder on a 5 s clip. + ``"full_then_blocks"`` (aliases ``"full_first"`` / + ``"one_shot"``) runs a fixed two-level ladder instead: the first + step plans the whole clip in one bidirectional chunk and every + later step reuses the uniform ``num_frame_per_block`` partition. min_num_frame_per_block: Block size the ladder stops refining at - every level-0 chunk is binary-split until it is at or below this size. ``1`` lets the pyramid reach fully causal (single-frame) @@ -189,30 +201,37 @@ class WanFlexForcingPipeline(WanSelfForcingPipeline): ``` """ # `denoise_mode` -> ladder depth, i.e. how many nested partitions. - # "fixed" -> 1: one partition held for every denoising step. - # "pyramid" -> one level per denoising step (3.2), so the caller never - # has to keep two numbers in sync. The ladder stops - # earlier at its fixed point; asking for more levels than - # the partition can yield is harmless. - # int -> an explicitly pinned depth, for callers that must run a - # truncated pyramid (e.g. the trainer's - # `--flex_pyramid_levels`). - if isinstance(denoise_mode, str): - mode = denoise_mode.strip().lower() - if mode == "fixed": + # "fixed" -> 1: one partition held for every denoising step. + # "pyramid" -> one level per denoising step (3.2), so the caller + # never has to keep two numbers in sync. The ladder + # stops earlier at its fixed point; asking for more + # levels than the partition can yield is harmless. + # "full_then_blocks" -> 2 levels, built specially below: the first step + # plans the whole clip in one bidirectional chunk, + # every later step reuses the uniform + # `num_frame_per_block` partition. + # int -> an explicitly pinned pyramid depth, for callers + # that must run a truncated pyramid (e.g. the + # trainer's `--flex_pyramid_levels`). + ladder_mode = (denoise_mode.strip().lower() + if isinstance(denoise_mode, str) else None) + if ladder_mode is not None: + if ladder_mode == "fixed": depth = 1 - elif mode == "pyramid": + elif ladder_mode == "pyramid": depth = max(1, int(num_inference_steps or 1)) + elif ladder_mode in FULL_THEN_BLOCKS_MODES: + depth = 2 else: raise ValueError( - "denoise_mode must be 'fixed', 'pyramid' or an int >= 1, " - f"got {denoise_mode!r}") + "denoise_mode must be 'fixed', 'pyramid', 'full_then_blocks' " + f"or an int >= 1, got {denoise_mode!r}") else: depth = int(denoise_mode) if depth < 1: raise ValueError( - "denoise_mode must be 'fixed', 'pyramid' or an int >= 1, " - f"got {denoise_mode!r}") + "denoise_mode must be 'fixed', 'pyramid', 'full_then_blocks' " + f"or an int >= 1, got {denoise_mode!r}") if chunk_spec is None and depth <= 1: # Nothing Flex-Forcing specific was asked for: hand the call to the @@ -286,12 +305,21 @@ class WanFlexForcingPipeline(WanSelfForcingPipeline): base = normalize_chunk_spec(chunk_spec, latent_frames) \ if chunk_spec is not None else None if depth > 1: - # Level 0 defaults to a single chunk over the whole clip - the - # paper's "high-level planning" step - unless the caller pinned it. - partitions = build_pyramid_partitions( - latent_frames, num_levels=int(depth), - min_num_frame_per_block=min_num_frame_per_block, base_chunks=base, - independent_first_frame=independent_first_frame) + if ladder_mode in FULL_THEN_BLOCKS_MODES: + # "First step full, every later step block-major": a fixed + # two-level ladder whose fine level is the uniform + # `num_frame_per_block` partition rather than a binary split, so + # it does not go through `build_pyramid_partitions`. + partitions = build_full_then_blocks_partitions( + latent_frames, num_frame_per_block, + independent_first_frame=independent_first_frame) + else: + # Level 0 defaults to a single chunk over the whole clip - the + # paper's "high-level planning" step - unless the caller pinned it. + partitions = build_pyramid_partitions( + latent_frames, num_levels=int(depth), + min_num_frame_per_block=min_num_frame_per_block, base_chunks=base, + independent_first_frame=independent_first_frame) validate_nested_partitions(partitions, latent_frames) chunk_sizes = partitions[0] else: diff --git a/videox_fun/utils/flex_chunking.py b/videox_fun/utils/flex_chunking.py index aaeaf8e..4d0d2c9 100644 --- a/videox_fun/utils/flex_chunking.py +++ b/videox_fun/utils/flex_chunking.py @@ -40,6 +40,7 @@ __all__ = [ "sample_flexible_chunks", "refine_partition", "build_pyramid_partitions", + "build_full_then_blocks_partitions", "validate_nested_partitions", "chunk_ends_tensor", "broadcast_chunk_sizes", @@ -293,6 +294,29 @@ def build_pyramid_partitions(num_frames: int, return partitions +def build_full_then_blocks_partitions(num_frames: int, + num_frame_per_block: int, + independent_first_frame: bool = False) -> List[ChunkSizes]: + """Two-level ladder: one full planning chunk, then block-major refinement. + + Level 0 is a single bidirectional chunk over the whole clip - the *first* + denoising step plans everything at once. Level 1 is the classic Self-Forcing + uniform ``num_frame_per_block`` partition, which the walk reuses for *every* + remaining step, so the rollout turns block-major right after the planning + pass. This is the coarse-to-fine special case of the pyramid where the fine + level is a fixed block size instead of a binary split. + + ``[F]`` -> blocks is always a valid refinement (the block partition keeps + the level-0 boundary at ``F``), so the ladder satisfies + :func:`validate_nested_partitions` and the KV cache written by the planning + step stays reusable by the block-major steps. + """ + num_frames = int(num_frames) + block = uniform_chunks(num_frames, num_frame_per_block, + independent_first_frame) + return [[num_frames], block] + + def validate_nested_partitions(partitions: Sequence[Sequence[int]], num_frames: int) -> Boundaries: """Check a pyramid ladder: same coverage + monotonically nested boundaries.