qwen image 21 control, update flex forcing and fix bug in qwen image 21
This commit is contained in:
@@ -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()
|
||||
@@ -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`
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user