Add Causal-Focing for Wan2.1 (#500)
This commit is contained in:
@@ -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
@@ -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 "."
|
||||
@@ -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:
|
||||
|
||||
|
||||
@@ -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 + (
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user