Update Self-Forcing and Ernie Image (#490)

This commit is contained in:
Bubbliiiing
2026-05-25 17:39:17 +08:00
committed by GitHub
parent 804a4258e2
commit 2b5596b8e6
27 changed files with 12946 additions and 25 deletions
+210
View File
@@ -0,0 +1,210 @@
import os
import sys
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
current_file_path = os.path.abspath(__file__)
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer,
ErnieImageTransformer2DModel, Mistral3Model)
from videox_fun.pipeline import ErnieImagePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/ERNIE-Image"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [1728, 992]
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
negative_prompt = "低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"
guidance_scale = 4.5
seed = 43
num_inference_steps = 40
lora_weight = 0.55
save_path = "samples/ernie-image-t2i"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = ErnieImageTransformer2DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, 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 = AutoencoderKLFlux2.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 tokenizer and text_encoder
tokenizer = AutoTokenizer.from_pretrained(
model_name, subfolder="tokenizer"
)
text_encoder = Mistral3Model.from_pretrained(
model_name, subfolder="text_encoder", torch_dtype=weight_dtype
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = ErnieImagePipeline(
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
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.layers))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if compile_dit:
for i in range(len(pipeline.transformer.layers)):
pipeline.transformer.layers[i] = torch.compile(pipeline.transformer.layers[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
sample = pipeline(
prompt,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
).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)
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0]
image.save(video_path)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+274
View File
@@ -0,0 +1,274 @@
import os
import sys
import numpy as np
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 (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, 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 = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# 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 and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
shift = 5
stochastic_sampling = True
# Load pretrained model if need
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Self-Forcing causal inference config
# Number of frames to generate per block (1 for standard causal, higher for faster but more memory)
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
local_attn_size = -1
# Others
independent_first_frame = False
context_noise = 0.0
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 1.0
seed = 43
num_inference_steps = 4
lora_weight = 0.55
save_path = "samples/wan-videos-self-forcing-t2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=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
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).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 Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
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)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
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=["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=["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)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = stochastic_sampling,
).videos
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)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+525
View File
@@ -0,0 +1,525 @@
# ERNIE-Image Full Parameter Training Guide
This document provides a complete workflow for full parameter training of ERNIE-Image Diffusion Transformer, including environment configuration, data preparation, distributed training, and inference testing.
---
## Table of Contents
- [1. Environment Configuration](#1-environment-configuration)
- [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. Full Parameter Training](#3-full-parameter-training)
- [3.1 Download Pretrained Model](#31-download-pretrained-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)
- [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 Configuration
**Method 1: Using requirements.txt**
```bash
pip install -r requirements.txt
```
**Method 2: Manual Dependency 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 deepspeed==0.17.0 numpy==1.26.4
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
```
**Method 3: Using Docker**
When using Docker, please ensure that the GPU driver and CUDA environment are correctly installed on your machine, 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
```
---
## 2. Data Preparation
### 2.1 Quick Test Dataset
We provide a test dataset containing several training samples.
```bash
# Download official demo dataset
modelscope download --dataset PAI/X-Fun-Images-Demo --local_dir ./datasets/X-Fun-Images-Demo
```
### 2.2 Dataset Structure
```
📦 datasets/
├── 📂 my_dataset/
│ ├── 📂 train/
│ │ ├── 📄 image001.jpg
│ │ ├── 📄 image002.png
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 2.3 metadata.json Format
**Relative Path Format** (example):
```json
[
{
"file_path": "train/image001.jpg",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"width": 1024,
"height": 1024
},
{
"file_path": "train/image002.png",
"text": "Portrait of a young woman, studio lighting, high quality",
"width": 1328,
"height": 1328
}
]
```
**Absolute Path Format**:
```json
[
{
"file_path": "/mnt/data/images/sunset.jpg",
"text": "A beautiful sunset over the ocean",
"width": 1024,
"height": 1024
}
]
```
**Key Fields Description**:
- `file_path`: Image path (relative or absolute)
- `text`: Image description (English prompt)
- `width` / `height`: Image dimensions (**recommended** to provide for bucket training; if not provided, they will be automatically read during training, which may slow down training when data is stored on slow systems like OSS)
- You can use `scripts/process_json_add_width_and_height.py` to add width and height fields to JSON files without these fields, supporting both images and videos
- Usage: `python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Images-Demo/metadata.json --output_file datasets/X-Fun-Images-Demo/metadata_add_width_height.json`
### 2.4 Relative vs Absolute Path Usage
**Relative Paths**:
If your data uses relative paths, configure the training script as follows:
```bash
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
```
**Absolute Paths**:
If your data uses absolute paths, configure the training script as follows:
```bash
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/metadata.json"
```
> 💡 **Recommendation**: If the dataset is small and stored locally, use relative paths. If the dataset is stored on external storage (e.g., NAS, OSS) or shared across multiple machines, use absolute paths.
---
## 3. Full Parameter Training
### 3.1 Download Pretrained Model
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download ERNIE-Image official weights
modelscope download --model PaddlePaddle/ERNIE-Image --local_dir models/Diffusion_Transformer/ERNIE-Image
```
### 3.2 Quick Start (DeepSpeed-Zero-2)
If you have downloaded the data as per **2.1 Quick Test Dataset** and the weights as per **3.1 Download Pretrained Model**, you can directly copy and run the quick start command.
DeepSpeed-Zero-2 and FSDP are recommended for training. Here we use DeepSpeed-Zero-2 as an example.
The difference between DeepSpeed-Zero-2 and FSDP lies in whether the model weights are sharded. **If VRAM is insufficient when using multiple GPUs with DeepSpeed-Zero-2**, you can switch to FSDP.
```bash
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/X-Fun-Images-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Images-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/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
### 3.3 Common Training Parameters
**Key Parameter Descriptions**:
| Parameter | Description | Example Value |
|-----|------|-------|
| `--pretrained_model_name_or_path` | Path to pretrained model | `models/Diffusion_Transformer/ERNIE-Image` |
| `--train_data_dir` | Training data directory | `datasets/internal_datasets/` |
| `--train_data_meta` | Training data metadata file | `datasets/internal_datasets/metadata.json` |
| `--train_batch_size` | Samples per batch | 1 |
| `--image_sample_size` | Maximum training resolution, auto bucketing | 1328 |
| `--gradient_accumulation_steps` | Gradient accumulation steps (equivalent to larger batch) | 1 |
| `--dataloader_num_workers` | DataLoader subprocesses | 8 |
| `--num_train_epochs` | Number of training epochs | 100 |
| `--checkpointing_steps` | Save checkpoint every N steps | 50 |
| `--learning_rate` | Initial learning rate | 2e-05 |
| `--lr_scheduler` | Learning rate scheduler | `constant_with_warmup` |
| `--lr_warmup_steps` | Learning rate warmup steps | 100 |
| `--seed` | Random seed | 42 |
| `--output_dir` | Output directory | `output_dir_ernie_image` |
| `--gradient_checkpointing` | Enable activation checkpointing | - |
| `--mixed_precision` | Mixed precision: `fp16/bf16` | `bf16` |
| `--adam_weight_decay` | AdamW weight decay | 3e-2 |
| `--adam_epsilon` | AdamW epsilon value | 1e-10 |
| `--vae_mini_batch` | Mini-batch size for VAE encoding | 1 |
| `--max_grad_norm` | Gradient clipping threshold | 0.05 |
| `--enable_bucket` | Enable bucket training: trains entire images grouped by resolution without center cropping | - |
| `--random_hw_adapt` | Auto-scale images to random size in range `[512, image_sample_size]` | - |
| `--resume_from_checkpoint` | Resume training from checkpoint path, use `"latest"` to auto-select latest | None |
| `--uniform_sampling` | Uniform timestep sampling | - |
| `--trainable_modules` | Trainable modules (`"."` means all modules) | `"."` |
| `--validation_steps` | Execute validation every N steps | 100 |
| `--validation_epochs` | Execute validation every N epochs | 100 |
| `--validation_prompts` | Prompts used during validation | `"a young girl..."` |
### 3.4 Training Validation
You can configure validation parameters to periodically generate test images during training, allowing you to monitor training progress and model quality.
**Validation Parameters**:
| Parameter | Description | Recommended Value |
|-----------|-------------|-------------------|
| `--validation_steps` | Execute validation every N steps | 100 |
| `--validation_epochs` | Execute validation every N epochs | 100 |
| `--validation_prompts` | Prompt for validation image generation. Use multiple space-separated prompt strings | Space-separated prompt strings |
**Example**:
```bash
--validation_steps=100 \
--validation_epochs=100 \
--validation_prompts="a young girl with flowing long hair, wearing a white halter dress"
```
**Notes**:
- Validation images will be saved to the `output_dir` directory
- For multi-prompt validation, use: `--validation_prompts "prompt1" "prompt2" "prompt3"`
### 3.5 Training with FSDP
**If VRAM is insufficient when using multiple GPUs with DeepSpeed-Zero-2**, you can switch to FSDP.
```sh
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
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 ErnieImageSharedAdaLNBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
### 3.6 Training Without DeepSpeed or FSDP
**This approach is not recommended as it lacks VRAM-saving backends and may easily cause out-of-memory errors**. This is provided for reference only.
```sh
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
### 3.7 Multi-Machine Distributed Training
**Suitable for**: Ultra-large-scale datasets, faster training speed
#### 3.7.1 Environment Configuration
Assuming 2 machines with 8 GPUs each:
**Machine 0 (Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/X-Fun-Images-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # Master machine IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # Total number of machines
export NUM_PROCESS=16 # Total processes = machines × 8
export RANK=0 # Current machine rank (0 or 1)
# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
**Machine 1 (Worker)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/X-Fun-Images-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # Same as Master
export MASTER_PORT=10086
export WORLD_SIZE=2
export NUM_PROCESS=16
export RANK=1 # Note this is 1
# 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
# Use the same accelerate launch command as Machine 0
```
#### 3.7.2 Multi-Machine Training Notes
- **Network Requirements**:
- RDMA/InfiniBand recommended (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)
## 4. Inference Testing
### 4.1 Inference Parameters
**Key Parameter Descriptions**:
| Parameter | Description | Example Value |
|------|------|-------|
| `GPU_memory_mode` | GPU memory mode, see table below for options | `model_cpu_offload` |
| `ulysses_degree` | Head dimension parallelization degree, 1 for single GPU | 1 |
| `ring_degree` | Sequence dimension parallelization degree, 1 for single GPU | 1 |
| `fsdp_dit` | Use FSDP for Transformer in multi-GPU inference to save VRAM | `False` |
| `fsdp_text_encoder` | Use FSDP for text encoder in multi-GPU inference | `False` |
| `compile_dit` | Compile Transformer to accelerate inference (effective at fixed resolution) | `False` |
| `model_name` | Model path | `models/Diffusion_Transformer/ERNIE-Image` |
| `sampler_name` | Sampler type: `Flow`, `Flow_Unipc`, `Flow_DPM++` | `Flow` |
| `transformer_path` | Path to trained Transformer weights | `None` |
| `vae_path` | Path to trained VAE weights | `None` |
| `lora_path` | LoRA weights path | `None` |
| `sample_size` | Generated image resolution `[height, width]` | `[1728, 992]` |
| `weight_dtype` | Model weight precision, use `torch.float16` for GPUs without bf16 support | `torch.bfloat16` |
| `prompt` | Positive prompt describing the content to generate | `"1girl, black_hair..."` |
| `negative_prompt` | Negative prompt for content to avoid | `"低分辨率,低画质..."` |
| `guidance_scale` | Guidance strength | 4.5 |
| `seed` | Random seed for reproducibility | 43 |
| `num_inference_steps` | Inference steps | 40 |
| `lora_weight` | LoRA weight strength | 0.55 |
| `save_path` | Generated image save path | `samples/ernie-image-t2i` |
**GPU Memory Mode Description**:
| Mode | Description | VRAM 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` | Layer group offload between CPU/CUDA | Low |
| `sequential_cpu_offload` | Offload each layer individually (slowest) | Lowest |
### 4.2 Single GPU Inference
Run single GPU inference with:
```bash
python examples/ernie_image/predict_t2i.py
```
Edit `examples/ernie_image/predict_t2i.py` according to your needs. For first-time inference, focus on these parameters. For other parameters, see the Inference Parameters section above.
```python
# Choose based on your GPU VRAM
GPU_memory_mode = "model_cpu_offload"
# Your actual model path
model_name = "models/Diffusion_Transformer/ERNIE-Image"
# Trained weights path, e.g. "output_dir_ernie_image/checkpoint-xxx/diffusion_pytorch_model.safetensors"
transformer_path = None
# Write based on content to generate
prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
# ...
```
### 4.3 Multi-GPU Parallel Inference
**Suitable for**: High-resolution generation, accelerated inference
#### Install Parallel Inference Dependencies
```bash
pip install xfuser==0.4.2 yunchang==0.6.2
```
#### Configure Parallel Strategy
Edit `examples/ernie_image/predict_t2i.py`:
```python
# Ensure ulysses_degree × ring_degree = number of GPUs
# For example, using 2 GPUs:
ulysses_degree = 2 # Head dimension parallelization
ring_degree = 1 # Sequence dimension parallelization
```
**Configuration Principles**:
- `ulysses_degree` must evenly divide the model's number of heads
- `ring_degree` splits on sequence dimension, affecting communication overhead; avoid using it when heads can be divided
**Example Configurations**:
| GPU Count | ulysses_degree | ring_degree | Description |
|---------|---------------|-------------|------|
| 1 | 1 | 1 | Single GPU |
| 4 | 4 | 1 | Head parallelization |
| 8 | 8 | 1 | Head parallelization |
| 8 | 4 | 2 | Hybrid parallelization |
#### Run Multi-GPU Inference
```bash
torchrun --nproc-per-node=2 examples/ernie_image/predict_t2i.py
```
## 5. Additional Resources
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
+525
View File
@@ -0,0 +1,525 @@
# ERNIE-Image 全量参数训练指南
本文档提供 ERNIE-Image Diffusion Transformer 全量参数训练的完整流程,包括环境配置、数据准备、分布式训练和推理测试。
---
## 目录
- [一、环境配置](#一环境配置)
- [二、数据准备](#二数据准备)
- [2.1 快速测试数据集](#21-快速测试数据集)
- [2.2 数据集结构](#22-数据集结构)
- [2.3 metadata.json 格式](#23-metadatajson-格式)
- [2.4 相对路径与绝对路径使用方案](#24-相对路径与绝对路径使用方案)
- [三、全量参数训练](#三全量参数训练)
- [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-多机分布式训练)
- [四、推理测试](#四推理测试)
- [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 deepspeed==0.17.0 numpy==1.26.4
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
```
**方式 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
```
---
## 二、数据准备
### 2.1 快速测试数据集
我们提供了一个测试的数据集,其中包含若干训练数据。
```bash
# 下载官方示例数据集
modelscope download --dataset PAI/X-Fun-Images-Demo --local_dir ./datasets/X-Fun-Images-Demo
```
### 2.2 数据集结构
```
📦 datasets/
├── 📂 my_dataset/
│ ├── 📂 train/
│ │ ├── 📄 image001.jpg
│ │ ├── 📄 image002.png
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 2.3 metadata.json 格式
**相对路径格式**(示例格式):
```json
[
{
"file_path": "train/image001.jpg",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"width": 1024,
"height": 1024
},
{
"file_path": "train/image002.png",
"text": "Portrait of a young woman, studio lighting, high quality",
"width": 1328,
"height": 1328
}
]
```
**绝对路径格式**:
```json
[
{
"file_path": "/mnt/data/images/sunset.jpg",
"text": "A beautiful sunset over the ocean",
"width": 1024,
"height": 1024
}
]
```
**关键字段说明**:
- `file_path`:图片路径(相对或绝对路径)
- `text`:图片描述(英文提示词)
- `width` / `height`:图片宽高(**最好提供**,用于分桶训练,如果不提供则自动在训练时读取,当数据存储在如oss这样的速度较慢的系统上时,可能会影响训练速度)。
- 可以使用`scripts/process_json_add_width_and_height.py`文件对无width与height字段的json进行提取,支持处理图片与视频。
- 使用方案为`python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Images-Demo/metadata.json --output_file datasets/X-Fun-Images-Demo/metadata_add_width_height.json`。
### 2.4 相对路径与绝对路径使用方案
**相对路径**:
如果数据的路径为相对路径,则在训练脚本中设置:
```bash
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
```
**绝对路径**:
如果数据的路径为绝对路径,则在训练脚本中设置:
```bash
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/metadata.json"
```
> 💡 **建议**:如果数据集较小且存储在本地,推荐使用相对路径;如果数据集存储在外部存储(如 NAS、OSS)或多个机器共享存储,推荐使用绝对路径。
---
## 三、全量参数训练
### 3.1 下载预训练模型
```bash
# 创建模型目录
mkdir -p models/Diffusion_Transformer
# 下载 ERNIE-Image 官方权重
modelscope download --model PaddlePaddle/ERNIE-Image --local_dir models/Diffusion_Transformer/ERNIE-Image
```
### 3.2 快速开始(DeepSpeed-Zero-2)
如果按照 **2.1 快速测试数据集下载数据** 与 **3.1 下载预训练模型下载权重**后,直接复制快速开始的启动指令进行启动。
推荐使用DeepSpeed-Zero-2与FSDP方案进行训练。这里使用DeepSpeed-Zero-2为例配置shell文件。
本文中DeepSpeed-Zero-2与FSDP的差别在于是否对模型权重进行分片,**如果使用多卡且使用DeepSpeed-Zero-2的情况下显存不足**,可以切换使用FSDP进行训练。
```bash
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/X-Fun-Images-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Images-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/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
### 3.3 训练常用参数解析
**关键参数说明**:
| 参数 | 说明 | 示例值 |
|-----|------|-------|
| `--pretrained_model_name_or_path` | 预训练模型路径 | `models/Diffusion_Transformer/ERNIE-Image` |
| `--train_data_dir` | 训练数据目录 | `datasets/internal_datasets/` |
| `--train_data_meta` | 训练数据元文件 | `datasets/internal_datasets/metadata.json` |
| `--train_batch_size` | 每批次样本数 | 1 |
| `--image_sample_size` | 最大训练分辨率,代码会自动分桶 | 1328 |
| `--gradient_accumulation_steps` | 梯度累积步数(等效增大 batch) | 1 |
| `--dataloader_num_workers` | DataLoader 子进程数 | 8 |
| `--num_train_epochs` | 训练 epoch 数 | 100 |
| `--checkpointing_steps` | 每 N 步保存 checkpoint | 50 |
| `--learning_rate` | 初始学习率 | 2e-05 |
| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` |
| `--lr_warmup_steps` | 学习率预热步数 | 100 |
| `--seed` | 随机种子 | 42 |
| `--output_dir` | 输出目录 | `output_dir_ernie_image` |
| `--gradient_checkpointing` | 激活重计算 | - |
| `--mixed_precision` | 混合精度:`fp16/bf16` | `bf16` |
| `--adam_weight_decay` | AdamW 权重衰减 | 3e-2 |
| `--adam_epsilon` | AdamW epsilon 值 | 1e-10 |
| `--vae_mini_batch` | VAE 编码时的迷你批次大小 | 1 |
| `--max_grad_norm` | 梯度裁剪阈值 | 0.05 |
| `--enable_bucket` | 启用分桶训练,不裁剪图片,按分辨率分组训练整个图像 | - |
| `--random_hw_adapt` | 自动缩放图片到 `[512, image_sample_size]` 范围内的随机尺寸 | - |
| `--resume_from_checkpoint` | 恢复训练路径,使用 `"latest"` 自动选择最新 checkpoint | None |
| `--uniform_sampling` | 均匀采样 timestep | - |
| `--trainable_modules` | 可训练模块(`"."` 表示所有模块) | `"."` |
| `--validation_steps` | 每 N 步执行一次验证 | 100 |
| `--validation_epochs` | 每 N 个epoch执行一次验证 | 100 |
| `--validation_prompts` | 验证图像生成的提示词 | `"一位年轻女子..."` |
### 3.4 训练验证
你可以配置验证参数,在训练过程中定期生成测试图像,以便监控训练进度和模型质量。
**验证参数说明**:
| 参数 | 说明 | 推荐值 |
|------|------|--------|
| `--validation_steps` | 每 N 步执行一次验证 | 100 |
| `--validation_epochs` | 每 N 个epoch执行一次验证 | 100 |
| `--validation_prompts` | 验证图像生成的提示词,可用空格分隔多个提示词 | 多个空格分隔的提示词 |
**示例**:
```bash
--validation_steps=100 \
--validation_epochs=100 \
--validation_prompts="一位年轻女子站在阳光明媚的海岸线上,白裙在轻拂的海风中微微飘动。"
```
**注意事项**:
- 验证图像会保存到 `output_dir` 目录中
- 多提示词验证格式:`--validation_prompts "prompt1" "prompt2" "prompt3"`
### 3.5 使用 FSDP 训练
**如果使用多卡且使用DeepSpeed-Zero-2的情况下显存不足**,可以切换使用FSDP进行训练。
```sh
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
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 ErnieImageSharedAdaLNBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
### 3.6 不使用 DeepSpeed 与 FSDP 训练
**该方案并不被推荐,因为没有显存节约后端,容易造成显存不足**。这里仅提供训练Shell用于参考训练。
```sh
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
### 3.7 多机分布式训练
**适合场景**:超大规模数据集、需要更快的训练速度
#### 3.7.1 环境配置
假设有 2 台机器,每台 8 张 GPU:
**机器 0(Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/X-Fun-Images-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # Master 机器 IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # 机器总数
export NUM_PROCESS=16 # 总进程数 = 机器数 × 8
export RANK=0 # 当前机器 rank(0 或 1)
# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/ernie_image/train.py \
--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_ernie_image" \
--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 "."
```
**机器 1(Worker)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
export DATASET_NAME="datasets/X-Fun-Images-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # 与 Master 相同
export MASTER_PORT=10086
export WORLD_SIZE=2
export NUM_PROCESS=16
export RANK=1 # 注意这里是 1
# 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
# 使用与机器 0 相同的 accelerate launch 命令
```
#### 3.7.2 多机训练注意事项
- **网络要求**:
- 推荐 RDMA/InfiniBand(高性能)
- 无 RDMA 时添加环境变量:
```bash
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
```
- **数据同步**:所有机器必须能够访问相同的数据路径(NFS/共享存储)
## 四、推理测试
### 4.1 推理参数解析
**关键参数说明**:
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `GPU_memory_mode` | 显存管理模式,可选值见下表 | `model_cpu_offload` |
| `ulysses_degree` | Head 维度并行度,单卡时为 1 | 1 |
| `ring_degree` | Sequence 维度并行度,单卡时为 1 | 1 |
| `fsdp_dit` | 多卡推理时对 Transformer 使用 FSDP 节省显存 | `False` |
| `fsdp_text_encoder` | 多卡推理时对文本编码器使用 FSDP | `False` |
| `compile_dit` | 编译 Transformer 加速推理(固定分辨率下有效) | `False` |
| `model_name` | 模型路径 | `models/Diffusion_Transformer/ERNIE-Image` |
| `sampler_name` | 采样器类型:`Flow`、`Flow_Unipc`、`Flow_DPM++` | `Flow` |
| `transformer_path` | 加载训练好的 Transformer 权重路径 | `None` |
| `vae_path` | 加载训练好的 VAE 权重路径 | `None` |
| `lora_path` | LoRA 权重路径 | `None` |
| `sample_size` | 生成图像分辨率 `[高度, 宽度]` | `[1728, 992]` |
| `weight_dtype` | 模型权重精度,不支持 bf16 的显卡使用 `torch.float16` | `torch.bfloat16` |
| `prompt` | 正向提示词,描述生成内容 | `"1girl, black_hair..."` |
| `negative_prompt` | 负向提示词,避免生成的内容 | `"低分辨率,低画质..."` |
| `guidance_scale` | 引导强度 | 4.5 |
| `seed` | 随机种子,用于复现结果 | 43 |
| `num_inference_steps` | 推理步数 | 40 |
| `lora_weight` | LoRA 权重强度 | 0.55 |
| `save_path` | 生成图像保存路径 | `samples/ernie-image-t2i` |
**显存管理模式说明**:
| 模式 | 说明 | 显存占用 |
|------|------|---------|
| `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/ernie_image/predict_t2i.py
```
根据需求修改编辑 `examples/ernie_image/predict_t2i.py`,初次推理重点关注如下参数,如果对其他参数感兴趣,请查看上方的推理参数解析。
```python
# 根据显卡显存选择
GPU_memory_mode = "model_cpu_offload"
# 根据实际模型路径
model_name = "models/Diffusion_Transformer/ERNIE-Image"
# 训练好的权重路径,如 "output_dir_ernie_image/checkpoint-xxx/diffusion_pytorch_model.safetensors"
transformer_path = None
# 根据生成内容编写
prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
# ...
```
### 4.3 多卡并行推理
**适合场景**:高分辨率生成、加速推理
#### 安装并行推理依赖
```bash
pip install xfuser==0.4.2 yunchang==0.6.2
```
#### 配置并行策略
编辑 `examples/ernie_image/predict_t2i.py`:
```python
# 确保 ulysses_degree × ring_degree = GPU 数量
# 例如使用 2 张 GPU:
ulysses_degree = 2 # Head 维度并行
ring_degree = 1 # Sequence 维度并行
```
**配置原则**:
- `ulysses_degree` 必须能整除模型的head数。
- `ring_degree` 会在sequence上切分,影响通信开销,在head数能切分的时候尽量不用。
**示例配置**:
| GPU 数量 | ulysses_degree | ring_degree | 说明 |
|---------|---------------|-------------|------|
| 1 | 1 | 1 | 单卡 |
| 4 | 4 | 1 | Head 并行 |
| 8 | 8 | 1 | Head 并行 |
| 8 | 4 | 2 | 混合并行 |
#### 运行多卡推理
```bash
torchrun --nproc-per-node=2 examples/ernie_image/predict_t2i.py
```
## 五、更多资源
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
File diff suppressed because it is too large Load Diff
+35
View File
@@ -0,0 +1,35 @@
export MODEL_NAME="models/Diffusion_Transformer/ERNIE-Image"
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 ErnieImageSharedAdaLNBlock --fsdp_sharding_strategy "FULL_SHARD" \
--fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/ernie_image/train.py \
--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=100 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ernie_image" \
--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 "."
+727
View File
@@ -0,0 +1,727 @@
# Wan2.1 Self-Forcing Distillation Training Guide
This document provides a complete workflow for Self-Forcing distillation of Wan2.1 including environment setup, data preparation, distributed training, and inference testing.
> **Note**: Wan2.1 Self-Forcing is a causal video generation model that supports text-to-video (T2V). Combined with distillation, this training code can reduce inference steps from 25-50 to 4-8 steps while enabling block-by-block causal generation with teacher forcing.
---
## 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. Distillation Training](#3-distillation-training)
- [3.1 Download Pretrained Models](#31-download-pretrained-models)
- [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-Node Distributed Training](#37-multi-node-distributed-training)
- [4. Inference Testing](#4-inference-testing)
- [4.1 Inference Parameters](#41-inference-parameters)
- [4.2 Text-to-Video (T2V) Inference](#42-text-to-video-t2v-inference)
- [4.3 Multi-GPU Parallel Inference](#43-multi-gpu-parallel-inference)
- [5. Additional Resources](#5-additional-resources)
---
## 1. Environment Setup
**Method 1: Using requirements.txt**
```bash
pip install -r requirements.txt
```
**Method 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 deepspeed==0.17.0 numpy==1.26.4
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
```
**Method 3: Using Docker**
When using Docker, please ensure that your machine has correctly installed GPU drivers and CUDA environment, 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
```
---
## 2. Data Preparation
### 2.1 Quick Test Dataset
We provide a test dataset that contains several training data samples.
```bash
# Download official example dataset
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 2.2 Dataset Structure
```
📦 datasets/
├── 📂 my_dataset/
│ ├── 📂 train/
│ │ ├── 📄 video001.mp4
│ │ ├── 📄 video002.mp4
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 2.3 metadata.json Format
**Relative Path Format** (example format):
```json
[
{
"file_path": "train/video001.mp4",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"type": "video",
"width": 1024,
"height": 1024
},
{
"file_path": "train/video002.mp4",
"text": "A person walking through a forest, cinematic view",
"type": "video",
"width": 1328,
"height": 1328
}
]
```
**Absolute Path Format**:
```json
[
{
"file_path": "/mnt/data/videos/sunset.mp4",
"text": "A beautiful sunset over the ocean",
"type": "video",
"width": 1024,
"height": 1024
}
]
```
**Key Field Descriptions**:
- `file_path`: Video path (relative or absolute path)
- `text`: Video description (English prompt)
- `type`: Data type, fixed as `"video"`
- `width` / `height`: Video dimensions (**recommended** to provide for bucket training. If not provided, it will be automatically read during training, which may affect training speed when data is stored on slower systems like OSS).
- You can use `scripts/process_json_add_width_and_height.py` to extract width and height fields for JSON files without these fields, supporting both images and videos.
- Usage: `python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Videos-Demo/metadata.json --output_file datasets/X-Fun-Videos-Demo/metadata_add_width_height.json`.
### 2.4 Relative vs Absolute Path Usage
**Relative Paths**:
If your data uses relative paths, configure in the training script:
```bash
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
```
**Absolute Paths**:
If your data uses absolute paths, configure in the training script:
```bash
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/metadata.json"
```
> 💡 **Recommendation**: If the dataset is small and stored locally, use relative paths. If the dataset is stored on external storage (such as NAS, OSS) or shared across multiple machines, use absolute paths.
---
## 3. Distillation Training
### 3.1 Download Pretrained Models
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download Wan2.1 official weights
# T2V model (text-to-video)
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
# Self-Forcing
hf download gdhe17/Self-Forcing --local-dir models/Diffusion_Transformer/Self-Forcing
```
### 3.2 Quick Start (DeepSpeed-Zero-2)
After downloading data according to **2.1 Quick Test Dataset** and downloading weights according to **3.1 Download Pretrained Models**, you can directly copy and run the quick start command.
We recommend using DeepSpeed-Zero-2 and FSDP for training. Here we use DeepSpeed-Zero-2 as an example to configure the shell file.
The difference between DeepSpeed-Zero-2 and FSDP lies in whether to shard model weights. **If you use multiple GPUs and encounter insufficient GPU memory with DeepSpeed-Zero-2**, you can switch to FSDP for training.
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
### 3.3 Common Training Parameters
**Key Parameter Descriptions**:
| Parameter | Description | Example Value |
|-----|------|-------|
| `--pretrained_model_name_or_path` | Pretrained model path | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B/` |
| `--train_data_dir` | Training data directory | `datasets/internal_datasets/` |
| `--train_data_meta` | Training data metadata file | `datasets/internal_datasets/metadata.json` |
| `--train_batch_size` | Batch size per GPU | 1 |
| `--image_sample_size` | Maximum image training resolution | 640 |
| `--video_sample_size` | Maximum video training resolution | 640 |
| `--token_sample_size` | Token sample size | 640 |
| `--video_sample_stride` | Video sampling stride | 2 |
| `--video_sample_n_frames` | Number of video frames | 81 |
| `--gradient_accumulation_steps` | Gradient accumulation steps (effectively increases batch) | 1 |
| `--dataloader_num_workers` | DataLoader worker processes | 8 |
| `--num_train_epochs` | Number of training epochs | 100 |
| `--checkpointing_steps` | Save checkpoint every N steps | 50 |
| `--learning_rate` | Initial learning rate (generator) | 2e-06 |
| `--learning_rate_critic` | Initial learning rate (critic) | 2e-07 |
| `--lr_scheduler` | Learning rate scheduler | `constant_with_warmup` |
| `--lr_warmup_steps` | Learning rate warmup steps | 100 |
| `--seed` | Random seed | 42 |
| `--output_dir` | Output directory | `output_dir_wan2.1_self_forcing_distill` |
| `--gradient_checkpointing` | Enable gradient checkpointing | - |
| `--mixed_precision` | Mixed precision: `fp16/bf16` | `bf16` |
| `--adam_weight_decay` | AdamW weight decay | 3e-2 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--vae_mini_batch` | VAE encoding mini-batch size | 1 |
| `--max_grad_norm` | Gradient clipping threshold | 0.05 |
| `--enable_bucket` | Enable bucket training, no cropping, group by resolution | - |
| `--random_hw_adapt` | Auto-scale images/videos to random sizes in `[min_size, max_size]` range | - |
| `--training_with_video_token_length` | Train based on token length, supports arbitrary resolutions | - |
| `--uniform_sampling` | Uniform timestep sampling | - |
| `--low_vram` | Low VRAM mode | - |
| `--train_mode` | Training mode: `normal` (T2V) | `normal` |
| `--resume_from_checkpoint` | Resume training path, use `"latest"` to auto-select latest checkpoint | None |
| `--validation_steps` | Run validation every N steps | 2000 |
| `--validation_epochs` | Run validation every N epochs | 5 |
| `--validation_prompts` | Prompts for video generation validation | `"A dog shaking head..."` |
| `--trainable_modules` | Trainable modules (`"."` means all modules) | `"."` |
**Distillation-Specific Parameters**:
| Parameter | Description | Example Value |
|-----|------|-------|
| `--denoising_step_indices_list` | Denoising step indices list (core distillation parameter) | `1000 750 500 250` |
| `--real_guidance_scale` | Real guidance scale for scoring | 6.0 |
| `--fake_guidance_scale` | Fake guidance scale for scoring | 0.0 |
| `--gen_update_interval` | Generator update interval | 5 |
| `--negative_prompt` | Negative prompt for distillation | Chinese negative prompt |
| `--train_sampling_steps` | Training sampling steps | 1000 |
| `--ode_transformer_path` | Path to ODE-trained weights to load into generator transformer3d | `models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt` |
**Self-Forcing-Specific Parameters**:
| Parameter | Description | Example Value |
|-----|------|-------|
| `--fix_sample_size` | Fixed sample size `[height, width]` for training | `480 832` |
| `--num_frame_per_block` | Number of frames per block for causal training | 3 |
| `--independent_first_frame` | Whether first frame is independent (`[1, N, N, ...]` pattern) | - |
| `--use_kv_cache_training` | Use KV cache block-by-block training (matches original Self-Forcing) | - |
| `--context_noise` | Context noise level for KV cache update | 0 |
| `--use_teacher_forcing` | Enable teacher forcing training (pass clean_x to transformer) | - |
| `--teacher_forcing_prob` | Probability of applying teacher forcing per step | 1.0 |
**Sample Size Configuration Guide**:
- `video_sample_size` represents the resolution size of videos; when `random_hw_adapt` is True, it represents the minimum value between video and image resolutions.
- `image_sample_size` represents the resolution size of images; when `random_hw_adapt` is True, it represents the maximum value between video and image resolutions.
- `token_sample_size` represents the resolution corresponding to the maximum token length when `training_with_video_token_length` is True.
- Due to potential confusion in configuration, **if you don't require arbitrary resolution for finetuning**, it is recommended to set `video_sample_size`, `image_sample_size`, and `token_sample_size` to the same fixed value, such as **(320, 480, 512, 640, 960)**.
- **All set to 320** represents **240P**.
- **All set to 480** represents **320P**.
- **All set to 640** represents **480P**.
- **All set to 960** represents **720P**.
**Token Length Training Guide**:
- When `training_with_video_token_length` is enabled, the model trains based on token length.
- For example: A video with 512x512 resolution and 49 frames has a token length of 13,312, requiring `token_sample_size = 512`.
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding frame counts allow the model to generate videos of different sizes.
### 3.4 Training Validation
You can configure validation parameters to periodically generate test videos during training, allowing you to monitor training progress and model quality.
**Validation Parameter Descriptions**:
| Parameter | Description | Recommended Value |
|------|------|--------|
| `--validation_steps` | Run validation every N steps | 2000 |
| `--validation_epochs` | Run validation every N epochs | 5 |
| `--validation_prompts` | Prompts for video generation validation | English prompts |
**T2V Validation Example**:
```bash
--validation_steps=2000 \
--validation_epochs=5 \
--validation_prompts="A brown dog shaking its head, sitting on a light-colored sofa in a cozy room. Behind the dog, there's a framed picture on the shelf, surrounded by pink flowers. The soft and warm lighting in the room creates a comfortable atmosphere."
```
**Notes**:
- Validation videos will be saved to the `output_dir` directory
- Multi-prompt validation format: `--validation_prompts "prompt1" "prompt2" "prompt3"`
### 3.5 Training with FSDP
**If you use multiple GPUs and encounter insufficient GPU memory with DeepSpeed-Zero-2**, you can switch to FSDP for training.
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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=CasualWanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
### 3.6 Other Backends
#### 3.6.1 Training with DeepSpeed-Zero-3
DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable.
DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model:
```bash
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
```
Training shell command is as follows:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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 --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
#### 3.6.2 Training without DeepSpeed and FSDP
**This approach is not recommended because there is no memory-saving backend, which can easily cause out-of-memory errors**. We only provide the training shell for reference.
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
### 3.7 Multi-Node Distributed Training
**Suitable for**: Ultra-large-scale datasets, faster training speed
#### 3.7.1 Environment Configuration
Assuming 2 machines, each with 8 GPUs:
**Machine 0 (Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # Master machine IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # Total number of machines
export NUM_PROCESS=16 # Total processes = machines × 8
export RANK=0 # Rank of this machine (0 or 1)
# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
**Machine 1 (Worker)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # Same as Master
export MASTER_PORT=10086
export WORLD_SIZE=2
export NUM_PROCESS=16
export RANK=1 # Note this is 1
# 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
# Use the same accelerate launch command as Machine 0
```
#### 3.7.2 Multi-Node Training Notes
- **Network Requirements**:
- RDMA/InfiniBand recommended (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)
---
## 4. Inference Testing
### 4.1 Inference Parameters
**Key Parameter Descriptions**:
| Parameter | Description | Example Value |
|------|------|-------|
| `GPU_memory_mode` | GPU memory mode, see options below | `sequential_cpu_offload` |
| `ulysses_degree` | Ulysses parallelism degree for multi-GPU inference | 1 |
| `ring_degree` | Ring parallelism degree for multi-GPU inference | 1 |
| `fsdp_dit` | Use FSDP for Transformer during multi-GPU inference to save memory | `False` |
| `fsdp_text_encoder` | Use FSDP for text encoder during multi-GPU inference | `True` |
| `compile_dit` | Compile Transformer for faster inference (effective at fixed resolution) | `False` |
| `model_name` | Model path | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B` |
| `sampler_name` | Sampler type: `Flow`, `Flow_Unipc`, `Flow_DPM++` | `Flow` |
| `transformer_path` | Trained Transformer weight path | `"models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"` |
| `vae_path` | Trained VAE weight path | `None` |
| `lora_path` | LoRA weight path | `None` |
| `sample_size` | Generated video resolution `[height, width]` | `[480, 832]` |
| `video_length` | Number of frames to generate | `81` |
| `fps` | Frames per second | `16` |
| `weight_dtype` | Model weight dtype, use `torch.float16` for GPUs that don't support bf16 | `torch.bfloat16` |
| `num_frame_per_block` | Number of frames to generate per block (1 for standard causal, higher for faster but more memory) | 3 |
| `local_attn_size` | Local attention window size (-1 for global attention) | -1 |
| `independent_first_frame` | Whether first frame is generated independently | `False` |
| `context_noise` | Context noise level for generation | 0.0 |
| `prompt` | Positive prompt describing what to generate | `"A stylish woman walks down a Tokyo street..."` |
| `negative_prompt` | Negative prompt to avoid certain content | Chinese negative prompt |
| `guidance_scale` | Guidance strength (distillation models typically use 1.0) | 1.0 |
| `seed` | Random seed for reproducibility | 43 |
| `num_inference_steps` | Number of inference steps (typically 4 for distillation models) | 4 |
| `lora_weight` | LoRA weight strength | 0.55 |
| `save_path` | Path to save generated videos | `samples/wan-videos-self-forcing-t2v` |
**GPU Memory Mode Descriptions**:
| Mode | Description | Memory Usage |
|------|------|---------|
| `model_full_load` | Entire model loaded 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` | Layer groups switch between CPU/CUDA | Low |
| `sequential_cpu_offload` | Layer-by-layer offload (slowest) | Lowest |
### 4.2 Text-to-Video (T2V) Inference
Run single GPU inference:
```bash
python examples/wan2.1_self_forcing/predict_t2v.py
```
Edit `examples/wan2.1_self_forcing/predict_t2v.py` according to your needs. For first-time inference, focus on the following key parameters. For other parameters, please refer to the inference parameter descriptions above.
```python
# Choose based on GPU memory
GPU_memory_mode = "sequential_cpu_offload"
# Your actual model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Trained weight path
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
# Distillation models typically use 4 steps
num_inference_steps = 4
# Distillation models guidance_scale is typically 1.0
guidance_scale = 1.0
# Self-Forcing causal inference config
num_frame_per_block = 3 # Number of frames to generate per block
local_attn_size = -1 # Local attention window size (-1 for global attention)
independent_first_frame = False
context_noise = 0.0
# Write according to your generated content
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage."
# ...
```
### 4.3 Multi-GPU Parallel Inference
**Suitable for**: High-resolution generation, accelerated inference
#### Install Parallel Inference Dependencies
```bash
pip install xfuser==0.4.2 yunchang==0.6.2
```
#### Configure Parallel Strategy
Edit `examples/wan2.1_self_forcing/predict_t2v.py`:
```python
# Ensure ulysses_degree × ring_degree = number of GPUs used
# For example, using 2 GPUs:
ulysses_degree = 2 # Head dimension parallelism
ring_degree = 1 # Sequence dimension parallelism
```
**Configuration Principles**:
- `ulysses_degree` must be divisible by the model's head count
- `ring_degree` splits on the sequence dimension, which affects communication overhead. Try to avoid using it when heads are evenly divisible.
**Configuration Examples**:
| GPU Count | ulysses_degree | ring_degree | Description |
|---------|---------------|-------------|------|
| 1 | 1 | 1 | Single GPU |
| 4 | 4 | 1 | Head parallelism |
| 8 | 8 | 1 | Head parallelism |
| 8 | 4 | 2 | Hybrid parallelism |
#### Run Multi-GPU Inference
```bash
torchrun --nproc-per-node=2 examples/wan2.1_self_forcing/predict_t2v.py
```
---
## 5. Additional Resources
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
+728
View File
@@ -0,0 +1,728 @@
# Wan2.1 Self-Forcing 蒸馏训练指南
本文档提供了将 Wan2.1 进行 Self-Forcing 蒸馏的完整工作流,包括环境配置、数据准备、分布式训练和推理测试。
> **说明**:Wan2.1 Self-Forcing 是一个支持文生视频(T2V)的因果视频生成模型。结合蒸馏训练,该代码可以将推理步数从 25-50 步减少到 4-8 步,同时支持逐块因果生成与 teacher forcing。
---
## 目录
- [一、环境配置](#一环境配置)
- [二、数据准备](#二数据准备)
- [2.1 快速测试数据集](#21-快速测试数据集)
- [2.2 数据集结构](#22-数据集结构)
- [2.3 metadata.json 格式](#23-metadatajson-格式)
- [2.4 相对路径与绝对路径使用方案](#24-相对路径与绝对路径使用方案)
- [三、蒸馏训练](#三蒸馏训练)
- [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-多机分布式训练)
- [四、推理测试](#四推理测试)
- [4.1 推理参数解析](#41-推理参数解析)
- [4.2 文生视频(T2V)推理](#42-文生视频t2v推理)
- [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 deepspeed==0.17.0 numpy==1.26.4
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
```
**方式 3:使用docker**
使用docker的情况下,请保证机器中已经正确安装显卡驱动与CUDA环境,然后以此执行以下命令:
```
# 拉取镜像
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
# 进入容器
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
```
---
## 二、数据准备
### 2.1 快速测试数据集
我们提供了一个测试的数据集,其中包含若干训练数据。
```bash
# 下载官方示例数据集
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 2.2 数据集结构
```
📦 datasets/
├── 📂 my_dataset/
│ ├── 📂 train/
│ │ ├── 📄 video001.mp4
│ │ ├── 📄 video002.mp4
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 2.3 metadata.json 格式
**相对路径格式**(示例格式):
```json
[
{
"file_path": "train/video001.mp4",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"type": "video",
"width": 1024,
"height": 1024
},
{
"file_path": "train/video002.mp4",
"text": "A person walking through a forest, cinematic view",
"type": "video",
"width": 1328,
"height": 1328
}
]
```
**绝对路径格式**:
```json
[
{
"file_path": "/mnt/data/videos/sunset.mp4",
"text": "A beautiful sunset over the ocean",
"type": "video",
"width": 1024,
"height": 1024
}
]
```
**关键字段说明**:
- `file_path`:视频路径(相对或绝对路径)
- `text`:视频描述(英文提示词)
- `type`:数据类型,固定为 `"video"`
- `width` / `height`:视频宽高(**最好提供**,用于分桶训练,如果不提供则自动在训练时读取,当数据存储在如oss这样的速度较慢的系统上时,可能会影响训练速度)。
- 可以使用`scripts/process_json_add_width_and_height.py`文件对无width与height字段的json进行提取,支持处理图片与视频。
- 使用方案为`python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Videos-Demo/metadata.json --output_file datasets/X-Fun-Videos-Demo/metadata_add_width_height.json`。
### 2.4 相对路径与绝对路径使用方案
**相对路径**:
如果数据的路径为相对路径,则在训练脚本中设置:
```bash
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
```
**绝对路径**:
如果数据的路径为绝对路径,则在训练脚本中设置:
```bash
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/metadata.json"
```
> 💡 **建议**:如果数据集较小且存储在本地,推荐使用相对路径;如果数据集存储在外部存储(如 NAS、OSS)或多个机器共享存储,推荐使用绝对路径。
---
## 三、蒸馏训练
### 3.1 下载预训练模型
```bash
# 创建模型目录
mkdir -p models/Diffusion_Transformer
# 下载 Wan2.1 官方权重
# T2V 模型(文生视频)
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
# Self-Forcing
hf download gdhe17/Self-Forcing --local-dir models/Diffusion_Transformer/Self-Forcing
```
### 3.2 快速开始(DeepSpeed-Zero-2)
如果按照 **2.1 快速测试数据集下载数据** 与 **3.1 下载预训练模型下载权重**后,直接复制快速开始的启动指令进行启动。
推荐使用DeepSpeed-Zero-2与FSDP方案进行训练。这里使用DeepSpeed-Zero-2为例配置shell文件。
本文中DeepSpeed-Zero-2与FSDP的差别在于是否对模型权重进行分片,**如果使用多卡且使用DeepSpeed-Zero-2的情况下显存不足**,可以切换使用FSDP进行训练。
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
### 3.3 训练常用参数解析
**关键参数说明**:
| 参数 | 说明 | 示例值 |
|-----|------|-------|
| `--pretrained_model_name_or_path` | 预训练模型路径 | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B/` |
| `--train_data_dir` | 训练数据目录 | `datasets/internal_datasets/` |
| `--train_data_meta` | 训练数据元文件 | `datasets/internal_datasets/metadata.json` |
| `--train_batch_size` | 每批次样本数 | 1 |
| `--image_sample_size` | 图像最大训练分辨率 | 640 |
| `--video_sample_size` | 视频最大训练分辨率 | 640 |
| `--token_sample_size` | Token 采样尺寸 | 640 |
| `--video_sample_stride` | 视频采样步幅 | 2 |
| `--video_sample_n_frames` | 视频采样帧数 | 81 |
| `--gradient_accumulation_steps` | 梯度累积步数(等效增大 batch) | 1 |
| `--dataloader_num_workers` | DataLoader 子进程数 | 8 |
| `--num_train_epochs` | 训练 epoch 数 | 100 |
| `--checkpointing_steps` | 每 N 步保存 checkpoint | 50 |
| `--learning_rate` | 初始学习率(生成器) | 2e-06 |
| `--learning_rate_critic` | 初始学习率(判别器) | 2e-07 |
| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` |
| `--lr_warmup_steps` | 学习率预热步数 | 100 |
| `--seed` | 随机种子 | 42 |
| `--output_dir` | 输出目录 | `output_dir_wan2.1_self_forcing_distill` |
| `--gradient_checkpointing` | 激活重计算 | - |
| `--mixed_precision` | 混合精度:`fp16/bf16` | `bf16` |
| `--adam_weight_decay` | AdamW 权重衰减 | 3e-2 |
| `--adam_epsilon` | AdamW epsilon 值 | 1e-10 |
| `--vae_mini_batch` | VAE 编码时的迷你批次大小 | 1 |
| `--max_grad_norm` | 梯度裁剪阈值 | 0.05 |
| `--enable_bucket` | 启用分桶训练,不裁剪图片/视频,按分辨率分组训练 | - |
| `--random_hw_adapt` | 自动缩放图片/视频到 `[min_size, max_size]` 范围内的随机尺寸 | - |
| `--training_with_video_token_length` | 根据 token 长度训练,支持任意分辨率 | - |
| `--uniform_sampling` | 均匀采样 timestep | - |
| `--low_vram` | 低显存模式 | - |
| `--train_mode` | 训练模式:`normal`(T2V) | `normal` |
| `--resume_from_checkpoint` | 恢复训练路径,使用 `"latest"` 自动选择最新 checkpoint | None |
| `--validation_steps` | 每 N 步执行一次验证 | 2000 |
| `--validation_epochs` | 每 N 个epoch执行一次验证 | 5 |
| `--validation_prompts` | 验证视频生成的提示词 | `"一只棕色的狗摇着头..."` |
| `--trainable_modules` | 可训练模块(`"."` 表示所有模块) | `"."` |
**蒸馏特有参数**:
| 参数 | 说明 | 示例值 |
|-----|------|-------|
| `--denoising_step_indices_list` | 去噪步骤列表(蒸馏核心参数) | `1000 750 500 250` |
| `--real_guidance_scale` | 用于评分的真实 guidance scale | 6.0 |
| `--fake_guidance_scale` | 用于评分的虚拟 guidance scale | 0.0 |
| `--gen_update_interval` | 生成器更新间隔 | 5 |
| `--negative_prompt` | 用于蒸馏的负向提示词 | 中文负向提示词 |
| `--train_sampling_steps` | 训练采样步数 | 1000 |
| `--ode_transformer_path` | ODE 训练权重路径,加载到 generator transformer3d 中 | `models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt` |
**Self-Forcing 特有参数**:
| 参数 | 说明 | 示例值 |
|-----|------|-------|
| `--fix_sample_size` | 固定训练尺寸 `[高度, 宽度]` | `480 832` |
| `--num_frame_per_block` | 每个块的帧数(用于因果训练) | 3 |
| `--independent_first_frame` | 第一帧是否独立生成(`[1, N, N, ...]` 模式) | - |
| `--use_kv_cache_training` | 使用 KV 缓存逐块训练(匹配原始 Self-Forcing) | - |
| `--context_noise` | KV 缓存更新的上下文噪声级别 | 0 |
| `--use_teacher_forcing` | 启用 teacher forcing 训练(将 clean_x 传给 transformer) | - |
| `--teacher_forcing_prob` | 每步应用 teacher forcing 的概率 | 1.0 |
**Sample Size 配置指南**:
- `video_sample_size` 表示视频的分辨率大小;当 `random_hw_adapt` 为 True 时,表示视频和图像分辨率的最小值。
- `image_sample_size` 表示图像的分辨率大小;当 `random_hw_adapt` 为 True 时,表示视频和图像分辨率的最大值。
- `token_sample_size` 表示当 `training_with_video_token_length` 为 True 时,最大 token 长度对应的分辨率。
- 由于配置可能产生混淆,**如果你不需要任意分辨率进行 finetuning**,建议将 `video_sample_size`、`image_sample_size` 和 `token_sample_size` 设置为相同的固定值,例如 **(320, 480, 512, 640, 960)**。
- **全部设置为 320** 代表 **240P**。
- **全部设置为 480** 代表 **320P**。
- **全部设置为 640** 代表 **480P**。
- **全部设置为 960** 代表 **720P**。
**Token Length 训练说明**:
- 当启用 `training_with_video_token_length` 时,模型根据 token 长度进行训练。
- 例如:512x512 分辨率、49 帧的视频,其 token 长度为 13,312,需要设置 `token_sample_size = 512`。
- 在 512x512 分辨率下,视频帧数为 49 (~= 512 * 512 * 49 / 512 / 512)。
- 在 768x768 分辨率下,视频帧数为 21 (~= 512 * 512 * 49 / 768 / 768)。
- 在 1024x1024 分辨率下,视频帧数为 9 (~= 512 * 512 * 49 / 1024 / 1024)。
- 这些分辨率与对应帧数的组合,使模型能够生成不同尺寸的视频。
### 3.4 训练验证
你可以配置验证参数,在训练过程中定期生成测试视频,以便监控训练进度和模型质量。
**验证参数说明**:
| 参数 | 说明 | 推荐值 |
|------|------|--------|
| `--validation_steps` | 每 N 步执行一次验证 | 2000 |
| `--validation_epochs` | 每 N 个epoch执行一次验证 | 5 |
| `--validation_prompts` | 验证视频生成的提示词 | 英文提示词 |
**T2V 验证示例**:
```bash
--validation_steps=2000 \
--validation_epochs=5 \
--validation_prompts="A brown dog shaking its head, sitting on a light-colored sofa in a cozy room. Behind the dog, there's a framed painting on a shelf, surrounded by pink flowers. The soft, warm lighting in the room creates a comfortable atmosphere."
```
**注意事项**:
- 验证视频会保存到 `output_dir` 目录中
- 多提示词验证格式:`--validation_prompts "prompt1" "prompt2" "prompt3"`
### 3.5 使用 FSDP 训练
**如果使用多卡且使用DeepSpeed-Zero-2的情况下显存不足**,可以切换使用FSDP进行训练。
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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=CasualWanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
### 3.6 其他后端
#### 3.6.1 使用DeepSpeed-Zero-3进行训练
目前不太推荐使用 DeepSpeed Zero-3。在本仓库中,使用 FSDP 出错更少且更稳定。
DeepSpeed Zero-3 适合高分辨率的 14B Wan。训练后,您可以使用以下命令获取最终模型:
```bash
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
```
训练 shell 命令如下:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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 --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
#### 3.6.2 不使用 DeepSpeed 与 FSDP 训练
**该方案并不被推荐,因为没有显存节约后端,容易造成显存不足**。这里仅提供训练Shell用于参考训练。
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-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/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
### 3.7 多机分布式训练
**适合场景**:超大规模数据集、需要更快的训练速度
#### 3.7.1 环境配置
假设有 2 台机器,每台 8 张 GPU:
**机器 0(Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # Master 机器 IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # 机器总数
export NUM_PROCESS=16 # 总进程数 = 机器数 × 8
export RANK=0 # 当前机器 rank(0 或 1)
# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
```
**机器 1(Worker)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json"
export MASTER_ADDR="192.168.1.100" # 与 Master 相同
export MASTER_PORT=10086
export WORLD_SIZE=2
export NUM_PROCESS=16
export RANK=1 # 注意这里是 1
# 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
# 使用与机器 0 相同的 accelerate launch 命令
```
#### 3.7.2 多机训练注意事项
- **网络要求**:
- 推荐 RDMA/InfiniBand(高性能)
- 无 RDMA 时添加环境变量:
```bash
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
```
- **数据同步**:所有机器必须能够访问相同的数据路径(NFS/共享存储)
---
## 四、推理测试
### 4.1 推理参数解析
**关键参数说明**:
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `GPU_memory_mode` | 显存管理模式,可选值见下表 | `sequential_cpu_offload` |
| `ulysses_degree` | Head 维度并行度,单卡时为 1 | 1 |
| `ring_degree` | Sequence 维度并行度,单卡时为 1 | 1 |
| `fsdp_dit` | 多卡推理时对 Transformer 使用 FSDP 节省显存 | `False` |
| `fsdp_text_encoder` | 多卡推理时对文本编码器使用 FSDP | `True` |
| `compile_dit` | 编译 Transformer 加速推理(固定分辨率下有效) | `False` |
| `model_name` | 模型路径 | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B` |
| `sampler_name` | 采样器类型:`Flow`、`Flow_Unipc`、`Flow_DPM++` | `Flow` |
| `transformer_path` | 加载训练好的 Transformer 权重路径 | `"models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"` |
| `vae_path` | 加载训练好的 VAE 权重路径 | `None` |
| `lora_path` | LoRA 权重路径 | `None` |
| `sample_size` | 生成视频分辨率 `[高度, 宽度]` | `[480, 832]` |
| `video_length` | 生成视频帧数 | `81` |
| `fps` | 每秒帧数 | `16` |
| `weight_dtype` | 模型权重精度,不支持 bf16 的显卡使用 `torch.float16` | `torch.bfloat16` |
| `num_frame_per_block` | 每个块生成的帧数(1 为标准因果,更高则更快但需更多显存) | 3 |
| `local_attn_size` | 局部注意力窗口大小(-1 为全局注意力) | -1 |
| `independent_first_frame` | 第一帧是否独立生成 | `False` |
| `context_noise` | 生成时的上下文噪声级别 | 0.0 |
| `prompt` | 正向提示词,描述生成内容 | `"A stylish woman walks down a Tokyo street..."` |
| `negative_prompt` | 负向提示词,避免生成的内容 | 中文负向提示词 |
| `guidance_scale` | 引导强度(蒸馏模型通常使用 1.0) | 1.0 |
| `seed` | 随机种子,用于复现结果 | 43 |
| `num_inference_steps` | 推理步数(蒸馏模型通常为 4) | 4 |
| `lora_weight` | LoRA 权重强度 | 0.55 |
| `save_path` | 生成视频保存路径 | `samples/wan-videos-self-forcing-t2v` |
**显存管理模式说明**:
| 模式 | 说明 | 显存占用 |
|------|------|---------|
| `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 文生视频(T2V)推理
单卡推理运行如下命令:
```bash
python examples/wan2.1_self_forcing/predict_t2v.py
```
根据需求修改编辑 `examples/wan2.1_self_forcing/predict_t2v.py`,初次推理重点关注如下参数,如果对其他参数感兴趣,请查看上方的推理参数解析。
```python
# 根据显卡显存选择
GPU_memory_mode = "sequential_cpu_offload"
# 根据实际模型路径
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# 训练好的权重路径
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
# 蒸馏模型通常使用 4 步
num_inference_steps = 4
# 蒸馏模型 guidance_scale 通常为 1.0
guidance_scale = 1.0
# Self-Forcing 因果推理配置
num_frame_per_block = 3 # 每个块生成的帧数
local_attn_size = -1 # 局部注意力窗口大小(-1 为全局注意力)
independent_first_frame = False
context_noise = 0.0
# 根据生成内容编写
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage."
# ...
```
### 4.3 多卡并行推理
**适合场景**:高分辨率生成、加速推理
#### 安装并行推理依赖
```bash
pip install xfuser==0.4.2 yunchang==0.6.2
```
#### 配置并行策略
编辑 `examples/wan2.1_self_forcing/predict_t2v.py`:
```python
# 确保 ulysses_degree × ring_degree = 使用的 GPU 数
# 例如使用 2 张 GPU:
ulysses_degree = 2 # Head 维度并行
ring_degree = 1 # Sequence 维度并行
```
**配置原则**:
- `ulysses_degree` 必须能整除模型的 head 数
- `ring_degree` 是在 sequence 维度切分,会影响通信开销,在 head 能整除的情况下尽量不要用
**配置示例**:
| GPU 数量 | ulysses_degree | ring_degree | 说明 |
|---------|---------------|-------------|------|
| 1 | 1 | 1 | 单 GPU |
| 4 | 4 | 1 | Head 并行 |
| 8 | 8 | 1 | Head 并行 |
| 8 | 4 | 2 | 混合并行 |
#### 运行多卡推理
```bash
torchrun --nproc-per-node=2 examples/wan2.1_self_forcing/predict_t2v.py
```
---
## 五、更多资源
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
+460
View File
@@ -0,0 +1,460 @@
# Wan2.1 Self-Forcing ODE Regression Training Guide
This document provides the complete workflow for **ODE regression pre-training** of Wan2.1 Self-Forcing, including environment setup, ODE trajectory pair generation, and ODE regression training.
> **What is ODE Regression Training?**
>
> ODE regression is the **pre-training stage** for Self-Forcing distillation. The pipeline is:
>
> 1. **Step 1 — Generate ODE pairs** (`generate_ode_pairs.py`): Use the **bidirectional teacher** (Wan2.1-T2V-1.3B) to perform full multi-step CFG denoising on a list of text prompts. Save the intermediate latents along the ODE trajectory together with the encoded prompt embeddings as `.safetensors` files.
> 2. **Step 2 — Train ODE regression** (`train_ode.py`): Load the generated ODE pairs, and train a **causal generator** to predict the clean endpoint `x0` at multiple sampled trajectory points. The output checkpoint serves as a strong initialization (typically saved as `ode_init.pt`) for the subsequent **Self-Forcing distillation** stage (`train_distill.py`, see [README_TRAIN.md](./README_TRAIN.md)).
---
## Table of Contents
- [1. Environment Setup](#1-environment-setup)
- [2. Download Pretrained Models](#2-download-pretrained-models)
- [3. Step 1 — Generate ODE Trajectory Pairs](#3-step-1--generate-ode-trajectory-pairs)
- [3.1 Download Prompt File](#31-download-prompt-file)
- [3.2 Run ODE Pair Generation](#32-run-ode-pair-generation)
- [3.3 Output Format](#33-output-format)
- [3.4 Generation Parameters](#34-generation-parameters)
- [3.5 Multi-GPU Generation](#35-multi-gpu-generation)
- [4. Step 2 — Train ODE Regression](#4-step-2--train-ode-regression)
- [4.1 Quick Start](#41-quick-start)
- [4.2 Common Training Parameters](#42-common-training-parameters)
- [4.3 Training with DeepSpeed-Zero-2 / FSDP](#43-training-with-deepspeed-zero-2--fsdp)
- [4.4 Multi-Node Distributed Training](#44-multi-node-distributed-training)
- [5. Use the Trained ODE Weights](#5-use-the-trained-ode-weights)
- [6. Additional Resources](#6-additional-resources)
---
## 1. Environment Setup
**Method 1: Using requirements.txt**
```bash
pip install -r requirements.txt
```
**Method 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 deepspeed==0.17.0 numpy==1.26.4
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
```
**Method 3: Using Docker**
When using Docker, please ensure that your machine has correctly installed GPU drivers and CUDA environment, then execute the following commands:
```bash
# 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
```
---
## 2. Download Pretrained Models
ODE generation uses the **bidirectional teacher** Wan2.1-T2V-1.3B as denoiser, and ODE training initializes the **causal generator** from the same base model.
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download Wan2.1 T2V base model (used as both teacher for generation and init for training)
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
```
---
## 3. Step 1 — Generate ODE Trajectory Pairs
This step uses the bidirectional teacher to run **48-step CFG denoising** on each prompt and saves the resulting ODE trajectory together with the prompt embedding into a `.safetensors` file. After all prompts are processed, an `outputs.json` annotation file is produced for the training stage to consume.
### 3.1 Download Prompt File
The official Self-Forcing prompt list is recommended:
```bash
mkdir -p datasets
# Download vidprom_filtered_extended.txt from the official Self-Forcing repo
hf download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir datasets/
# Final path: datasets/vidprom_filtered_extended.txt
```
You can also use any plain-text file with one prompt per line.
### 3.2 Run ODE Pair Generation
The ready-to-use launcher is [scripts/wan2.1_self_forcing/generate_ode_pairs.sh](./generate_ode_pairs.sh):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/generate_ode_pairs.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/wan2.1/wan_civitai.yaml" \
--video_sample_n_frames=81 \
--height=480 \
--width=832 \
--guidance_scale=6.0 \
--shift=8.0 \
--num_inference_steps=48 \
--caption_path="datasets/vidprom_filtered_extended.txt" \
--output_folder="datasets/ode_pairs_output" \
--sample_every_n_prompts=50
```
Or simply run the shell script:
```bash
bash scripts/wan2.1_self_forcing/generate_ode_pairs.sh
```
### 3.3 Output Format
After generation, `--output_folder` will contain:
```
📦 datasets/ode_pairs_output/
├── 📄 00000.safetensors # Per-prompt ODE trajectory + prompt embeds
├── 📄 00001.safetensors
├── 📄 ...
├── 📂 sample/ # Optional preview videos (when sample_every_n_prompts > 0)
│ └── 📄 00000_clean.mp4
└── 📄 outputs.json # Annotation file consumed by train_ode.py
```
Each `.safetensors` file contains:
| Key | Shape | Description |
|-----|-------|-------------|
| `latents` | `[5, C, F, H, W]` | Sparse 5-point sampling of the 48-step ODE trajectory: indices `[0, 12, 24, 36, -1]` (initial noise → 3 mid-points → clean endpoint) |
| `prompt_embeds` | `[512, D]` | Padded T5 prompt embeddings (max length 512) |
| `prompt_attention_mask` | `[512]` | Attention mask for the prompt embeddings |
The auto-generated `outputs.json` follows the same format as a standard `metadata.json`:
```json
[
{ "file_path": "datasets/ode_pairs_output/00000.safetensors" },
{ "file_path": "datasets/ode_pairs_output/00001.safetensors" }
]
```
### 3.4 Generation Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--pretrained_model_name_or_path` | Path to Wan2.1-T2V-1.3B teacher | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B` |
| `--config_path` | Model config YAML | `config/wan2.1/wan_civitai.yaml` |
| `--caption_path` | Plain-text file with one prompt per line | `datasets/vidprom_filtered_extended.txt` |
| `--output_folder` | Output directory for `.safetensors` files and `outputs.json` | `datasets/ode_pairs_output` |
| `--guidance_scale` | CFG guidance scale used by the teacher | 6.0 |
| `--num_inference_steps` | Number of teacher denoising steps (must be ≥ 37 because indices `[0,12,24,36,-1]` are sampled) | 48 |
| `--shift` | Shift value for `FlowMatchEulerDiscreteScheduler` (must match training!) | 8.0 |
| `--video_sample_n_frames` | Pixel frame count of generated video | 81 |
| `--height` / `--width` | Video resolution in pixels | 480 / 832 |
| `--negative_prompt` | Negative prompt used for CFG | (Chinese default) |
| `--sample_every_n_prompts` | Decode and save preview MP4 every N prompts (0 to disable) | 50 |
| `--mixed_precision` | `no` / `fp16` / `bf16` | `bf16` |
> ⚠️ Keep `--shift` **identical** between generation and ODE training. Both default to `8.0` in the provided scripts.
### 3.5 Multi-GPU Generation
`generate_ode_pairs.py` is built on `accelerate`. Each rank automatically processes an interleaved subset of prompts (`prompt_index = index * world_size + rank`) and skips already-existing files, so the job is **resumable** and **parallelizable** out of the box:
```bash
# 8-GPU generation
accelerate launch --multi_gpu --num_processes=8 --mixed_precision="bf16" \
scripts/wan2.1_self_forcing/generate_ode_pairs.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/wan2.1/wan_civitai.yaml" \
--caption_path="datasets/vidprom_filtered_extended.txt" \
--output_folder="datasets/ode_pairs_output" \
--num_inference_steps=48 --guidance_scale=6.0 --shift=8.0 \
--height=480 --width=832 --video_sample_n_frames=81
```
Only the main process writes the final `outputs.json`.
---
## 4. Step 2 — Train ODE Regression
After Step 1 completes and `datasets/ode_pairs_output/outputs.json` is generated, train the causal generator (`WanTransformer3DModel_SelfForcing`) to regress the ODE trajectory.
For each training sample the script:
1. Loads the saved sparse trajectory (5 points) and prompt embedding from one `.safetensors` file.
2. Randomly picks one trajectory point per **block** (with `--num_frame_per_block` frames sharing the same timestep), feeds the noisy latent and the per-frame timestep through the causal generator.
3. Converts the predicted flow into an `x0` prediction and computes MSE loss against the **clean endpoint** of the trajectory.
### 4.1 Quick Start
The ready-to-use launcher is [scripts/wan2.1_self_forcing/train_ode.sh](./train_ode.sh):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.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/wan2.1_self_forcing/train_ode.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$ODE_DATA_META \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=500 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_ode_regression" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.05 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--denoising_step_indices_list 1000 750 500 250 \
--shift=8.0 \
--resume_from_checkpoint="latest" \
--trainable_modules "."
```
Or simply:
```bash
bash scripts/wan2.1_self_forcing/train_ode.sh
```
> 💡 Because the ODE trajectory and prompt embeddings are pre-computed in Step 1, **no VAE / text encoder is invoked during ODE training** — training is fast, memory-efficient, and `train_data_dir` can be left empty when `outputs.json` already contains absolute paths.
### 4.2 Common Training Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--pretrained_model_name_or_path` | Base model used to initialize the causal generator | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B` |
| `--config_path` | Model config YAML | `config/wan2.1/wan_civitai.yaml` |
| `--train_data_dir` | Optional root prepended to `file_path`; can be empty when `outputs.json` stores absolute paths | `""` |
| `--train_data_meta` | Annotation JSON produced by Step 1 | `datasets/ode_pairs_output/outputs.json` |
| `--train_batch_size` | Per-GPU batch size | 1 |
| `--gradient_accumulation_steps` | Gradient accumulation steps | 1 |
| `--dataloader_num_workers` | DataLoader workers | 8 |
| `--num_train_epochs` | Number of training epochs | 100 |
| `--checkpointing_steps` | Save checkpoint every N steps | 500 |
| `--learning_rate` | Initial learning rate | 2e-06 |
| `--lr_scheduler` | LR scheduler type | `constant_with_warmup` |
| `--lr_warmup_steps` | LR warmup steps | 100 |
| `--seed` | Random seed | 42 |
| `--output_dir` | Output directory | `output_dir_wan2.1_self_forcing_ode_regression` |
| `--gradient_checkpointing` | Enable gradient checkpointing | - |
| `--mixed_precision` | `fp16` / `bf16` | `bf16` |
| `--adam_weight_decay` | AdamW weight decay | 3e-2 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--max_grad_norm` | Gradient clipping threshold | 0.05 |
| `--trainable_modules` | Trainable modules (`"."` = all) | `"."` |
| `--resume_from_checkpoint` | Resume path or `"latest"` | `latest` |
**ODE-specific parameters** (must match Step 1 unless you understand the consequences):
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--train_sampling_steps` | Total scheduler timesteps from which `denoising_step_indices_list` is sampled | 1000 |
| `--denoising_step_indices_list` | Discrete timestep indices used during ODE regression (corresponds to the 5 sparse points sampled in Step 1) | `1000 750 500 250` |
| `--shift` | Shift for `FlowMatchEulerDiscreteScheduler` — **must match `--shift` used in Step 1** | 8.0 |
| `--num_frame_per_block` | Number of frames per causal block (frames in a block share the same timestep) | 3 |
| `--independent_first_frame` | First frame is independent (`[1, N, N, ...]` block pattern) | - |
| `--context_noise` | Context noise level (matches downstream Self-Forcing distillation config) | 0 |
### 4.3 Training with DeepSpeed-Zero-2 / FSDP
For multi-GPU training, the same memory-saving backends as the distillation stage are supported.
**DeepSpeed-Zero-2** (recommended default):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.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/wan2.1_self_forcing/train_ode.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$ODE_DATA_META \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=500 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_ode_regression" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.05 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--denoising_step_indices_list 1000 750 500 250 \
--shift=8.0 \
--resume_from_checkpoint="latest" \
--trainable_modules "."
```
**FSDP** (use when DeepSpeed-Zero-2 runs out of memory):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.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=CasualWanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.1_self_forcing/train_ode.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$ODE_DATA_META \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=500 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_ode_regression" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.05 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--denoising_step_indices_list 1000 750 500 250 \
--shift=8.0 \
--resume_from_checkpoint="latest" \
--trainable_modules "."
```
### 4.4 Multi-Node Distributed Training
Assuming 2 machines × 8 GPUs:
**Machine 0 (Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.json"
export MASTER_ADDR="192.168.1.100" # Master machine IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # Total number of machines
export NUM_PROCESS=16 # Total processes = machines × 8
export RANK=0 # Rank of this machine (0 or 1)
# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_self_forcing/train_ode.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$ODE_DATA_META \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=500 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_ode_regression" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.05 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--denoising_step_indices_list 1000 750 500 250 \
--shift=8.0 \
--resume_from_checkpoint="latest" \
--trainable_modules "."
```
**Machine 1 (Worker)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.json"
export MASTER_ADDR="192.168.1.100" # Same as Master
export MASTER_PORT=10086
export WORLD_SIZE=2
export NUM_PROCESS=16
export RANK=1 # Note this is 1
# 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
# Use the same accelerate launch command as Machine 0
```
**Notes**:
- Use RDMA / InfiniBand whenever possible. Without RDMA, set `NCCL_IB_DISABLE=1` and `NCCL_P2P_DISABLE=1`.
- All machines must share the same `outputs.json` and the underlying `.safetensors` files (NFS / shared storage).
---
## 5. Use the Trained ODE Weights
The ODE-init checkpoint produced under `output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N}/` is intended to bootstrap Self-Forcing distillation. Pass its path to `train_distill.py` via `--ode_transformer_path`:
```bash
# Example: pick the saved weight file (e.g. diffusion_pytorch_model.safetensors)
--ode_transformer_path="output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N}/diffusion_pytorch_model.safetensors"
```
The official released equivalent is `models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt`. See [README_TRAIN.md](./README_TRAIN.md) for the full distillation workflow.
---
## 6. Additional Resources
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
+394
View File
@@ -0,0 +1,394 @@
# Wan2.1 Self-Forcing ODE 回归预训练指南
本文档介绍 Wan2.1 Self-Forcing 的 **ODE 回归预训练** 完整流程,涵盖环境配置、ODE 轨迹对生成、ODE 回归训练。
> **什么是 ODE 回归训练?**
>
> ODE 回归是 Self-Forcing 蒸馏的 **预训练阶段**,整体流程分两步:
>
> 1. **第一步 — 生成 ODE 对**(`generate_ode_pairs.py`):使用 **双向教师模型** Wan2.1-T2V-1.3B,对一组文本提示词执行完整的多步 CFG 去噪,将 ODE 轨迹上的中间 latent 与编码后的 prompt embedding 一起保存为 `.safetensors` 文件。
> 2. **第二步 — ODE 回归训练**(`train_ode.py`):加载第一步生成的 ODE 对,训练一个 **因果生成器**,在轨迹上随机抽样多个噪声等级,预测干净的终点 `x0`。训练得到的权重(通常保存为 `ode_init.pt`)作为 **Self-Forcing 蒸馏阶段**(`train_distill.py`,参见 [README_TRAIN.md](./README_TRAIN.md))的强初始化。
---
## 目录
- [一、环境配置](#一环境配置)
- [二、下载预训练模型](#二下载预训练模型)
- [三、第一步 — 生成 ODE 轨迹对](#三第一步--生成-ode-轨迹对)
- [3.1 下载提示词文件](#31-下载提示词文件)
- [3.2 运行 ODE 对生成](#32-运行-ode-对生成)
- [3.3 输出格式](#33-输出格式)
- [3.4 生成参数说明](#34-生成参数说明)
- [3.5 多卡生成](#35-多卡生成)
- [四、第二步 — ODE 回归训练](#四第二步--ode-回归训练)
- [4.1 快速开始](#41-快速开始)
- [4.2 训练常用参数](#42-训练常用参数)
- [4.3 使用 DeepSpeed-Zero-2 / FSDP 训练](#43-使用-deepspeed-zero-2--fsdp-训练)
- [4.4 多机分布式训练](#44-多机分布式训练)
- [五、使用训练好的 ODE 权重](#五使用训练好的-ode-权重)
- [六、更多资源](#六更多资源)
---
## 一、环境配置
**方式 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 deepspeed==0.17.0 numpy==1.26.4
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
```
**方式 3:使用 docker**
使用 docker 时,请确保机器中已正确安装显卡驱动与 CUDA 环境,然后依次执行以下命令:
```bash
# 拉取镜像
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
# 进入容器
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
```
---
## 二、下载预训练模型
ODE 生成阶段使用 **双向教师模型** Wan2.1-T2V-1.3B 进行去噪;ODE 训练阶段同样以该基础模型来初始化 **因果生成器**。
```bash
# 创建模型目录
mkdir -p models/Diffusion_Transformer
# 下载 Wan2.1 T2V 基础模型(生成时作为教师,训练时作为初始化)
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
```
---
## 三、第一步 — 生成 ODE 轨迹对
该步骤使用双向教师模型对每条提示词执行 **48 步 CFG 去噪**,并将得到的 ODE 轨迹与对应的 prompt embedding 一起保存为 `.safetensors` 文件。所有提示词处理完成后,会自动生成一个 `outputs.json` 标注文件,供后续训练阶段使用。
### 3.1 下载提示词文件
推荐使用 Self-Forcing 官方提供的提示词列表:
```bash
mkdir -p datasets
# 从 Self-Forcing 官方仓库下载 vidprom_filtered_extended.txt
hf download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir datasets/
# 最终路径:datasets/vidprom_filtered_extended.txt
```
也可以使用任意纯文本文件,每行一条提示词。
### 3.2 运行 ODE 对生成
直接复用启动脚本 [scripts/wan2.1_self_forcing/generate_ode_pairs.sh](./generate_ode_pairs.sh):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/generate_ode_pairs.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/wan2.1/wan_civitai.yaml" \
--video_sample_n_frames=81 \
--height=480 \
--width=832 \
--guidance_scale=6.0 \
--shift=8.0 \
--num_inference_steps=48 \
--caption_path="datasets/vidprom_filtered_extended.txt" \
--output_folder="datasets/ode_pairs_output" \
--sample_every_n_prompts=50
```
或者直接执行 shell 脚本:
```bash
bash scripts/wan2.1_self_forcing/generate_ode_pairs.sh
```
### 3.3 输出格式
生成完成后,`--output_folder` 中包含以下内容:
```
📦 datasets/ode_pairs_output/
├── 📄 00000.safetensors # 单条提示词对应的 ODE 轨迹与 prompt embedding
├── 📄 00001.safetensors
├── 📄 ...
├── 📂 sample/ # 可选预览视频(当 sample_every_n_prompts > 0 时)
│ └── 📄 00000_clean.mp4
└── 📄 outputs.json # 由 train_ode.py 读取的标注文件
```
每个 `.safetensors` 文件包含以下字段:
| 字段 | 形状 | 说明 |
|------|------|------|
| `latents` | `[5, C, F, H, W]` | 对 48 步 ODE 轨迹的稀疏 5 点采样:索引 `[0, 12, 24, 36, -1]`(初始噪声 → 3 个中间点 → 干净终点) |
| `prompt_embeds` | `[512, D]` | 经 padding 的 T5 prompt embedding(最大长度 512) |
| `prompt_attention_mask` | `[512]` | prompt embedding 的注意力掩码 |
自动生成的 `outputs.json` 与标准 `metadata.json` 格式一致:
```json
[
{ "file_path": "datasets/ode_pairs_output/00000.safetensors" },
{ "file_path": "datasets/ode_pairs_output/00001.safetensors" }
]
```
### 3.4 生成参数说明
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--pretrained_model_name_or_path` | Wan2.1-T2V-1.3B 教师模型路径 | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B` |
| `--config_path` | 模型配置 YAML | `config/wan2.1/wan_civitai.yaml` |
| `--caption_path` | 每行一条提示词的纯文本文件 | `datasets/vidprom_filtered_extended.txt` |
| `--output_folder` | `.safetensors` 与 `outputs.json` 的输出目录 | `datasets/ode_pairs_output` |
| `--guidance_scale` | 教师模型使用的 CFG 引导强度 | 6.0 |
| `--num_inference_steps` | 教师去噪步数(必须 ≥ 37,因为代码采样的索引为 `[0,12,24,36,-1]`) | 48 |
| `--shift` | `FlowMatchEulerDiscreteScheduler` 的 shift 值(**必须与训练阶段一致**) | 8.0 |
| `--video_sample_n_frames` | 生成视频的像素帧数 | 81 |
| `--height` / `--width` | 视频分辨率(像素) | 480 / 832 |
| `--negative_prompt` | CFG 使用的负向提示词 | (默认中文负向提示词) |
| `--sample_every_n_prompts` | 每 N 条提示词解码并保存一次预览 MP4(0 表示关闭) | 50 |
| `--mixed_precision` | `no` / `fp16` / `bf16` | `bf16` |
> ⚠️ **生成与训练阶段必须使用相同的 `--shift` 值**,提供的脚本均默认为 `8.0`。
### 3.5 多卡生成
`generate_ode_pairs.py` 基于 `accelerate` 实现,每个 rank 自动按 `prompt_index = index * world_size + rank` 交替处理提示词,并自动跳过已存在的文件,因此天然 **可断点续跑、可多卡并行**:
```bash
# 8 卡生成
accelerate launch --multi_gpu --num_processes=8 --mixed_precision="bf16" \
scripts/wan2.1_self_forcing/generate_ode_pairs.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/wan2.1/wan_civitai.yaml" \
--caption_path="datasets/vidprom_filtered_extended.txt" \
--output_folder="datasets/ode_pairs_output" \
--num_inference_steps=48 --guidance_scale=6.0 --shift=8.0 \
--height=480 --width=832 --video_sample_n_frames=81
```
最终 `outputs.json` 仅由主进程写入。
---
## 四、第二步 — ODE 回归训练
第一步完成、`datasets/ode_pairs_output/outputs.json` 生成后,即可训练因果生成器(`WanTransformer3DModel_SelfForcing`)来回归 ODE 轨迹。
每个训练样本上,训练脚本会:
1. 从一个 `.safetensors` 文件中加载稀疏的 5 点轨迹与 prompt embedding;
2. 按 **块**(每 `--num_frame_per_block` 帧共享同一时间步)随机选取一个轨迹点,将带噪 latent 与逐帧时间步送入因果生成器;
3. 将生成器输出的 flow 转换为 `x0` 预测,与轨迹的 **干净终点** 计算 MSE 损失。
### 4.1 快速开始
直接复用启动脚本 [scripts/wan2.1_self_forcing/train_ode.sh](./train_ode.sh):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.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/wan2.1_self_forcing/train_ode.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$ODE_DATA_META \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=500 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_ode_regression" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.05 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--denoising_step_indices_list 1000 750 500 250 \
--shift=8.0 \
--resume_from_checkpoint="latest" \
--trainable_modules "."
```
或者直接执行:
```bash
bash scripts/wan2.1_self_forcing/train_ode.sh
```
> 💡 因为 ODE 轨迹与 prompt embedding 已在第一步预先计算完毕,**ODE 训练阶段不会再调用 VAE / 文本编码器**,训练速度快、显存占用低。当 `outputs.json` 中已经使用绝对路径时,`train_data_dir` 可以留空。
### 4.2 训练常用参数
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--pretrained_model_name_or_path` | 用于初始化因果生成器的基础模型 | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B` |
| `--config_path` | 模型配置 YAML | `config/wan2.1/wan_civitai.yaml` |
| `--train_data_dir` | 拼接到 `file_path` 之前的可选根目录;若 `outputs.json` 已使用绝对路径可留空 | `""` |
| `--train_data_meta` | 第一步生成的标注 JSON | `datasets/ode_pairs_output/outputs.json` |
| `--train_batch_size` | 每卡 batch size | 1 |
| `--gradient_accumulation_steps` | 梯度累积步数 | 1 |
| `--dataloader_num_workers` | DataLoader 子进程数 | 8 |
| `--num_train_epochs` | 训练 epoch 数 | 100 |
| `--checkpointing_steps` | 每 N 步保存一次 checkpoint | 500 |
| `--learning_rate` | 初始学习率 | 2e-06 |
| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` |
| `--lr_warmup_steps` | 学习率预热步数 | 100 |
| `--seed` | 随机种子 | 42 |
| `--output_dir` | 输出目录 | `output_dir_wan2.1_self_forcing_ode_regression` |
| `--gradient_checkpointing` | 启用激活重计算 | - |
| `--mixed_precision` | `fp16` / `bf16` | `bf16` |
| `--adam_weight_decay` | AdamW 权重衰减 | 3e-2 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--max_grad_norm` | 梯度裁剪阈值 | 0.05 |
| `--trainable_modules` | 可训练模块(`"."` 表示全量) | `"."` |
| `--resume_from_checkpoint` | 恢复训练路径或 `"latest"` | `latest` |
**ODE 特有参数**(除非清楚后果,否则需与第一步保持一致):
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--train_sampling_steps` | 调度器总时间步数,从中按 `denoising_step_indices_list` 抽样 | 1000 |
| `--denoising_step_indices_list` | ODE 回归使用的离散时间步索引(与第一步抽样的 5 个稀疏轨迹点对应) | `1000 750 500 250` |
| `--shift` | `FlowMatchEulerDiscreteScheduler` 的 shift —— **必须与第一步生成时使用的 `--shift` 一致** | 8.0 |
| `--num_frame_per_block` | 每个因果块包含的帧数(同一块内的帧共享同一时间步) | 3 |
| `--independent_first_frame` | 第一帧是否独立(`[1, N, N, ...]` 块模式) | - |
| `--context_noise` | 上下文噪声等级(与下游 Self-Forcing 蒸馏配置匹配) | 0 |
**验证参数(可选)**:
| 参数 | 说明 | 示例 |
|------|------|------|
| `--validation_steps` | 每 N 步执行一次验证 | 2000 |
| `--validation_epochs` | 每 N 个 epoch 执行一次验证 | 5 |
| `--validation_prompts` | 验证视频生成使用的提示词 | 英文提示词 |
| `--video_sample_size` | 验证采样尺寸 | 640 |
| `--video_sample_n_frames` | 验证生成的视频帧数 | 81 |
| `--fix_sample_size` | 验证使用的固定 `[高度, 宽度]` | `480 832` |
### 4.3 使用 DeepSpeed-Zero-2 / FSDP 训练
多卡训练支持与蒸馏阶段相同的显存节约后端。将 4.1 中 `accelerate launch` 前缀替换为以下任意一种即可:
**DeepSpeed-Zero-2**(推荐默认):
```bash
accelerate launch \
--use_deepspeed --deepspeed_config_file config/zero_stage2_config.json \
--deepspeed_multinode_launcher standard \
scripts/wan2.1_self_forcing/train_ode.py \
... # 训练参数与 4.1 相同
```
**FSDP**(DeepSpeed-Zero-2 显存不足时使用):
```bash
accelerate launch --mixed_precision="bf16" \
--use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \
--fsdp_transformer_layer_cls_to_wrap=CasualWanAttentionBlock \
--fsdp_sharding_strategy "FULL_SHARD" \
--fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" \
--fsdp_cpu_ram_efficient_loading False \
scripts/wan2.1_self_forcing/train_ode.py \
... # 训练参数与 4.1 相同
```
**DeepSpeed-Zero-3**(适用于超大模型,1.3B 通常不需要):
```bash
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true \
--use_deepspeed --deepspeed_config_file config/zero_stage3_config.json \
--deepspeed_multinode_launcher standard \
scripts/wan2.1_self_forcing/train_ode.py \
... # 训练参数与 4.1 相同
# 训练完成后将分片 checkpoint 转为单文件 bf16:
python scripts/zero_to_bf16.py \
output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N} \
output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N}-outputs \
--max_shard_size 80GB --safe_serialization
```
### 4.4 多机分布式训练
假设 2 台机器、每台 8 卡:
**机器 0(Master)**:
```bash
export MASTER_ADDR="192.168.1.100"
export MASTER_PORT=10086
export WORLD_SIZE=2
export NUM_PROCESS=16
export RANK=0
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
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 \
--use_deepspeed --deepspeed_config_file config/zero_stage2_config.json \
--deepspeed_multinode_launcher standard \
scripts/wan2.1_self_forcing/train_ode.py \
... # 训练参数与 4.1 相同
```
**机器 1(Worker)**:与 Master 完全相同,仅将 `export RANK=1`。
**注意事项**:
- 优先使用 RDMA / InfiniBand。无 RDMA 时需设置 `NCCL_IB_DISABLE=1` 与 `NCCL_P2P_DISABLE=1`。
- 所有机器必须共享同一份 `outputs.json` 与对应的 `.safetensors` 文件(NFS / 共享存储)。
---
## 五、使用训练好的 ODE 权重
`output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N}/` 中保存的 ODE-init 权重作为 Self-Forcing 蒸馏的初始化。在 `train_distill.py` 中通过 `--ode_transformer_path` 指定即可:
```bash
# 例:保存的权重文件(如 diffusion_pytorch_model.safetensors)
--ode_transformer_path="output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N}/diffusion_pytorch_model.safetensors"
```
官方发布的对应权重为 `models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt`。完整的蒸馏流程参见 [README_TRAIN.md](./README_TRAIN.md)。
---
## 六、更多资源
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
@@ -0,0 +1,341 @@
# Based on https://github.com/guandeh17/Self-Forcing
import argparse
import gc
import json
import math
import os
import sys
import torch
from accelerate import Accelerator
from diffusers import FlowMatchEulerDiscreteScheduler
from einops import rearrange
from omegaconf import OmegaConf
from safetensors.torch import save_file
from tqdm import tqdm
from transformers import AutoTokenizer
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.models import (AutoencoderKLWan, WanT5EncoderModel,
WanTransformer3DModel)
from videox_fun.utils.utils import save_videos_grid
def filter_kwargs(cls, kwargs):
import inspect
sig = inspect.signature(cls.__init__)
valid_params = set(sig.parameters.keys()) - {'self', 'cls'}
return {k: v for k, v in kwargs.items() if k in valid_params}
def load_prompts(caption_path):
with open(caption_path, encoding="utf-8") as f:
return [line.rstrip() for line in f if line.strip()]
def main():
parser = argparse.ArgumentParser(description="Generate ODE trajectory pairs for ODE regression training.")
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
required=True,
help="Path to pretrained model or model identifier from huggingface.co/models.",
)
parser.add_argument(
"--config_path",
type=str,
required=True,
help="Path to the model config YAML file (e.g. config/wan2.1/wan_civitai.yaml).",
)
parser.add_argument(
"--caption_path",
type=str,
required=True,
help="Path to a text file containing prompts, one per line. Download at https://huggingface.co/gdhe17/Self-Forcing/blob/main/vidprom_filtered_extended.txt",
)
parser.add_argument(
"--output_folder",
type=str,
required=True,
help="The output directory where per-prompt .safetensors ODE trajectory files will be saved.",
)
parser.add_argument(
"--guidance_scale",
type=float,
default=6.0,
help="Classifier-free guidance scale for ODE denoising. Default: 6.0.",
)
parser.add_argument(
"--num_inference_steps",
type=int,
default=48,
help="Number of ODE denoising steps for the teacher model. Default: 48.",
)
parser.add_argument(
"--shift",
type=float,
default=8.0,
help="Shift value for FlowMatchEulerDiscreteScheduler. Default: 8.0.",
)
parser.add_argument(
"--video_sample_n_frames",
type=int,
default=81,
help="Number of pixel frames for the generated video. Default: 81.",
)
parser.add_argument(
"--height",
type=int,
default=480,
help="Video height in pixels. Will be divided by VAE spatial ratio (8) for latent size. Default: 480.",
)
parser.add_argument(
"--width",
type=int,
default=832,
help="Video width in pixels. Will be divided by VAE spatial ratio (8) for latent size. Default: 832.",
)
parser.add_argument(
"--negative_prompt",
type=str,
default="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
help="The negative prompt for classifier-free guidance.",
)
parser.add_argument(
"--vae_mini_batch",
type=int,
default=1,
help="Mini batch size for VAE decode. Default: 1.",
)
parser.add_argument(
"--sample_every_n_prompts",
type=int,
default=0,
help="Decode and save sample video every N prompts for visualization. 0 to disable. Default: 0.",
)
parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank.")
parser.add_argument(
"--mixed_precision",
type=str,
default="bf16",
choices=["no", "fp16", "bf16"],
help="Whether to use mixed precision. Default: bf16.",
)
args = parser.parse_args()
# Initialize accelerator for distributed generation
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device = accelerator.device
world_size = accelerator.num_processes
rank = accelerator.process_index
# Disable gradients globally since this is inference-only
torch.set_grad_enabled(False)
# Enable TF32 for faster computation on Ampere GPUs
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
config = OmegaConf.load(args.config_path)
# For mixed precision we cast all weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
weight_dtype = torch.float32
if accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
# Load tokenizer and text encoder
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(args.pretrained_model_name_or_path,
config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer'))
)
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(device).eval()
text_encoder.requires_grad_(False)
# Load VAE
vae = AutoencoderKLWan.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(device, dtype=weight_dtype).eval()
vae.requires_grad_(False)
# Load bidirectional transformer (teacher)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path,
config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
).to(device, dtype=weight_dtype).eval()
transformer.requires_grad_(False)
# Load scheduler and configure shift
scheduler_kwargs = OmegaConf.to_container(config['scheduler_kwargs'])
scheduler_kwargs['shift'] = args.shift
noise_scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, scheduler_kwargs)
)
noise_scheduler.set_timesteps(args.num_inference_steps, device=device)
timesteps = noise_scheduler.timesteps
# Compute latent shapes from VAE config
# latent_h/w = pixel_h/w / spatial_compression_ratio
# num_frames = (pixel_frames - 1) / temporal_compression_ratio + 1
latent_h = args.height // vae.spatial_compression_ratio
latent_w = args.width // vae.spatial_compression_ratio
latent_channels = vae.latent_channels
num_frames = (args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio + 1
# Compute seq_len for transformer
patch_size = transformer.config.patch_size
seq_len = math.ceil((latent_h * latent_w) / (patch_size[1] * patch_size[2]) * num_frames)
# Load prompts and distribute across ranks
prompts = load_prompts(args.caption_path)
os.makedirs(args.output_folder, exist_ok=True)
total_per_rank = int(math.ceil(len(prompts) / world_size))
# Negative prompt embedding (unconditional)
with torch.no_grad():
neg_inputs = tokenizer(
[args.negative_prompt], padding="max_length", max_length=512,
truncation=True, add_special_tokens=True, return_tensors="pt"
)
neg_seq_lens = neg_inputs.attention_mask.gt(0).sum(dim=1).long()
neg_embeds = text_encoder(neg_inputs.input_ids.to(device), attention_mask=neg_inputs.attention_mask.to(device))[0]
neg_prompt_embeds = [neg_embeds[i, :neg_seq_lens[i]] for i in range(neg_embeds.shape[0])]
# Main generation loop: each rank processes interleaved prompts
for index in tqdm(range(total_per_rank), disable=rank != 0, desc="Generating ODE pairs"):
prompt_index = index * world_size + rank
if prompt_index >= len(prompts):
continue
prompt = prompts[prompt_index]
output_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
print(rank, output_path)
if os.path.exists(output_path):
continue
# Encode prompt (keep padded [512, D] for saving to safetensors)
text_inputs = tokenizer(
[prompt],
padding="max_length",
max_length=512,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
prompt_attention_mask = text_inputs.attention_mask # [1, 512]
text_seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
text_embeds = text_encoder(text_inputs.input_ids.to(device), attention_mask=prompt_attention_mask.to(device))[0] # [1, 512, D]
prompt_embeds = [text_embeds[i, :text_seq_lens[i]] for i in range(text_embeds.shape[0])]
# Sample initial noise: [B, C, F, H, W]
latents = torch.randn(
[1, latent_channels, num_frames, latent_h, latent_w],
dtype=weight_dtype, device=device
)
# Run full ODE denoising with CFG, collecting intermediate latents
noisy_inputs = []
# Reset scheduler state for each prompt to avoid stale `_step_index`
# leaking across iterations and causing IndexError on `self.sigmas[sigma_idx + 1]`.
noise_scheduler._step_index = None
if hasattr(noise_scheduler, 'model_outputs'):
noise_scheduler.model_outputs = []
for progress_id, t in enumerate(timesteps):
timestep = t.expand(latents.shape[0]) # [B]
noisy_inputs.append(latents.clone())
# Conditional prediction
with torch.cuda.amp.autocast(dtype=weight_dtype):
flow_pred_cond = transformer(
x=latents,
context=prompt_embeds,
t=timestep,
seq_len=seq_len,
)
# Unconditional prediction
flow_pred_uncond = transformer(
x=latents,
context=neg_prompt_embeds,
t=timestep,
seq_len=seq_len,
)
# CFG
flow_pred = flow_pred_uncond + args.guidance_scale * (flow_pred_cond - flow_pred_uncond)
# Scheduler step
latents = noise_scheduler.step(flow_pred, t, latents, return_dict=False)[0]
# Append final clean latent
noisy_inputs.append(latents.clone())
# Stack all intermediate + final latents: [1, num_steps+1, C, F, H, W]
noisy_inputs_tensor = torch.stack(noisy_inputs, dim=1)
# Sparse sample 5 points along the ODE trajectory: [0, 12, 24, 36, -1]
# This reduces storage while preserving the trajectory shape
noisy_inputs_tensor = noisy_inputs_tensor[:, [0, 12, 24, 36, -1]]
# Save as safetensors with latents, prompt_embeds, prompt_attention_mask
save_file(
{
"latents": noisy_inputs_tensor.squeeze(0).cpu(),
"prompt_embeds": text_embeds.squeeze(0).cpu(),
"prompt_attention_mask": prompt_attention_mask.squeeze(0).cpu(),
},
output_path,
metadata={"prompt": prompt},
)
# Decode and save sample video for visualization
if args.sample_every_n_prompts > 0 and prompt_index % args.sample_every_n_prompts == 0:
sample_dir = os.path.join(args.output_folder, "sample")
os.makedirs(sample_dir, exist_ok=True)
with torch.no_grad():
# Decode the final clean latent (last sparse point)
clean_latent = noisy_inputs_tensor[:, -1] # [1, C, F, H, W]
video = vae.decode(clean_latent.to(vae.dtype)).sample
video = (video / 2 + 0.5).clamp(0, 1)
save_videos_grid(
video.cpu().float(),
os.path.join(sample_dir, f"{prompt_index:05d}_clean.mp4"),
)
gc.collect()
torch.cuda.empty_cache()
accelerator.wait_for_everyone()
# Write outputs.json annotation file for ImageVideoSafetensorsDataset
# This JSON lists all generated safetensors files so they can be loaded by train_ode.py
if accelerator.is_main_process:
safe_json_writer = []
for i in range(len(prompts)):
safetensor_path = os.path.join(args.output_folder, f"{i:05d}.safetensors")
if os.path.exists(safetensor_path):
safe_json_writer.append({"file_path": safetensor_path})
json_path = os.path.join(args.output_folder, "outputs.json")
with open(json_path, "w", encoding="utf-8") as f:
json.dump(safe_json_writer, f, ensure_ascii=False, indent=4)
print(f"Done. Generated {len(safe_json_writer)} ODE pairs, saved to {args.output_folder}")
print(f"Annotation JSON: {json_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,20 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# 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
# Download vidprom_filtered_extended.txt from:
# https://huggingface.co/gdhe17/Self-Forcing/blob/main/vidprom_filtered_extended.txt
accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/generate_ode_pairs.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/wan2.1/wan_civitai.yaml" \
--video_sample_n_frames=81 \
--height=480 \
--width=832 \
--guidance_scale=6.0 \
--shift=8.0 \
--num_inference_steps=48 \
--caption_path="datasets/vidprom_filtered_extended.txt" \
--output_folder="datasets/ode_pairs_output" \
--sample_every_n_prompts=50
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,47 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/train_distill.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--fix_sample_size 480 832 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--learning_rate_critic=4e-07 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_kv_cache_training \
--num_frame_per_block=3 \
--train_mode="normal" \
--trainable_modules "." \
--ode_transformer_path="models/Diffusion_Transformer/Self-Forcing/checkpoints/ode_init.pt" \
--low_vram
File diff suppressed because it is too large Load Diff
+34
View File
@@ -0,0 +1,34 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
export DATASET_NAME=""
export ODE_DATA_META="datasets/ode_pairs_output/outputs.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/wan2.1_self_forcing/train_ode.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$ODE_DATA_META \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=500 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_self_forcing_ode_regression" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.05 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--denoising_step_indices_list 1000 750 500 250 \
--shift=8.0 \
--resume_from_checkpoint="latest" \
--trainable_modules "."
+4 -2
View File
@@ -494,7 +494,8 @@ class ImageVideoControlDataset(Dataset):
shuffle(subject_id)
subject_images = []
for i in range(min(len(subject_id), 4)):
subject_image = Image.open(subject_id[i])
subject_image_path = subject_id[i] if self.data_root is None else os.path.join(self.data_root, subject_id[i])
subject_image = Image.open(subject_image_path)
if self.padding_subject_info:
img = padding_image(subject_image, visual_width, visual_height)
@@ -547,7 +548,8 @@ class ImageVideoControlDataset(Dataset):
shuffle(subject_id)
subject_images = []
for i in range(min(len(subject_id), 4)):
subject_image = Image.open(subject_id[i]).convert('RGB')
subject_image_path = subject_id[i] if self.data_root is None else os.path.join(self.data_root, subject_id[i])
subject_image = Image.open(subject_image_path).convert('RGB')
if self.padding_subject_info:
img = padding_image(subject_image, visual_width, visual_height)
+6 -4
View File
@@ -1,6 +1,7 @@
import importlib.util
from .cogvideox_xfuser import CogVideoXMultiGPUsAttnProcessor2_0
from .ernie_image_xfuser import ErnieImageMultiGPUsAttnProcessor
from .flashhead_xfuser import usp_attn_flashhead_forward
from .flux2_xfuser import Flux2MultiGPUsAttnProcessor2_0
from .flux_xfuser import FluxMultiGPUsAttnProcessor2_0
@@ -8,9 +9,9 @@ from .fsdp import shard_model
from .fuser import (get_sequence_parallel_rank,
get_sequence_parallel_world_size, get_sp_group,
get_world_group, init_distributed_environment,
initialize_model_parallel, sequence_parallel_all_gather,
sequence_parallel_chunk, set_multi_gpus_devices,
xFuserLongContextAttention)
initialize_model_parallel, model_parallel_is_initialized,
sequence_parallel_all_gather, sequence_parallel_chunk,
set_multi_gpus_devices, xFuserLongContextAttention)
from .hunyuanvideo_xfuser import HunyuanVideoMultiGPUsAttnProcessor2_0
from .infinitalk_xfuser import usp_attn_infinitetalk_forward
from .longcatvideo_xfuser import (usp_attn_longcatvideo_avatar_forward,
@@ -20,7 +21,8 @@ from .longcatvideo_xfuser import (usp_attn_longcatvideo_avatar_forward,
from .ltx2_xfuser import (LTX2MultiGPUsAttnProcessor,
LTX2PerturbedMultiGPUsAttnProcessor)
from .qwen_xfuser import QwenImageMultiGPUsAttnProcessor2_0
from .wan_xfuser import usp_attn_forward, usp_attn_s2v_forward
from .wan_xfuser import (usp_attn_forward, usp_attn_s2v_forward,
usp_attn_self_forcing_forward)
from .z_image_xfuser import ZMultiGPUsSingleStreamAttnProcessor
# The pai_fuser is an internally developed acceleration package, which can be used on PAI.
+98
View File
@@ -0,0 +1,98 @@
from typing import Optional
import torch
import torch.nn.functional as F
from diffusers.models.attention import Attention
from .fuser import xFuserLongContextAttention
class ErnieImageMultiGPUsAttnProcessor:
"""
Processor for Ernie-Image multi-GPU inference using sequence parallel attention.
This processor adapts the single-stream attention mechanism to work with
xFuserLongContextAttention for distributed inference across multiple GPUs.
"""
_attention_backend = None
_parallel_config = None
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"ErnieImageMultiGPUsAttnProcessor requires PyTorch 2.0. "
"To use it, please upgrade PyTorch to version 2.0 or higher."
)
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
freqs_cis: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# Step 1: QKV projections
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
# Reshape to [batch, seq_len, heads, head_dim]
query = query.unflatten(-1, (attn.heads, -1))
key = key.unflatten(-1, (attn.heads, -1))
value = value.unflatten(-1, (attn.heads, -1))
# Step 2: Apply QK normalization
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# Step 3: Apply rotary positional embeddings (RoPE)
# Same rotate_half logic as ErnieImageSingleStreamAttnProcessor (rotary_interleaved=False)
if freqs_cis is not None:
def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
rot_dim = freqs_cis.shape[-1]
x, x_pass = x_in[..., :rot_dim], x_in[..., rot_dim:]
cos_ = torch.cos(freqs_cis).to(x.dtype)
sin_ = torch.sin(freqs_cis).to(x.dtype)
# Non-interleaved rotate_half: [-x2, x1]
x1, x2 = x.chunk(2, dim=-1)
x_rotated = torch.cat((-x2, x1), dim=-1)
return torch.cat((x * cos_ + x_rotated * sin_, x_pass), dim=-1)
query = apply_rotary_emb(query, freqs_cis)
key = apply_rotary_emb(key, freqs_cis)
# Step 4: Cast to correct dtype
dtype = query.dtype
query, key = query.to(dtype), key.to(dtype)
# Step 5: Handle attention mask format conversion if needed
# From [batch, seq_len] to [batch, 1, 1, seq_len] -> broadcast to [batch, heads, seq_len, seq_len]
if attention_mask is not None and attention_mask.ndim == 2:
attention_mask = attention_mask[:, None, None, :]
# Step 6: Perform distributed attention using xFuserLongContextAttention
# This handles sequence parallelism automatically
half_dtypes = (torch.float16, torch.bfloat16)
def half(x):
return x if x.dtype in half_dtypes else x.to(torch.bfloat16)
hidden_states = xFuserLongContextAttention()(
None,
half(query),
half(key),
half(value),
dropout_p=0.0,
causal=False,
)
# Step 7: Reshape back and project output
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(dtype)
output = attn.to_out[0](hidden_states)
return output
-18
View File
@@ -1,21 +1,3 @@
# Copyright 2025 The VideoX-Fun 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.
"""
Multi-GPU sequence parallel attention processors for LTX2 transformer.
"""
from typing import Tuple
import os
+154
View File
@@ -1,3 +1,5 @@
import math
import torch
import torch.cuda.amp as amp
@@ -61,6 +63,55 @@ def rope_apply(x, grid_sizes, freqs):
output.append(x_i)
return torch.stack(output).to(dtype)
@amp.autocast(enabled=False)
@torch.compiler.disable()
def causal_rope_apply(x, grid_sizes, freqs, start_frame=0):
"""
Apply causal rotary positional embedding with frame offset support.
This function applies RoPE with a starting frame offset, enabling causal
inference where different frames can have different positional indices.
Args:
x: Input tensor with shape (batch, seq_len, n_channels, c*2)
grid_sizes: Grid dimensions (f, h, w) for each sample
freqs: Precomputed frequency parameters
start_frame: Starting frame index for causal positioning
Returns:
Tensor with causal RoPE applied
"""
n, c = x.size(2), x.size(3) // 2
# Split freqs into temporal, height, and width components
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# Process each sample in the batch
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# Reshape and convert to complex numbers
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
seq_len, n, -1, 2))
# Broadcast frequencies with start_frame offset for temporal dimension
freqs_i = torch.cat([
freqs[0][start_frame:start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1),
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
],
dim=-1).reshape(seq_len, 1, -1)
# Apply rotation: x * exp(i*freq)
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
# Concatenate with padding tokens (if any)
x_i = torch.cat([x_i, x[i, seq_len:]])
# Append to collection
output.append(x_i)
return torch.stack(output).type_as(x)
def rope_apply_qk(q, k, grid_sizes, freqs):
q = rope_apply(q, grid_sizes, freqs)
k = rope_apply(k, grid_sizes, freqs)
@@ -178,4 +229,107 @@ def usp_attn_s2v_forward(self,
# output
x = x.flatten(2)
x = self.o(x)
return x
def usp_attn_self_forcing_forward(
self,
x,
seq_lens,
grid_sizes,
freqs,
block_mask,
kv_cache=None,
current_start=0,
cache_start=None,
dtype=torch.bfloat16,
t=0
):
"""
USP attention forward for Self-Forcing with KV cache support.
Combines sequence parallelism with causal KV cache inference.
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
half_dtypes = (torch.float16, torch.bfloat16)
sp_size = get_sequence_parallel_world_size()
sp_rank = get_sequence_parallel_rank()
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
if cache_start is None:
cache_start = current_start
# QKV computation
def qkv_fn(x):
q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
return q, k, v
q, k, v = qkv_fn(x)
# Inference mode with KV cache
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
current_start_frame = current_start // frame_seqlen
# Step 1: all_gather QKV to restore full sequence
q_full = get_sp_group().all_gather(q, dim=1) # [B, L_full, H, D]
k_full = get_sp_group().all_gather(k, dim=1)
v_full = get_sp_group().all_gather(v, dim=1)
# Step 2: apply causal RoPE on full sequence with frame offset
roped_query_full = causal_rope_apply(q_full, grid_sizes, freqs,
start_frame=current_start_frame).type_as(v_full)
roped_key_full = causal_rope_apply(k_full, grid_sizes, freqs,
start_frame=current_start_frame).type_as(v_full)
current_end = current_start + roped_query_full.shape[1]
sink_tokens = self.sink_size * frame_seqlen
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query_full.shape[1]
# Step 3: KV cache update logic with full keys
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and \
(num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key_full
kv_cache["v"][:, local_start_index:local_end_index] = v_full
else:
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key_full
kv_cache["v"][:, local_start_index:local_end_index] = v_full
# Step 4: chunk back to SP distribution for attention computation
roped_query = torch.chunk(roped_query_full, sp_size, dim=1)[sp_rank]
# Step 5: compute attention using xFuserLongContextAttention for sequence parallelism
# Chunk KV cache window to match SP distribution
kv_k_full = kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
kv_v_full = kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
kv_k = torch.chunk(kv_k_full, sp_size, dim=1)[sp_rank]
kv_v = torch.chunk(kv_v_full, sp_size, dim=1)[sp_rank]
x = xFuserLongContextAttention()(
None,
query=half(roped_query),
key=half(kv_k),
value=kv_v,
window_size=self.window_size
)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
# Output projection
x = x.flatten(2)
x = self.o(x)
return x
+10 -1
View File
@@ -25,10 +25,18 @@ try:
from transformers import Qwen3VLForConditionalGeneration
except:
Qwen3VLForConditionalGeneration = None
print("Your transformers version is too old to load Qwen3VLForConditionalGeneration. If you wish to use QwenImage, please upgrade your transformers package to the latest version.")
print("Your transformers version is too old to load Qwen3VLForConditionalGeneration. If you wish to use Qwen3VLForConditionalGeneration, please upgrade your transformers package to the latest version.")
try:
from transformers import Mistral3Model, Ministral3ForCausalLM
except:
Mistral3Model = None
Ministral3ForCausalLM = None
print("Your transformers version is too old to load Mistral3Model and Ministral3ForCausalLM. If you wish to use ErnieImage, please upgrade your transformers package to the latest version.")
from .cogvideox_transformer3d import CogVideoXTransformer3DModel
from .cogvideox_vae import AutoencoderKLCogVideoX
from .ernie_image_transformer import ErnieImageTransformer2DModel
from .fantasytalking_audio_encoder import FantasyTalkingAudioEncoder
from .fantasytalking_transformer3d import FantasyTalkingTransformer3DModel
from .flashhead_audio_encoder import FlashHeadAudioEncoder
@@ -69,6 +77,7 @@ from .wan_transformer3d import (Wan2_2Transformer3DModel, WanRMSNorm,
WanSelfAttention, WanTransformer3DModel)
from .wan_transformer3d_animate import Wan2_2Transformer3DModel_Animate
from .wan_transformer3d_s2v import Wan2_2Transformer3DModel_S2V
from .wan_transformer3d_self_forcing import WanTransformer3DModel_SelfForcing
from .wan_transformer3d_vace import VaceWanTransformer3DModel
from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_
from .wan_vae3_8 import AutoencoderKLWan2_2_, AutoencoderKLWan3_8
@@ -0,0 +1,501 @@
# Modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/transformers/transformer_ernie_image.py
# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
Ernie-Image Transformer2DModel for HuggingFace Diffusers.
"""
import inspect
from dataclasses import dataclass
from typing import Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import PeftAdapterMixin
from diffusers.models.attention_processor import Attention
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import RMSNorm
from diffusers.utils import BaseOutput, logging
from .attention_utils import attention
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class ErnieImageTransformer2DModelOutput(BaseOutput):
sample: torch.Tensor
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
assert dim % 2 == 0
scale = torch.arange(0, dim, 2, dtype=torch.float32, device=pos.device) / dim
omega = 1.0 / (theta**scale)
out = torch.einsum("...n,d->...nd", pos, omega)
return out.float()
class ErnieImageEmbedND3(nn.Module):
def __init__(self, dim: int, theta: int, axes_dim: Tuple[int, int, int]):
super().__init__()
self.dim = dim
self.theta = theta
self.axes_dim = list(axes_dim)
def forward(self, ids: torch.Tensor) -> torch.Tensor:
emb = torch.cat([rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(3)], dim=-1)
emb = emb.unsqueeze(2) # [B, S, 1, head_dim//2]
return torch.stack([emb, emb], dim=-1).reshape(*emb.shape[:-1], -1) # [B, S, 1, head_dim]
class ErnieImagePatchEmbedDynamic(nn.Module):
def __init__(self, in_channels: int, embed_dim: int, patch_size: int):
super().__init__()
self.patch_size = patch_size
self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size, bias=True)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.proj(x)
batch_size, dim, height, width = x.shape
return x.reshape(batch_size, dim, height * width).transpose(1, 2).contiguous()
class ErnieImageSingleStreamAttnProcessor:
_attention_backend = None
_parallel_config = None
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"ErnieImageSingleStreamAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher."
)
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
freqs_cis: torch.Tensor | None = None,
) -> torch.Tensor:
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
query = query.unflatten(-1, (attn.heads, -1))
key = key.unflatten(-1, (attn.heads, -1))
value = value.unflatten(-1, (attn.heads, -1))
# Apply Norms
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# Apply RoPE: same rotate_half logic as Megatron _apply_rotary_pos_emb_bshd (rotary_interleaved=False)
# x_in: [B, S, heads, head_dim], freqs_cis: [B, S, 1, head_dim] with angles [θ0,θ0,θ1,θ1,...]
def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
rot_dim = freqs_cis.shape[-1]
x, x_pass = x_in[..., :rot_dim], x_in[..., rot_dim:]
cos_ = torch.cos(freqs_cis).to(x.dtype)
sin_ = torch.sin(freqs_cis).to(x.dtype)
# Non-interleaved rotate_half: [-x2, x1]
x1, x2 = x.chunk(2, dim=-1)
x_rotated = torch.cat((-x2, x1), dim=-1)
return torch.cat((x * cos_ + x_rotated * sin_, x_pass), dim=-1)
if freqs_cis is not None:
query = apply_rotary_emb(query, freqs_cis)
key = apply_rotary_emb(key, freqs_cis)
# Cast to correct dtype
dtype = query.dtype
query, key = query.to(dtype), key.to(dtype)
# From [batch, seq_len] to [batch, 1, 1, seq_len] -> broadcast to [batch, heads, seq_len, seq_len]
if attention_mask is not None and attention_mask.ndim == 2:
attention_mask = attention_mask[:, None, None, :]
# Compute joint attention
hidden_states = attention(
query,
key,
value,
attn_mask=attention_mask,
)
# Reshape back
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(dtype)
output = attn.to_out[0](hidden_states)
return output
class ErnieImageAttention(nn.Module):
_default_processor_cls = ErnieImageSingleStreamAttnProcessor
def __init__(
self,
query_dim: int,
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias: bool = False,
qk_norm: str = "rms_norm",
added_proj_bias: bool | None = True,
out_bias: bool = True,
eps: float = 1e-5,
out_dim: int = None,
elementwise_affine: bool = True,
processor=None,
):
super().__init__()
self.head_dim = dim_head
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.query_dim = query_dim
self.out_dim = out_dim if out_dim is not None else query_dim
self.heads = out_dim // dim_head if out_dim is not None else heads
self.use_bias = bias
self.dropout = dropout
self.added_proj_bias = added_proj_bias
self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
# QK Norm
if qk_norm == "layer_norm":
self.norm_q = torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
self.norm_k = torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
elif qk_norm == "rms_norm":
self.norm_q = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
self.norm_k = torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
else:
raise ValueError(
f"unknown qk_norm: {qk_norm}. Should be one of None, 'layer_norm', 'fp32_layer_norm', 'layer_norm_across_heads', 'rms_norm', 'rms_norm_across_heads', 'l2'."
)
self.to_out = torch.nn.ModuleList([])
self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
if processor is None:
processor = self._default_processor_cls()
self.set_processor(processor)
def set_processor(self, processor) -> None:
"""
Set the attention processor to use.
Args:
processor: The attention processor to use.
"""
if (
hasattr(self, "processor")
and isinstance(self.processor, torch.nn.Module)
and not isinstance(processor, torch.nn.Module)
):
logger.info(f"You are removing possibly trained weights of {self.processor} with {processor}")
self._modules.pop("processor")
self.processor = processor
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
image_rotary_emb: torch.Tensor | None = None,
**kwargs,
) -> torch.Tensor:
attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys())
unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters]
if len(unused_kwargs) > 0:
logger.warning(
f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored."
)
kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters}
return self.processor(self, hidden_states, attention_mask, image_rotary_emb, **kwargs)
class ErnieImageFeedForward(nn.Module):
def __init__(self, hidden_size: int, ffn_hidden_size: int):
super().__init__()
# Separate gate and up projections (matches converted weights)
self.gate_proj = nn.Linear(hidden_size, ffn_hidden_size, bias=False)
self.up_proj = nn.Linear(hidden_size, ffn_hidden_size, bias=False)
self.linear_fc2 = nn.Linear(ffn_hidden_size, hidden_size, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear_fc2(self.up_proj(x) * F.gelu(self.gate_proj(x)))
class ErnieImageSharedAdaLNBlock(nn.Module):
def __init__(
self, hidden_size: int, num_heads: int, ffn_hidden_size: int, eps: float = 1e-6, qk_layernorm: bool = True
):
super().__init__()
self.adaLN_sa_ln = RMSNorm(hidden_size, eps=eps)
self.self_attention = ErnieImageAttention(
query_dim=hidden_size,
dim_head=hidden_size // num_heads,
heads=num_heads,
qk_norm="rms_norm" if qk_layernorm else None,
eps=eps,
bias=False,
out_bias=False,
processor=ErnieImageSingleStreamAttnProcessor(),
)
self.adaLN_mlp_ln = RMSNorm(hidden_size, eps=eps)
self.mlp = ErnieImageFeedForward(hidden_size, ffn_hidden_size)
def forward(
self,
x,
rotary_pos_emb,
temb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None = None,
):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = temb
residual = x
x = self.adaLN_sa_ln(x)
x = (x.float() * (1 + scale_msa.float()) + shift_msa.float()).to(x.dtype)
x_bsh = x.permute(1, 0, 2) # [S, B, H] → [B, S, H] for diffusers Attention (batch-first)
attn_out = self.self_attention(x_bsh, attention_mask=attention_mask, image_rotary_emb=rotary_pos_emb)
attn_out = attn_out.permute(1, 0, 2) # [B, S, H] → [S, B, H]
x = residual + (gate_msa.float() * attn_out.float()).to(x.dtype)
residual = x
x = self.adaLN_mlp_ln(x)
x = (x.float() * (1 + scale_mlp.float()) + shift_mlp.float()).to(x.dtype)
return residual + (gate_mlp.float() * self.mlp(x).float()).to(x.dtype)
class ErnieImageAdaLNContinuous(nn.Module):
def __init__(self, hidden_size: int, eps: float = 1e-6):
super().__init__()
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=eps)
self.linear = nn.Linear(hidden_size, hidden_size * 2)
def forward(self, x: torch.Tensor, conditioning: torch.Tensor) -> torch.Tensor:
scale, shift = self.linear(conditioning).chunk(2, dim=-1)
x = self.norm(x)
# Broadcast conditioning to sequence dimension
x = x * (1 + scale.unsqueeze(0)) + shift.unsqueeze(0)
return x
class ErnieImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
_supports_gradient_checkpointing = True
_repeated_blocks = ["ErnieImageSharedAdaLNBlock"]
@register_to_config
def __init__(
self,
hidden_size: int = 3072,
num_attention_heads: int = 24,
num_layers: int = 24,
ffn_hidden_size: int = 8192,
in_channels: int = 128,
out_channels: int = 128,
patch_size: int = 1,
text_in_dim: int = 2560,
rope_theta: int = 256,
rope_axes_dim: Tuple[int, int, int] = (32, 48, 48),
eps: float = 1e-6,
qk_layernorm: bool = True,
):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_attention_heads
self.head_dim = hidden_size // num_attention_heads
self.num_layers = num_layers
self.patch_size = patch_size
self.in_channels = in_channels
self.out_channels = out_channels
self.text_in_dim = text_in_dim
self.x_embedder = ErnieImagePatchEmbedDynamic(in_channels, hidden_size, patch_size)
self.text_proj = nn.Linear(text_in_dim, hidden_size, bias=False) if text_in_dim != hidden_size else None
self.time_proj = Timesteps(hidden_size, flip_sin_to_cos=False, downscale_freq_shift=0)
self.time_embedding = TimestepEmbedding(hidden_size, hidden_size)
self.pos_embed = ErnieImageEmbedND3(dim=self.head_dim, theta=rope_theta, axes_dim=rope_axes_dim)
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size))
nn.init.zeros_(self.adaLN_modulation[-1].weight)
nn.init.zeros_(self.adaLN_modulation[-1].bias)
self.layers = nn.ModuleList(
[
ErnieImageSharedAdaLNBlock(
hidden_size, num_attention_heads, ffn_hidden_size, eps, qk_layernorm=qk_layernorm
)
for _ in range(num_layers)
]
)
self.final_norm = ErnieImageAdaLNContinuous(hidden_size, eps)
self.final_linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels)
nn.init.zeros_(self.final_linear.weight)
nn.init.zeros_(self.final_linear.bias)
self.gradient_checkpointing = False
# Multi-GPU inference support
self.sp_world_size = 1
self.sp_world_rank = 0
def enable_multi_gpus_inference(self):
"""Enable multi-GPU inference using sequence parallelism."""
from ..dist import (ErnieImageMultiGPUsAttnProcessor,
get_sequence_parallel_rank,
get_sequence_parallel_world_size, get_sp_group)
self.sp_world_size = get_sequence_parallel_world_size()
self.sp_world_rank = get_sequence_parallel_rank()
self.all_gather = get_sp_group().all_gather
self.set_attn_processor(ErnieImageMultiGPUsAttnProcessor())
def set_attn_processor(self, processor):
"""Set attention processor for all attention layers.
Args:
processor: The attention processor to use for all attention layers.
"""
for name, module in self.named_modules():
if hasattr(module, "set_processor"):
module.set_processor(processor)
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
# encoder_hidden_states: List[torch.Tensor],
text_bth: torch.Tensor,
text_lens: torch.Tensor,
return_dict: bool = True,
):
device, dtype = hidden_states.device, hidden_states.dtype
B, C, H, W = hidden_states.shape
p, Hp, Wp = self.patch_size, H // self.patch_size, W // self.patch_size
N_img = Hp * Wp
# Store original N_img for sequence parallel
N_img_full = N_img
img_sbh = self.x_embedder(hidden_states).transpose(0, 1).contiguous()
# text_bth, text_lens = self._pad_text(encoder_hidden_states, device, dtype)
if self.text_proj is not None and text_bth.numel() > 0:
text_bth = self.text_proj(text_bth)
Tmax = text_bth.shape[1]
text_sbh = text_bth.transpose(0, 1).contiguous()
# Sequence parallel: chunk image tokens across GPUs
if self.sp_world_size > 1:
N_img = N_img // self.sp_world_size
img_sbh = torch.chunk(img_sbh, self.sp_world_size, dim=0)[self.sp_world_rank]
x = torch.cat([img_sbh, text_sbh], dim=0)
S = x.shape[0]
# Position IDs
text_ids = (
torch.cat(
[
torch.arange(Tmax, device=device, dtype=torch.float32).view(1, Tmax, 1).expand(B, -1, -1),
torch.zeros((B, Tmax, 2), device=device),
],
dim=-1,
)
if Tmax > 0
else torch.zeros((B, 0, 3), device=device)
)
grid_yx = torch.stack(
torch.meshgrid(
torch.arange(Hp, device=device, dtype=torch.float32),
torch.arange(Wp, device=device, dtype=torch.float32),
indexing="ij",
),
dim=-1,
).reshape(-1, 2)
# Sequence parallel: use only the image_ids chunk for this GPU
if self.sp_world_size > 1:
chunk_start = self.sp_world_rank * N_img
chunk_end = chunk_start + N_img
image_ids = torch.cat(
[text_lens.float().view(B, 1, 1).expand(-1, N_img, -1),
grid_yx[chunk_start:chunk_end].view(1, N_img, 2).expand(B, -1, -1)],
dim=-1,
)
else:
image_ids = torch.cat(
[text_lens.float().view(B, 1, 1).expand(-1, N_img, -1), grid_yx.view(1, N_img, 2).expand(B, -1, -1)],
dim=-1,
)
rotary_pos_emb = self.pos_embed(torch.cat([image_ids, text_ids], dim=1))
# Attention mask: True = valid (attend), False = padding (mask out), matches sdpa bool convention
valid_text = (
torch.arange(Tmax, device=device).view(1, Tmax) < text_lens.view(B, 1)
if Tmax > 0
else torch.zeros((B, 0), device=device, dtype=torch.bool)
)
attention_mask = torch.cat([torch.ones((B, N_img), device=device, dtype=torch.bool), valid_text], dim=1)[
:, None, None, :
]
# AdaLN
sample = self.time_proj(timestep)
sample = sample.to(dtype=dtype)
c = self.time_embedding(sample)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = [
t.unsqueeze(0).expand(S, -1, -1).contiguous() for t in self.adaLN_modulation(c).chunk(6, dim=-1)
]
for layer in self.layers:
temb = [shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp]
if torch.is_grad_enabled() and self.gradient_checkpointing:
x = self._gradient_checkpointing_func(
layer,
x,
rotary_pos_emb,
temb,
attention_mask,
)
else:
x = layer(x, rotary_pos_emb, temb, attention_mask)
x = self.final_norm(x, c).type_as(x)
patches = self.final_linear(x)
# Sequence parallel: gather image patches from all GPUs
if self.sp_world_size > 1:
# Only gather the image part (first N_img tokens)
img_patches = patches[:N_img]
img_patches = self.all_gather(img_patches, dim=0)
# Reconstruct full patches: [full_img_tokens, text_tokens]
patches = torch.cat([img_patches, patches[N_img:]], dim=0)
# Use full N_img for output reshape
N_img = N_img_full
output = (
patches[:N_img].transpose(0, 1).contiguous()
.view(B, Hp, Wp, p, p, self.out_channels)
.permute(0, 5, 1, 3, 2, 4)
.contiguous()
.view(B, self.out_channels, H, W)
)
return ErnieImageTransformer2DModelOutput(sample=output) if return_dict else (output,)
File diff suppressed because it is too large Load Diff
+2
View File
@@ -1,6 +1,7 @@
from .pipeline_cogvideox_fun import CogVideoXFunPipeline
from .pipeline_cogvideox_fun_control import CogVideoXFunControlPipeline
from .pipeline_cogvideox_fun_inpaint import CogVideoXFunInpaintPipeline
from .pipeline_ernie_image import ErnieImagePipeline
from .pipeline_fantasytalking import FantasyTalkingPipeline
from .pipeline_flashhead import FlashHeadPipeline
from .pipeline_flux import FluxPipeline
@@ -31,6 +32,7 @@ from .pipeline_wan2_2_vace_fun import Wan2_2VaceFunPipeline
from .pipeline_wan_fun_control import WanFunControlPipeline
from .pipeline_wan_fun_inpaint import WanFunInpaintPipeline
from .pipeline_wan_phantom import WanFunPhantomPipeline
from .pipeline_wan_self_forcing import WanSelfForcingPipeline
from .pipeline_wan_vace import WanVacePipeline
from .pipeline_z_image import ZImagePipeline
from .pipeline_z_image_control import ZImageControlPipeline
+415
View File
@@ -0,0 +1,415 @@
# Modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/ernie_image/pipeline_ernie_image.py
# Copyright 2025 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
Ernie-Image Pipeline for HuggingFace Diffusers.
"""
import json
from dataclasses import dataclass
from typing import Callable, List, Optional, Union
import numpy as np
import PIL.Image
import torch
from diffusers.image_processor import VaeImageProcessor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import BaseOutput
from diffusers.utils.torch_utils import randn_tensor
from ..models import (AutoencoderKLFlux2, AutoTokenizer,
ErnieImageTransformer2DModel, Ministral3ForCausalLM,
Mistral3Model)
if not hasattr(PIL.Image, "Image"):
raise ImportError("`ErnieImagePipeline` requires `PIL.Image`. Please install it with: `pip install Pillow`.")
@dataclass
class ErnieImagePipelineOutput(BaseOutput):
"""
Output class for Ernie-Image pipelines.
Args:
images (`List[PIL.Image.Image]` or `np.ndarray`)
List of denoised PIL images of length `batch_size` or numpy array of shape `(batch_size, height, width,
num_channels)`. PIL images or numpy array present the denoised images of the diffusion pipeline.
revised_prompts (`List[str]`, *optional*):
List of revised prompts after PE enhancement.
"""
images: Union[List[PIL.Image.Image], np.ndarray]
revised_prompts: Optional[List[str]] = None
class ErnieImagePipeline(DiffusionPipeline):
"""
Pipeline for text-to-image generation using ErnieImageTransformer2DModel.
This pipeline uses:
- A custom DiT transformer model
- A Flux2-style VAE for encoding/decoding latents
- A text encoder (e.g., Qwen) for text conditioning
- Flow Matching Euler Discrete Scheduler
"""
model_cpu_offload_seq = "pe->text_encoder->transformer->vae"
# For SGLang fallback ...
_optional_components = ["pe", "pe_tokenizer"]
_callback_tensor_inputs = ["latents"]
def __init__(
self,
transformer: ErnieImageTransformer2DModel,
vae: AutoencoderKLFlux2,
text_encoder: Mistral3Model,
tokenizer: AutoTokenizer,
scheduler: FlowMatchEulerDiscreteScheduler,
pe: Optional[Ministral3ForCausalLM] = None,
pe_tokenizer: Optional[AutoTokenizer] = None,
):
super().__init__()
self.register_modules(
transformer=transformer,
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
scheduler=scheduler,
pe=pe,
pe_tokenizer=pe_tokenizer,
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels)) if getattr(self, "vae", None) else 16
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
@property
def guidance_scale(self):
return self._guidance_scale
@property
def do_classifier_free_guidance(self):
return self._guidance_scale > 1.0
@torch.no_grad()
def _enhance_prompt_with_pe(
self,
prompt: str,
device: torch.device,
width: int = 1024,
height: int = 1024,
system_prompt: Optional[str] = None,
temperature: float = 0.6,
top_p: float = 0.95,
) -> str:
"""Use PE model to rewrite/enhance a short prompt via chat_template."""
# Build user message as JSON carrying prompt text and target resolution
user_content = json.dumps(
{"prompt": prompt, "width": width, "height": height},
ensure_ascii=False,
)
messages = []
if system_prompt is not None:
messages.append({"role": "system", "content": system_prompt})
messages.append({"role": "user", "content": user_content})
# apply_chat_template picks up the chat_template.jinja loaded with pe_tokenizer
input_text = self.pe_tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=False, # "Output:" is already in the user block
)
inputs = self.pe_tokenizer(input_text, return_tensors="pt").to(device)
output_ids = self.pe.generate(
**inputs,
max_new_tokens=self.pe_tokenizer.model_max_length,
do_sample=temperature != 1.0 or top_p != 1.0,
temperature=temperature,
top_p=top_p,
pad_token_id=self.pe_tokenizer.pad_token_id,
eos_token_id=self.pe_tokenizer.eos_token_id,
)
# Decode only newly generated tokens
generated_ids = output_ids[0][inputs["input_ids"].shape[1] :]
return self.pe_tokenizer.decode(generated_ids, skip_special_tokens=True).strip()
@torch.no_grad()
def encode_prompt(
self,
prompt: Union[str, List[str]],
device: torch.device,
num_images_per_prompt: int = 1,
) -> List[torch.Tensor]:
"""Encode text prompts to embeddings."""
if isinstance(prompt, str):
prompt = [prompt]
text_hiddens = []
for p in prompt:
ids = self.tokenizer(
p,
add_special_tokens=True,
truncation=True,
padding=False,
)["input_ids"]
if len(ids) == 0:
if self.tokenizer.bos_token_id is not None:
ids = [self.tokenizer.bos_token_id]
else:
ids = [0]
input_ids = torch.tensor([ids], device=device)
with torch.no_grad():
outputs = self.text_encoder(
input_ids=input_ids,
output_hidden_states=True,
)
# Use second to last hidden state (matches training)
hidden = outputs.hidden_states[-2][0] # [T, H]
# Repeat for num_images_per_prompt
for _ in range(num_images_per_prompt):
text_hiddens.append(hidden)
return text_hiddens
@staticmethod
def _patchify_latents(latents: torch.Tensor) -> torch.Tensor:
"""2x2 patchify: [B, 32, H, W] -> [B, 128, H/2, W/2]"""
b, c, h, w = latents.shape
latents = latents.view(b, c, h // 2, 2, w // 2, 2)
latents = latents.permute(0, 1, 3, 5, 2, 4)
return latents.reshape(b, c * 4, h // 2, w // 2)
@staticmethod
def _unpatchify_latents(latents: torch.Tensor) -> torch.Tensor:
"""Reverse patchify: [B, 128, H/2, W/2] -> [B, 32, H, W]"""
b, c, h, w = latents.shape
latents = latents.reshape(b, c // 4, 2, 2, h, w)
latents = latents.permute(0, 1, 4, 2, 5, 3)
return latents.reshape(b, c // 4, h * 2, w * 2)
@staticmethod
def _pad_text(text_hiddens: List[torch.Tensor], device: torch.device, dtype: torch.dtype, text_in_dim: int):
B = len(text_hiddens)
if B == 0:
return torch.zeros((0, 0, text_in_dim), device=device, dtype=dtype), torch.zeros(
(0,), device=device, dtype=torch.long
)
normalized = [
th.squeeze(1).to(device).to(dtype) if th.dim() == 3 else th.to(device).to(dtype) for th in text_hiddens
]
lens = torch.tensor([t.shape[0] for t in normalized], device=device, dtype=torch.long)
Tmax = int(lens.max().item())
text_bth = torch.zeros((B, Tmax, text_in_dim), device=device, dtype=dtype)
for i, t in enumerate(normalized):
text_bth[i, : t.shape[0], :] = t
return text_bth, lens
@torch.no_grad()
def __call__(
self,
prompt: Optional[Union[str, List[str]]] = None,
negative_prompt: Optional[Union[str, List[str]]] = "",
height: int = 1024,
width: int = 1024,
num_inference_steps: int = 50,
guidance_scale: float = 4.0,
num_images_per_prompt: int = 1,
generator: Optional[torch.Generator] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: list[torch.FloatTensor] | None = None,
negative_prompt_embeds: list[torch.FloatTensor] | None = None,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[Callable[[int, int, dict], None]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
use_pe: bool = True, # 默认使用PE进行改写
):
"""
Generate images from text prompts.
Args:
prompt: Text prompt(s)
negative_prompt: Negative prompt(s) for CFG. Default is "".
height: Image height in pixels (must be divisible by 16). Default: 1024.
width: Image width in pixels (must be divisible by 16). Default: 1024.
num_inference_steps: Number of denoising steps
guidance_scale: CFG scale (1.0 = no guidance). Default: 4.0.
num_images_per_prompt: Number of images per prompt
generator: Random generator for reproducibility
latents: Pre-generated latents (optional)
prompt_embeds: Pre-computed text embeddings for positive prompts (optional).
If provided, `encode_prompt` is skipped for positive prompts.
negative_prompt_embeds: Pre-computed text embeddings for negative prompts (optional).
If provided, `encode_prompt` is skipped for negative prompts.
output_type: "pil" or "latent"
return_dict: Whether to return a dataclass
callback_on_step_end: Optional callback invoked at the end of each denoising step.
Called as `callback_on_step_end(pipeline, step, timestep, callback_kwargs)` where `callback_kwargs`
contains the tensors listed in `callback_on_step_end_tensor_inputs`. The callback may return a dict to
override those tensors for subsequent steps.
callback_on_step_end_tensor_inputs: List of tensor names passed into the callback kwargs.
Must be a subset of `_callback_tensor_inputs` (default: `["latents"]`).
use_pe: Whether to use the PE model to enhance prompts before generation.
Returns:
:class:`ErnieImagePipelineOutput` with `images` and `revised_prompts`.
"""
device = self._execution_device
dtype = self.transformer.dtype
self._guidance_scale = guidance_scale
# Validate prompt / prompt_embeds
if prompt is None and prompt_embeds is None:
raise ValueError("Must provide either `prompt` or `prompt_embeds`.")
if prompt is not None and prompt_embeds is not None:
raise ValueError("Cannot provide both `prompt` and `prompt_embeds` at the same time.")
# Validate dimensions
if height % self.vae_scale_factor != 0 or width % self.vae_scale_factor != 0:
raise ValueError(f"Height and width must be divisible by {self.vae_scale_factor}")
# Handle prompts
if prompt is not None:
if isinstance(prompt, str):
prompt = [prompt]
# [Phase 1] PE: enhance prompts
revised_prompts: Optional[List[str]] = None
if prompt is not None and use_pe and self.pe is not None and self.pe_tokenizer is not None:
prompt = [self._enhance_prompt_with_pe(p, device, width=width, height=height) for p in prompt]
revised_prompts = list(prompt)
if prompt is not None:
batch_size = len(prompt)
else:
batch_size = len(prompt_embeds)
total_batch_size = batch_size * num_images_per_prompt
# Handle negative prompt
if negative_prompt is None:
negative_prompt = ""
if isinstance(negative_prompt, str):
negative_prompt = [negative_prompt] * batch_size
if len(negative_prompt) != batch_size:
raise ValueError(f"negative_prompt must have same length as prompt ({batch_size})")
# [Phase 2] Text encoding
if prompt_embeds is not None:
text_hiddens = [h for h in prompt_embeds for _ in range(num_images_per_prompt)]
else:
text_hiddens = self.encode_prompt(prompt, device, num_images_per_prompt)
# CFG with negative prompt
if self.do_classifier_free_guidance:
if negative_prompt_embeds is not None:
uncond_text_hiddens = [h for h in negative_prompt_embeds for _ in range(num_images_per_prompt)]
else:
uncond_text_hiddens = self.encode_prompt(negative_prompt, device, num_images_per_prompt)
# Latent dimensions
latent_h = height // self.vae_scale_factor
latent_w = width // self.vae_scale_factor
latent_channels = self.transformer.config.in_channels # After patchify
# Initialize latents
if latents is None:
latents = randn_tensor(
(total_batch_size, latent_channels, latent_h, latent_w),
generator=generator,
device=device,
dtype=dtype,
)
# Setup scheduler
sigmas = torch.linspace(1.0, 0.0, num_inference_steps + 1)
self.scheduler.set_timesteps(sigmas=sigmas[:-1], device=device)
# Denoising loop
if self.do_classifier_free_guidance:
cfg_text_hiddens = list(uncond_text_hiddens) + list(text_hiddens)
else:
cfg_text_hiddens = text_hiddens
text_bth, text_lens = self._pad_text(
text_hiddens=cfg_text_hiddens, device=device, dtype=dtype, text_in_dim=self.transformer.config.text_in_dim
)
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(self.scheduler.timesteps):
if self.do_classifier_free_guidance:
latent_model_input = torch.cat([latents, latents], dim=0)
t_batch = torch.full((total_batch_size * 2,), t.item(), device=device, dtype=dtype)
else:
latent_model_input = latents
t_batch = torch.full((total_batch_size,), t.item(), device=device, dtype=dtype)
# Model prediction
pred = self.transformer(
hidden_states=latent_model_input,
timestep=t_batch,
text_bth=text_bth,
text_lens=text_lens,
return_dict=False,
)[0]
# Apply CFG
if self.do_classifier_free_guidance:
pred_uncond, pred_cond = pred.chunk(2, dim=0)
pred = pred_uncond + guidance_scale * (pred_cond - pred_uncond)
# Scheduler step
latents = self.scheduler.step(pred, t, latents).prev_sample
# Callback
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)
progress_bar.update()
if output_type == "latent":
images = latents
else:
# Decode latents to images
# Unnormalize latents using VAE's BN stats
# TODO: switch to `self.vae.config.batch_norm_eps` once the hub config is updated to match the trained value (1e-5).
bn_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(device=device, dtype=latents.dtype)
bn_std = torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) + 1e-5).to(
device=device, dtype=latents.dtype
)
latents = latents * bn_std + bn_mean
# Unpatchify
latents = self._unpatchify_latents(latents)
# Decode
images = self.vae.decode(latents, return_dict=False)[0]
# Post-process
images = self.image_processor.postprocess(images, output_type=output_type)
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return (images,)
return ErnieImagePipelineOutput(images=images, revised_prompts=revised_prompts)
@@ -0,0 +1,865 @@
# Modified from https://github.com/guandeh17/Self-Forcing/blob/main/pipeline/causal_diffusion_inference.py
import inspect
import math
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from ..models import (AutoencoderKLWan, AutoTokenizer, WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from ..utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
get_sampling_sigmas)
from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def stochastic_sampling_timesteps(num_inference_steps, shift, device, num_timesteps=1000):
"""Official FlashHead timestep schedule with shift transform."""
if num_inference_steps == 4:
timesteps = [1000, 750, 500, 250]
else:
timesteps = np.linspace(num_timesteps, 1, num_inference_steps, dtype=np.float32).tolist()
timesteps = torch.tensor(timesteps + [0.0], dtype=torch.float32, device=device)
t = timesteps / num_timesteps
return shift * t / (1 + (shift - 1) * t) * num_timesteps
EXAMPLE_DOC_STRING = """
Examples:
```python
pass
```
"""
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
def retrieve_timesteps(
scheduler,
num_inference_steps: Optional[int] = None,
device: Optional[Union[str, torch.device]] = None,
timesteps: Optional[List[int]] = None,
sigmas: Optional[List[float]] = None,
**kwargs,
):
"""
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
Args:
scheduler (`SchedulerMixin`):
The scheduler to get timesteps from.
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
must be `None`.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
timesteps (`List[int]`, *optional*):
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
`num_inference_steps` and `sigmas` must be `None`.
sigmas (`List[float]`, *optional*):
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
`num_inference_steps` and `timesteps` must be `None`.
Returns:
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
if timesteps is not None:
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
timesteps = scheduler.timesteps
return timesteps, num_inference_steps
@dataclass
class WanSelfForcingPipelineOutput(BaseOutput):
r"""
Output class for CogVideo pipelines.
Args:
video (`torch.Tensor`, `np.ndarray`, or List[List[PIL.Image.Image]]):
List of video outputs - It can be a nested list of length `batch_size,` with each sub-list containing
denoised PIL image sequences of length `num_frames.` It can also be a NumPy array or Torch tensor of shape
`(batch_size, num_frames, channels, height, width)`.
"""
videos: torch.Tensor
class WanSelfForcingPipeline(DiffusionPipeline):
r"""
Pipeline for text-to-video generation using Wan.
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the
library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)
"""
_optional_components = []
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = [
"latents",
"prompt_embeds",
"negative_prompt_embeds",
]
def __init__(
self,
tokenizer: AutoTokenizer,
text_encoder: WanT5EncoderModel,
vae: AutoencoderKLWan,
transformer: WanTransformer3DModel_SelfForcing,
scheduler: FlowMatchEulerDiscreteScheduler,
):
super().__init__()
self.register_modules(
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer, scheduler=scheduler
)
self.video_processor = VideoProcessor(vae_scale_factor=self.vae.spatial_compression_ratio)
self.kv_cache_pos = None
self.kv_cache_neg = None
self.crossattn_cache_pos = None
self.crossattn_cache_neg = None
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_videos_per_prompt: int = 1,
max_sequence_length: int = 512,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1])
logger.warning(
"The following part of your input was truncated because `max_sequence_length` is set to "
f" {max_sequence_length} tokens: {removed_text}"
)
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask.to(device))[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
return [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
def encode_prompt(
self,
prompt: Union[str, List[str]],
negative_prompt: Optional[Union[str, List[str]]] = None,
do_classifier_free_guidance: bool = True,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
max_sequence_length: int = 512,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
r"""
Encodes the prompt into text encoder hidden states.
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
less than `1`).
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
Whether to use classifier free guidance or not.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
device: (`torch.device`, *optional*):
torch device
dtype: (`torch.dtype`, *optional*):
torch dtype
"""
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
if prompt is not None:
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds = self._get_t5_prompt_embeds(
prompt=prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
if do_classifier_free_guidance and negative_prompt_embeds is None:
negative_prompt = negative_prompt or ""
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
if prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
" the batch size of `prompt`."
)
negative_prompt_embeds = self._get_t5_prompt_embeds(
prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
return prompt_embeds, negative_prompt_embeds
def prepare_latents(
self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None
):
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."
)
# Shape: [B, C, F, H, W] (standard PyTorch format)
shape = (
batch_size,
num_channels_latents,
(num_frames - 1) // self.vae.temporal_compression_ratio + 1,
height // self.vae.spatial_compression_ratio,
width // self.vae.spatial_compression_ratio,
)
if latents is None:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
# scale the initial noise by the standard deviation required by the scheduler
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
return latents
def decode_latents(self, latents: torch.Tensor) -> torch.Tensor:
frames = self.vae.decode(latents.to(self.vae.dtype)).sample
frames = (frames / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
frames = frames.cpu().float().numpy()
return frames
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs
def prepare_extra_step_kwargs(self, generator, eta):
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
# and should be between [0, 1]
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
extra_step_kwargs = {}
if accepts_eta:
extra_step_kwargs["eta"] = eta
# check if the scheduler accepts generator
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
# Copied from diffusers.pipelines.latte.pipeline_latte.LattePipeline.check_inputs
def check_inputs(
self,
prompt,
height,
width,
negative_prompt,
callback_on_step_end_tensor_inputs,
prompt_embeds=None,
negative_prompt_embeds=None,
):
if height % 8 != 0 or width % 8 != 0:
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
f" {negative_prompt_embeds.shape}."
)
@property
def guidance_scale(self):
return self._guidance_scale
@property
def num_timesteps(self):
return self._num_timesteps
@property
def attention_kwargs(self):
return self._attention_kwargs
@property
def interrupt(self):
return self._interrupt
def _initialize_kv_cache(self, batch_size, dtype, device, frame_seq_length, num_latent_frames):
"""
Initialize KV cache for causal self-attention.
"""
kv_cache_pos = []
kv_cache_neg = []
# Compute KV cache size based on actual resolution and frame count
local_attn_size = getattr(self.transformer.config, 'local_attn_size', -1)
if local_attn_size != -1:
kv_cache_size = local_attn_size * frame_seq_length
else:
kv_cache_size = num_latent_frames * frame_seq_length
num_heads = self.transformer.config.num_heads
head_dim = self.transformer.config.dim // num_heads
for _ in range(self.transformer.config.num_layers):
kv_cache_pos.append({
"k": torch.zeros([batch_size, kv_cache_size, num_heads, head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, kv_cache_size, num_heads, head_dim], dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device)
})
kv_cache_neg.append({
"k": torch.zeros([batch_size, kv_cache_size, num_heads, head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, kv_cache_size, num_heads, head_dim], dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device)
})
self.kv_cache_pos = kv_cache_pos
self.kv_cache_neg = kv_cache_neg
def _initialize_crossattn_cache(self, batch_size, dtype, device):
"""
Initialize cross-attention cache.
"""
crossattn_cache_pos = []
crossattn_cache_neg = []
text_len = self.transformer.config.text_len
num_heads = self.transformer.config.num_heads
head_dim = self.transformer.config.dim // num_heads
for _ in range(self.transformer.config.num_layers):
crossattn_cache_pos.append({
"k": torch.zeros([batch_size, text_len, num_heads, head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, text_len, num_heads, head_dim], dtype=dtype, device=device),
"is_init": False
})
crossattn_cache_neg.append({
"k": torch.zeros([batch_size, text_len, num_heads, head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, text_len, num_heads, head_dim], dtype=dtype, device=device),
"is_init": False
})
self.crossattn_cache_pos = crossattn_cache_pos
self.crossattn_cache_neg = crossattn_cache_neg
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Optional[Union[str, List[str]]] = None,
negative_prompt: Optional[Union[str, List[str]]] = None,
height: int = 480,
width: int = 720,
num_frames: int = 49,
num_inference_steps: int = 50,
timesteps: Optional[List[int]] = None,
guidance_scale: float = 6,
num_videos_per_prompt: int = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 512,
comfyui_progressbar: bool = False,
shift: float = 5.0,
initial_latent: Optional[torch.FloatTensor] = None,
start_frame_index: int = 0,
num_frame_per_block: int = 1,
independent_first_frame: bool = True,
context_noise: int = 0,
stochastic_sampling: bool = True,
) -> Union[WanSelfForcingPipelineOutput, Tuple]:
r"""
Function invoked when calling the pipeline for Self-Forcing causal generation.
Args:
initial_latent: Optional initial latent frames for I2V/video extension.
Shape: (batch_size, num_input_frames, channels, height, width)
start_frame_index: Starting frame index for long video generation.
Used when continuing generation from a previous segment.
num_frame_per_block: Number of frames to generate per block.
independent_first_frame: Whether to generate the first frame independently (T2V mode).
context_noise: Context noise level for KV cache update (matches training config).
Examples:
```python
pass
```
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
num_videos_per_prompt = 1
# 1. Check inputs
self.check_inputs(
prompt,
height,
width,
negative_prompt,
callback_on_step_end_tensor_inputs,
prompt_embeds,
negative_prompt_embeds,
)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._interrupt = False
# 2. Default call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
weight_dtype = self.text_encoder.dtype
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
# corresponds to doing no classifier free guidance.
do_classifier_free_guidance = guidance_scale > 1.0
# 3. Encode input prompt
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
prompt,
negative_prompt,
do_classifier_free_guidance,
num_videos_per_prompt=num_videos_per_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
max_sequence_length=max_sequence_length,
device=device,
)
if do_classifier_free_guidance:
in_prompt_embeds = negative_prompt_embeds + prompt_embeds
else:
in_prompt_embeds = prompt_embeds
# 4. Prepare timesteps
if stochastic_sampling:
timesteps = stochastic_sampling_timesteps(num_inference_steps, shift, device)
elif isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps)
elif isinstance(self.scheduler, FlowUniPCMultistepScheduler):
self.scheduler.set_timesteps(num_inference_steps, device=device, shift=shift)
timesteps = self.scheduler.timesteps
elif isinstance(self.scheduler, FlowDPMSolverMultistepScheduler):
sampling_sigmas = get_sampling_sigmas(num_inference_steps, shift)
timesteps, _ = retrieve_timesteps(
self.scheduler,
device=device,
sigmas=sampling_sigmas)
else:
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps)
self._num_timesteps = len(timesteps)
# 5. Prepare latents (noise) and output buffer separately
latent_channels = self.transformer.config.in_channels
# For I2V: num_frames_to_generate is frames to generate (not including input frames)
num_frames_to_generate = num_frames
if initial_latent is not None:
# In I2V mode, num_frames is total frames, but noise should only be the new frames
# VAE compression: num_latent_frames = (num_frames - 1) // temporal_compression + 1
num_input_frames_temp = initial_latent.shape[2] # [B, C, F, H, W]
total_latent_frames = (num_frames - 1) // self.vae.temporal_compression_ratio + 1
input_latent_frames = num_input_frames_temp
num_frames_to_generate = total_latent_frames - input_latent_frames
# Prepare noise (only for frames to generate)
noise = self.prepare_latents(
batch_size,
latent_channels,
num_frames_to_generate,
height,
width,
weight_dtype,
device,
generator,
latents,
)
# Calculate total output frames (input + generated)
num_input_frames = initial_latent.shape[2] if initial_latent is not None else 0 # [B, C, F, H, W]
num_output_frames = num_frames_to_generate + num_input_frames
# Allocate output buffer: [B, C, F_total, H, W]
output = torch.zeros_like(
noise,
device=device,
dtype=weight_dtype
)
# 6. Calculate sequence length and frame_seq_length
target_shape = (
self.vae.latent_channels,
(num_frames - 1) // self.vae.temporal_compression_ratio + 1,
width // self.vae.spatial_compression_ratio,
height // self.vae.spatial_compression_ratio,
)
seq_len = math.ceil(
(target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2])
* target_shape[1]
)
# Calculate frame_seq_length: tokens per frame
frame_seq_length = (target_shape[2] * target_shape[3]) // (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2])
# 7. Causal generation loop - block by block
# num_latent_frames is the number of frames after VAE compression
num_latent_frames = target_shape[1]
# Determine num_blocks based on mode (T2V vs I2V)
# Reference: causal_inference.py line 70-78
if not independent_first_frame or (independent_first_frame and initial_latent is not None):
# I2V mode: even with independent_first_frame, if initial_latent is provided, frames should be divisible
assert num_latent_frames % num_frame_per_block == 0, \
f"num_latent_frames ({num_latent_frames}) must be divisible by num_frame_per_block ({num_frame_per_block})"
num_blocks = num_latent_frames // num_frame_per_block
else:
# T2V mode: no initial_latent, use [1, 4, 4, ...] pattern
assert (num_latent_frames - 1) % num_frame_per_block == 0, \
f"num_latent_frames-1 ({num_latent_frames - 1}) must be divisible by num_frame_per_block ({num_frame_per_block})"
num_blocks = (num_latent_frames - 1) // num_frame_per_block
# Initialize ComfyUI progress bar after calculating num_blocks
if comfyui_progressbar:
from comfy.utils import ProgressBar
# Total steps = num_blocks * num_inference_steps + 1 (for latent preparation)
pbar = ProgressBar(num_blocks * num_inference_steps + 1)
pbar.update(1)
# Self-Forcing causal state (reset per call)
current_start_frame = start_frame_index
cache_start_frame = 0
# 8. Initialize KV cache and cross-attention cache
# Reset caches if they exist (for multiple inference calls)
required_kv_size = num_latent_frames * frame_seq_length
if self.kv_cache_pos is not None and self.kv_cache_pos[0]["k"].shape[1] >= required_kv_size:
for block_index in range(len(self.kv_cache_pos)):
self.kv_cache_pos[block_index]["global_end_index"] = torch.tensor(
[0], dtype=torch.long, device=device)
self.kv_cache_pos[block_index]["local_end_index"] = torch.tensor(
[0], dtype=torch.long, device=device)
self.kv_cache_neg[block_index]["global_end_index"] = torch.tensor(
[0], dtype=torch.long, device=device)
self.kv_cache_neg[block_index]["local_end_index"] = torch.tensor(
[0], dtype=torch.long, device=device)
for block_index in range(len(self.crossattn_cache_pos)):
self.crossattn_cache_pos[block_index]["is_init"] = False
self.crossattn_cache_neg[block_index]["is_init"] = False
else:
self._initialize_kv_cache(batch_size=batch_size, dtype=weight_dtype, device=device, frame_seq_length=frame_seq_length, num_latent_frames=num_latent_frames)
self._initialize_crossattn_cache(batch_size=batch_size, dtype=weight_dtype, device=device)
# Build all_num_frames list
# Self-Forcing: T2V with independent_first_frame uses [1, 4, 4, 4, ...] pattern
# I2V mode uses [4, 4, 4, ...] pattern (first frame is provided)
all_num_frames = [num_frame_per_block] * num_blocks
if independent_first_frame and initial_latent is None:
# First frame is generated independently (standard Self-Forcing T2V pattern)
all_num_frames = [1] + all_num_frames
for block_idx, current_num_frames in enumerate(all_num_frames):
# Extract noise for current block and convert to list format
# noise only contains frames to generate, indexed from 0
# current_start_frame tracks global position (including input frames for I2V)
# Need to offset by num_input_frames to get index in noise
start_idx = current_start_frame - num_input_frames
end_idx = start_idx + current_num_frames
noisy_input = noise[:, :, start_idx:end_idx]
# Denoising loop for current block
# Reset scheduler state for each block (required for causal generation)
# For Euler scheduler, resetting _step_index is sufficient.
# For multi-step schedulers (UniPC, DPM++), also clear accumulated model outputs.
self.scheduler._step_index = None
if hasattr(self.scheduler, 'model_outputs'):
self.scheduler.model_outputs = []
denoise_timesteps = timesteps[:-1] if stochastic_sampling else timesteps
with self.progress_bar(total=len(denoise_timesteps)) as progress_bar:
for step_idx, t in enumerate(denoise_timesteps):
# Per-frame timesteps for causal generation
timestep = torch.ones([batch_size, current_num_frames], device=device, dtype=weight_dtype) * t
if comfyui_progressbar:
pbar.update(1)
if do_classifier_free_guidance:
# Conditional path
with torch.cuda.amp.autocast(dtype=weight_dtype):
flow_pred_cond = self.transformer(
x=noisy_input,
context=prompt_embeds,
t=timestep,
seq_len=seq_len,
kv_cache=self.kv_cache_pos,
crossattn_cache=self.crossattn_cache_pos,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
)
# Unconditional path
with torch.cuda.amp.autocast(dtype=weight_dtype):
flow_pred_uncond = self.transformer(
x=noisy_input,
context=negative_prompt_embeds,
t=timestep,
seq_len=seq_len,
kv_cache=self.kv_cache_neg,
crossattn_cache=self.crossattn_cache_neg,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
)
# CFG guidance
# Transformer output shape check
if flow_pred_cond.dim() == 5:
# Already [B, C, F, H, W]
flow_pred = flow_pred_uncond + guidance_scale * (flow_pred_cond - flow_pred_uncond)
elif flow_pred_cond.dim() == 4:
# [F, C, H, W], need to add batch dim
flow_pred_cond = flow_pred_cond.unsqueeze(0).permute(0, 2, 1, 3, 4)
flow_pred_uncond = flow_pred_uncond.unsqueeze(0).permute(0, 2, 1, 3, 4)
flow_pred = flow_pred_uncond + guidance_scale * (flow_pred_cond - flow_pred_uncond)
else:
raise ValueError(f"Unexpected flow_pred_cond dim: {flow_pred_cond.dim()}, shape: {flow_pred_cond.shape}")
else:
# Forward pass with KV cache
with torch.cuda.amp.autocast(dtype=weight_dtype):
flow_pred = self.transformer(
x=noisy_input,
context=in_prompt_embeds,
t=timestep,
seq_len=seq_len,
kv_cache=self.kv_cache_pos,
crossattn_cache=self.crossattn_cache_pos,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
)
# Transformer output shape check
if flow_pred.dim() == 4:
# [F, C, H, W], need to add batch dim and permute
flow_pred = flow_pred.unsqueeze(0).permute(0, 2, 1, 3, 4)
# If already 5D [B, C, F, H, W], no need to permute
# compute the previous noisy sample x_t -> x_t-1
if stochastic_sampling:
t_i = (timesteps[step_idx] / 1000).to(weight_dtype)
t_i_1 = (timesteps[step_idx + 1] / 1000).to(weight_dtype)
denoised_pred = noisy_input - flow_pred * t_i
noisy_input = (1 - t_i_1) * denoised_pred + t_i_1 * torch.randn(
denoised_pred.shape, dtype=denoised_pred.dtype, device=device, generator=generator
)
else:
# Get current sigma for x0 conversion
sigma_t = self.scheduler.sigmas[step_idx]
# Convert to x0: x0 = x_t - sigma_t * flow_pred
denoised_pred = noisy_input - sigma_t * flow_pred
if step_idx < len(denoise_timesteps) - 1:
# Not the last step: add noise for next timestep
next_sigma = self.scheduler.sigmas[step_idx + 1]
local_noise = torch.randn(denoised_pred.shape, device=denoised_pred.device, dtype=denoised_pred.dtype, generator=generator)
noisy_input = (1 - next_sigma) * denoised_pred + next_sigma * local_noise
else:
noisy_input = denoised_pred
progress_bar.update()
# Update output with denoised block
output[:, :, cache_start_frame:cache_start_frame + current_num_frames] = denoised_pred
# Update KV cache with clean context (timestep=context_noise) for next block
# Reference: causal_inference.py line 227 - uses context_noise for KV cache update
if block_idx < len(all_num_frames) - 1:
context_timestep = torch.ones([batch_size, current_num_frames], device=device, dtype=torch.long) * context_noise
if do_classifier_free_guidance:
# Update both positive and negative caches
with torch.cuda.amp.autocast(dtype=weight_dtype):
self.transformer(
x=denoised_pred,
context=prompt_embeds,
t=context_timestep,
seq_len=seq_len,
kv_cache=self.kv_cache_pos,
crossattn_cache=self.crossattn_cache_pos,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
)
self.transformer(
x=denoised_pred,
context=negative_prompt_embeds,
t=context_timestep,
seq_len=seq_len,
kv_cache=self.kv_cache_neg,
crossattn_cache=self.crossattn_cache_neg,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
)
else:
with torch.cuda.amp.autocast(dtype=weight_dtype):
self.transformer(
x=denoised_pred,
context=in_prompt_embeds,
t=context_timestep,
seq_len=seq_len,
kv_cache=self.kv_cache_pos,
crossattn_cache=self.crossattn_cache_pos,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
)
current_start_frame += current_num_frames
cache_start_frame += current_num_frames
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, block_idx, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
# 9. Decode output
if output_type == "pil":
video = self.decode_latents(output)
video = torch.from_numpy(video)
else:
video = output
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return (video,)
return WanSelfForcingPipelineOutput(videos=video)