Compare commits

...
Author SHA1 Message Date
JerryZhou54 38f9c41e46 refactor hy15 sf distill 2026-01-26 20:40:11 +00:00
JerryZhou54 b1f88ad5d5 refactor hy15 causal denoising 2026-01-26 15:54:20 +00:00
JerryZhou54 90a598bd9e fix lint 2026-01-26 03:09:37 +00:00
JerryZhou54 90d86d5a79 compatible with wan sf 2026-01-26 01:25:34 +00:00
JerryZhou54 7d373cd2c4 Add context forcing, ode_init to sf training 2026-01-26 01:01:47 +00:00
JerryZhou54 0806218156 small change 2026-01-26 01:01:47 +00:00
JerryZhou54 8c6056fbe2 ckpt 2026-01-26 01:01:44 +00:00
JerryZhou54 e4705349d0 Ode init running for hy15 2026-01-26 00:57:22 +00:00
JerryZhou54 50e63840c6 Ode runnable for hy15 2026-01-26 00:57:18 +00:00
JerryZhou54 f1ec0cde18 Add support for ode_init inference for hy15 & support multiple timesteps for hy15 2026-01-26 00:51:50 +00:00
45 changed files with 3943 additions and 408 deletions
@@ -0,0 +1,169 @@
#!/bin/bash
#SBATCH --job-name=dmd_3333
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_3333_output/dmd_3333_%j.out
#SBATCH --error=dmd_3333_output/dmd_3333_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv-distill
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation with Hunyuan1.5:
GENERATOR_MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v" # Teacher model
FAKE_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
# REAL_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled"
# FAKE_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
# DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
DATA_DIR="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/vidprom_16k_text_embed"
DATA_DIR_2="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/fv-ode-preprocessing-16k-hy15-121"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFhy1.5_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_3333_1e-5_1e-5_cfg6_corrected_scheduler"
--wandb_run_name "self_forcing_3333_context_forcing"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 848
--enable_gradient_checkpointing_type "full"
--simulate_generator_forward
--log_visualization
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
# --resume-from-checkpoint "/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_1333/checkpoint-300"
--init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_worldplay_df_init_3333/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_distilled_ode_init_3333/checkpoint-2400/transformer/diffusion_pytorch_model.safetensors
)
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
# --data_path_2 "$DATA_DIR_2"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "8"
--validation_guidance_scale "1.0" # not used for dmd inference
--text-encoder-cpu-offload
# --vae_cpu_offload True
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 100
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
)
dmd_args=(
--warp_denoising_step
--dmd_denoising_steps '1000,875,750,625,500,375,250,125'
# --dmd_denoising_steps '1000,760,520,280'
# --dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.5
--fake_score_learning_rate 5e-6
--fake_score_betas '0.0,0.999'
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
# --use-context-forcing True
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/hy15_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
+24 -4
View File
@@ -2,14 +2,14 @@ from fastvideo import VideoGenerator
import json
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_hy15"
OUTPUT_PATH = "video_samples_hy15_t2v_distilled"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
@@ -18,15 +18,35 @@ def main():
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
# init_weights_from_safetensors="/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_1333_5e-5_2e-6_cfg3.5/checkpoint-500/ema/generator_ema.safetensors"
)
# json_path = "/mnt/weka/home/hao.zhang/wei/FastVideo/data/mixkit_i2v_full_720p.json"
# with open(json_path, 'r') as f:
# data_list = json.load(f)["data"]
# # Now you can index into data_list however you like
# # For example: data_list[0], data_list[1:3], etc.
# for data in data_list:
# prompt = data["prompt"]
# image_path = data["image_path"]
# generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path, num_frames=121, fps=24)
# return
# prompt = (
# "In a brightly lit studio, a photographer wearing a denim jacket focuses intently, capturing shots with a professional camera. Facing him, a model stands gracefully, adjusting her long, flowing hair with delicate movements. The scene is characterized by strong contrasts; the model's soft pink attire and gentle gestures complement the rugged, precise demeanor of the photographer. Positioned against a minimalist backdrop, the pair work seamlessly, with the camera\u2019s lens pointed directly at the model, capturing her elegance. The soft, diffused lighting casts a gentle glow on both subjects, creating an airy and ethereal atmosphere perfect for a high-fashion photo shoot."
# )
# video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=121, fps=24, image_path="data/1.png")
# return
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
@@ -35,7 +55,7 @@ def main():
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
@@ -1,36 +1,38 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_MODE=offline
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
DATA_DIR="data/ode-preprocessing-hy15-test/"
VALIDATION_DATASET_FILE="data/validation_64.json"
NUM_GPUS=1
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--output_dir "ode_init_hy15_test"
--wandb_run_name "vidprom_bz128_1e-5"
# --resume_from_checkpoint "ode_init_diffusers/"
# --warp_denoising_step
# --log_visualization
--max_train_steps 6001
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_latent_t 19
--num_height 480
--num_width 832
--num_frames 77
--warp_denoising_step
--num_frames 73
--dmd_denoising_steps "1000,750,500,250"
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--num_gpus 1
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
@@ -51,18 +53,17 @@ dataset_args=(
# Validation arguments
validation_args=(
--log_validation
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
@@ -78,6 +79,7 @@ miscellaneous_args=(
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
# --enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
@@ -0,0 +1,132 @@
#!/bin/bash
#SBATCH --job-name=hy15_ode_3333
#SBATCH --partition=main
#SBATCH --nodes=8
#SBATCH --ntasks=8
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=hy15_ode_3333_output/hy15_ode_3333_%j.out
#SBATCH --error=hy15_ode_3333_output/hy15_ode_3333_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv-distill
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v"
DATA_DIR="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/fv-ode-preprocessing-16k-hy15-121"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "hy15_ode_init"
--output_dir "/mnt/weka/home/hao.zhang/wei/hy15_ode_init_1333_new"
--wandb_run_name "hy15_ode_init_1333_new"
--warp_denoising_step
--dmd_denoising_steps '1000,760,520,280,0'
--log_visualization
--visualization_steps 100
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 31
--num_height 480
--num_width 848
--num_frames 121
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--text-encoder-cpu-offload
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 400
--weight_decay 0.01
--max_grad_norm 1.0
)
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 10
--validation_sampling_steps "4"
--validation_guidance_scale "1.0" # not used for dmd inference
# --vae_cpu_offload True
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 6
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
# --enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -26,6 +26,9 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
param_names_mapping: dict = field(
default_factory=lambda: {
r"^cond_type_embed\.(.*)$":
r"cond_type_embed.\1",
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
@@ -55,6 +58,16 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
r"^image_embedder\.linear_1\.(.*)$":
r"image_embedder.linear_1.\1",
r"^image_embedder\.linear_2\.(.*)$":
r"image_embedder.linear_2.\1",
r"^image_embedder\.norm_in\.(.*)$":
r"image_embedder.norm_in.\1",
r"^image_embedder\.norm_out\.(.*)$":
r"image_embedder.norm_out.\1",
# 2. txt_in_2 mapping:
r"^context_embedder_2\.(.*)$":
@@ -144,6 +157,12 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
# Causal HunyuanVideo1.5
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 31
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
+17 -1
View File
@@ -127,13 +127,29 @@ class Hunyuan15T2V480PConfig(PipelineConfig):
vae_tiling: bool = True
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Hunyuan15DistilledI2V480PConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 7
@dataclass
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 9
@dataclass
class SelfForcingHunyuan15T2V480PConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 5
is_causal: bool = True
# dmd_denoising_steps: list[int] | None = field(
# default_factory=lambda: [1000, 750, 500, 250])
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 875, 750, 625, 500, 375, 250, 125])
warp_denoising_step: bool = True
+5 -1
View File
@@ -8,9 +8,9 @@ from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig, SelfForcingHunyuan15T2V480PConfig, Hunyuan15DistilledI2V480PConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.turbodiffusion import (
@@ -35,6 +35,10 @@ logger = init_logger(__name__)
PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v":
SelfForcingHunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled":
Hunyuan15DistilledI2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
+10 -2
View File
@@ -17,8 +17,8 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
# sigmas: list[float] | None = field(
# default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
@@ -28,6 +28,14 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
@dataclass
class Hunyuan15_480P_Distilled_SamplingParam(Hunyuan15_480P_SamplingParam):
num_inference_steps: int = 8
fps: int = 24
guidance_scale: float = 1.0
@dataclass
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
height: int = 720
+5 -1
View File
@@ -5,8 +5,8 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam, Hunyuan15_480P_Distilled_SamplingParam
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
@@ -44,6 +44,10 @@ logger = init_logger(__name__)
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled":
Hunyuan15_480P_Distilled_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
+24 -2
View File
@@ -127,7 +127,10 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
def ode_text_only_record_creator(
video_name: str, text_embedding: np.ndarray, caption: str,
trajectory_latents: np.ndarray,
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
trajectory_timesteps: np.ndarray,
text_mask: np.ndarray | None = None,
text_embedding_2: np.ndarray | None = None,
text_mask_2: np.ndarray | None = None) -> dict[str, Any]:
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
Args:
@@ -165,6 +168,25 @@ def ode_text_only_record_creator(
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
if text_embedding_2 is not None:
record.update({
"text_embedding_2_bytes": text_embedding_2.tobytes(),
"text_embedding_2_shape": list(text_embedding_2.shape),
"text_embedding_2_dtype": str(text_embedding_2.dtype),
})
if text_mask is not None:
record.update({
"text_mask_bytes": text_mask.tobytes(),
"text_mask_shape": list(text_mask.shape),
"text_mask_dtype": str(text_mask.dtype),
})
if text_mask_2 is not None:
record.update({
"text_mask_2_bytes": text_mask_2.tobytes(),
"text_mask_2_shape": list(text_mask_2.shape),
"text_mask_2_dtype": str(text_mask_2.dtype),
})
return record
@@ -187,4 +209,4 @@ def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
"text_embedding_dtype": str(text_embedding.dtype),
"caption": caption,
}
return record
return record
+10 -1
View File
@@ -90,6 +90,15 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_embedding_2_bytes", pa.binary()),
pa.field("text_embedding_2_shape", pa.list_(pa.int64())),
pa.field("text_embedding_2_dtype", pa.string()),
pa.field("text_mask_bytes", pa.binary()),
pa.field("text_mask_shape", pa.list_(pa.int64())),
pa.field("text_mask_dtype", pa.string()),
pa.field("text_mask_2_bytes", pa.binary()),
pa.field("text_mask_2_shape", pa.list_(pa.int64())),
pa.field("text_mask_2_dtype", pa.string()),
# --- ODE Trajectory ---
pa.field("trajectory_latents_bytes", pa.binary()),
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
@@ -115,4 +124,4 @@ pyarrow_schema_text_only = pa.schema([
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("caption", pa.string()),
])
])
+3 -2
View File
@@ -17,6 +17,7 @@ from PIL import Image
from transformers import AutoTokenizer
from fastvideo.logger import init_logger
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -655,7 +656,7 @@ class TextDataset(torch.utils.data.IterableDataset,
self.seed = seed
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer_path = os.path.join(maybe_download_model(args.model_path), "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
@@ -758,4 +759,4 @@ class TextDataset(torch.utils.data.IterableDataset,
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
self.processed_batches = state_dict["processed_batches"]
+9 -3
View File
@@ -154,8 +154,14 @@ def collate_rows_from_parquet_schema(rows,
) if rng else random.random()) < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
else:
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
if row[f"{tensor_name}_dtype"] == "float32":
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
elif row[f"{tensor_name}_dtype"] == "int64":
data = np.frombuffer(
bytes_data, dtype=np.int64).reshape(shape).copy()
else:
raise ValueError(f"Unsupported dtype: {row[f"{tensor_name}_dtype"]}")
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
@@ -168,7 +174,7 @@ def collate_rows_from_parquet_schema(rows,
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_name == 'text_embedding':
if tensor_name == 'null':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
+22
View File
@@ -830,6 +830,7 @@ class TrainingArgs(FastVideoArgs):
precedence.
"""
data_path: str = ""
data_path_2: str | None = None
dataloader_num_workers: int = 0
num_height: int = 0
num_width: int = 0
@@ -859,6 +860,7 @@ class TrainingArgs(FastVideoArgs):
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
visualization_steps: float = 0.0
log_validation: bool = False
trackers: list[str] = dataclasses.field(default_factory=list)
tracker_project_name: str = ""
@@ -874,6 +876,7 @@ class TrainingArgs(FastVideoArgs):
num_train_epochs: int = 0
max_train_steps: int = 0
gradient_accumulation_steps: int = 0
optimizer_type: str = "adamw"
learning_rate: float = 0.0
scale_lr: bool = False
lr_scheduler: str = "constant"
@@ -943,6 +946,8 @@ class TrainingArgs(FastVideoArgs):
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
use_context_forcing: bool = False
use_ode_init: bool = False
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -998,6 +1003,10 @@ class TrainingArgs(FastVideoArgs):
type=str,
required=True,
help="Path to parquet files")
parser.add_argument("--data-path-2",
type=str,
required=False,
help="Path to parquet files")
parser.add_argument("--dataloader-num-workers",
type=int,
required=True,
@@ -1091,6 +1100,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--validation-steps",
type=float,
help="Number of validation steps")
parser.add_argument("--visualization-steps",
type=float,
help="Number of visualization steps")
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
@@ -1139,6 +1151,10 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--gradient-accumulation-steps",
type=int,
help="Number of steps to accumulate gradients")
parser.add_argument("--optimizer-type",
type=str,
choices=["adamw", "muon"],
help="Optimizer type")
parser.add_argument("--learning-rate",
type=float,
required=True,
@@ -1365,6 +1381,12 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
parser.add_argument("--use-context-forcing",
action=StoreBoolean,
help="Whether to use context forcing")
parser.add_argument("--use-ode-init",
action=StoreBoolean,
help="Whether to use ODE init")
return parser
@@ -0,0 +1,808 @@
# Copyright 2025 The Hunyuan 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.
from typing import Any, Dict, Optional, List
import math
import torch
# import torch._dynamo
# torch._dynamo.config.cache_size_limit = 128
# try:
# torch._dynamo.config.recompile_limit = 128
# except AttributeError:
# pass
import torch.nn as nn
import torch.nn.functional as F
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
from torch.nn.attention.flex_attention import BlockMask
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="default")
import torch.distributed as dist
from fastvideo.attention import LocalAttention
from fastvideo.forward_context import set_forward_context
from fastvideo.configs.models.dits import HunyuanVideo15Config
from fastvideo.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
# TODO(will-PY-refactor): RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed, _apply_rotary_emb
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
unpatchify)
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.models.dits.hunyuanvideo15 import (
HunyuanRMSNorm,
HunyuanVideo15TimeEmbedding,
HunyuanVideo15ByT5TextProjection,
HunyuanVideo15ImageProjection,
SingleTokenRefiner,
FinalLayer)
logger = init_logger(__name__)
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
local_attn_size: int = -1,
sink_size: int = 0,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
self.head_dim = hidden_size // num_attention_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
self.max_attention_size = 21 * 1590 if local_attn_size == -1 else local_attn_size * 1590
self.attn = LocalAttention(
num_heads=self.num_attention_heads,
head_size=self.head_dim,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn"
)
def forward_txt(
self,
txt: torch.Tensor,
vec: torch.Tensor,
cache_txt: bool = False,
):
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
t_kv = {}
if cache_txt:
t_kv["k_txt"] = txt_k
t_kv["v_txt"] = txt_v
txt_attn = self.attn(txt_q, txt_k, txt_v)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return txt, t_kv
def forward_vision(
self,
img: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple,
block_mask: BlockMask,
kv_cache: dict | None = None,
txt_kv_cache: list | None = None,
current_start: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
if vec.dim() == 3:
img_mod_outputs = self.img_mod(vec).unflatten(dim=-1, sizes=(6, -1))
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=2)
else:
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings
cos, sin = freqs_cis
img_q = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False)
img_k = _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
# Apply flex_attention
# Does not support SP padding for now
if kv_cache is None:
q = img_q
k = torch.cat([img_k, txt_kv_cache["k_txt"]], dim=1)
v = torch.cat([img_v, txt_kv_cache["v_txt"]], dim=1)
# Padding for flex attention
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_kv_length = math.ceil(k.shape[1] / 128) * 128 - k.shape[1]
padded_roped_query = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_roped_key = torch.cat(
[k, torch.zeros([k.shape[0], padded_kv_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_kv_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
img_attn = flex_attention(
query=padded_roped_query.transpose(2, 1),
key=padded_roped_key.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask
)[:, :, :-padded_length].transpose(2, 1)
assert img_attn.shape[1] == image_seq_len
updated_kv_cache = None
else:
current_end = current_start + img_q.shape[1]
num_new_tokens = img_q.shape[1]
sink_tokens = self.sink_size * 1590
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = self.max_attention_size
# Clone cache to avoid in-place modification during gradient checkpointing
k_cache = kv_cache["k"].clone()
v_cache = kv_cache["v"].clone()
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):
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_size - num_new_tokens - sink_tokens
k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
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
assert local_end_index == self.max_attention_size
else:
# Assign new keys/values directly up to current_end
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
assert local_start_index >= 0
q = img_q
k = torch.cat([k_cache[:, :local_start_index], img_k, txt_kv_cache["k_txt"]], dim=1)
v = torch.cat([v_cache[:, :local_start_index], img_v, txt_kv_cache["v_txt"]], dim=1)
img_attn = self.attn(q, k, v)
k_cache[:, local_start_index:local_end_index] = img_k
v_cache[:, local_start_index:local_end_index] = img_v
updated_kv_cache = {
"k": k_cache,
"v": v_cache,
"global_end_index": torch.tensor([current_end], dtype=torch.long, device=k_cache.device),
"local_end_index": torch.tensor([local_end_index], dtype=torch.long, device=k_cache.device)
}
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
return img, updated_kv_cache
def forward(
self,
txt_inference=False,
vision_inference=False,
**kwargs
):
if txt_inference:
return self.forward_txt(**kwargs)
elif vision_inference:
return self.forward_vision(**kwargs)
else:
raise ValueError("txt_inference and vision_inference cannot be both False")
class CausalHunyuanVideo15Transformer3DModel(CachableDiT):
r"""
A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5).
"""
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = HunyuanVideo15Config()._fsdp_shard_conditions
_compile_conditions = HunyuanVideo15Config()._compile_conditions
_supported_attention_backends = HunyuanVideo15Config(
)._supported_attention_backends
param_names_mapping = HunyuanVideo15Config().param_names_mapping
reverse_param_names_mapping = HunyuanVideo15Config(
).reverse_param_names_mapping
lora_param_names_mapping = HunyuanVideo15Config().lora_param_names_mapping
def __init__(
self,
config: HunyuanVideo15Config,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
self.out_channels = config.out_channels or config.in_channels
self.patch_size = (config.patch_size_t, config.patch_size, config.patch_size)
# 1. Latent and condition embedders
self.img_in = PatchEmbed(self.patch_size,
config.in_channels,
self.hidden_size,
prefix=f"{config.prefix}.img_in")
self.image_embedder = HunyuanVideo15ImageProjection(config.image_embed_dim, self.hidden_size)
self.txt_in = SingleTokenRefiner(config.text_embed_dim,
self.hidden_size,
config.num_attention_heads,
depth=config.num_refiner_layers,
dtype=None,
prefix=f"{config.prefix}.txt_in")
self.txt_in_2 = HunyuanVideo15ByT5TextProjection(config.text_embed_2_dim, 2048, self.hidden_size)
self.time_in = HunyuanVideo15TimeEmbedding(self.hidden_size, use_meanflow=config.use_meanflow)
self.cond_type_embed = nn.Embedding(3, self.hidden_size)
# 3. Dual stream transformer blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
local_attn_size=config.local_attn_size,
sink_size=config.sink_size,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.num_layers)
]
)
# 5. Output projection
self.final_layer = FinalLayer(self.hidden_size,
self.patch_size,
self.out_channels,
prefix=f"{config.prefix}.final_layer")
self.gradient_checkpointing = False
self.num_frame_per_block = config.num_frames_per_block
self.local_attn_size = config.local_attn_size
self.block_mask = None
self.__post_init__()
def get_text_and_mask(
self,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_2: torch.Tensor,
encoder_attention_mask: torch.Tensor,
encoder_attention_mask_2: torch.Tensor,
encoder_hidden_states_image: torch.Tensor,
timestep: torch.Tensor,
):
batch_size, txt_seq_len = encoder_hidden_states.shape[0], encoder_hidden_states.shape[1]
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
assert encoder_hidden_states.shape[0] == 1
return encoder_hidden_states, encoder_attention_mask
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1,
text_seq_len: int = 0
) -> BlockMask:
"""
we will divide the token sequence into the following format
[1 latent frame] [1 latent frame] ... [1 latent frame]
We use flexattention to construct the attention mask
"""
total_length = num_frames * frame_seqlen
total_kv_length = total_length + text_seq_len
total_length_tensor = torch.tensor(total_length, device=device)
total_kv_length_tensor = torch.tensor(total_kv_length, device=device)
# we do right padding to get to a multiple of 128
padded_length = math.ceil(total_length / 128) * 128 - total_length
kv_padded_length = math.ceil(total_kv_length / 128) * 128 - total_kv_length
ends = torch.zeros(total_length + padded_length,
device=device, dtype=torch.long)
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
frame_indices = torch.arange(
# start=frame_seqlen,
start=0,
end=total_length,
step=frame_seqlen * num_frame_per_block,
device=device
)
# frame_indices = torch.cat([torch.tensor([0], device=device), frame_indices])
for i, tmp in enumerate(frame_indices):
# if i == 0:
# ends[tmp:tmp + frame_seqlen] = tmp + frame_seqlen
# else:
ends[tmp:tmp + frame_seqlen * num_frame_per_block] = tmp + \
frame_seqlen * num_frame_per_block
def attention_mask(b, h, q_idx, kv_idx):
if local_attn_size == -1:
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx) | ((kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor))
else:
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx) | ((kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor))
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
KV_LEN=total_kv_length + kv_padded_length, _compile=False, device=device)
# if not dist.is_initialized() or dist.get_rank() == 0:
# print(
# f" cache a block wise causal mask with block size of {num_frame_per_block} frames")
# print(block_mask)
# import imageio
# import numpy as np
# from torch.nn.attention.flex_attention import create_mask
# mask = create_mask(attention_mask, B=None, H=None, Q_LEN=total_length +
# padded_length, KV_LEN=total_length + padded_length, device=device)
# import cv2
# mask = cv2.resize(mask[0, 0].cpu().float().numpy(), (1024, 1024))
# imageio.imwrite("mask_%d.jpg" % (0), np.uint8(255. * mask))
return block_mask
def forward_txt(
self,
encoder_hidden_states: List[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[List[torch.Tensor]] = None,
encoder_attention_mask: Optional[List[torch.Tensor]] = None,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
cache_txt: bool = False,
**kwargs
):
# Check that the timestep is only consisted of 0s
assert torch.all(timestep == 0), "Timestep for txt must be only consisted of 0s"
if cache_txt:
_kv_cache_new = []
transformer_num_layers = len(self.double_blocks)
for _ in range(transformer_num_layers):
_kv_cache_new.append(
{"k_vision": None, "v_vision": None, "k_txt": None, "v_txt": None}
)
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
# 2. Conditional embeddings
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
encoder_hidden_states, encoder_attention_mask = self.get_text_and_mask(
encoder_hidden_states,
encoder_hidden_states_2,
encoder_attention_mask,
encoder_attention_mask_2,
encoder_hidden_states_image,
timestep
)
encoder_hidden_states = encoder_hidden_states[encoder_attention_mask.bool().to(encoder_hidden_states.device)].unsqueeze(0)
# 4. Transformer blocks
for index, block in enumerate(self.double_blocks):
encoder_hidden_states, t_kv = block(
txt_inference=True,
vision_inference=False,
txt=encoder_hidden_states,
vec=temb,
cache_txt=cache_txt,
)
if cache_txt:
_kv_cache_new[index]["k_txt"] = t_kv["k_txt"]
_kv_cache_new[index]["v_txt"] = t_kv["v_txt"]
if cache_txt:
return _kv_cache_new
def forward_vision(
self,
hidden_states: torch.Tensor,
timestep: torch.LongTensor,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
kv_cache: dict | None = None,
txt_kv_cache: list | None = None,
current_start: int = 0,
rope_start_idx: int = 0,
):
assert txt_kv_cache is not None, "txt_kv_cache must be provided"
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width), self.hidden_size,
self.num_attention_heads, self.config.rope_axes_dim, self.config.rope_theta, start_frame=rope_start_idx)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
hidden_states = self.img_in(hidden_states)
# Prepare block-wise causal attention mask
if kv_cache is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
text_seq_len=txt_kv_cache[0]["k_txt"].shape[1]
)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block_index, block in enumerate(self.double_blocks):
hidden_states, new_cache = self._gradient_checkpointing_func(
block,
txt_inference=False,
vision_inference=True,
img=hidden_states,
vec=temb,
freqs_cis=freqs_cis,
block_mask=self.block_mask,
kv_cache=kv_cache[block_index] if kv_cache is not None else None,
txt_kv_cache=txt_kv_cache[block_index],
current_start=current_start
)
if new_cache is not None and kv_cache is not None:
for k in new_cache.keys():
kv_cache[block_index][k] = new_cache[k].clone()
else:
for block_index, block in enumerate(self.double_blocks):
hidden_states, new_cache = block(
txt_inference=False,
vision_inference=True,
img=hidden_states,
vec=temb,
freqs_cis=freqs_cis,
block_mask=self.block_mask,
kv_cache=kv_cache[block_index] if kv_cache is not None else None,
txt_kv_cache=txt_kv_cache[block_index],
current_start=current_start
)
if new_cache is not None and kv_cache is not None:
for k in new_cache.keys():
kv_cache[block_index][k] = new_cache[k].clone()
# Final layer processing
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states, kv_cache
def forward(
self,
txt_inference=False,
vision_inference=False,
**kwargs,
):
if txt_inference:
return self.forward_txt(**kwargs)
elif vision_inference:
return self.forward_vision(**kwargs)
else:
raise ValueError("txt_inference and vision_inference cannot be both False")
+18 -3
View File
@@ -127,8 +127,9 @@ class HunyuanVideo15TimeEmbedding(nn.Module):
self,
timestep: torch.Tensor,
timestep_r: Optional[torch.Tensor] = None,
timestep_seq_len: int | None = None,
) -> torch.Tensor:
timesteps_emb = self.timestep_embedder(timestep)
timesteps_emb = self.timestep_embedder(timestep, timestep_seq_len)
if timestep_r is not None:
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
@@ -473,6 +474,7 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
**kwargs
):
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
@@ -494,7 +496,12 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
@@ -698,6 +705,7 @@ class SingleTokenRefiner(nn.Module):
else:
mask_float = mask.float().unsqueeze(-1) # [B, L, 1]
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
context_aware_representations = context_aware_representations.to(original_dtype)
context_aware_representations = self.c_embedder(
context_aware_representations)
@@ -850,6 +858,13 @@ class FinalLayer(nn.Module):
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
if c.dim() == 3:
# [bs, seq_len, inner_dim]
num_frames = scale.shape[1]
frame_seqlen = x.shape[1] // num_frames
x = (self.norm_final(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale.unsqueeze(2)) + shift.unsqueeze(2)).flatten(1, 2)
else:
# [bs, inner_dim]
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
+7 -5
View File
@@ -391,6 +391,8 @@ class TextEncoderLoader(ComponentLoader):
from fastvideo.platforms import current_platform
logger.info("Loading text encoder with cpu_offload: %s", use_cpu_offload)
if use_cpu_offload:
pin_cpu_memory = fastvideo_args.pin_cpu_memory and is_pin_memory_available()
# Disable FSDP for MPS as it's not compatible
@@ -558,10 +560,9 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
class_name = config.get("_class_name")
assert class_name is not None, (
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
)
config.pop("_name_or_path", None)
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
fastvideo_args.model_paths["vae"] = model_path
from fastvideo.platforms import current_platform
@@ -729,6 +730,7 @@ class TransformerLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
config.pop("_name_or_path", None)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
if cls_name is None:
@@ -829,7 +831,7 @@ class TransformerLoader(ComponentLoader):
)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
logger.info("Loaded model with %.2fB parameters, with cpu_offload: %s", total_params / 1e9, fastvideo_args.dit_cpu_offload)
assert next(model.parameters()).dtype == default_dtype, (
"Model dtype does not match default dtype"
+4 -3
View File
@@ -78,9 +78,10 @@ def hf_to_custom_state_dict(
for source_param_name, full_tensor in hf_param_sd: # type: ignore
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is None:
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
+2
View File
@@ -28,6 +28,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
"HYWorldTransformer3DModel":
("dits", "hyworld", "HYWorldTransformer3DModel"),
"CausalHunyuanVideo15Transformer3DModel":
("dits", "causal_hunyuanvideo15", "CausalHunyuanVideo15Transformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
@@ -155,8 +155,8 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
self.sigmas = sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
self.sigma_min = sigma_min if sigma_min is not None else self.sigmas[-1].item()
self.sigma_max = sigma_max if sigma_max is not None else self.sigmas[0].item()
BaseScheduler.__init__(self)
@@ -289,6 +289,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
sigmas: list[float] | None = None,
mu: float | None = None,
timesteps: list[float] | None = None,
extra_one_step: bool = False,
) -> None:
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -350,7 +351,10 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
if timesteps_array is None:
t_max = self._sigma_to_t(self.sigma_max)
t_min = self._sigma_to_t(self.sigma_min)
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
if extra_one_step:
timesteps_array = np.linspace(t_max, t_min, num_inference_steps + 1)[:-1]
else:
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
sigmas_array = timesteps_array / self.config.num_train_timesteps
else:
sigmas_array = np.array(sigmas).astype(np.float32)
@@ -644,7 +648,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
self,
clean_latent: torch.Tensor,
noise: torch.Tensor,
timestep: torch.IntTensor,
timestep: torch.Tensor,
) -> torch.Tensor:
"""
+3
View File
@@ -5,6 +5,9 @@ from typing import Any
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move it elsewhere
def auto_attributes(init_func):
+1 -1
View File
@@ -663,7 +663,7 @@ class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = False
self.use_tiling = True
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
+228
View File
@@ -0,0 +1,228 @@
import math
import torch
try:
from torch.distributed.tensor import DTensor
except ImportError:
# handle old pytorch versions
Dtensor = None
# This code is modified from the GitHub repository of KellerJordan:
# https://github.com/KellerJordan/Muon/blob/master/muon.py
def zeropower_via_newtonschulz5(G, steps=5):
"""
Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a
quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose
of minimizing steps, it turns out to be empirically effective to keep increasing the slope at
zero even beyond the point where the iteration no longer converges all the way to one everywhere
on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T
where S' is diagonal with S_{ii}' ~ Uniform(0.5, 1.5), which turns out not to hurt model
performance at all relative to UV^T, where USV^T = G is the SVD.
"""
if isinstance(G, DTensor):
device_mesh = G.device_mesh
G = G.full_tensor()
else:
device_mesh = None
assert len(G.shape) >= 2
a, b, c = (3.4445, -4.7750, 2.0315)
X = G
if G.size(-2) > G.size(-1):
X = X.mT
# Ensure spectral norm is at most 1
X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7)
# Perform the NS iterations
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A # quintic computation strategy adapted from suggestion by @jxbz, @leloykun, and @YouJiacheng
X = a * X + B @ X
if G.size(-2) > G.size(-1):
X = X.mT
if device_mesh is not None:
return DTensor.from_local(X, device_mesh)
else:
return X
class Muon(torch.optim.Optimizer):
"""
Muon - MomentUm Orthogonalized by Newton-schulz
Arguments:
muon_params: The parameters to be optimized by Muon.
lr: The learning rate. The updates will have spectral norm of `lr`. (0.02 is a good default)
momentum: The momentum used by the internal SGD. (0.95 is a good default)
nesterov: Whether to use Nesterov-style momentum in the internal SGD. (recommended)
ns_steps: The number of Newton-Schulz iterations to run. (6 is probably always enough)
adamw_params: The parameters to be optimized by AdamW. Any parameters in `muon_params` which are
{0, 1}-D or are detected as being the embed or lm_head will be optimized by AdamW as well.
adamw_lr: The learning rate for the internal AdamW.
adamw_betas: The betas for the internal AdamW.
adamw_eps: The epsilon for the internal AdamW.
adamw_wd: The weight decay for the internal AdamW.
"""
def __init__(
self,
lr=1e-3,
wd=0.1,
muon_params=None,
momentum=0.95,
nesterov=True,
ns_steps=5,
adamw_params=None,
adamw_betas=(0.95, 0.95),
adamw_eps=1e-8,
):
defaults = dict(
lr=lr,
wd=wd,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
)
params = list(muon_params)
adamw_params = list(adamw_params) if adamw_params is not None else []
params.extend(adamw_params)
super().__init__(params, defaults)
# Sort parameters into those for which we will use Muon, and those for which we will not
for p in muon_params:
# Use Muon for every parameter in muon_params which is >= 2D and doesn't look like an embedding or head layer
assert p.ndim >= 2, p.ndim
self.state[p]["use_muon"] = True
for p in adamw_params:
# Do not use Muon for parameters in adamw_params
self.state[p]["use_muon"] = False
def adjust_lr_for_muon(self, lr, param_shape):
A, B = param_shape[:2]
# We adjust the learning rate and weight decay based on the size of the parameter matrix
# as describted in the paper
adjusted_ratio = 0.2 * math.sqrt(max(A, B))
adjusted_lr = lr * adjusted_ratio
return adjusted_lr
def step(self, closure=None):
"""Perform a single optimization step.
Args:
closure (Callable, optional): A closure that reevaluates the model
and returns the loss.
"""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
############################
# Muon #
############################
params = [p for p in group["params"] if self.state[p]["use_muon"]]
lr = group["lr"]
wd = group["wd"]
momentum = group["momentum"]
# generate weight updates in distributed fashion
for p in params:
# sanity check
g = p.grad
if g is None:
continue
if g.ndim > 2:
g = g.view(g.size(0), -1)
assert g is not None
# calc update
state = self.state[p]
if "momentum_buffer" not in state:
state["momentum_buffer"] = torch.zeros_like(g)
buf = state["momentum_buffer"]
buf.mul_(momentum).add_(g)
g = g.add(buf, alpha=momentum) if group["nesterov"] else buf
g = g.bfloat16()
u = zeropower_via_newtonschulz5(g, steps=group["ns_steps"])
# scale update
adjusted_lr = self.adjust_lr_for_muon(lr, p.shape)
# apply weight decay
p.data.mul_(1 - lr * wd)
# apply update
p.data.add_(u.view(p.shape), alpha=-adjusted_lr)
############################
# AdamW backup #
############################
params = [
p for p in group["params"] if not self.state[p]["use_muon"]
]
lr = group['lr']
beta1, beta2 = group["adamw_betas"]
eps = group["adamw_eps"]
weight_decay = group["wd"]
for p in params:
g = p.grad
if g is None:
continue
state = self.state[p]
if "step" not in state:
state["step"] = 0
state["moment1"] = torch.zeros_like(g)
state["moment2"] = torch.zeros_like(g)
state["step"] += 1
step = state["step"]
buf1 = state["moment1"]
buf2 = state["moment2"]
buf1.lerp_(g, 1 - beta1)
buf2.lerp_(g.square(), 1 - beta2)
g = buf1 / (eps + buf2.sqrt())
bias_correction1 = 1 - beta1**step
bias_correction2 = 1 - beta2**step
scale = bias_correction1 / bias_correction2**0.5
p.data.mul_(1 - lr * weight_decay)
p.data.add_(g, alpha=-lr / scale)
return loss
# help function to create the Muon optimizer
def get_muon_optimizer(model,
lr=1e-3,
weight_decay=0.1,
momentum=0.95,
adamw_betas=(0.95, 0.95),
adamw_eps=1e-8):
muon_params = [
p for name, p in model.named_parameters()
if p.requires_grad and p.ndim >= 2
]
adamw_params = [
p for name, p in model.named_parameters()
if p.requires_grad and not (p.ndim >= 2)
]
return Muon(
lr=lr,
wd=weight_decay,
muon_params=muon_params,
momentum=momentum,
adamw_params=adamw_params,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
)
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan causal DMD pipeline implementation.
This module wires the causal DMD denoising stage into the modular pipeline.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
Hy15CausalDMDDenosingStage,
InputValidationStage,
Hy15ImageEncodingStage,
LatentPreparationStage,
TextEncodingStage)
# isort: on
logger = init_logger(__name__)
class Hy15CausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=Hy15CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = Hy15CausalDMDPipeline
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
Hy15ImageEncodingStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15ImageToVideoPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HunyuanVideo15ImageToVideoPipeline
@@ -84,6 +84,7 @@ class ComposedPipelineBase(ABC):
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
# fastvideo_args.dit_cpu_offload = False
with self.profiler_controller.region("profiler_region_model_loading"):
self.modules = self.load_modules(fastvideo_args, loaded_modules)
@@ -287,6 +288,7 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
model_index.pop("_name_or_path", None)
model_index.pop("workload_type", None)
if "boundary_ratio" in model_index and model_index[
"boundary_ratio"] is not None:
@@ -232,6 +232,7 @@ class TrainingBatch:
image_latents: torch.Tensor | None = None
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
video_latent: torch.Tensor | None = None
# ODE trajectory supervision
trajectory_latents: torch.Tensor | None = None
@@ -240,6 +241,9 @@ class TrainingBatch:
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
timesteps: torch.Tensor | None = None
use_gt_trajectory: bool = False
trajectory_timesteps: torch.Tensor | None = None
start_timestep_index: int | None = None
sigmas: torch.Tensor | None = None
noise: torch.Tensor | None = None
+2
View File
@@ -29,6 +29,8 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
"HYWorldPipeline": "hyworld",
"Hy15CausalDMDPipeline": "hunyuan15",
"HunyuanVideo15ImageToVideoPipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"Cosmos2_5Pipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
@@ -303,7 +303,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
@@ -422,4 +422,4 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
num_processed_samples = 0
num_processed_samples = 0
@@ -28,16 +28,12 @@ from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages import (
DecodingStage, DenoisingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage, Hy15ImageEncodingStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
logger = init_logger(__name__)
@@ -47,7 +43,8 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
@@ -61,19 +58,19 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
assert fastvideo_args.pipeline_config.flow_shift == 5
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
denoising_strength=1.0)
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
@@ -82,6 +79,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
@@ -95,11 +95,13 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
args):
"""Preprocess text-only data and generate trajectory information."""
num_encoders = len(self.prompt_encoding_stage.text_encoders)
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
@@ -130,12 +132,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
@@ -144,61 +143,48 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
negative_prompt_embeds_list = []
negative_prompt_masks_list = []
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = prompt_embeds_list
batch.prompt_attention_mask = prompt_masks_list
batch.negative_prompt_embeds = negative_prompt_embeds_list
batch.negative_attention_mask = negative_prompt_masks_list
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.num_frames = args.num_frames
batch.fps = args.train_fps
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.num_inference_steps = 48
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.image_encoding_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch, fastvideo_args)
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
trajectory_latents.append(result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
extra_features = {
@@ -209,10 +195,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
if j in [5, 7]:
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
@@ -227,7 +214,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
text_embedding = prompt_embeds_list[0].float().cpu().numpy()
text_mask = prompt_masks_list[0].cpu().numpy()
text_embedding_2 = prompt_embeds_list[1].float().cpu(
).numpy()
text_mask_2 = prompt_masks_list[1].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
@@ -253,6 +244,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"trajectory_latents"],
trajectory_timesteps=sample_extra_features[
"trajectory_timesteps"],
text_embedding_2=text_embedding_2,
text_mask=text_mask,
text_mask_2=text_mask_2,
)
batch_data.append(record)
@@ -58,7 +58,7 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
@@ -181,4 +181,4 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text
EntryClass = PreprocessPipeline_Text
+18 -20
View File
@@ -1,9 +1,7 @@
import argparse
import os
from typing import Any
from fastvideo import PipelineConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.distributed import (
get_world_size, maybe_init_distributed_environment_and_model_parallel)
from fastvideo.fastvideo_args import FastVideoArgs
@@ -16,38 +14,38 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
def main(args) -> None:
args.model_path = maybe_download_model(args.model_path)
# args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
print(pipeline_config.__class__.__name__)
kwargs: dict[str, Any] = {}
if args.preprocess_task == "text_only":
kwargs = {
"text_encoder_cpu_offload": False,
}
else:
# Full config for video/image processing
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
pipeline_config.update_config_from_dict(kwargs)
# kwargs: dict[str, Any] = {}
# if args.preprocess_task == "text_only":
# kwargs = {
# "text_encoder_cpu_offload": False,
# }
# else:
# # Full config for video/image processing
# kwargs = {
# "vae_precision": "fp32",
# "vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
# }
# pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pipeline_config=pipeline_config,
)
if args.preprocess_task == "t2v":
@@ -134,4 +132,4 @@ if __name__ == "__main__":
)
args = parser.parse_args()
main(args)
main(args)
@@ -24,4 +24,4 @@ if __name__ == "__main__":
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
fastvideo_args = FastVideoArgs.from_cli_args(args)
main(fastvideo_args)
main(fastvideo_args)
+2
View File
@@ -8,6 +8,7 @@ complete diffusion pipelines.
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.hy15_causal_denoising import Hy15CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
@@ -52,6 +53,7 @@ __all__ = [
"Cosmos25LatentPreparationStage",
"LTX2LatentPreparationStage",
"LTX2AudioDecodingStage",
"Hy15CausalDMDDenosingStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
@@ -45,12 +45,12 @@ class CausalDMDDenosingStage(DenoisingStage):
self.transformer_2 = transformer_2
self.vae = vae
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
self.num_transformer_blocks = self.transformer.config.num_layers
self.num_frames_per_block = self.transformer.config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.sliding_window_num_frames
try:
self.local_attn_size = getattr(self.transformer.model,
self.local_attn_size = getattr(self.transformer.config,
"local_attn_size",
-1) # type: ignore
except Exception:
@@ -412,8 +412,8 @@ class CausalDMDDenosingStage(DenoisingStage):
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = self.transformer.attention_head_dim
num_attention_heads = self.transformer.config.num_attention_heads
attention_head_dim = self.transformer.config.attention_head_dim
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
+18 -10
View File
@@ -26,7 +26,7 @@ from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import dict_to_3d_list, masks_like
from fastvideo.utils import dict_to_3d_list, masks_like, PRECISION_TO_TYPE
try:
from fastvideo.attention.backends.sliding_tile_attn import (
@@ -210,7 +210,10 @@ class DenoisingStage(PipelineStage):
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
assert self.vae is not None, "VAE is not provided for TI2V task"
z = self.vae.encode(batch.pil_image).mean.float()
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
self.vae = self.vae.to(get_local_torch_device())
z = self.vae.encode(batch.pil_image.to(vae_dtype)).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
@@ -223,6 +226,9 @@ class DenoisingStage(PipelineStage):
else:
z = z * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to('cpu')
latent_model_input = latent_model_input.squeeze(0)
_, mask2 = masks_like([latent_model_input], zero=True)
@@ -232,18 +238,18 @@ class DenoisingStage(PipelineStage):
latent_model_input = latent_model_input.to(get_local_torch_device())
latents = latent_model_input
F = batch.num_frames
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
seq_len = ((F - 1) // temporal_scale +
1) * (batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size[1] *
patch_size[2])
temporal_scale = fastvideo_args.pipeline_config.vae_config.temporal_compression_ratio
spatial_scale = fastvideo_args.pipeline_config.vae_config.spatial_compression_ratio
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
seq_len = ((F - 1) // temporal_scale + 1) * (
batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size * patch_size)
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = [latents]
logger.info("timesteps: %s", timesteps)
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -408,6 +414,7 @@ class DenoisingStage(PipelineStage):
**action_kwargs,
)
assert batch.do_classifier_free_guidance, "do_classifier_free_guidance is not supported"
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
@@ -462,6 +469,7 @@ class DenoisingStage(PipelineStage):
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_timesteps.append(torch.zeros_like(t))
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
@@ -0,0 +1,361 @@
import math
import torch # type: ignore
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.utils import PRECISION_TO_TYPE
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
logger = init_logger(__name__)
class Hy15CausalDMDDenosingStage(CausalDMDDenosingStage):
"""
Denoising stage for causal diffusion.
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
if isinstance(self.transformer.config.patch_size, tuple):
patch_ratio = self.transformer.config.patch_size[
1] * self.transformer.config.patch_size[2]
elif isinstance(self.transformer.config.patch_size, int):
patch_ratio = self.transformer.config.patch_size**2
else:
raise ValueError(
f"Unsupported patch size type: {type(self.transformer.config.patch_size)}"
)
self.frame_seq_length = latent_seq_length // patch_ratio
# TODO(will): make this a parameter once we add i2v support
independent_first_frame = self.transformer.independent_first_frame if hasattr(
self.transformer, 'independent_first_frame') else False
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
self.scheduler.set_timesteps(num_inference_steps=1000,
extra_one_step=True,
device=get_local_torch_device())
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
logger.info("[causal_denoising] timesteps: %s", timesteps)
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
# Image kwargs (kept empty unless caller provides compatible args)
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# STA
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Latents and prompts
assert batch.latents is not None, "latents must be provided"
latents = batch.latents # [B, C, T, H, W]
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
pos_start_base = 0
num_blocks = math.ceil(t / self.num_frames_per_block)
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
# Initialize txt kv cache
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
txt_kv_cache = self.transformer(
txt_inference=True,
vision_inference=False,
encoder_hidden_states=prompt_embeds,
encoder_hidden_states_image=image_embeds,
encoder_attention_mask=batch.prompt_attention_mask,
timestep=torch.zeros([latents.shape[0]], device=latents.device),
cache_txt=True,
)
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert self.vae is not None, "VAE is not provided for causal video gen task"
self.vae = self.vae.to(get_local_torch_device())
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
first_frame_latent = self.vae.encode(
batch.pil_image.to(vae_dtype)).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent = first_frame_latent * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
t_zero = torch.zeros([latents.shape[0], 1],
device=latents.device,
dtype=torch.long)
if batch.video_latent is not None:
video_latent_chunk = batch.video_latent[:, :, start_index:
start_index + 1, :, :]
first_frame_input = torch.cat([
first_frame_latent,
video_latent_chunk,
torch.zeros_like(first_frame_latent),
],
dim=1)
else:
first_frame_input = first_frame_latent.clone()
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.transformer(
txt_inference=False,
vision_inference=True,
hidden_states=first_frame_input.to(target_dtype),
timestep=t_zero,
kv_cache=kv_cache1,
txt_kv_cache=txt_kv_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
rope_start_idx=start_index,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
vision_input_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"vision_inference": True,
"txt_inference": False,
},
)
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for current_num_frames in block_sizes:
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# use BTCHW for DMD conversion routines
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
# Copy for pred conversion
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
if batch.video_latent is not None:
video_latent_chunk = batch.video_latent[:, :,
start_index:
start_index +
current_num_frames, :, :]
latent_model_input = torch.cat([
latent_model_input,
video_latent_chunk,
torch.zeros_like(current_latents),
],
dim=1)
elif batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(target_dtype)
],
dim=2)
# Prepare inputs
t_expand = t_cur.repeat(latent_model_input.shape[0])
# Attention metadata if needed
if (vsa_available and self.attn_backend
== VideoSparseAttentionBackend):
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=i, # type: ignore
raw_latent_shape=(current_num_frames, h,
w), # type: ignore
patch_size=fastvideo_args.pipeline_config.
dit_config.patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(), # type: ignore
) # type: ignore
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw, kv_cache1 = self.transformer(
hidden_states=latent_model_input,
timestep=t_expanded_noise,
kv_cache=kv_cache1,
txt_kv_cache=txt_kv_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
rope_start_idx=start_index,
**vision_input_kwargs,
)
pred_noise_btchw = pred_noise_btchw.permute(
0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1],
dtype=torch.long,
device=pred_video_btchw.device)
noise = torch.randn(
video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(
0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(
0, 2, 1, 3, 4)
if progress_bar is not None:
progress_bar.update()
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Re-run with context timestep to update KV cache using clean context
context_noise = getattr(fastvideo_args.pipeline_config,
"context_noise", 0)
t_context = torch.ones([latents.shape[0], current_num_frames],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
if batch.video_latent is not None:
video_latent_chunk = batch.video_latent[:, :, start_index:
start_index +
current_num_frames, :, :]
context_bcthw = torch.cat([
context_bcthw,
video_latent_chunk,
torch.zeros_like(current_latents),
],
dim=1)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
_, kv_cache1 = self.transformer(
hidden_states=context_bcthw,
timestep=t_context,
kv_cache=kv_cache1,
txt_kv_cache=txt_kv_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
rope_start_idx=start_index,
**vision_input_kwargs,
)
start_index += current_num_frames
batch.latents = latents
return batch
+4 -4
View File
@@ -115,10 +115,10 @@ class Hy15ImageEncodingStage(ImageEncodingStage):
"""
Encode the prompt into image encoder hidden states.
"""
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
# if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
raw_latent_shape = list(batch.raw_latent_shape)
raw_latent_shape[1] = 1
@@ -120,9 +120,9 @@ class InputValidationStage(PipelineStage):
else:
# Standard Wan logic
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
dh, dw = patch_size * vae_stride, patch_size * vae_stride
max_area = 480 * 848
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
+76 -48
View File
@@ -28,8 +28,6 @@ from fastvideo.distributed import (cleanup_dist_env_and_memory,
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
@@ -42,6 +40,7 @@ from fastvideo.training.training_utils import (
shift_timestep)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
from fastvideo.optim.muon import get_muon_optimizer
vsa_available = is_vsa_available()
@@ -85,7 +84,6 @@ class DistillationPipeline(TrainingPipeline):
if training_args.real_score_model_path:
logger.info("Loading real score transformer from: %s",
training_args.real_score_model_path)
training_args.override_transformer_cls_name = "WanTransformer3DModel"
# TODO(will): can use deepcopy instead if the model is the same
self.real_score_transformer = self.load_module_from_path(
training_args.real_score_model_path, "transformer",
@@ -111,7 +109,6 @@ class DistillationPipeline(TrainingPipeline):
if training_args.fake_score_model_path:
logger.info("Loading fake score transformer from: %s",
training_args.fake_score_model_path)
training_args.override_transformer_cls_name = "WanTransformer3DModel"
self.fake_score_transformer = self.load_module_from_path(
training_args.fake_score_model_path, "transformer",
training_args)
@@ -146,8 +143,8 @@ class DistillationPipeline(TrainingPipeline):
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
# self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
# shift=self.timestep_shift)
if self.training_args.boundary_ratio is not None:
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
@@ -192,16 +189,24 @@ class DistillationPipeline(TrainingPipeline):
if fake_score_lr == 0.0:
fake_score_lr = training_args.learning_rate
betas_str = training_args.fake_score_betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
betas_str = training_args.fake_score_betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.fake_score_optimizer = get_muon_optimizer(
self.fake_score_transformer,
lr=fake_score_lr,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
self.fake_score_lr_scheduler = get_scheduler(
training_args.fake_score_lr_scheduler,
@@ -218,13 +223,23 @@ class DistillationPipeline(TrainingPipeline):
fake_score_params_2 = list(
filter(lambda p: p.requires_grad,
self.fake_score_transformer_2.parameters()))
self.fake_score_optimizer_2 = torch.optim.AdamW(
fake_score_params_2,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
self.fake_score_optimizer_2 = torch.optim.AdamW(
fake_score_params_2,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.fake_score_optimizer_2 = get_muon_optimizer(
self.fake_score_transformer_2,
lr=fake_score_lr,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
self.fake_score_lr_scheduler_2 = get_scheduler(
training_args.fake_score_lr_scheduler,
optimizer=self.fake_score_optimizer_2,
@@ -272,21 +287,6 @@ class DistillationPipeline(TrainingPipeline):
self.generator_ema: EMA_FSDP | None = None
self.generator_ema_2: EMA_FSDP | None = None
if (self.training_args.ema_decay
is not None) and (self.training_args.ema_decay > 0.0):
self.generator_ema = EMA_FSDP(self.transformer,
decay=self.training_args.ema_decay)
logger.info("Initialized generator EMA with decay=%s",
self.training_args.ema_decay)
# Initialize EMA for transformer_2 if it exists
if self.transformer_2 is not None:
self.generator_ema_2 = EMA_FSDP(
self.transformer_2, decay=self.training_args.ema_decay)
logger.info("Initialized generator EMA_2 with decay=%s",
self.training_args.ema_decay)
else:
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
def load_module_from_path(self, model_path: str, module_type: str,
training_args: "TrainingArgs"):
@@ -572,6 +572,7 @@ class DistillationPipeline(TrainingPipeline):
training_batch.input_kwargs = {
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_hidden_states_image": training_batch.image_embeds,
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep,
"return_dict": False,
@@ -723,10 +724,18 @@ class DistillationPipeline(TrainingPipeline):
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).detach().unflatten(0,
(1, generator_pred_video.shape[1]))
noisy_latent_copy = noisy_latent.clone()
if training_batch.video_latent is not None:
noisy_latent_copy = torch.cat([
noisy_latent_copy,
training_batch.video_latent,
torch.zeros_like(noisy_latent_copy),
],
dim=2)
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
noisy_latent_copy, timestep, training_batch.conditional_dict,
training_batch)
current_fake_score_transformer = self._get_fake_score_transformer(
timestep)
@@ -742,7 +751,7 @@ class DistillationPipeline(TrainingPipeline):
# real_score_transformer cond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
noisy_latent_copy, timestep, training_batch.conditional_dict,
training_batch)
current_real_score_transformer = self._get_real_score_transformer(
timestep)
@@ -758,7 +767,7 @@ class DistillationPipeline(TrainingPipeline):
# real_score_transformer uncond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.unconditional_dict,
noisy_latent_copy, timestep, training_batch.unconditional_dict,
training_batch)
# Use same transformer as conditional forward for consistency
real_score_pred_noise_uncond = current_real_score_transformer(
@@ -779,9 +788,16 @@ class DistillationPipeline(TrainingPipeline):
original_latent - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
original_latent.float(),
(original_latent.float() - grad.float()).detach())
if self.training_args.use_context_forcing and training_batch.trajectory_latents is not None:
context_forcing_length = training_batch.trajectory_latents.shape[1]
dmd_loss = 0.5 * F.mse_loss(
original_latent.float()[:, context_forcing_length:],
(original_latent.float()[:, context_forcing_length:] -
grad.float()[:, context_forcing_length:]).detach())
else:
dmd_loss = 0.5 * F.mse_loss(
original_latent.float(),
(original_latent.float() - grad.float()).detach())
training_batch.dmd_latent_vis_dict.update({
"training_batch_dmd_fwd_clean_latent":
@@ -831,6 +847,13 @@ class DistillationPipeline(TrainingPipeline):
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
fake_score_timestep).unflatten(0,
(1, generator_pred_video.shape[1]))
if training_batch.video_latent is not None:
noisy_generator_pred_video = torch.cat([
noisy_generator_pred_video,
training_batch.video_latent,
torch.zeros_like(noisy_generator_pred_video),
],
dim=2)
with set_forward_context(current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
@@ -877,7 +900,7 @@ class DistillationPipeline(TrainingPipeline):
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
super()._prepare_dit_inputs(training_batch)
# super()._prepare_dit_inputs(training_batch)
conditional_dict = {
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
@@ -898,7 +921,7 @@ class DistillationPipeline(TrainingPipeline):
self.video_latent_shape = training_batch.latents.shape
self.video_latent_shape_sp = training_batch.latents.shape
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
@@ -1239,6 +1262,7 @@ class DistillationPipeline(TrainingPipeline):
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch, steps)
batch.prompt_attention_mask = []
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
@@ -1249,8 +1273,11 @@ class DistillationPipeline(TrainingPipeline):
)
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
batch_negative, training_args)
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
if len(result_batch.prompt_embeds) == 1:
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
else:
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds, result_batch.prompt_attention_mask
logger.info(
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
@@ -1354,6 +1381,7 @@ class DistillationPipeline(TrainingPipeline):
if hasattr(self, 'transformer_2') and self.transformer_2 is not None:
self.transformer_2.train()
gc.collect()
torch.cuda.empty_cache()
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
File diff suppressed because it is too large Load Diff
+222 -53
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import sys
from copy import deepcopy
from typing import Any, cast
@@ -13,14 +14,15 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
# from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
# WanCausalDMDPipeline)
from fastvideo.pipelines.basic.hunyuan15.hunyuan15_causal_dmd_pipeline import Hy15CausalDMDPipeline
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.decoding import DecodingStage
logger = init_logger(__name__)
@@ -35,16 +37,18 @@ class ODEInitTrainingPipeline(TrainingPipeline):
- minimizing MSE to the stored next latent at timestep t_next
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
_required_config_modules = [
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Match the preprocess/generation scheduler for consistent stepping
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
training=True)
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# # Match the preprocess/generation scheduler for consistent stepping
# self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
# shift=fastvideo_args.pipeline_config.flow_shift,
# sigma_min=0.0,
# extra_one_step=True)
# self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
# training=True)
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
@@ -56,18 +60,36 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.text_encoder = self.get_module("text_encoder")
self.text_encoder.requires_grad_(False)
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer")
],
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
# self.noise_scheduler = SelfForcingFlowMatchScheduler(
# shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
extra_one_step=True,
device=get_local_torch_device())
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
dtype=torch.long,
device=get_local_torch_device())
self.dmd_denoising_steps = torch.tensor(
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
torch.tensor([0],
@@ -78,7 +100,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
logger.info("warped self.dmd_denoising_steps: %s",
self.dmd_denoising_steps)
else:
raise ValueError("warp_denoising_step must be true")
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32))).cuda()
self.dmd_denoising_steps = timesteps[self.dmd_denoising_steps]
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
get_local_torch_device())
@@ -98,7 +123,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
# Warm start validation with current transformer
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
self.validation_pipeline = Hy15CausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
@@ -109,7 +134,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
dit_cpu_offload=False)
def _get_next_batch(
self,
@@ -117,15 +142,40 @@ class ODEInitTrainingPipeline(TrainingPipeline):
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
device = get_local_torch_device()
encoder_hidden_states = batch['text_embedding'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_hidden_states_2 = batch['text_embedding_2'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_attention_mask = batch['text_mask'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_attention_mask_2 = batch['text_mask_2'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_hidden_states_image = [
torch.zeros(1,
729,
1152,
device=get_local_torch_device(),
dtype=torch.bfloat16)
]
infos = batch['info_list']
if encoder_hidden_states.dim() < 3:
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
infos[0]["caption"],
self.training_args,
encoder_index=[0],
return_attention_mask=True,
)
encoder_hidden_states = prompt_embeds_list[0].to(
device, dtype=torch.bfloat16)
encoder_attention_mask = prompt_masks_list[0].to(
device, dtype=torch.bfloat16)
# Trajectory tensors may include a leading singleton batch dim per row
trajectory_latents = batch['trajectory_latents']
if trajectory_latents.dim() == 7:
@@ -154,11 +204,13 @@ class ODEInitTrainingPipeline(TrainingPipeline):
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
# Move to device
device = get_local_torch_device()
training_batch.encoder_hidden_states = encoder_hidden_states.to(
device, dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
device, dtype=torch.bfloat16)
training_batch.encoder_hidden_states = [
encoder_hidden_states, encoder_hidden_states_2
]
training_batch.encoder_attention_mask = [
encoder_attention_mask, encoder_attention_mask_2
]
training_batch.encoder_hidden_states_image = encoder_hidden_states_image
training_batch.infos = infos
return training_batch, trajectory_latents.to(
@@ -197,6 +249,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
dtype=torch.long).repeat(1, num_frame)
return timestep
else:
num_pad_frames = 0
if num_frame % num_frame_per_block != 0:
# Pad num_frame to be divisible by num_frame_per_block
num_pad_frames = num_frame_per_block - (num_frame %
num_frame_per_block)
num_frame += num_pad_frames
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, num_frame],
device=self.device,
@@ -207,12 +265,15 @@ class ODEInitTrainingPipeline(TrainingPipeline):
num_frame_per_block)
timestep[:, :, 1:] = timestep[:, :, 0:1]
timestep = timestep.reshape(timestep.shape[0], -1)
if num_pad_frames > 0:
timestep = timestep[:, num_pad_frames:]
return timestep
def _step_predict_next_latent(
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor
encoder_attention_mask: torch.Tensor,
encoder_hidden_states_image: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
torch.Tensor]]:
latent_vis_dict: dict[str, torch.Tensor] = {}
@@ -225,7 +286,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36], dtype=torch.long).cpu()
[0, 12, 24, 36, 50], dtype=torch.long).cpu()
# [0, 1, 2, 3], dtype=torch.long).cpu()
logger.info("self._cached_closest_idx_per_dmd: %s",
self._cached_closest_idx_per_dmd)
@@ -241,6 +302,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
traj_latents,
dim=1,
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
# relevant_traj_latents = traj_latents
logger.info("relevant_traj_latents: %s", relevant_traj_latents.shape)
# assert relevant_traj_latents.shape[0] == 1
@@ -251,51 +313,149 @@ class ODEInitTrainingPipeline(TrainingPipeline):
num_frames,
3,
uniform_timestep=False)
logger.info("indexes: %s", indexes.shape)
logger.info("indexes: %s", indexes)
# noisy_input = relevant_traj_latents[indexes]
noisy_input = torch.gather(
latents = torch.gather(
relevant_traj_latents,
dim=1,
index=indexes.reshape(B, 1, num_frames, 1, 1,
1).expand(-1, -1, -1, num_channels, height,
width).to(self.device)).squeeze(1)
noisy_input = torch.cat([
latents,
torch.zeros_like(latents),
torch.zeros_like(latents[:, :, 0:1])
],
dim=2)
timestep = self.dmd_denoising_steps[indexes]
logger.info("selected timestep for rank %s: %s",
self.global_rank,
timestep,
local_main_process_only=False)
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.permute(
latent_vis_dict["noisy_input"] = latents.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
4).detach().clone().cpu()
input_kwargs = {
logger.info("timestep: %s", timestep)
txt_input_kwargs = {
"txt_inference":
True,
"vision_inference":
False,
"encoder_hidden_states":
encoder_hidden_states,
"encoder_hidden_states_image":
encoder_hidden_states_image,
"encoder_attention_mask":
encoder_attention_mask,
"timestep":
torch.zeros([latents.shape[0]],
device=latents.device,
dtype=torch.bfloat16),
"cache_txt":
True,
}
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
txt_kv_cache = self.transformer(**txt_input_kwargs)
vision_input_kwargs = {
"txt_inference": False,
"vision_inference": True,
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep.to(device, dtype=torch.bfloat16),
"return_dict": False,
"txt_kv_cache": txt_kv_cache,
}
# Predict noise and step the scheduler to obtain next latent
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
noise_pred = self.transformer(**vision_input_kwargs).permute(
0, 2, 1, 3, 4)
from fastvideo.models.utils import pred_noise_to_pred_video
pred_video = pred_noise_to_pred_video(
pred_noise=noise_pred.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
noise_input_latent=latents.flatten(0, 1),
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
scheduler=self.modules["scheduler"]).unflatten(
0, noise_pred.shape[:2])
scheduler=self.noise_scheduler).unflatten(0, noise_pred.shape[:2])
latent_vis_dict["pred_video"] = pred_video.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
return pred_video, target_latent, timestep, latent_vis_dict
# def _step_predict_next_latent(
# self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
# encoder_hidden_states: torch.Tensor,
# encoder_attention_mask: torch.Tensor,
# encoder_hidden_states_image: torch.Tensor
# ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
# torch.Tensor]]:
# latent_vis_dict: dict[str, torch.Tensor] = {}
# device = get_local_torch_device()
# target_latent = traj_latents[:, -1]
# del traj_latents
# # Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
# B, num_frames, num_channels, height, width = target_latent.shape
# indexes = self._get_timestep( # [B, num_frames]
# 0,
# 1000,
# B,
# num_frames,
# 3,
# uniform_timestep=False)
# timestep = self.noise_scheduler.timesteps[indexes.cpu()].to(device)
# latents = self.noise_scheduler.add_noise(target_latent.flatten(0, 1), torch.randn_like(target_latent.flatten(0, 1)), timestep.flatten(0, 1)).unflatten(0, (B, num_frames))
# noisy_input = torch.cat([latents, torch.zeros_like(latents), torch.zeros_like(latents[:, :, 0:1])], dim=2)
# # Prepare inputs for transformer
# latent_vis_dict["noisy_input"] = latents.permute(
# 0, 2, 1, 3, 4).detach().clone().cpu()
# latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
# 4).detach().clone().cpu()
# logger.info("timestep: %s", timestep)
# txt_input_kwargs = {
# "txt_inference": True,
# "vision_inference": False,
# "encoder_hidden_states": encoder_hidden_states,
# "encoder_hidden_states_image": encoder_hidden_states_image,
# "encoder_attention_mask": encoder_attention_mask,
# "timestep": torch.zeros([latents.shape[0]], device=latents.device, dtype=torch.bfloat16),
# "cache_txt": True,
# }
# with set_forward_context(current_timestep=timestep,
# attn_metadata=None,
# forward_batch=None):
# txt_kv_cache = self.transformer(**txt_input_kwargs)
# vision_input_kwargs = {
# "txt_inference": False,
# "vision_inference": True,
# "hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
# "timestep": timestep.to(device, dtype=torch.bfloat16),
# "txt_kv_cache": txt_kv_cache,
# }
# # Predict noise and step the scheduler to obtain next latent
# with set_forward_context(current_timestep=timestep,
# attn_metadata=None,
# forward_batch=None):
# noise_pred = self.transformer(**vision_input_kwargs).permute(0, 2, 1, 3, 4)
# from fastvideo.models.utils import pred_noise_to_pred_video
# pred_video = pred_noise_to_pred_video(
# pred_noise=noise_pred.flatten(0, 1),
# noise_input_latent=latents.flatten(0, 1),
# timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
# scheduler=self.noise_scheduler).unflatten(
# 0, noise_pred.shape[:2])
# latent_vis_dict["pred_video"] = pred_video.permute(
# 0, 2, 1, 3, 4).detach().clone().cpu()
# return pred_video, target_latent, timestep, latent_vis_dict
def train_one_step(self, training_batch): # type: ignore[override]
self.transformer.train()
self.optimizer.zero_grad()
@@ -307,8 +467,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
for _ in range(args.gradient_accumulation_steps):
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
training_batch)
# traj_latents = traj_latents[:, :, :21]
text_embeds = training_batch.encoder_hidden_states
text_attention_mask = training_batch.encoder_attention_mask
image_embeds = training_batch.encoder_hidden_states_image
assert traj_latents.shape[0] == 1
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
@@ -318,7 +480,8 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
traj_latents, traj_timesteps, text_embeds, text_attention_mask,
image_embeds)
training_batch.latent_vis_dict.update(latent_vis_dict)
@@ -356,6 +519,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
except Exception:
grad_value = 0.0
training_batch.grad_norm = grad_value
if training_batch.current_timestep % 10 == 0:
gc.collect()
torch.cuda.empty_cache()
return training_batch
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
@@ -367,14 +534,13 @@ class ODEInitTrainingPipeline(TrainingPipeline):
assert latent_key in latents_vis_dict and latents_vis_dict[
latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
pixel_latent = self.decoding_stage.decode(latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
video_artifact = self.tracker.video(
video, fps=16, format="mp4") # change to 16 for Wan2.1
video, fps=24, format="mp4") # change to 16 for Wan2.1
if video_artifact is not None:
tracker_loss_dict[latent_key] = video_artifact
# Clean up references
@@ -383,6 +549,9 @@ class ODEInitTrainingPipeline(TrainingPipeline):
if self.global_rank == 0 and tracker_loss_dict:
self.tracker.log_artifacts(tracker_loss_dict, step)
gc.collect()
torch.cuda.empty_cache()
def main(args) -> None:
logger.info("Starting ODE-init training pipeline...")
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import os
import gc
import time
from collections import deque
from typing import Any
@@ -8,6 +9,7 @@ from typing import Any
import numpy as np
import torch
import torch.distributed as dist
import torch.nn.functional as F
from einops import rearrange
from tqdm.auto import tqdm
@@ -19,8 +21,8 @@ from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import SelfForcingFlowMatchScheduler
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import SelfForcingFlowMatchScheduler
from fastvideo.pipelines import TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.training.training_utils import (EMA_FSDP,
@@ -84,35 +86,23 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self.last_step_only = getattr(training_args, 'last_step_only', False)
self.context_noise = getattr(training_args, 'context_noise', 0)
self.kv_cache1: list[dict[str, Any]] | None = None
self.crossattn_cache: list[dict[str, Any]] | None = None
logger.info("Self-forcing generator update ratio: %s",
self.dfake_gen_update_ratio)
logger.info("RANK: %s, exiting initialize_training_pipeline",
self.global_rank,
local_main_process_only=False)
def generate_and_sync_list(self, num_blocks: int, num_denoising_steps: int,
def generate_and_sync_list(self, num_blocks: int, start_timestep_index: int,
end_timestep_index: int,
device: torch.device) -> list[int]:
"""Generate and synchronize random exit flags across distributed processes."""
logger.info(
"RANK: %s, enter generate_and_sync_list blocks=%s steps=%s device=%s",
self.global_rank,
num_blocks,
num_denoising_steps,
str(device),
local_main_process_only=False)
rank = dist.get_rank() if dist.is_initialized() else 0
if rank == 0:
# Generate random indices
indices = torch.randint(low=0,
high=num_denoising_steps,
indices = torch.randint(low=start_timestep_index,
high=end_timestep_index,
size=(num_blocks, ),
device=device)
if self.last_step_only:
indices = torch.ones_like(indices) * (num_denoising_steps - 1)
indices = torch.ones_like(indices) * (end_timestep_index - 1)
else:
indices = torch.empty(num_blocks, dtype=torch.long, device=device)
@@ -120,12 +110,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
dist.broadcast(indices,
src=0) # Broadcast the random indices to all ranks
flags = indices.tolist()
logger.info(
"RANK: %s, exit generate_and_sync_list flags_len=%s first=%s",
self.global_rank,
len(flags),
flags[0] if len(flags) > 0 else None,
local_main_process_only=False)
return flags
def generator_loss(
@@ -216,7 +200,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
batch_size, num_generated_frames, *self.video_latent_shape[2:]
]
noise = torch.randn(noise_shape, device=self.device, dtype=dtype)
if training_batch.use_gt_trajectory and training_batch.trajectory_latents is not None:
noise = training_batch.trajectory_latents.to(self.device,
dtype=dtype)
else:
noise = torch.randn(noise_shape, device=self.device, dtype=dtype)
if self.sp_world_size > 1:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
@@ -252,7 +240,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
# Step 1: Initialize KV cache to all zeros
cache_frames = num_generated_frames + num_input_frames
self.kv_cache1, self.crossattn_cache = self._initialize_simulation_caches(
kv_cache1, crossattn_cache = self._initialize_simulation_caches(
batch_size, dtype, self.device, max_num_frames=cache_frames)
# Step 2: Cache context feature
@@ -286,9 +274,16 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
all_num_frames = [self.num_frame_per_block] * num_blocks
if self.independent_first_frame and initial_latent is None:
all_num_frames = [1] + all_num_frames
num_denoising_steps = len(self.denoising_step_list)
if training_batch.use_gt_trajectory and training_batch.trajectory_latents is not None:
start_timestep_index = training_batch.start_timestep_index
end_timestep_index = len(self.denoising_step_list)
else:
start_timestep_index = 0
end_timestep_index = len(self.denoising_step_list)
exit_flags = self.generate_and_sync_list(len(all_num_frames),
num_denoising_steps,
start_timestep_index,
end_timestep_index,
device=noise.device)
start_gradient_frame_index = max(0, num_output_frames - 21)
@@ -297,8 +292,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
num_input_frames:current_start_frame +
current_num_frames - num_input_frames]
index = start_timestep_index
current_timestep = self.denoising_step_list[index]
assert index < len(
self.denoising_step_list
), "Index is greater than the number of denoising steps"
# Step 3.1: Spatial denoising loop
for index, current_timestep in enumerate(self.denoising_step_list):
while index < len(self.denoising_step_list):
if self.same_step_across_blocks:
exit_flag = (index == exit_flags[0])
else:
@@ -329,8 +330,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=current_start_frame).permute(
@@ -369,8 +370,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=current_start_frame).permute(
@@ -389,8 +390,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=current_start_frame).permute(
@@ -404,6 +405,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
0, pred_flow.shape[:2])
break
index += 1
current_timestep = self.denoising_step_list[index]
# Step 3.2: record the model's output
output[:, current_start_frame:current_start_frame +
current_num_frames] = denoised_pred
@@ -416,6 +420,17 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
context_timestep).unflatten(0, denoised_pred.shape[:2])
with torch.no_grad():
if training_batch.video_latent is not None:
denoised_pred = torch.cat([
denoised_pred,
training_batch.
video_latent[:,
current_start_frame:current_start_frame +
current_num_frames],
torch.zeros_like(denoised_pred),
],
dim=2)
training_batch_temp = self._build_distill_input_kwargs(
denoised_pred, context_timestep,
training_batch.conditional_dict, training_batch)
@@ -430,8 +445,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
timestep=training_batch_temp.input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame * self.frame_seq_length,
start_frame=current_start_frame)
@@ -523,9 +538,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
min_num_frames, dtype=torch.float32, device=self.device)
# Clean up caches
assert self.kv_cache1 is not None
assert self.crossattn_cache is not None
self._reset_simulation_caches(self.kv_cache1, self.crossattn_cache)
assert kv_cache1 is not None
assert crossattn_cache is not None
self._reset_simulation_caches(kv_cache1, crossattn_cache)
return final_output if gradient_mask is not None else pred_image_or_video
@@ -538,28 +553,36 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
max_num_frames: int | None = None,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
num_transformer_blocks = len(self.transformer.blocks)
num_transformer_blocks = self.transformer.config.num_layers
latent_shape = self.video_latent_shape_sp
_, num_frames, _, height, width = latent_shape
_, p_h, p_w = self.transformer.patch_size
post_patch_height = height // p_h
post_patch_width = width // p_w
if isinstance(self.transformer.config.patch_size, tuple):
ph, pw = self.transformer.config.patch_size[
1], self.transformer.config.patch_size[2]
elif isinstance(self.transformer.config.patch_size, int):
ph, pw = self.transformer.config.patch_size, self.transformer.config.patch_size
else:
raise ValueError(
f"Unsupported patch size type: {type(self.transformer.config.patch_size)}"
)
post_patch_height = height // ph
post_patch_width = width // pw
frame_seq_length = post_patch_height * post_patch_width
self.frame_seq_length = frame_seq_length
# Get model configuration parameters - handle FSDP wrapping
num_attention_heads = getattr(self.transformer, 'num_attention_heads',
None)
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
None)
text_len = getattr(self.transformer, 'text_len', None)
num_attention_heads = self.transformer.config.num_attention_heads
attention_head_dim = self.transformer.config.attention_head_dim
text_len = getattr(self.transformer.config, 'text_len', None)
if max_num_frames is None:
max_num_frames = num_frames
num_max_frames = max(max_num_frames, num_frames)
kv_cache_size = num_max_frames * frame_seq_length
local_attn_size = getattr(self.transformer.config, 'local_attn_size',
-1)
kv_cache_size = num_max_frames * frame_seq_length if local_attn_size == -1 else local_attn_size * frame_seq_length
kv_cache = []
for _ in range(num_transformer_blocks):
@@ -979,6 +1002,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
logger.info("Starting training from scratch")
self.train_loader_iter = iter(self.train_dataloader)
if getattr(self, "train_dataloader_2", None) is not None:
self.train_loader_iter_2 = iter(self.train_dataloader_2)
step_times: deque[float] = deque(maxlen=100)
+80 -39
View File
@@ -2,6 +2,7 @@
from dataclasses import asdict
import math
import os
import gc
import time
from abc import ABC, abstractmethod
from collections import deque
@@ -25,7 +26,7 @@ from fastvideo.attention.backends.video_sparse_attn import (
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory_text_only, pyarrow_schema_t2v
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_local_torch_device, get_sp_group,
@@ -47,6 +48,7 @@ from fastvideo.training.training_utils import (
shard_latents_across_sp)
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
from fastvideo.optim.muon import get_muon_optimizer
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
@@ -89,6 +91,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
def set_schemas(self) -> None:
self.train_dataset_schema = pyarrow_schema_t2v
self.train_dataset_schema_2 = pyarrow_schema_ode_trajectory_text_only
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
@@ -108,7 +111,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
set_random_seed(self.seed + self.global_rank)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
@@ -127,16 +130,24 @@ class TrainingPipeline(LoRAPipeline, ABC):
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
# Parse betas from string format "beta1,beta2"
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.optimizer = get_muon_optimizer(
self.transformer,
lr=training_args.learning_rate,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
self.init_steps = 0
logger.info("optimizer: %s", self.optimizer)
@@ -156,13 +167,23 @@ class TrainingPipeline(LoRAPipeline, ABC):
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.optimizer_2 = get_muon_optimizer(
self.transformer_2,
lr=training_args.learning_rate,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
@@ -186,6 +207,19 @@ class TrainingPipeline(LoRAPipeline, ABC):
text_len, # type: ignore[attr-defined]
seed=self.seed)
if getattr(training_args, 'data_path_2', None) is not None:
self.train_dataset_2, self.train_dataloader_2 = build_parquet_map_style_dataloader(
training_args.data_path_2,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema_2,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=self.seed)
self.noise_scheduler = noise_scheduler
if self.training_args.boundary_ratio is not None:
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
@@ -575,13 +609,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed)
device=current_platform.device_name).manual_seed(self.seed +
self.global_rank)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
self.seed + self.global_rank)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
@@ -648,22 +684,25 @@ class TrainingPipeline(LoRAPipeline, ABC):
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
try:
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
context_len = int(training_batch.encoder_hidden_states.shape[1])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
context_len = int(training_batch.encoder_hidden_states.shape[1])
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
arch_config = self.training_args.pipeline_config.dit_config.arch_config
arch_config = self.training_args.pipeline_config.dit_config.arch_config
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
except:
pass
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
@@ -676,12 +715,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_visualization and step % self.training_args.visualization_steps == 0:
self.visualize_intermediate_latents(training_batch,
self.training_args, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
with self.profiler_controller.region(
"profiler_region_training_validation"):
if self.training_args.log_visualization:
self.visualize_intermediate_latents(
training_batch, self.training_args, step)
self._log_validation(self.transformer, self.training_args,
step)
gpu_memory_usage = current_platform.get_torch_device(
+191 -50
View File
@@ -256,8 +256,6 @@ def save_distillation_checkpoint(
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -289,8 +287,6 @@ def save_distillation_checkpoint(
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
if generator_ema_2 is not None:
generator_2_states["ema"] = generator_ema_2.state_dict()
generator_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
@@ -416,6 +412,67 @@ def save_distillation_checkpoint(
rank,
local_main_process_only=False)
# Persist EMA separately to avoid shape mismatches across ranks.
# Supports:
# - mode="rank0_full": save consolidated EMA only on rank 0
# - mode="local_shard": save per-rank EMA shard for each rank
try:
if generator_ema is not None and getattr(generator_ema, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema,
generator_transformer, rank,
save_dir, "generator_ema")
elif generator_ema is not None and getattr(generator_ema, "mode",
None) == "local_shard":
# Save per-rank shard
ema_dir_shard = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard, exist_ok=True)
ema_shard_path = os.path.join(ema_dir_shard,
f"generator_ema_rank{rank}.pt")
torch.save(generator_ema.state_dict(), ema_shard_path)
logger.info(
"rank: %s, saved generator EMA shard (local_shard) to %s",
rank,
ema_shard_path,
local_main_process_only=False)
# Also consolidate EMA to a single full-state file on rank 0 by applying EMA to the model and gathering
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema, generator_transformer, rank, save_dir,
"generator_ema")
except Exception as e:
logger.warning("rank: %s, failed saving EMA separately: %s", rank,
str(e))
try:
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema_2,
generator_transformer_2, rank,
save_dir, "generator_ema_2")
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "local_shard":
# Save per-rank shard for EMA_2
ema_dir_shard_2 = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard_2, exist_ok=True)
ema2_shard_path = os.path.join(ema_dir_shard_2,
f"generator_ema_2_rank{rank}.pt")
torch.save(generator_ema_2.state_dict(), ema2_shard_path)
logger.info(
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
rank,
ema2_shard_path,
local_main_process_only=False)
# Also consolidate EMA_2 to a single full-state file on rank 0
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema_2, generator_transformer_2, rank, save_dir,
"generator_ema_2")
except Exception as e:
logger.warning("rank: %s, failed saving EMA_2 separately: %s", rank,
str(e))
# Save generator model weights (consolidated) for inference
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
device=None)
@@ -453,46 +510,45 @@ def save_distillation_checkpoint(
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(
generator_transformer_2, device=None)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(generator_transformer_2,
device=None)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(inference_save_dir_2,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2,
"config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2, "config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
def load_checkpoint(transformer,
@@ -643,18 +699,38 @@ def load_distillation_checkpoint(
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
# Load EMA separately if saved in rank0_full mode
if generator_ema is not None:
try:
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully",
rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
if getattr(generator_ema, "mode", None) == "rank0_full":
ema_path = os.path.join(checkpoint_path, "ema",
"generator_ema.pt")
if rank == 0 and os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA (rank0_full) loaded from %s",
rank, ema_path)
elif rank == 0:
logger.info(
"rank: %s, generator EMA file not found at %s; skipping",
rank, ema_path)
elif getattr(generator_ema, "mode", None) == "local_shard":
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
f"generator_ema_rank{rank}.pt")
if os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA shard (local_shard) loaded from %s",
rank, ema_path)
else:
logger.info(
"rank: %s, generator EMA shard file not found at %s; skipping",
rank, ema_path)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank,
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
str(e))
# Load generator_2 distributed checkpoint (MoE support)
@@ -849,7 +925,7 @@ def load_distillation_checkpoint(
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * 0.476986
return latents * vae.config.scaling_factor
elif model_type == "wan":
latents_mean = torch.tensor(vae.latents_mean)
latents_std = 1.0 / torch.tensor(vae.latents_std)
@@ -1153,6 +1229,71 @@ def custom_to_hf_state_dict(
return new_state_dict
def _save_full_ema_safetensors_from_state(
state_dict: dict[str, Any],
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
output_path: str,
) -> None:
"""
Convert a training-format state_dict to HF format and save as safetensors.
"""
diffusers_state_dict = custom_to_hf_state_dict(state_dict,
reverse_param_names_mapping)
save_file(diffusers_state_dict, output_path)
def _save_rank0_full_ema_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
if rank != 0:
return
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
ema_state = ema.state_dict()
_save_full_ema_safetensors_from_state(ema_state,
module.reverse_param_names_mapping,
output_path)
logger.info("rank: %s, saved %s as consolidated EMA safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
def _consolidate_local_shard_ema_and_save_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
try:
# Temporarily apply EMA to the live (sharded) module and gather full CPU state on rank 0
with ema.apply_to_model(module):
cpu_state_full = gather_state_dict_on_cpu_rank0(module, device=None)
if rank == 0:
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
_save_full_ema_safetensors_from_state(
cpu_state_full, module.reverse_param_names_mapping, output_path)
logger.info(
"rank: %s, saved consolidated %s EMA (from local_shard) as safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
except Exception as ce:
logger.warning(
"rank: %s, failed consolidating %s EMA (local_shard): %s", rank,
base_name, str(ce))
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1: