qwen image 21 control, update flex forcing and fix bug in qwen image 21

This commit is contained in:
bubbliiiing
2026-09-23 13:53:03 +08:00
parent bce42b301f
commit 261e1d85d3
17 changed files with 6828 additions and 84 deletions
@@ -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()
@@ -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()
+7 -7
View File
@@ -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`
+3 -3
View File
@@ -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:
+548
View File
@@ -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/<ts>/checkpoint-<step>/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
@@ -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/<ts>/checkpoint-<step>/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
@@ -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
# `<checkpoint>/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 `<checkpoint>` 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()
File diff suppressed because it is too large Load Diff
+39
View File
@@ -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"
File diff suppressed because it is too large Load Diff
@@ -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
+96 -43
View File
@@ -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
@@ -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)
+6 -7
View File
@@ -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
@@ -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<image1><|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 = "<image1><|vision_start|><|image_pad|><|vision_end|>"
for i in range(2, n_imgs + 1):
replace += f" <image{i}><|vision_start|><|image_pad|><|vision_end|>"
template = self.prompt_template_ti2i.replace(
"<image1><|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)
@@ -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:
+24
View File
@@ -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.