Add Causal-Focing for Wan2.1 (#500)

This commit is contained in:
hkz
2026-07-22 17:37:05 +08:00
committed by GitHub
parent 403f1f7b78
commit 248ab0ac0e
24 changed files with 9606 additions and 102 deletions
@@ -0,0 +1,328 @@
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 = "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 = 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.0
# `stochastic_sampling`: False: ar
# True : ccd and dmd
stochastic_sampling = True
# Causal-Forcing checkpoint to overlay on top of the Wan2.1 base model.
transformer_path = "output_dir_wan2.1_causal_forcing_dmd/checkpoint-2000/diffusion_pytorch_model.safetensors"
use_ema = False
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Causal-Forcing causal inference config
# `num_frame_per_block`: 3 = chunk-wise:
# 1 = frame-wise:
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
# Some 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压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# -------- Stage selector (uncomment ONE block) --------
# All few-step distilled stages (2, 3) bake CFG into the student weights, so
# inference MUST use guidance_scale=1.0 — CF's CausalInferencePipeline does
# zero CFG (grep "unconditional/cfg/guidance" in pipeline/causal_inference.py
# returns 0). Using gs>1 stacks CFG on top of a CFG-baked model and produces
# over-saturated, AR-unstable outputs.
#
# Stage 1 — AR diffusion (`ar_diffusion.pt`): 50-step UniPC + CFG.
# guidance_scale = 3.0
# num_inference_steps = 50
#
# Stage 2 — CCD (`causal_cd.pt`): 4-step consistency-distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
#
# Stage 3 — DMD (`causal_forcing.pt`): 4-step distribution-matching distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
guidance_scale = 1.0
num_inference_steps = 4
seed = 43
lora_weight = 0.55
save_path = "samples/wan-videos-causal-forcing"
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:
def _resolve_transformer_path(raw_path: str, prefer_ema: bool) -> str:
"""Resolve a checkpoint path to the actual weights file to load.
File paths are returned unchanged so external `.pt` ckpts (CF official) keep
working. For trainer-output dirs, prefer EMA when asked and available,
otherwise fall back to live `transformer/` weights.
"""
if os.path.isfile(raw_path):
return raw_path
if os.path.isdir(raw_path):
candidates = []
if prefer_ema:
candidates.append(os.path.join(raw_path, "ema_transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "diffusion_pytorch_model.safetensors"))
for c in candidates:
if os.path.isfile(c):
return c
raise FileNotFoundError(
f"transformer_path={raw_path!r} is neither a file nor a checkpoint dir "
f"with a known safetensors layout (transformer/ or ema_transformer/)."
)
_raw_transformer_path = transformer_path
transformer_path = _resolve_transformer_path(transformer_path, prefer_ema=use_ema)
if transformer_path != _raw_transformer_path:
print(f"use_ema={use_ema}: resolved {_raw_transformer_path} -> {transformer_path}")
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
# Causal-Forcing's FSDP-saved ckpts (causal_cd.pt / causal_forcing.pt) keep the
# `model._fsdp_wrapped_module.` prefix; strip it before the generic `model.` strip
# so both kinds of ckpt land at bare parameter names.
if any("._fsdp_wrapped_module." in k for k in state_dict.keys()):
state_dict = {k.replace("model._fsdp_wrapped_module.", "model.", 1) if k.startswith("model._fsdp_wrapped_module.") else k: v for k, v in state_dict.items()}
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()
@@ -0,0 +1,373 @@
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, StreamVideoSaver,
SegmentVideoSaver)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_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.0
# `stochastic_sampling`: False: ar
# True : ccd and dmd
stochastic_sampling = True
# Causal-Forcing checkpoint to overlay on top of the Wan2.1 base model.
transformer_path = "output_dir_wan2.1_causal_forcing_dmd/checkpoint-2000/diffusion_pytorch_model.safetensors"
use_ema = False
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Causal-Forcing causal inference config
# `num_frame_per_block`: 3 = chunk-wise:
# 1 = frame-wise:
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
# Streaming decode: decode and write each causal block to disk as soon as it is
# generated ("generate a chunk, save a chunk"), instead of holding all latents
# and decoding once at the end. The VAE causal cache is threaded across blocks,
# so the saved video is seam-free and identical to a full decode.
streaming = True
# Streaming save mode (only used when streaming=True):
# "stream" : append every block into one continuous mp4 (finalized on close).
# "segments" : save each decoded block as its own standalone mp4, flushed
# immediately after decoding ("decode a chunk -> save it now").
# Robust to interruption; segments can be concatenated later.
save_mode = "segments"
# Use torch.float16 if GPU does not support torch.bfloat16
# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
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压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# -------- Stage selector (uncomment ONE block) --------
# All few-step distilled stages (2, 3) bake CFG into the student weights, so
# inference MUST use guidance_scale=1.0 — CF's CausalInferencePipeline does
# zero CFG (grep "unconditional/cfg/guidance" in pipeline/causal_inference.py
# returns 0). Using gs>1 stacks CFG on top of a CFG-baked model and produces
# over-saturated, AR-unstable outputs.
#
# Stage 1 — AR diffusion (`ar_diffusion.pt`): 50-step UniPC + CFG.
# guidance_scale = 3.0
# num_inference_steps = 50
#
# Stage 2 — CCD (`causal_cd.pt`): 4-step consistency-distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
#
# Stage 3 — DMD (`causal_forcing.pt`): 4-step distribution-matching distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
guidance_scale = 1.0
num_inference_steps = 4
seed = 43
lora_weight = 0.55
save_path = "samples/wan-videos-causal-forcing"
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:
def _resolve_transformer_path(raw_path: str, prefer_ema: bool) -> str:
"""Resolve a checkpoint path to the actual weights file to load.
File paths are returned unchanged so external `.pt` ckpts (CF official) keep
working. For trainer-output dirs, prefer EMA when asked and available,
otherwise fall back to live `transformer/` weights.
"""
if os.path.isfile(raw_path):
return raw_path
if os.path.isdir(raw_path):
candidates = []
if prefer_ema:
candidates.append(os.path.join(raw_path, "ema_transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "diffusion_pytorch_model.safetensors"))
for c in candidates:
if os.path.isfile(c):
return c
raise FileNotFoundError(
f"transformer_path={raw_path!r} is neither a file nor a checkpoint dir "
f"with a known safetensors layout (transformer/ or ema_transformer/)."
)
_raw_transformer_path = transformer_path
transformer_path = _resolve_transformer_path(transformer_path, prefer_ema=use_ema)
if transformer_path != _raw_transformer_path:
print(f"use_ema={use_ema}: resolved {_raw_transformer_path} -> {transformer_path}")
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
# Causal-Forcing's FSDP-saved ckpts (causal_cd.pt / causal_forcing.pt) keep the
# `model._fsdp_wrapped_module.` prefix; strip it before the generic `model.` strip
# so both kinds of ckpt land at bare parameter names.
if any("._fsdp_wrapped_module." in k for k in state_dict.keys()):
state_dict = {k.replace("model._fsdp_wrapped_module.", "model.", 1) if k.startswith("model._fsdp_wrapped_module.") else k: v for k, v in state_dict.items()}
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)
# Only the main process (rank 0, or single-GPU) writes files to disk.
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
is_main_process = dist.get_rank() == 0
else:
is_main_process = True
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
# In streaming mode, set up an incremental writer up-front so each block
# can be flushed to disk right after it is decoded.
saver = None
decode_callback = None
if streaming:
if is_main_process:
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 save_mode == "segments":
# One standalone mp4 per decoded block under a per-prompt dir.
saver = SegmentVideoSaver(os.path.join(save_path, prefix), fps)
else:
# One continuous mp4, appended block by block.
saver = StreamVideoSaver(os.path.join(save_path, prefix + ".mp4"), fps)
decode_callback = saver
else:
# Non-main ranks still decode locally (matching non-streaming
# behaviour) but must not write; a no-op avoids accumulation.
decode_callback = lambda video_chunk, block_idx: None
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,
streaming = streaming,
decode_callback = decode_callback,
).videos
if saver is not None:
saver.close()
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)
# Streaming already wrote the video incrementally above; only the
# non-streaming path needs the one-shot save below.
if not streaming and is_main_process:
save_results()
@@ -0,0 +1,319 @@
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, StreamVideoSaver,
SegmentVideoSaver)
# 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
# Streaming decode: decode and write each causal block to disk as soon as it is
# generated ("generate a chunk, save a chunk"), instead of holding all latents
# and decoding once at the end. The VAE causal cache is threaded across blocks,
# so the saved video is seam-free and identical to a full decode.
streaming = True
# Streaming save mode (only used when streaming=True):
# "stream" : append every block into one continuous mp4 (finalized on close).
# "segments" : save each decoded block as its own standalone mp4, flushed
# immediately after decoding ("decode a chunk -> save it now").
# Robust to interruption; segments can be concatenated later.
save_mode = "segments"
# 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)
# Only the main process (rank 0, or single-GPU) writes files to disk.
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
is_main_process = dist.get_rank() == 0
else:
is_main_process = True
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
# In streaming mode, set up an incremental writer up-front so each block
# can be flushed to disk right after it is decoded.
saver = None
decode_callback = None
if streaming:
if is_main_process:
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 save_mode == "segments":
# One standalone mp4 per decoded block under a per-prompt dir.
saver = SegmentVideoSaver(os.path.join(save_path, prefix), fps)
else:
# One continuous mp4, appended block by block.
saver = StreamVideoSaver(os.path.join(save_path, prefix + ".mp4"), fps)
decode_callback = saver
else:
# Non-main ranks still decode locally (matching non-streaming
# behaviour) but must not write; a no-op avoids accumulation.
decode_callback = lambda video_chunk, block_idx: None
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,
streaming = streaming,
decode_callback = decode_callback,
).videos
if saver is not None:
saver.close()
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)
# Streaming already wrote the video incrementally above; only the
# non-streaming path needs the one-shot save below.
if not streaming and is_main_process:
save_results()
@@ -0,0 +1,276 @@
# Wan2.1 Causal-Forcing Stage 1: Autoregressive Diffusion Training Guide
This document provides the complete workflow for **Causal-Forcing Stage 1 — Autoregressive Diffusion (AR Diffusion) training** on Wan2.1-T2V-1.3B.
> **What is Causal-Forcing?**
>
> Causal-Forcing is a three-stage pipeline that compresses a video diffusion model into a **few-step causal autoregressive generator**:
>
> 1. **Stage 1 — AR Diffusion** (`train_ar_diffusion.py`): Train the model with **teacher forcing** on clean video latents, learning to denoise chunk-by-chunk (or frame-by-frame) in a causal autoregressive manner. This produces a strong AR backbone.
> 2. **Stage 2 — Causal Consistency Distillation** (`train_causal_consistency_distill.py`): Distill the multi-step AR model into a one-step-per-block consistency model.
> 3. **Stage 3 — Causal DMD** (`train_causal_dmd.py`): Further distill to a **2-step** generator using distribution matching with a 14B teacher.
>
> This README covers **Stage 1 only**.
---
## Table of Contents
- [1. Environment Setup](#1-environment-setup)
- [2. Download Pretrained Models](#2-download-pretrained-models)
- [3. Prepare Training Data](#3-prepare-training-data)
- [3.1 Quick Demo Dataset](#31-quick-demo-dataset)
- [3.2 Dataset Structure](#32-dataset-structure)
- [3.3 metadata.json Format](#33-metadatajson-format)
- [4. Training](#4-training)
- [4.1 Quick Start](#41-quick-start)
- [4.2 Key Parameters](#42-key-parameters)
- [4.3 Causal-Forcing-Specific Parameters](#43-causal-forcing-specific-parameters)
- [4.4 Training with FSDP](#44-training-with-fsdp)
- [5. Use the Trained Checkpoint](#5-use-the-trained-checkpoint)
- [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
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**Method 3: Using Docker**
```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
Stage 1 initializes the Wan2.1-T2V-1.3B model and trains it with causal teacher forcing.
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download Wan2.1 T2V base model
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
```
---
## 3. Prepare Training Data
Stage 1 trains directly on **raw videos** with online VAE encoding — no ODE trajectory pairs are needed.
### 3.1 Quick Demo Dataset
We provide a small demo dataset for quick testing:
```bash
# Download official demo dataset
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 3.2 Dataset Structure
```
📦 datasets/
├── 📂 X-Fun-Videos-Demo/
│ ├── 📂 train/
│ │ ├── 📄 video001.mp4
│ │ ├── 📄 video002.mp4
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 3.3 metadata.json Format
**Relative paths** (recommended):
```json
[
{
"file_path": "train/video001.mp4",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"type": "video"
},
{
"file_path": "train/video002.mp4",
"text": "A person walking through a forest, cinematic view",
"type": "video"
}
]
```
**Absolute paths**:
```json
[
{
"file_path": "/mnt/data/videos/sunset.mp4",
"text": "A beautiful sunset over the ocean",
"type": "video"
}
]
```
---
## 4. Training
### 4.1 Quick Start
The ready-to-use launcher is [train_ar_diffusion.sh](./train_ar_diffusion.sh):
```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"
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_causal_forcing/train_ar_diffusion.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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_ar_diffusion" \
--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 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--shift=5.0 \
--use_timestep_weight \
--trainable_modules "."
```
Or simply:
```bash
bash scripts/wan2.1_causal_forcing/train_ar_diffusion.sh
```
### 4.2 Key Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--pretrained_model_name_or_path` | Base model to initialize | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B/` |
| `--config_path` | Model config YAML | `config/wan2.1/wan_civitai.yaml` |
| `--train_data_dir` | Video root directory (prepended to relative `file_path`) | `""` |
| `--train_data_meta` | Path to `metadata.json` | `datasets/X-Fun-Videos-Demo/metadata.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 | 200 |
| `--learning_rate` | Initial learning rate | 2e-06 |
| `--lr_scheduler` | LR scheduler type | `constant_with_warmup` |
| `--lr_warmup_steps` | LR warmup steps | 100 |
| `--output_dir` | Output directory | `output_dir_wan2.1_causal_forcing_ar_diffusion` |
| `--gradient_checkpointing` | Enable activation 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) | `"."` |
**Video Sampling Parameters**:
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--image_sample_size` | Image sampling size | 640 |
| `--video_sample_size` | Video sampling size | 640 |
| `--token_sample_size` | Token sampling size | 640 |
| `--fix_sample_size` | Fixed `[height, width]` for output | `480 832` |
| `--video_sample_stride` | Frame sampling stride | 2 |
| `--video_sample_n_frames` | Number of video frames | 81 |
| `--random_hw_adapt` | Enable random resolution adaptation | - |
| `--training_with_video_token_length` | Enable token-length-based training | - |
| `--enable_bucket` | Enable aspect-ratio bucket sampling | - |
| `--vae_mini_batch` | VAE encoding mini-batch size (1 to avoid OOM) | 1 |
### 4.3 Causal-Forcing-Specific Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--num_frame_per_block` | Frames per causal block. `3` = chunkwise, `1` = framewise | 3 |
| `--independent_first_frame` | First frame is independent (`[1, N, N, ...]` block pattern, useful for I2V) | - |
| `--shift` | `FlowMatchEulerDiscreteScheduler` shift (Causal-Forcing default: 5.0) | 5.0 |
| `--train_sampling_steps` | Total scheduler timesteps for flow matching | 1000 |
| `--use_timestep_weight` | Apply per-timestep Gaussian loss weight (centered at T/2) | - |
| `--no_teacher_forcing` | Disable teacher forcing (use diffusion forcing instead) | - |
| `--noise_augmentation_max_timestep` | Add light noise to clean context tokens during teacher forcing (0 = off) | 0 |
> **Teacher Forcing**: By default, Stage 1 uses teacher forcing — the model receives the **clean GT latent** as context (`clean_x`) when predicting the current block. This stabilizes early training. Disable with `--no_teacher_forcing` to train under diffusion forcing instead.
### 4.4 Training with FSDP
The script above already uses FSDP with `CasualWanAttentionBlock` auto-wrapping. For single-GPU training without FSDP:
```bash
accelerate launch --mixed_precision="bf16" \
scripts/wan2.1_causal_forcing/train_ar_diffusion.py \
...
```
---
## 5. Use the Trained Checkpoint
The Stage 1 checkpoint is used to initialize **Stage 2 (Causal Consistency Distillation)**. Pass its path to `train_causal_consistency_distill.py` via `--transformer_path` and `--teacher_transformer_path`:
```bash
--transformer_path="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-{N}/diffusion_pytorch_model.safetensors" \
--teacher_transformer_path="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-{N}/diffusion_pytorch_model.safetensors"
```
See [train_causal_consistency_distill.sh](./train_causal_consistency_distill.sh) for the full Stage 2 workflow.
---
## 6. Additional Resources
- **Causal-Forcing Paper**: https://github.com/thu-ml/Causal-Forcing
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
@@ -0,0 +1,276 @@
# Wan2.1 Causal-Forcing 第一阶段:自回归扩散(AR Diffusion)训练指南
本文档介绍 **Causal-Forcing 第一阶段 — 自回归扩散(AR Diffusion)训练** 的完整流程。
> **什么是 Causal-Forcing?**
>
> Causal-Forcing 是一个三阶段流水线,将视频扩散模型压缩为 **少步因果自回归生成器**:
>
> 1. **第一阶段 — AR Diffusion**(`train_ar_diffusion.py`):在干净视频 latent 上使用 **teacher forcing** 训练,让模型学会按块(或按帧)因果自回归地去噪,得到一个强大的 AR 骨架。
> 2. **第二阶段 — 因果一致性蒸馏**(`train_causal_consistency_distill.py`):将多步 AR 模型蒸馏为每块一步的一致性模型。
> 3. **第三阶段 — 因果 DMD**(`train_causal_dmd.py`):利用 14B 教师模型做分布匹配,进一步蒸馏为 **2 步** 生成器。
>
> 本文档仅覆盖 **第一阶段**。
---
## 目录
- [一、环境配置](#一环境配置)
- [二、下载预训练模型](#二下载预训练模型)
- [三、准备训练数据](#三准备训练数据)
- [3.1 快速测试数据集](#31-快速测试数据集)
- [3.2 数据集结构](#32-数据集结构)
- [3.3 metadata.json 格式](#33-metadatajson-格式)
- [四、训练](#四训练)
- [4.1 快速开始](#41-快速开始)
- [4.2 主要参数说明](#42-主要参数说明)
- [4.3 Causal-Forcing 特有参数](#43-causal-forcing-特有参数)
- [4.4 使用 FSDP 训练](#44-使用-fsdp-训练)
- [五、使用训练好的 Checkpoint](#五使用训练好的-checkpoint)
- [六、更多资源](#六更多资源)
---
## 一、环境配置
**方式 1:使用 requirements.txt**
```bash
pip install -r requirements.txt
```
**方式 2:手动安装依赖**
```bash
pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image
pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime
pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2"
pip install yunchang xfuser modelscope openpyxl
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**方式 3:使用 docker**
```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
```
---
## 二、下载预训练模型
第一阶段以 Wan2.1-T2V-1.3B 为基础模型,在其上进行因果 teacher forcing 训练。
```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
```
---
## 三、准备训练数据
第一阶段直接在 **原始视频** 上训练,VAE 在线编码 — **不需要** ODE 轨迹对或预计算 latent。
### 3.1 快速测试数据集
我们提供了一个小型 demo 数据集用于快速测试:
```bash
# 下载官方示例数据集
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 3.2 数据集结构
```
📦 datasets/
├── 📂 X-Fun-Videos-Demo/
│ ├── 📂 train/
│ │ ├── 📄 video001.mp4
│ │ ├── 📄 video002.mp4
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 3.3 metadata.json 格式
**相对路径**(推荐):
```json
[
{
"file_path": "train/video001.mp4",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"type": "video"
},
{
"file_path": "train/video002.mp4",
"text": "A person walking through a forest, cinematic view",
"type": "video"
}
]
```
**绝对路径**:
```json
[
{
"file_path": "/mnt/data/videos/sunset.mp4",
"text": "A beautiful sunset over the ocean",
"type": "video"
}
]
```
---
## 四、训练
### 4.1 快速开始
直接执行启动脚本 [train_ar_diffusion.sh](./train_ar_diffusion.sh):
```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"
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_causal_forcing/train_ar_diffusion.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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_ar_diffusion" \
--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 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--shift=5.0 \
--use_timestep_weight \
--trainable_modules "."
```
或者直接执行:
```bash
bash scripts/wan2.1_causal_forcing/train_ar_diffusion.sh
```
### 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` 之前) | `""` |
| `--train_data_meta` | `metadata.json` 路径 | `datasets/X-Fun-Videos-Demo/metadata.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 | 200 |
| `--learning_rate` | 初始学习率 | 2e-06 |
| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` |
| `--lr_warmup_steps` | 学习率预热步数 | 100 |
| `--output_dir` | 输出目录 | `output_dir_wan2.1_causal_forcing_ar_diffusion` |
| `--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` | 可训练模块(`"."` 表示全量) | `"."` |
**视频采样参数**:
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--image_sample_size` | 图像采样尺寸 | 640 |
| `--video_sample_size` | 视频采样尺寸 | 640 |
| `--token_sample_size` | Token 采样尺寸 | 640 |
| `--fix_sample_size` | 固定输出 `[高度, 宽度]` | `480 832` |
| `--video_sample_stride` | 帧采样步长 | 2 |
| `--video_sample_n_frames` | 视频帧数 | 81 |
| `--random_hw_adapt` | 启用随机分辨率自适应 | - |
| `--training_with_video_token_length` | 启用基于 token 长度的训练 | - |
| `--enable_bucket` | 启用宽高比 bucket 采样 | - |
| `--vae_mini_batch` | VAE 编码 mini-batch 大小(设为 1 可避免显存溢出) | 1 |
### 4.3 Causal-Forcing 特有参数
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--num_frame_per_block` | 每个因果块的帧数。`3` = chunkwise,`1` = framewise | 3 |
| `--independent_first_frame` | 第一帧独立(`[1, N, N, ...]` 块模式,I2V 场景适用) | - |
| `--shift` | `FlowMatchEulerDiscreteScheduler` 的 shift 值(Causal-Forcing 默认 5.0) | 5.0 |
| `--train_sampling_steps` | flow matching 调度器总时间步数 | 1000 |
| `--use_timestep_weight` | 启用逐时间步高斯损失权重(以 T/2 为中心) | - |
| `--no_teacher_forcing` | 关闭 teacher forcing(改为 diffusion forcing) | - |
| `--noise_augmentation_max_timestep` | teacher forcing 时向干净上下文 token 添加轻微噪声(0 = 关闭) | 0 |
> **Teacher Forcing 说明**:默认情况下,第一阶段使用 teacher forcing — 模型在预测当前块时接收 **干净的 GT latent** 作为上下文(`clean_x`)。这可以稳定早期训练。如需使用 diffusion forcing,可通过 `--no_teacher_forcing` 关闭。
### 4.4 使用 FSDP 训练
上述脚本已配置 FSDP,使用 `CasualWanAttentionBlock` 自动包装。单卡训练无需 FSDP:
```bash
accelerate launch --mixed_precision="bf16" \
scripts/wan2.1_causal_forcing/train_ar_diffusion.py \
...
```
---
## 五、使用训练好的 Checkpoint
第一阶段输出的 checkpoint 用于初始化 **第二阶段(因果一致性蒸馏)**。在 `train_causal_consistency_distill.py` 中通过 `--transformer_path` 和 `--teacher_transformer_path` 指定:
```bash
--transformer_path="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-{N}/diffusion_pytorch_model.safetensors" \
--teacher_transformer_path="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-{N}/diffusion_pytorch_model.safetensors"
```
完整的第二阶段流程参见 [train_causal_consistency_distill.sh](./train_causal_consistency_distill.sh)。
---
## 六、更多资源
- **Causal-Forcing 论文**:https://github.com/thu-ml/Causal-Forcing
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
@@ -0,0 +1,293 @@
# Wan2.1 Causal-Forcing Stage 2: Causal Consistency Distillation (CCD) Training Guide
This document provides the complete workflow for **Causal-Forcing Stage 2 — Causal Consistency Distillation (CCD)** on Wan2.1-T2V-1.3B.
> **What is Causal Consistency Distillation?**
>
> CCD is the **second stage** of the Causal-Forcing pipeline, which compresses a video diffusion model into a **few-step causal autoregressive generator**:
>
> 1. **Stage 1 — AR Diffusion** (`train_ar_diffusion.py`): Train the model with teacher forcing on clean video latents, producing a strong AR backbone.
> 2. **Stage 2 — Causal Consistency Distillation** (`train_causal_consistency_distill.py`): Distill the multi-step AR model into a **one-step-per-block** consistency model using an EMA teacher with CFG guidance.
> 3. **Stage 3 — Causal DMD** (`train_causal_dmd.py`): Further distill to a **2-step** generator using distribution matching with a 14B teacher.
>
> This README covers **Stage 2 only**. See [README_TRAIN_AR_DIFFUSION.md](./README_TRAIN_AR_DIFFUSION.md) for Stage 1.
---
## Table of Contents
- [1. Prerequisites](#1-prerequisites)
- [2. Environment Setup](#2-environment-setup)
- [3. Download Pretrained Models](#3-download-pretrained-models)
- [4. Prepare Training Data](#4-prepare-training-data)
- [4.1 Quick Demo Dataset](#41-quick-demo-dataset)
- [4.2 Dataset Structure](#42-dataset-structure)
- [4.3 metadata.json Format](#43-metadatajson-format)
- [5. Training](#5-training)
- [5.1 Quick Start](#51-quick-start)
- [5.2 Key Parameters](#52-key-parameters)
- [5.3 CCD-Specific Parameters](#53-ccd-specific-parameters)
- [6. Use the Trained Checkpoint](#6-use-the-trained-checkpoint)
- [7. Additional Resources](#7-additional-resources)
---
## 1. Prerequisites
Stage 2 requires a **Stage 1 AR Diffusion checkpoint** to initialize both the generator and the EMA teacher.
```bash
# Example: Stage 1 checkpoint from AR Diffusion training
export STAGE1_CKPT="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-8000/diffusion_pytorch_model.safetensors"
```
See [README_TRAIN_AR_DIFFUSION.md](./README_TRAIN_AR_DIFFUSION.md) for how to produce this checkpoint.
---
## 2. 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
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**Method 3: Using Docker**
```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
```
---
## 3. Download Pretrained Models
CCD initializes the generator and teacher from Wan2.1-T2V-1.3B, then loads the Stage 1 checkpoint on top.
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download Wan2.1 T2V base model
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
```
---
## 4. Prepare Training Data
CCD trains directly on **raw videos** with online VAE encoding — same data format as Stage 1.
### 4.1 Quick Demo Dataset
We provide a small demo dataset for quick testing:
```bash
# Download official demo dataset
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 4.2 Dataset Structure
```
📦 datasets/
├── 📂 X-Fun-Videos-Demo/
│ ├── 📂 train/
│ │ ├── 📄 video001.mp4
│ │ ├── 📄 video002.mp4
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 4.3 metadata.json Format
**Relative paths** (recommended):
```json
[
{
"file_path": "train/video001.mp4",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"type": "video"
}
]
```
**Absolute paths**:
```json
[
{
"file_path": "/mnt/data/videos/sunset.mp4",
"text": "A beautiful sunset over the ocean",
"type": "video"
}
]
```
---
## 5. Training
### 5.1 Quick Start
The ready-to-use launcher is [train_causal_consistency_distill.sh](./train_causal_consistency_distill.sh):
```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 STAGE1_CKPT="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-8000/diffusion_pytorch_model.safetensors"
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_causal_forcing/train_causal_consistency_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 \
--transformer_path=$STAGE1_CKPT \
--teacher_transformer_path=$STAGE1_CKPT \
--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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2.0e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_ccd" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=0.0 \
--adam_beta1=0.0 \
--adam_beta2=0.999 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=10.0 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--num_frame_per_block=3 \
--shift=5.0 \
--discrete_cd_N=48 \
--guidance_scale=3.0 \
--ema_weight=0.99 \
--ema_start_step=200 \
--trainable_modules "."
```
Or simply:
```bash
bash scripts/wan2.1_causal_forcing/train_causal_consistency_distill.sh
```
### 5.2 Key Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--pretrained_model_name_or_path` | Base model to initialize | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B/` |
| `--config_path` | Model config YAML | `config/wan2.1/wan_civitai.yaml` |
| `--train_data_dir` | Video root directory (prepended to relative `file_path`) | `""` |
| `--train_data_meta` | Path to `metadata.json` | `datasets/X-Fun-Videos-Demo/metadata.json` |
| `--transformer_path` | Stage 1 checkpoint for generator init | `$STAGE1_CKPT` |
| `--teacher_transformer_path` | Stage 1 checkpoint for teacher init (defaults to generator) | `$STAGE1_CKPT` |
| `--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 | 200 |
| `--learning_rate` | Initial learning rate | 2e-06 |
| `--lr_scheduler` | LR scheduler type | `constant_with_warmup` |
| `--lr_warmup_steps` | LR warmup steps | 100 |
| `--output_dir` | Output directory | `output_dir_wan2.1_causal_forcing_ccd` |
| `--gradient_checkpointing` | Enable activation checkpointing | - |
| `--mixed_precision` | `fp16` / `bf16` | `bf16` |
| `--adam_weight_decay` | AdamW weight decay | 0.0 |
| `--adam_beta1` | AdamW beta1 (CCD uses 0.0) | 0.0 |
| `--adam_beta2` | AdamW beta2 | 0.999 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--max_grad_norm` | Gradient clipping threshold | 10.0 |
| `--trainable_modules` | Trainable modules (`"."` = all) | `"."` |
| `--low_vram` | Enable low VRAM mode (offload VAE/text encoder) | - |
**Video Sampling Parameters**:
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--image_sample_size` | Image sampling size | 640 |
| `--video_sample_size` | Video sampling size | 640 |
| `--token_sample_size` | Token sampling size | 640 |
| `--fix_sample_size` | Fixed `[height, width]` for output | `480 832` |
| `--video_sample_stride` | Frame sampling stride | 2 |
| `--video_sample_n_frames` | Number of video frames | 81 |
| `--random_hw_adapt` | Enable random resolution adaptation | - |
| `--training_with_video_token_length` | Enable token-length-based training | - |
| `--enable_bucket` | Enable aspect-ratio bucket sampling | - |
| `--vae_mini_batch` | VAE encoding mini-batch size (1 to avoid OOM) | 1 |
### 5.3 CCD-Specific Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--num_frame_per_block` | Frames per causal block. `3` = chunkwise, `1` = framewise | 3 |
| `--independent_first_frame` | First frame is independent (`[1, N, N, ...]` block pattern, useful for I2V) | - |
| `--shift` | `FlowMatchEulerDiscreteScheduler` shift (Causal-Forcing default: 5.0) | 5.0 |
| `--discrete_cd_N` | Number of discrete timesteps for the consistency schedule | 48 |
| `--guidance_scale` | CFG guidance scale used by the EMA teacher | 3.0 |
| `--ema_weight` | EMA decay for the consistency-target generator copy (set <=0 to disable) | 0.99 |
| `--ema_start_step` | Steps to wait before EMA tracking starts | 200 |
> **How CCD works**: For each training sample, CCD:
> 1. Loads clean video latents (via online VAE encoding).
> 2. Samples a timestep index from the discrete consistency schedule `[0, N-2]`.
> 3. Adds noise to the clean latents at timestep `t` and `t_next`.
> 4. Runs the **generator** on `x_t` to predict `x0`.
> 5. Runs the **EMA teacher** (with CFG guidance) on `x_{t_next}` to produce the consistency target.
> 6. Minimizes the L2 loss between the generator prediction and the teacher target.
> **EMA Teacher**: The EMA copy tracks the generator with polyak updates (`ema = decay*ema + (1-decay)*gen`). Before `--ema_start_step`, the EMA mirrors the live generator. The teacher uses CFG with `--guidance_scale=3.0` to produce higher-quality targets.
---
## 6. Use the Trained Checkpoint
The Stage 2 CCD checkpoint is used to initialize **Stage 3 (Causal DMD)**. Pass its path to `train_causal_dmd.py`:
```bash
--ode_transformer_path="output_dir_wan2.1_causal_forcing_ccd/checkpoint-{N}/diffusion_pytorch_model.safetensors"
```
See `train_causal_dmd.sh` for the full Stage 3 workflow.
---
## 7. Additional Resources
- **Causal-Forcing Paper**: https://github.com/thu-ml/Causal-Forcing
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
@@ -0,0 +1,293 @@
# Wan2.1 Causal-Forcing 第二阶段:因果一致性蒸馏(CCD)训练指南
本文档介绍 **Causal-Forcing 第二阶段 — 因果一致性蒸馏(CCD)** 的完整流程。
> **什么是因果一致性蒸馏?**
>
> CCD 是 Causal-Forcing 流水线的 **第二阶段**,将视频扩散模型压缩为 **少步因果自回归生成器**:
>
> 1. **第一阶段 — AR Diffusion**(`train_ar_diffusion.py`):在干净视频 latent 上使用 teacher forcing 训练,得到强大的 AR 骨架。
> 2. **第二阶段 — 因果一致性蒸馏**(`train_causal_consistency_distill.py`):利用 EMA teacher 和 CFG 引导,将多步 AR 模型蒸馏为 **每块一步** 的一致性模型。
> 3. **第三阶段 — 因果 DMD**(`train_causal_dmd.py`):利用 14B 教师模型做分布匹配,进一步蒸馏为 **2 步** 生成器。
>
> 本文档仅覆盖 **第二阶段**。第一阶段参见 [README_TRAIN_AR_DIFFUSION_zh-CN.md](./README_TRAIN_AR_DIFFUSION_zh-CN.md)。
---
## 目录
- [一、前置条件](#一前置条件)
- [二、环境配置](#二环境配置)
- [三、下载预训练模型](#三下载预训练模型)
- [四、准备训练数据](#四准备训练数据)
- [4.1 快速测试数据集](#41-快速测试数据集)
- [4.2 数据集结构](#42-数据集结构)
- [4.3 metadata.json 格式](#43-metadatajson-格式)
- [五、训练](#五训练)
- [5.1 快速开始](#51-快速开始)
- [5.2 主要参数说明](#52-主要参数说明)
- [5.3 CCD 特有参数](#53-ccd-特有参数)
- [六、使用训练好的 Checkpoint](#六使用训练好的-checkpoint)
- [七、更多资源](#七更多资源)
---
## 一、前置条件
第二阶段需要 **第一阶段 AR Diffusion 的 checkpoint** 来初始化 generator 和 EMA teacher。
```bash
# 示例:第一阶段 AR Diffusion 训练输出的 checkpoint
export STAGE1_CKPT="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-8000/diffusion_pytorch_model.safetensors"
```
第一阶段的训练方法参见 [README_TRAIN_AR_DIFFUSION_zh-CN.md](./README_TRAIN_AR_DIFFUSION_zh-CN.md)。
---
## 二、环境配置
**方式 1:使用 requirements.txt**
```bash
pip install -r requirements.txt
```
**方式 2:手动安装依赖**
```bash
pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image
pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime
pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2"
pip install yunchang xfuser modelscope openpyxl
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**方式 3:使用 docker**
```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
```
---
## 三、下载预训练模型
CCD 以 Wan2.1-T2V-1.3B 初始化 generator 和 teacher,然后加载第一阶段 checkpoint。
```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
```
---
## 四、准备训练数据
CCD 直接在 **原始视频** 上训练,VAE 在线编码 — 数据格式与第一阶段相同。
### 4.1 快速测试数据集
我们提供了一个小型 demo 数据集用于快速测试:
```bash
# 下载官方示例数据集
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 4.2 数据集结构
```
📦 datasets/
├── 📂 X-Fun-Videos-Demo/
│ ├── 📂 train/
│ │ ├── 📄 video001.mp4
│ │ ├── 📄 video002.mp4
│ │ └── 📄 ...
│ └── 📄 metadata.json
```
### 4.3 metadata.json 格式
**相对路径**(推荐):
```json
[
{
"file_path": "train/video001.mp4",
"text": "A beautiful sunset over the ocean, golden hour lighting",
"type": "video"
}
]
```
**绝对路径**:
```json
[
{
"file_path": "/mnt/data/videos/sunset.mp4",
"text": "A beautiful sunset over the ocean",
"type": "video"
}
]
```
---
## 五、训练
### 5.1 快速开始
直接执行启动脚本 [train_causal_consistency_distill.sh](./train_causal_consistency_distill.sh):
```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 STAGE1_CKPT="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-8000/diffusion_pytorch_model.safetensors"
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_causal_forcing/train_causal_consistency_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 \
--transformer_path=$STAGE1_CKPT \
--teacher_transformer_path=$STAGE1_CKPT \
--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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2.0e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_ccd" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=0.0 \
--adam_beta1=0.0 \
--adam_beta2=0.999 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=10.0 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--num_frame_per_block=3 \
--shift=5.0 \
--discrete_cd_N=48 \
--guidance_scale=3.0 \
--ema_weight=0.99 \
--ema_start_step=200 \
--trainable_modules "."
```
或者直接执行:
```bash
bash scripts/wan2.1_causal_forcing/train_causal_consistency_distill.sh
```
### 5.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` 之前) | `""` |
| `--train_data_meta` | `metadata.json` 路径 | `datasets/X-Fun-Videos-Demo/metadata.json` |
| `--transformer_path` | 第一阶段 checkpoint(generator 初始化) | `$STAGE1_CKPT` |
| `--teacher_transformer_path` | 第一阶段 checkpoint(teacher 初始化,默认同 generator) | `$STAGE1_CKPT` |
| `--train_batch_size` | 每卡 batch size | 1 |
| `--gradient_accumulation_steps` | 梯度累积步数 | 1 |
| `--dataloader_num_workers` | DataLoader 子进程数 | 8 |
| `--num_train_epochs` | 训练 epoch 数 | 100 |
| `--checkpointing_steps` | 每 N 步保存 checkpoint | 200 |
| `--learning_rate` | 初始学习率 | 2e-06 |
| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` |
| `--lr_warmup_steps` | 学习率预热步数 | 100 |
| `--output_dir` | 输出目录 | `output_dir_wan2.1_causal_forcing_ccd` |
| `--gradient_checkpointing` | 启用激活重计算 | - |
| `--mixed_precision` | `fp16` / `bf16` | `bf16` |
| `--adam_weight_decay` | AdamW 权重衰减 | 0.0 |
| `--adam_beta1` | AdamW beta1(CCD 使用 0.0) | 0.0 |
| `--adam_beta2` | AdamW beta2 | 0.999 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--max_grad_norm` | 梯度裁剪阈值 | 10.0 |
| `--trainable_modules` | 可训练模块(`"."` 表示全量) | `"."` |
| `--low_vram` | 启用低显存模式(VAE/文本编码器分时加载) | - |
**视频采样参数**:
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--image_sample_size` | 图像采样尺寸 | 640 |
| `--video_sample_size` | 视频采样尺寸 | 640 |
| `--token_sample_size` | Token 采样尺寸 | 640 |
| `--fix_sample_size` | 固定输出 `[高度, 宽度]` | `480 832` |
| `--video_sample_stride` | 帧采样步长 | 2 |
| `--video_sample_n_frames` | 视频帧数 | 81 |
| `--random_hw_adapt` | 启用随机分辨率自适应 | - |
| `--training_with_video_token_length` | 启用基于 token 长度的训练 | - |
| `--enable_bucket` | 启用宽高比 bucket 采样 | - |
| `--vae_mini_batch` | VAE 编码 mini-batch 大小(设为 1 可避免显存溢出) | 1 |
### 5.3 CCD 特有参数
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `--num_frame_per_block` | 每个因果块的帧数。`3` = chunkwise,`1` = framewise | 3 |
| `--independent_first_frame` | 第一帧独立(`[1, N, N, ...]` 块模式,I2V 场景适用) | - |
| `--shift` | `FlowMatchEulerDiscreteScheduler` 的 shift 值(Causal-Forcing 默认 5.0) | 5.0 |
| `--discrete_cd_N` | 一致性调度器的离散时间步数量 | 48 |
| `--guidance_scale` | EMA teacher 使用的 CFG 引导强度 | 3.0 |
| `--ema_weight` | EMA 衰减率(一致性目标 generator 副本,设为 <=0 禁用) | 0.99 |
| `--ema_start_step` | EMA 跟踪启动前的等待步数 | 200 |
> **CCD 工作原理**:对每个训练样本,CCD 执行以下步骤:
> 1. 加载干净视频 latent(通过 VAE 在线编码)。
> 2. 从离散一致性调度 `[0, N-2]` 中采样一个时间步索引。
> 3. 在时间步 `t` 和 `t_next` 分别给干净 latent 加噪。
> 4. **Generator** 在 `x_t` 上预测 `x0`。
> 5. **EMA teacher**(带 CFG 引导)在 `x_{t_next}` 上生成一致性目标。
> 6. 最小化 generator 预测与 teacher 目标之间的 L2 损失。
> **EMA Teacher**:EMA 副本通过 polyak 更新跟踪 generator(`ema = decay*ema + (1-decay)*gen`)。在 `--ema_start_step` 之前,EMA 直接镜像 generator。Teacher 使用 `--guidance_scale=3.0` 的 CFG 来生成更高质量的目标。
---
## 六、使用训练好的 Checkpoint
第二阶段 CCD 输出的 checkpoint 用于初始化 **第三阶段(因果 DMD)**。在 `train_causal_dmd.py` 中指定路径:
```bash
--ode_transformer_path="output_dir_wan2.1_causal_forcing_ccd/checkpoint-{N}/diffusion_pytorch_model.safetensors"
```
完整的第三阶段流程参见 `train_causal_dmd.sh`。
---
## 七、更多资源
- **Causal-Forcing 论文**:https://github.com/thu-ml/Causal-Forcing
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
@@ -0,0 +1,303 @@
# Wan2.1 Causal-Forcing Stage 3: Distribution Matching Distillation (DMD) Training Guide
This document provides the complete workflow for **Causal-Forcing Stage 3 — Distribution Matching Distillation (DMD)** on Wan2.1-T2V-1.3B.
> **What is Distribution Matching Distillation?**
>
> DMD is the **third and final stage** of the Causal-Forcing pipeline, which further compresses a CCD model into a **few-step causal autoregressive generator** using a large (14B) teacher:
>
> 1. **Stage 1 — AR Diffusion** (`train_ar_diffusion.py`): Train the model with teacher forcing on clean video latents, producing a strong AR backbone.
> 2. **Stage 2 — Causal Consistency Distillation** (`train_causal_consistency_distill.py`): Distill the multi-step AR model into a **one-step-per-block** consistency model using an EMA teacher with CFG guidance.
> 3. **Stage 3 — Distribution Matching Distillation** (`train_causal_dmd.py`): Further distill to a **few-step** generator using distribution matching with a **14B real-score teacher**.
>
> This README covers **Stage 3 only**. See [README_TRAIN_CAUSAL_CONSISTENCY_DISTILL.md](./README_TRAIN_CAUSAL_CONSISTENCY_DISTILL.md) for Stage 2 and [README_TRAIN_AR_DIFFUSION.md](./README_TRAIN_AR_DIFFUSION.md) for Stage 1.
---
## Table of Contents
- [1. Prerequisites](#1-prerequisites)
- [2. Environment Setup](#2-environment-setup)
- [3. Download Pretrained Models](#3-download-pretrained-models)
- [4. Prepare Training Data](#4-prepare-training-data)
- [4.1 Quick Demo Dataset](#41-quick-demo-dataset)
- [4.2 Dataset Structure](#42-dataset-structure)
- [4.3 metadata.json Format](#43-metadatajson-format)
- [5. Training](#5-training)
- [5.1 Quick Start](#51-quick-start)
- [5.2 Key Parameters](#52-key-parameters)
- [5.3 DMD-Specific Parameters](#53-dmd-specific-parameters)
- [6. Use the Trained Checkpoint](#6-use-the-trained-checkpoint)
- [7. Additional Resources](#7-additional-resources)
---
## 1. Prerequisites
Stage 3 requires:
1. A **Stage 2 CCD checkpoint** to initialize the generator (and critic).
2. A **Wan2.1-T2V-14B** model as the DMD real-score teacher.
```bash
# Example: Stage 2 checkpoint from CCD training
export STAGE2_CKPT="output_dir_wan2.1_causal_forcing_ccd/checkpoint-5000/transformer/diffusion_pytorch_model.safetensors"
```
See [README_TRAIN_CAUSAL_CONSISTENCY_DISTILL.md](./README_TRAIN_CAUSAL_CONSISTENCY_DISTILL.md) for how to produce this checkpoint.
---
## 2. 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
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**Method 3: Using Docker**
```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
```
---
## 3. Download Pretrained Models
DMD requires **two** pretrained models:
- **Wan2.1-T2V-1.3B**: base model for the generator/critic.
- **Wan2.1-T2V-14B**: the non-causal real-score teacher used by DMD to compute the real distribution score.
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download Wan2.1 T2V 1.3B (student base model)
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
# Download Wan2.1 T2V 14B (DMD real-score teacher)
modelscope download --model Wan-AI/Wan2.1-T2V-14B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-14B
```
The Stage 2 CCD checkpoint (generator/critic init) is loaded via `--ode_transformer_path`.
---
## 4. Prepare Training Data
DMD uses a **TextDataset** (`train_mode="normal"`) — it only needs prompts, not video data, because the generator creates its own training samples via autoregressive rollout. However, a `metadata.json` with prompts is still required.
### 4.1 Quick Demo Dataset
```bash
# Download official demo dataset
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 4.2 Dataset Structure
Same as Stage 1/2. See [README_TRAIN_CAUSAL_CONSISTENCY_DISTILL.md](./README_TRAIN_CAUSAL_CONSISTENCY_DISTILL.md#42-dataset-structure) for details.
### 4.3 metadata.json Format
DMD only uses the `"text"` field from each entry. The `file_path` and other fields are ignored when `--train_mode="normal"` (TextDataset mode).
```json
[
{
"text": "A beautiful sunset over the ocean, golden hour lighting"
},
{
"text": "A person walking through a forest, cinematic view"
}
]
```
> **Note**: You can reuse any video metadata.json — only the `text` field matters for DMD.
---
## 5. Training
### 5.1 Quick Start
The ready-to-use launcher is [train_causal_dmd.sh](./train_causal_dmd.sh):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export REAL_SCORE_MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-14B"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json"
export STAGE2_CKPT="output_dir_wan2.1_causal_forcing_ccd/checkpoint-5000/transformer/diffusion_pytorch_model.safetensors"
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_causal_forcing/train_causal_dmd.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--real_score_pretrained_model_name_or_path=$REAL_SCORE_MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--ode_transformer_path=$STAGE2_CKPT \
--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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2.0e-06 \
--learning_rate_critic=2.0e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_dmd" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=0.0 \
--adam_beta1=0.0 \
--adam_beta2=0.999 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=10.0 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--num_frame_per_block=3 \
--use_kv_cache_training \
--denoising_step_indices_list 1000 667 334 1 \
--real_guidance_scale=6.0 \
--randomize_step_indices \
--fake_guidance_scale=0.0 \
--gen_update_interval=5 \
--trainable_modules "."
```
Or simply:
```bash
bash scripts/wan2.1_causal_forcing/train_causal_dmd.sh
```
### 5.2 Key Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--pretrained_model_name_or_path` | Base model (1.3B) for generator/critic init | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B/` |
| `--real_score_pretrained_model_name_or_path` | 14B real-score teacher for DMD | `models/Diffusion_Transformer/Wan2.1-T2V-14B` |
| `--config_path` | Model config YAML | `config/wan2.1/wan_civitai.yaml` |
| `--train_data_dir` | Data root directory | `""` |
| `--train_data_meta` | Path to `metadata.json` (prompts only) | `datasets/X-Fun-Videos-Demo/metadata.json` |
| `--ode_transformer_path` | Stage 2 CCD checkpoint for generator/critic init | `$STAGE2_CKPT` |
| `--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 | 200 |
| `--learning_rate` | Generator learning rate | 2e-06 |
| `--learning_rate_critic` | Critic learning rate | 2e-06 |
| `--lr_scheduler` | LR scheduler type | `constant_with_warmup` |
| `--lr_warmup_steps` | LR warmup steps | 100 |
| `--output_dir` | Output directory | `output_dir_wan2.1_causal_forcing_dmd` |
| `--gradient_checkpointing` | Enable activation checkpointing | - |
| `--mixed_precision` | `fp16` / `bf16` | `bf16` |
| `--adam_weight_decay` | AdamW weight decay | 0.0 |
| `--adam_beta1` | AdamW beta1 (DMD uses 0.0) | 0.0 |
| `--adam_beta2` | AdamW beta2 | 0.999 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--max_grad_norm` | Gradient clipping threshold | 10.0 |
| `--trainable_modules` | Trainable modules (`"."` = all) | `"."` |
| `--low_vram` | Enable low VRAM mode (offload VAE/text encoder) | - |
**Video Sampling Parameters**:
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--image_sample_size` | Image sampling size | 640 |
| `--video_sample_size` | Video sampling size | 640 |
| `--token_sample_size` | Token sampling size | 640 |
| `--fix_sample_size` | Fixed `[height, width]` for output | `480 832` |
| `--video_sample_stride` | Frame sampling stride | 2 |
| `--video_sample_n_frames` | Number of video frames | 81 |
| `--random_hw_adapt` | Enable random resolution adaptation | - |
| `--training_with_video_token_length` | Enable token-length-based training | - |
| `--enable_bucket` | Enable aspect-ratio bucket sampling | - |
| `--vae_mini_batch` | VAE encoding mini-batch size (1 to avoid OOM) | 1 |
### 5.3 DMD-Specific Parameters
| Parameter | Description | Example Value |
|-----------|-------------|---------------|
| `--denoising_step_indices_list` | Denoising step indices (DMD core param). Example runs 4-step; parser default `[1000, 500]` = 2-step | `1000 667 334 1` |
| `--real_guidance_scale` | CFG scale for the real-score (14B teacher) | 6.0 |
| `--randomize_step_indices` | Randomize the denoising step indices during training | - |
| `--fake_guidance_scale` | CFG scale for the fake-score (generator). 0.0 = no CFG | 0.0 |
| `--gen_update_interval` | Generator update interval (generator updates every N critic steps) | 5 |
| `--num_frame_per_block` | Frames per causal block. `3` = chunkwise (default), `1` = frame-wise | 3 |
| `--use_kv_cache_training` | Use KV cache block-by-block training (matches original Self-Forcing) | - |
| `--independent_first_frame` | First frame is independent (`[1, N, N, ...]` block pattern, useful for I2V) | - |
| `--context_noise` | Context noise level for KV cache update | 0 |
| `--use_teacher_forcing` | Enable teacher forcing (pass clean_x to transformer) | - |
| `--teacher_forcing_prob` | Probability of applying teacher forcing per step (1.0 = always) | 1.0 |
| `--train_mode` | Training mode: `normal` (TextDataset, prompt-only) or `i2v` | `normal` |
| `--resume_from_checkpoint` | Resume from checkpoint. Use `"latest"` to auto-select | `"latest"` |
---
## 6. Use the Trained Checkpoint
The Stage 3 DMD checkpoint is the **final model** in the Causal-Forcing pipeline. Use it for inference:
```python
# In examples/wan2.1_causal_forcing/predict_t2v.py
transformer_path = "output_dir_wan2.1_causal_forcing_dmd/checkpoint-{N}/diffusion_pytorch_model.safetensors"
# DMD Stage 3 inference config
guidance_scale = 1.0 # CFG is baked into distilled weights
num_inference_steps = 4 # 4-step DMD
stochastic_sampling = True
num_frame_per_block = 3 # Chunk-wise generation
```
Or run:
```bash
python examples/wan2.1_causal_forcing/predict_t2v.py
```
> **Stage Selector** in `predict_t2v.py` provides preset configs for all stages:
> - **Stage 1 (AR Diffusion)**: `guidance_scale=3.0`, `num_inference_steps=50`, `stochastic_sampling=False`
> - **Stage 2 (CCD)**: `guidance_scale=1.0`, `num_inference_steps=4`, `stochastic_sampling=True`
> - **Stage 3 (DMD)**: `guidance_scale=1.0`, `num_inference_steps=4`, `stochastic_sampling=True`
---
## 7. Additional Resources
- **Causal-Forcing Paper**: https://github.com/thu-ml/Causal-Forcing
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
@@ -0,0 +1,302 @@
# Wan2.1 Causal-Forcing Stage 3: 分布匹配蒸馏 (DMD) 训练指南
本文档提供了将 Wan2.1 进行 **Causal-Forcing Stage 3 — 分布匹配蒸馏 (DMD)** 的完整工作流。
> **什么是分布匹配蒸馏?**
>
> DMD 是 Causal-Forcing 流水线的**第三阶段(最终阶段)**,使用大规模(14B)teacher 将 CCD 模型进一步蒸馏为 **少步因果自回归生成器**:
>
> 1. **Stage 1 — AR Diffusion** (`train_ar_diffusion.py`):在干净视频 latent 上以 teacher forcing 训练模型,产出强 AR 基座。
> 2. **Stage 2 — 因果一致性蒸馏 (CCD)** (`train_causal_consistency_distill.py`):使用 EMA teacher + CFG 将多步 AR 模型蒸馏为**每块一步**的一致性模型。
> 3. **Stage 3 — 分布匹配蒸馏 (DMD)** (`train_causal_dmd.py`):使用 **14B real-score teacher** 的分布匹配进一步蒸馏为 **少步**生成器。
>
> 本文档仅覆盖 **Stage 3**。Stage 2 请参见 [README_TRAIN_CAUSAL_CONSISTENCY_DISTILL_zh-CN.md](./README_TRAIN_CAUSAL_CONSISTENCY_DISTILL_zh-CN.md),Stage 1 请参见 [README_TRAIN_AR_DIFFUSION_zh-CN.md](./README_TRAIN_AR_DIFFUSION_zh-CN.md)。
---
## 目录
- [一、前置条件](#一前置条件)
- [二、环境配置](#二环境配置)
- [三、下载预训练模型](#三下载预训练模型)
- [四、数据准备](#四数据准备)
- [4.1 快速测试数据集](#41-快速测试数据集)
- [4.2 数据集结构](#42-数据集结构)
- [4.3 metadata.json 格式](#43-metadatajson-格式)
- [五、训练](#五训练)
- [5.1 快速开始](#51-快速开始)
- [5.2 关键参数](#52-关键参数)
- [5.3 DMD 特有参数](#53-dmd-特有参数)
- [六、使用训练好的 Checkpoint](#六使用训练好的-checkpoint)
- [七、更多资源](#七更多资源)
---
## 一、前置条件
Stage 3 需要:
1. **Stage 2 CCD checkpoint**:用于初始化生成器(和判别器)。
2. **Wan2.1-T2V-14B** 模型:作为 DMD real-score teacher。
```bash
# 示例:Stage 2 CCD 训练的 checkpoint
export STAGE2_CKPT="output_dir_wan2.1_causal_forcing_ccd/checkpoint-5000/transformer/diffusion_pytorch_model.safetensors"
```
如何产出该 checkpoint 请参见 [README_TRAIN_CAUSAL_CONSISTENCY_DISTILL_zh-CN.md](./README_TRAIN_CAUSAL_CONSISTENCY_DISTILL_zh-CN.md)。
---
## 二、环境配置
**方式 1:使用 requirements.txt**
```bash
pip install -r requirements.txt
```
**方式 2:手动安装依赖**
```bash
pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image
pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime
pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2"
pip install yunchang xfuser modelscope openpyxl
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**方式 3:使用 Docker**
```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
```
---
## 三、下载预训练模型
DMD 需要**两个**预训练模型:
- **Wan2.1-T2V-1.3B**:生成器/判别器的基础模型。
- **Wan2.1-T2V-14B**:非因果的 real-score teacher,用于计算真实分布得分。
```bash
# 创建模型目录
mkdir -p models/Diffusion_Transformer
# 下载 Wan2.1 T2V 1.3B(student 基础模型)
modelscope download --model Wan-AI/Wan2.1-T2V-1.3B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-1.3B
# 下载 Wan2.1 T2V 14B(DMD real-score teacher)
modelscope download --model Wan-AI/Wan2.1-T2V-14B --local_dir models/Diffusion_Transformer/Wan2.1-T2V-14B
```
Stage 2 CCD checkpoint(生成器/判别器初始化)通过 `--ode_transformer_path` 加载。
---
## 四、数据准备
DMD 使用 **TextDataset**(`train_mode="normal"`)— 只需要提示词,不需要视频数据,因为生成器通过自回归 rollout 自行创建训练样本。但仍需提供包含提示词的 `metadata.json`。
### 4.1 快速测试数据集
```bash
# 下载官方示例数据集
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
### 4.2 数据集结构
与 Stage 1/2 相同。详见 [README_TRAIN_CAUSAL_CONSISTENCY_DISTILL_zh-CN.md](./README_TRAIN_CAUSAL_CONSISTENCY_DISTILL_zh-CN.md#42-数据集结构)。
### 4.3 metadata.json 格式
DMD 仅使用每条记录的 `"text"` 字段。当 `--train_mode="normal"`(TextDataset 模式)时,`file_path` 和其他字段会被忽略。
```json
[
{
"text": "A beautiful sunset over the ocean, golden hour lighting"
},
{
"text": "A person walking through a forest, cinematic view"
}
]
```
> **说明**:你可以复用任何视频 metadata.json — DMD 只关心 `text` 字段。
---
## 五、训练
### 5.1 快速开始
可直接使用的启动脚本为 [train_causal_dmd.sh](./train_causal_dmd.sh):
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export REAL_SCORE_MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-14B"
export DATASET_NAME="datasets/X-Fun-Videos-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json"
export STAGE2_CKPT="output_dir_wan2.1_causal_forcing_ccd/checkpoint-5000/transformer/diffusion_pytorch_model.safetensors"
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_causal_forcing/train_causal_dmd.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--real_score_pretrained_model_name_or_path=$REAL_SCORE_MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--ode_transformer_path=$STAGE2_CKPT \
--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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2.0e-06 \
--learning_rate_critic=2.0e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_dmd" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=0.0 \
--adam_beta1=0.0 \
--adam_beta2=0.999 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=10.0 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--num_frame_per_block=3 \
--use_kv_cache_training \
--denoising_step_indices_list 1000 667 334 1 \
--real_guidance_scale=6.0 \
--randomize_step_indices \
--fake_guidance_scale=0.0 \
--gen_update_interval=5 \
--trainable_modules "."
```
或直接运行:
```bash
bash scripts/wan2.1_causal_forcing/train_causal_dmd.sh
```
### 5.2 关键参数
| 参数 | 说明 | 示例值 |
|------|------|--------|
| `--pretrained_model_name_or_path` | 基础模型(1.3B),用于生成器/判别器初始化 | `models/Diffusion_Transformer/Wan2.1-T2V-1.3B/` |
| `--real_score_pretrained_model_name_or_path` | 14B real-score teacher | `models/Diffusion_Transformer/Wan2.1-T2V-14B` |
| `--config_path` | 模型配置 YAML | `config/wan2.1/wan_civitai.yaml` |
| `--train_data_dir` | 数据根目录 | `""` |
| `--train_data_meta` | `metadata.json` 路径(仅含提示词) | `datasets/X-Fun-Videos-Demo/metadata.json` |
| `--ode_transformer_path` | Stage 2 CCD checkpoint(生成器/判别器初始化) | `$STAGE2_CKPT` |
| `--train_batch_size` | 每 GPU batch 大小 | 1 |
| `--gradient_accumulation_steps` | 梯度累积步数 | 1 |
| `--dataloader_num_workers` | DataLoader 子进程数 | 8 |
| `--num_train_epochs` | 训练 epoch 数 | 100 |
| `--checkpointing_steps` | 每 N 步保存 checkpoint | 200 |
| `--learning_rate` | 生成器学习率 | 2e-06 |
| `--learning_rate_critic` | 判别器学习率 | 2e-06 |
| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` |
| `--lr_warmup_steps` | 学习率预热步数 | 100 |
| `--output_dir` | 输出目录 | `output_dir_wan2.1_causal_forcing_dmd` |
| `--gradient_checkpointing` | 激活重计算 | - |
| `--mixed_precision` | 混合精度:`fp16/bf16` | `bf16` |
| `--adam_weight_decay` | AdamW 权重衰减 | 0.0 |
| `--adam_beta1` | AdamW beta1(DMD 使用 0.0) | 0.0 |
| `--adam_beta2` | AdamW beta2 | 0.999 |
| `--adam_epsilon` | AdamW epsilon | 1e-10 |
| `--max_grad_norm` | 梯度裁剪阈值 | 10.0 |
| `--trainable_modules` | 可训练模块(`"."` = 全部) | `"."` |
| `--low_vram` | 低显存模式(卸载 VAE/文本编码器) | - |
**视频采样参数**:
| 参数 | 说明 | 示例值 |
|------|------|--------|
| `--image_sample_size` | 图像采样尺寸 | 640 |
| `--video_sample_size` | 视频采样尺寸 | 640 |
| `--token_sample_size` | Token 采样尺寸 | 640 |
| `--fix_sample_size` | 固定输出 `[高度, 宽度]` | `480 832` |
| `--video_sample_stride` | 帧采样步幅 | 2 |
| `--video_sample_n_frames` | 视频帧数 | 81 |
| `--random_hw_adapt` | 启用随机分辨率适配 | - |
| `--training_with_video_token_length` | 启用基于 token 长度的训练 | - |
| `--enable_bucket` | 启用宽高比分桶采样 | - |
| `--vae_mini_batch` | VAE 编码迷你批次大小(设为 1 避免 OOM) | 1 |
### 5.3 DMD 特有参数
| 参数 | 说明 | 示例值 |
|------|------|--------|
| `--denoising_step_indices_list` | 去噪步骤索引(DMD 核心参数)。示例为 4 步;解析器默认 `[1000, 500]` = 2 步 | `1000 667 334 1` |
| `--real_guidance_scale` | real-score(14B teacher)的 CFG scale | 6.0 |
| `--randomize_step_indices` | 训练时随机化去噪步骤索引 | - |
| `--fake_guidance_scale` | fake-score(生成器)的 CFG scale。0.0 = 无 CFG | 0.0 |
| `--gen_update_interval` | 生成器更新间隔(每 N 步判别器更新后更新 1 次生成器) | 5 |
| `--num_frame_per_block` | 每个因果块的帧数。`3` = chunkwise(默认),`1` = 逐帧 | 3 |
| `--use_kv_cache_training` | 使用 KV 缓存逐块训练(匹配原始 Self-Forcing) | - |
| `--independent_first_frame` | 第一帧是否独立(`[1, N, N, ...]` 模式,适用于 I2V) | - |
| `--context_noise` | KV 缓存更新的上下文噪声级别 | 0 |
| `--use_teacher_forcing` | 启用 teacher forcing(将 clean_x 传给 transformer) | - |
| `--teacher_forcing_prob` | 每步应用 teacher forcing 的概率(1.0 = 始终) | 1.0 |
| `--train_mode` | 训练模式:`normal`(TextDataset,仅提示词)或 `i2v` | `normal` |
| `--resume_from_checkpoint` | 恢复训练。使用 `"latest"` 自动选择最新 checkpoint | `"latest"` |
---
## 六、使用训练好的 Checkpoint
Stage 3 DMD checkpoint 是 Causal-Forcing 流水线的**最终模型**,直接用于推理:
```python
# 在 examples/wan2.1_causal_forcing/predict_t2v.py 中
transformer_path = "output_dir_wan2.1_causal_forcing_dmd/checkpoint-{N}/diffusion_pytorch_model.safetensors"
# DMD Stage 3 推理配置
guidance_scale = 1.0 # CFG 已烘焦到蒸馏权重中
num_inference_steps = 4 # 4 步 DMD
stochastic_sampling = True
num_frame_per_block = 3 # 分块生成
```
或运行:
```bash
python examples/wan2.1_causal_forcing/predict_t2v.py
```
> **阶段选择器**:`predict_t2v.py` 中包含阶段选择器部分,提供不同阶段的预设配置:
> - **Stage 1 (AR Diffusion)**:`guidance_scale=3.0`、`num_inference_steps=50`、`stochastic_sampling=False`
> - **Stage 2 (CCD)**:`guidance_scale=1.0`、`num_inference_steps=4`、`stochastic_sampling=True`
> - **Stage 3 (DMD)**:`guidance_scale=1.0`、`num_inference_steps=4`、`stochastic_sampling=True`
---
## 七、更多资源
- **Causal-Forcing 论文**:https://github.com/thu-ml/Causal-Forcing
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,51 @@
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
# Causal-Forcing Stage 1: Autoregressive Diffusion Training (teacher forcing).
# - num_frame_per_block=3 -> chunkwise; set to 1 for the framewise variant.
# - shift=5.0 mirrors the Causal-Forcing default scheduler shift.
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_causal_forcing/train_ar_diffusion.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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_ar_diffusion" \
--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 \
--num_frame_per_block=3 \
--train_sampling_steps=1000 \
--shift=5.0 \
--use_timestep_weight \
--trainable_modules "."
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,59 @@
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"
export STAGE1_CKPT="output_dir_wan2.1_causal_forcing_ar_diffusion/checkpoint-8000/diffusion_pytorch_model.safetensors"
# 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
# Causal-Forcing Stage 2: Causal Consistency Distillation (CCD).
# - --transformer_path / --teacher_transformer_path point to the Stage 1 AR-diffusion ckpt.
# - num_frame_per_block=3 -> chunkwise; set to 1 for the framewise variant.
# - discrete_cd_N=48 mirrors the official `causal_cd_chunkwise.yaml`.
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_causal_forcing/train_causal_consistency_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 \
--transformer_path=$STAGE1_CKPT \
--teacher_transformer_path=$STAGE1_CKPT \
--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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2.0e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_ccd" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=0.0 \
--adam_beta1=0.0 \
--adam_beta2=0.999 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=10.0 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--num_frame_per_block=3 \
--shift=5.0 \
--discrete_cd_N=48 \
--guidance_scale=3.0 \
--ema_weight=0.99 \
--ema_start_step=200 \
--trainable_modules "."
File diff suppressed because it is too large Load Diff
+64
View File
@@ -0,0 +1,64 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/"
export REAL_SCORE_MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-14B"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export STAGE2_CKPT="output_dir_wan2.1_causal_forcing_ccd/checkpoint-5000/transformer/diffusion_pytorch_model.safetensors"
# 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
# Causal-Forcing Stage 3: Distribution Matching Distillation (DMD), frame-wise 2-step variant.
# - --real_score_pretrained_model_name_or_path -> Wan2.1-T2V-14B non-causal teacher (DMD real_score).
# - --transformer_path -> Stage 2 CCD ckpt (generator/critic init).
# - num_frame_per_block=3 -> frame-wise; --denoising_step_indices_list 1000 500 -> 2-step DMD.
# - train_mode="normal" (default) -> TextDataset (prompt-only); generation shape from video_sample_* / fix_sample_size.
# - Note: DMD parser has no --shift; the flow scheduler shift is fixed inside the training loop.
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_causal_forcing/train_causal_dmd.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--real_score_pretrained_model_name_or_path=$REAL_SCORE_MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--ode_transformer_path=$STAGE2_CKPT \
--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 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=200 \
--learning_rate=2.0e-06 \
--learning_rate_critic=2.0e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_wan2.1_causal_forcing_dmd" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=0.0 \
--adam_beta1=0.0 \
--adam_beta2=0.999 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=10.0 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--num_frame_per_block=3 \
--use_kv_cache_training \
--denoising_step_indices_list 1000 667 334 1 \
--real_guidance_scale=6.0 \
--randomize_step_indices \
--fake_guidance_scale=0.0 \
--gen_update_interval=5 \
--trainable_modules "."
+21 -11
View File
@@ -76,7 +76,8 @@ from videox_fun.data import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512,
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
CLIPModel, Wan2_2Transformer3DModel,
WanT5EncoderModel)
from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline
from videox_fun.pipeline import (Wan2_2I2VPipeline, Wan2_2Pipeline,
Wan2_2TI2VPipeline)
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions, get_image_latent,
get_image_to_video_latent,
@@ -198,7 +199,16 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac
transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d
if args.train_mode != "normal":
if args.train_mode == "ti2v":
pipeline = Wan2_2TI2VPipeline(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer3d_1,
transformer_2=transformer3d_2,
scheduler=scheduler,
)
elif args.train_mode != "normal":
pipeline = Wan2_2I2VPipeline(
vae=vae,
text_encoder=text_encoder,
@@ -2165,8 +2175,8 @@ def main():
fake_score_main_cond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=fake_score_main_cond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
if args.fake_guidance_scale != 0.0:
@@ -2180,8 +2190,8 @@ def main():
fake_score_main_uncond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=fake_score_main_uncond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
fake_score_main = fake_score_main_uncond + (
fake_score_main_cond - fake_score_main_uncond
@@ -2200,8 +2210,8 @@ def main():
real_score_main_cond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=real_score_main_cond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
real_score_main_uncond = real_score_transformer3d(
@@ -2214,8 +2224,8 @@ def main():
real_score_main_uncond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=real_score_main_uncond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
real_score_main = real_score_main_uncond + (
@@ -2371,7 +2381,7 @@ def main():
if args.low_vram:
fake_score_transformer3d = fake_score_transformer3d.to(accelerator.device)
generator_transformer3d = generator_transformer3d.to(accelerator.device)
# Checks if the accelerator has performed an optimization step behind the scenes
if accelerator.sync_gradients:
+20 -10
View File
@@ -76,7 +76,8 @@ from videox_fun.data import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512,
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
CLIPModel, Wan2_2Transformer3DModel,
WanT5EncoderModel)
from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline
from videox_fun.pipeline import (Wan2_2I2VPipeline, Wan2_2Pipeline,
Wan2_2TI2VPipeline)
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
create_network, merge_lora,
@@ -201,7 +202,16 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, c
transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d
if args.train_mode != "normal":
if args.train_mode == "ti2v":
pipeline = Wan2_2TI2VPipeline(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer3d_1,
transformer_2=transformer3d_2,
scheduler=scheduler,
)
elif args.train_mode != "normal":
pipeline = Wan2_2I2VPipeline(
vae=vae,
text_encoder=text_encoder,
@@ -2210,8 +2220,8 @@ def main():
fake_score_main_cond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=fake_score_main_cond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
if args.fake_guidance_scale != 0.0:
@@ -2225,8 +2235,8 @@ def main():
fake_score_main_uncond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=fake_score_main_uncond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
fake_score_main = fake_score_main_uncond + (
fake_score_main_cond - fake_score_main_uncond
@@ -2245,8 +2255,8 @@ def main():
real_score_main_cond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=real_score_main_cond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
real_score_main_uncond = real_score_transformer3d(
@@ -2259,8 +2269,8 @@ def main():
real_score_main_uncond = convert_flow_pred_to_x0(
scheduler=noise_scheduler,
flow_pred=real_score_main_uncond,
xt=generator_denoised_input,
timestep=generator_timestep
xt=_generator_denoised_input,
timestep=generator_timestep,
)
real_score_main = real_score_main_uncond + (
+10
View File
@@ -182,6 +182,11 @@ class ImageVideoDataset(Dataset):
)
if min_sample_n_frames == 0:
raise ValueError(f"No Frames in video.")
min_video_sample_n_frames = getattr(self, "min_video_sample_n_frames", None)
if min_video_sample_n_frames is not None and min_sample_n_frames < min_video_sample_n_frames:
raise ValueError(
f"Video too short: sampled {min_sample_n_frames} frames < required {min_video_sample_n_frames}."
)
# Select contiguous clip with random start position
video_length = int(self.video_length_drop_end * len(video_reader))
@@ -392,6 +397,11 @@ class ImageVideoControlDataset(Dataset):
)
if min_sample_n_frames == 0:
raise ValueError(f"No Frames in video.")
min_video_sample_n_frames = getattr(self, "min_video_sample_n_frames", None)
if min_video_sample_n_frames is not None and min_sample_n_frames < min_video_sample_n_frames:
raise ValueError(
f"Video too short: sampled {min_sample_n_frames} frames < required {min_video_sample_n_frames}."
)
# Select contiguous clip with random start position
video_length = int(self.video_length_drop_end * len(video_reader))
+169 -12
View File
@@ -13,7 +13,7 @@ from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils.accelerate_utils import apply_forward_hook
from einops import rearrange
from torch.utils.checkpoint import checkpoint as torch_checkpoint
CACHE_T = 2
@@ -549,6 +549,76 @@ class AutoencoderKLWan_(nn.Module):
self.clear_cache()
return x
def encode_stream(self, x, scale=None, is_first_chunk=True):
# Streaming variant of `encode`. Unlike `encode`, it does NOT call
# `clear_cache()` at the start/end, so the causal temporal feature
# cache (`self._enc_feat_map`, mutated in place by the encoder's
# CausalConv3d layers) is threaded across successive pixel chunks. The
# caller is responsible for calling `clear_cache()` exactly once before
# the first chunk and once after the last chunk. Encoding the pixel
# sequence chunk-by-chunk this way is mathematically equivalent to a
# single `encode` call over the whole sequence (no boundary seams).
#
# Temporal alignment (mirrors the VAE's 1,4,4,... compression):
# - is_first_chunk=True : chunk contains the very first pixel frame;
# `t` must be `1 + 4*k` (k>=0) -> yields `1 + k` latent frames.
# - is_first_chunk=False: continuation chunk; `t` must be a multiple
# of 4 -> yields `t // 4` latent frames.
# x: [b,c,t,h,w]
t = x.shape[2]
if scale is not None:
scale = [item.to(x.device, x.dtype) for item in scale]
out = None
if is_first_chunk:
assert (t - 1) % 4 == 0, \
f"first streaming encode chunk needs t=1+4*k frames, got t={t}"
iter_ = 1 + (t - 1) // 4
for i in range(iter_):
self._enc_conv_idx = [0]
if i == 0:
cur = self.encoder(
x[:, :, :1, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
else:
cur = self.encoder(
x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = cur if out is None else torch.cat([out, cur], 2)
else:
assert t % 4 == 0, \
f"continuation streaming encode chunk needs t=4*k frames, got t={t}"
iter_ = t // 4
for i in range(iter_):
self._enc_conv_idx = [0]
cur = self.encoder(
x[:, :, 4 * i:4 * (i + 1), :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = cur if out is None else torch.cat([out, cur], 2)
# conv1 has temporal kernel size 1, so applying it per chunk is
# identical to applying it once over the concatenated sequence.
mu, log_var = self.conv1(out).chunk(2, dim=1)
if scale is not None:
if isinstance(scale[0], torch.Tensor):
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
mu = (mu - scale[0]) * scale[1]
x = torch.cat([mu, log_var], dim=1)
return x
def _decode_frame(self, frame, feat_map_in):
# Decode a single latent frame with a functional (non-mutating) cache
# so it can be wrapped by gradient checkpointing safely. The input
# cache list is copied and never mutated in place, and the updated
# cache is returned explicitly to be threaded to the next frame.
conv_idx = [0]
feat_map = list(feat_map_in)
out = self.decoder(frame, feat_cache=feat_map, feat_idx=conv_idx)
return out, feat_map
def decode(self, z, scale=None):
self.clear_cache()
# z: [b,c,t,h,w]
@@ -561,22 +631,63 @@ class AutoencoderKLWan_(nn.Module):
z = z / scale[1] + scale[0]
iter_ = z.shape[2]
x = self.conv2(z)
# Per-frame gradient checkpointing: each latent frame is an independent
# recompute unit, so backward only materializes one frame's decoder
# activations at a time instead of all frames at once.
use_ckpt = getattr(self, "gradient_checkpointing", False) \
and torch.is_grad_enabled()
feat_map = self._feat_map
outs = []
for i in range(iter_):
self._conv_idx = [0]
if i == 0:
out = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
frame = x[:, :, i:i + 1, :, :]
if use_ckpt:
out_, feat_map = torch_checkpoint(
self._decode_frame, frame, feat_map,
use_reentrant=False)
else:
out_ = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2)
out_, feat_map = self._decode_frame(frame, feat_map)
outs.append(out_)
out = torch.cat(outs, 2)
self.clear_cache()
return out
def decode_stream(self, z, scale=None):
# Streaming variant of `decode`. Unlike `decode`, it does NOT call
# `clear_cache()` at the start/end, so the causal temporal feature
# cache (`self._feat_map`) is threaded across successive chunks. The
# caller is responsible for calling `clear_cache()` exactly once before
# the first chunk and once after the last chunk. Decoding the latent
# sequence chunk-by-chunk this way is mathematically equivalent to a
# single `decode` call over the whole sequence (no boundary seams),
# while enabling "generate a chunk, decode a chunk" streaming.
# z: [b,c,t,h,w]
if scale is not None:
scale = [item.to(z.device, z.dtype) for item in scale]
if isinstance(scale[0], torch.Tensor):
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
1, self.z_dim, 1, 1, 1)
else:
z = z / scale[1] + scale[0]
iter_ = z.shape[2]
x = self.conv2(z)
use_ckpt = getattr(self, "gradient_checkpointing", False) \
and torch.is_grad_enabled()
feat_map = self._feat_map
outs = []
for i in range(iter_):
frame = x[:, :, i:i + 1, :, :]
if use_ckpt:
out_, feat_map = torch_checkpoint(
self._decode_frame, frame, feat_map,
use_reentrant=False)
else:
out_, feat_map = self._decode_frame(frame, feat_map)
outs.append(out_)
# Persist the updated cache so the next chunk continues seamlessly.
self._feat_map = feat_map
out = torch.cat(outs, 2)
return out
def reparameterize(self, mu, log_var):
std = torch.exp(0.5 * log_var)
eps = torch.randn_like(std)
@@ -657,6 +768,7 @@ class AutoencoderKLWan(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.gradient_checkpointing = kwargs["enable"]
else:
raise ValueError("Invalid set gradient checkpointing")
self.model.gradient_checkpointing = self.gradient_checkpointing
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = [
@@ -678,7 +790,30 @@ class AutoencoderKLWan(ModelMixin, ConfigMixin, FromOriginalModelMixin):
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def _encode_stream(self, x, is_first_chunk=True):
# Streaming encode of one pixel chunk over the full batch at once. The
# whole batch must be encoded in a single call so that the persistent
# `self.model._enc_feat_map` stays consistent across chunks (a per-sample
# loop would clobber the cache between samples/chunks).
return self.model.encode_stream(x, self.scale, is_first_chunk=is_first_chunk)
@apply_forward_hook
def encode_stream(
self, x: torch.Tensor, is_first_chunk: bool = True, return_dict: bool = True
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
h = self._encode_stream(x, is_first_chunk=is_first_chunk)
posterior = DiagonalGaussianDistribution(h)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(self, zs):
# Gradient checkpointing is applied per-frame inside self.model.decode
# (see AutoencoderKLWan_.decode), which lowers the backward peak to a
# single latent frame's decoder activations.
self.model.gradient_checkpointing = self.gradient_checkpointing
dec = [
self.model.decode(u.unsqueeze(0), self.scale).clamp_(-1, 1).squeeze(0)
for u in zs
@@ -695,6 +830,28 @@ class AutoencoderKLWan(ModelMixin, ConfigMixin, FromOriginalModelMixin):
return (decoded,)
return DecoderOutput(sample=decoded)
def clear_cache(self):
# Reset the causal temporal feature cache. Must be called once before
# the first streaming chunk and once after the last one.
self.model.clear_cache()
def _decode_stream(self, zs):
# Streaming decode of one chunk over the full batch at once. The whole
# batch must be decoded in a single call so that the persistent
# `self.model._feat_map` stays consistent across chunks (a per-sample
# loop would clobber the cache between samples/chunks).
self.model.gradient_checkpointing = self.gradient_checkpointing
dec = self.model.decode_stream(zs, self.scale).clamp_(-1, 1)
return DecoderOutput(sample=dec)
@apply_forward_hook
def decode_stream(self, z: torch.Tensor, return_dict: bool = True) -> Union[DecoderOutput, torch.Tensor]:
decoded = self._decode_stream(z).sample
if not return_dict:
return (decoded,)
return DecoderOutput(sample=decoded)
@classmethod
def from_pretrained(cls, pretrained_model_path, additional_kwargs={}):
def filter_kwargs(cls, kwargs):
+40 -22
View File
@@ -14,7 +14,7 @@ from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils.accelerate_utils import apply_forward_hook
from einops import rearrange
from torch.utils.checkpoint import checkpoint as torch_checkpoint
CACHE_T = 2
@@ -817,6 +817,17 @@ class AutoencoderKLWan2_2_(nn.Module):
self.clear_cache()
return x
def _decode_frame(self, frame, feat_map_in, first_chunk=False):
# Decode a single latent frame with a functional (non-mutating) cache
# so it can be wrapped by gradient checkpointing safely. The input
# cache list is copied and never mutated in place, and the updated
# cache is returned explicitly to be threaded to the next frame.
conv_idx = [0]
feat_map = list(feat_map_in)
out = self.decoder(frame, feat_cache=feat_map, feat_idx=conv_idx,
first_chunk=first_chunk)
return out, feat_map
def decode(self, z, scale):
self.clear_cache()
# z: [b,c,t,h,w]
@@ -828,22 +839,24 @@ class AutoencoderKLWan2_2_(nn.Module):
z = z / scale[1] + scale[0]
iter_ = z.shape[2]
x = self.conv2(z)
# Per-frame gradient checkpointing: each latent frame is an independent
# recompute unit, so backward only materializes one frame's decoder
# activations at a time instead of all frames at once.
use_ckpt = getattr(self, "gradient_checkpointing", False) \
and torch.is_grad_enabled()
feat_map = self._feat_map
outs = []
for i in range(iter_):
self._conv_idx = [0]
if i == 0:
out = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx,
first_chunk=True,
)
frame = x[:, :, i:i + 1, :, :]
first_chunk = (i == 0)
if use_ckpt:
out_, feat_map = torch_checkpoint(
self._decode_frame, frame, feat_map, first_chunk,
use_reentrant=False)
else:
out_ = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx,
)
out = torch.cat([out, out_], 2)
out_, feat_map = self._decode_frame(frame, feat_map, first_chunk)
outs.append(out_)
out = torch.cat(outs, 2)
out = unpatchify(out, patch_size=2)
self.clear_cache()
return out
@@ -901,7 +914,7 @@ class AutoencoderKLWan3_8(ModelMixin, ConfigMixin, FromOriginalModelMixin):
dim_mult=[1, 2, 4, 4],
temperal_downsample=[False, True, True],
temporal_compression_ratio=4,
spatial_compression_ratio=8
spatial_compression_ratio=16
):
super().__init__()
mean = torch.tensor(
@@ -1012,12 +1025,12 @@ class AutoencoderKLWan3_8(ModelMixin, ConfigMixin, FromOriginalModelMixin):
# init model
self.model = _video_vae(
pretrained_path=vae_pth,
z_dim=latent_channels,
dim=c_dim,
dim_mult=dim_mult,
temperal_downsample=temperal_downsample,
).eval().requires_grad_(False)
pretrained_path=vae_pth,
z_dim=latent_channels,
dim=c_dim,
dim_mult=dim_mult,
temperal_downsample=temperal_downsample,
).eval().requires_grad_(False)
self.gradient_checkpointing = False
@@ -1028,6 +1041,7 @@ class AutoencoderKLWan3_8(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.gradient_checkpointing = kwargs["enable"]
else:
raise ValueError("Invalid set gradient checkpointing")
self.model.gradient_checkpointing = self.gradient_checkpointing
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = [
@@ -1050,6 +1064,10 @@ class AutoencoderKLWan3_8(ModelMixin, ConfigMixin, FromOriginalModelMixin):
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(self, zs):
# Gradient checkpointing is applied per-frame inside self.model.decode
# (see AutoencoderKLWan2_2_.decode), which lowers the backward peak to a
# single latent frame's decoder activations.
self.model.gradient_checkpointing = self.gradient_checkpointing
dec = [
self.model.decode(u.unsqueeze(0), self.scale).clamp_(-1, 1).squeeze(0)
for u in zs
+1 -31
View File
@@ -43,34 +43,4 @@ WanFunPipeline = WanPipeline
WanI2VPipeline = WanFunInpaintPipeline
Wan2_2FunPipeline = Wan2_2Pipeline
Wan2_2I2VPipeline = Wan2_2FunInpaintPipeline
import importlib.util
if importlib.util.find_spec("paifuser") is not None:
# --------------------------------------------------------------- #
# Sparse Attention
# --------------------------------------------------------------- #
from paifuser.ops import sparse_reset
# Wan2.1
WanFunInpaintPipeline.__call__ = sparse_reset(WanFunInpaintPipeline.__call__)
WanFunPipeline.__call__ = sparse_reset(WanFunPipeline.__call__)
WanFunControlPipeline.__call__ = sparse_reset(WanFunControlPipeline.__call__)
WanI2VPipeline.__call__ = sparse_reset(WanI2VPipeline.__call__)
WanPipeline.__call__ = sparse_reset(WanPipeline.__call__)
WanVacePipeline.__call__ = sparse_reset(WanVacePipeline.__call__)
# Phantom
WanFunPhantomPipeline.__call__ = sparse_reset(WanFunPhantomPipeline.__call__)
# Wan2.2
Wan2_2FunInpaintPipeline.__call__ = sparse_reset(Wan2_2FunInpaintPipeline.__call__)
Wan2_2FunPipeline.__call__ = sparse_reset(Wan2_2FunPipeline.__call__)
Wan2_2FunControlPipeline.__call__ = sparse_reset(Wan2_2FunControlPipeline.__call__)
Wan2_2Pipeline.__call__ = sparse_reset(Wan2_2Pipeline.__call__)
Wan2_2I2VPipeline.__call__ = sparse_reset(Wan2_2I2VPipeline.__call__)
Wan2_2TI2VPipeline.__call__ = sparse_reset(Wan2_2TI2VPipeline.__call__)
Wan2_2S2VPipeline.__call__ = sparse_reset(Wan2_2S2VPipeline.__call__)
Wan2_2VaceFunPipeline.__call__ = sparse_reset(Wan2_2VaceFunPipeline.__call__)
Wan2_2AnimatePipeline.__call__ = sparse_reset(Wan2_2AnimatePipeline.__call__)
Wan2_2I2VPipeline = Wan2_2FunInpaintPipeline
@@ -23,9 +23,17 @@ 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."""
"""Official FlashHead timestep schedule with shift transform.
Hardcoded schedules match the CF reference for distilled few-step inference:
- 4-step: [1000, 750, 500, 250] — Stage 2 CCD (`causal_cd_framewise.yaml`)
- 2-step: [1000, 500] — Stage 3 DMD (`causal_forcing_dmd_framewise_2step.yaml`)
Anything else falls back to linspace.
"""
if num_inference_steps == 4:
timesteps = [1000, 750, 500, 250]
elif num_inference_steps == 2:
timesteps = [1000, 500]
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)
@@ -312,6 +320,18 @@ class WanSelfForcingPipeline(DiffusionPipeline):
frames = frames.cpu().float().numpy()
return frames
def decode_latents_stream(self, latents: torch.Tensor) -> torch.Tensor:
# Decode a single chunk of latent frames while preserving the VAE's
# causal temporal feature cache across chunks (see
# AutoencoderKLWan.decode_stream). `self.vae.clear_cache()` must be
# called once before the first chunk and once after the last one.
# Returns a CPU float tensor of shape [B, C, F_pixels, H, W] in [0, 1],
# matching the layout produced by `decode_latents`.
frames = self.vae.decode_stream(latents.to(self.vae.dtype)).sample
frames = (frames / 2 + 0.5).clamp(0, 1)
frames = frames.cpu().float()
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
@@ -490,6 +510,8 @@ class WanSelfForcingPipeline(DiffusionPipeline):
independent_first_frame: bool = True,
context_noise: int = 0,
stochastic_sampling: bool = True,
streaming: bool = False,
decode_callback: Optional[Callable[[torch.Tensor, int], None]] = None,
) -> Union[WanSelfForcingPipelineOutput, Tuple]:
r"""
Function invoked when calling the pipeline for Self-Forcing causal generation.
@@ -502,6 +524,16 @@ class WanSelfForcingPipeline(DiffusionPipeline):
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).
streaming: If True, decode each causal block into pixels right after it is
generated (instead of accumulating all latents and decoding once at the end).
The VAE's causal temporal cache is threaded across blocks so the result is
seam-free and identical to a single full decode, while lowering peak decode
memory and enabling incremental output.
decode_callback: Optional callable invoked as `decode_callback(video_chunk, block_idx)`
after each block is decoded in streaming mode, where `video_chunk` is a CPU float
tensor of shape [B, C, F_pixels, H, W] in [0, 1]. When provided, chunks are NOT
accumulated in memory and the returned `videos` is an empty tensor (the caller is
responsible for consuming/saving each chunk). Only used when `streaming` is True.
Examples:
```python
@@ -684,6 +716,12 @@ class WanSelfForcingPipeline(DiffusionPipeline):
# First frame is generated independently (standard Self-Forcing T2V pattern)
all_num_frames = [1] + all_num_frames
# Streaming decode state: reset the VAE causal cache once before the
# first block so the per-block decodes can be threaded seamlessly.
streamed_videos = []
if streaming:
self.vae.clear_cache()
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
@@ -778,25 +816,36 @@ class WanSelfForcingPipeline(DiffusionPipeline):
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
# Delegate to the scheduler's own step() so each sampler
# (Flow Euler, UniPC, DPM++) applies its correct
# multi-step formula. The previous "predict-x0 + resample
# fresh noise" hybrid did NOT match CF's UniPC reference
# (CF's CausalDiffusionInferencePipeline calls
# sample_scheduler.step(...)). For framewise AR rollouts
# this mismatch caused subtle drift on top of CF's own
# output for the same checkpoint.
t = denoise_timesteps[step_idx]
step_out = self.scheduler.step(
flow_pred, t, noisy_input, return_dict=False
)
noisy_input = step_out[0]
denoised_pred = noisy_input
progress_bar.update()
# Update output with denoised block
output[:, :, cache_start_frame:cache_start_frame + current_num_frames] = denoised_pred
# Streaming decode: turn this block's latents into pixels right away.
# The VAE cache is preserved across blocks, so this is seam-free and
# matches a single full decode of `output`.
if streaming:
video_chunk = self.decode_latents_stream(denoised_pred)
if decode_callback is not None:
decode_callback(video_chunk, block_idx)
else:
streamed_videos.append(video_chunk)
# 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:
@@ -850,7 +899,16 @@ class WanSelfForcingPipeline(DiffusionPipeline):
# 9. Decode output
if output_type == "pil":
if streaming:
# Close the VAE causal cache opened before the block loop.
self.vae.clear_cache()
if decode_callback is not None:
# Frames were streamed out via the callback; nothing to return.
video = output.new_zeros(0)
else:
# Concatenate the per-block pixel chunks along the frame axis.
video = torch.cat(streamed_videos, dim=2)
elif output_type == "pil":
video = self.decode_latents(output)
video = torch.from_numpy(video)
else:
+2 -1
View File
@@ -16,4 +16,5 @@ from .trigflow_sampler import (RectifiedFlow_TrigFlowWrapper,
from .utils import (calculate_dimensions, filter_kwargs, get_autocast_dtype,
get_image_latent, get_image_to_video_latent,
get_video_to_video_latent, save_videos_grid,
save_videos_with_audio_grid)
save_videos_with_audio_grid, StreamVideoSaver,
SegmentVideoSaver)
+67
View File
@@ -84,6 +84,73 @@ def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, f
outputs[0].save(path, format='GIF', append_images=outputs, save_all=True, duration=100, loop=0)
print(f"Saved video to: {path}")
class StreamVideoSaver:
"""Incrementally write video frames to an mp4 as each block is decoded.
Used as a streaming pipeline `decode_callback`: it receives one pixel chunk
per causal block ([B, C, F, H, W] float in [0, 1]) and appends its frames to
an open imageio writer, so a long video never has to be held in memory. The
file is only finalized on `close()`.
"""
def __init__(self, video_path, fps):
os.makedirs(os.path.dirname(video_path), exist_ok=True)
self.video_path = video_path
self.writer = imageio.get_writer(video_path, fps=fps)
self.num_frames = 0
def __call__(self, video_chunk, block_idx):
# video_chunk: [B, C, F, H, W] in [0, 1]; take the first sample.
chunk = video_chunk[0].permute(1, 2, 3, 0) # [F, H, W, C]
chunk = (chunk.clamp(0, 1) * 255).numpy().astype(np.uint8)
for frame in chunk:
self.writer.append_data(frame)
self.num_frames += 1
def close(self):
self.writer.close()
print(f"Saved video to: {self.video_path} ({self.num_frames} frames)")
class SegmentVideoSaver:
"""Save each decoded block as its own standalone mp4, flushed immediately,
while also assembling one complete continuous mp4.
Unlike `StreamVideoSaver` (only one continuous file finalized on close),
every block is written to a separate, fully-closed mp4 the moment it is
decoded (so partial results survive an interruption), AND its frames are
appended to a single `full.mp4` so a ready-to-play complete video is also
produced. Segments are named by block index.
"""
def __init__(self, out_dir, fps, full_name="full.mp4"):
os.makedirs(out_dir, exist_ok=True)
self.out_dir = out_dir
self.fps = fps
self.segments = []
# Continuous writer for the assembled full video.
self.full_path = os.path.join(out_dir, full_name)
self.full_writer = imageio.get_writer(self.full_path, fps=fps)
self.num_frames = 0
def __call__(self, video_chunk, block_idx):
# video_chunk: [B, C, F, H, W] in [0, 1]; take the first sample.
chunk = video_chunk[0].permute(1, 2, 3, 0) # [F, H, W, C]
chunk = (chunk.clamp(0, 1) * 255).numpy().astype(np.uint8)
# 1) Standalone, fully-flushed segment file.
seg_path = os.path.join(self.out_dir, f"seg_{block_idx:04d}.mp4")
with imageio.get_writer(seg_path, fps=self.fps) as writer:
for frame in chunk:
writer.append_data(frame)
self.segments.append(seg_path)
# 2) Append the same frames to the continuous full video.
for frame in chunk:
self.full_writer.append_data(frame)
self.num_frames += 1
print(f"Saved segment to: {seg_path} ({len(chunk)} frames)")
def close(self):
self.full_writer.close()
print(f"Saved {len(self.segments)} segments to: {self.out_dir}")
print(f"Saved full video to: {self.full_path} ({self.num_frames} frames)")
def save_videos_with_audio_grid(
videos: torch.Tensor,
audio: torch.Tensor,