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