Compare commits

...
Author SHA1 Message Date
JerryZhou54 32711250ad checkpoint 2025-09-09 02:48:02 +00:00
JerryZhou54 3d6cac57fc Fix runtime errors during training 2025-09-09 01:15:18 +00:00
Matthew Noto 57fd3d8159 new branch 2025-09-09 01:15:18 +00:00
JerryZhou54 1eacdd80de new branch 2025-09-09 01:15:18 +00:00
Matthew Noto 2ca24b3288 new branch 2025-09-08 19:20:32 +00:00
RandNMR73 328eb611c4 training refactor 2025-09-07 02:40:19 +00:00
RandNMR73 3408e20d7a Merge branch 'matthew/causal' of github.com:hao-ai-lab/FastVideo into matthew/causal 2025-09-07 00:17:20 +00:00
SolitaryThinker 4328fe1ebf fix warping 2025-09-05 10:20:10 +00:00
SolitaryThinker 6d0eba5789 warp timestep 2025-09-05 09:54:35 +00:00
RandNMR73 b3e76ae7dd Merge branch 'matthew/causal' of github.com:hao-ai-lab/FastVideo into matthew/causal 2025-09-02 00:57:20 +00:00
Matthew Noto 1ffd80ee51 training experiments in progress 2025-09-02 00:03:58 +00:00
Matthew Noto e84fdaedde add single example preprocessing 2025-09-01 03:49:33 +00:00
Matthew Noto dd0fe401c9 fix kv-cache + training loop 2025-09-01 01:52:11 +00:00
Matthew Noto 60eac9f18b clean-up + kv cache debugging 2025-08-31 08:31:19 +00:00
Matthew Noto a75bccb75d debugging kv cache 2025-08-31 05:30:21 +00:00
RandNMR73 c52ff91747 Merge branch 'matthew/causal' of github.com:hao-ai-lab/FastVideo into matthew/causal 2025-08-28 12:21:18 +00:00
RandNMR73 36bb0935a8 val data path fix 2025-08-28 12:18:41 +00:00
RandNMR73 d1e26abd63 Merge branch 'matthew/causal' of github.com:hao-ai-lab/FastVideo into matthew/causal 2025-08-28 11:53:46 +00:00
RandNMR73 b6f187f338 working training script 2025-08-28 11:32:53 +00:00
RandNMR73 de4938c3a7 working training script 2025-08-28 11:27:18 +00:00
RandNMR73 a9f7407228 debugging training 2025-08-28 08:06:31 +00:00
RandNMR73 258d1da0d3 training script testing 2025-08-28 05:03:12 +00:00
RandNMR73 17f6dff632 training branch 2025-08-28 04:49:51 +00:00
RandNMR73 e66f16057f training script testing 2025-08-28 04:42:51 +00:00
JerryZhou54 8acc7c3655 small fix 2025-08-28 04:10:51 +00:00
JerryZhou54 8dd0b6536d Fix pre-commit tests 2025-08-28 04:07:11 +00:00
JerryZhou54 1a79b30ea4 Fix inter-block issues 2025-08-28 03:56:37 +00:00
JerryZhou54 772ead0d34 Wan2.1 causal, few-step inference runnable 2025-08-28 03:56:37 +00:00
SolitaryThinker 1575102965 debugging model 2025-08-28 03:56:37 +00:00
SolitaryThinker b198ba5607 checkpoint 2025-08-28 03:56:34 +00:00
William Lin cf1942fd47 [dev] Will/causal infer (#764) 2025-08-28 03:55:16 +00:00
Wei Zhou b5519f1f91 Delete fastvideo/configs/models/dits/causal_wanvideo.py 2025-08-28 03:55:16 +00:00
JerryZhou54 a464f96b95 Fix file dir 2025-08-28 03:55:16 +00:00
JerryZhou54 363cf0d173 First version of causalwan 2025-08-28 03:55:15 +00:00
JerryZhou54 36371c5689 checkpoint 2025-08-28 03:55:15 +00:00
SolitaryThinker 8fea7c02b5 debugging model 2025-08-27 02:56:31 -07:00
SolitaryThinker e2b6f49879 checkpoint 2025-08-27 02:12:43 -07:00
William Lin 9f24aef7cf [dev] Will/causal infer (#764) 2025-08-26 23:13:20 -07:00
Wei Zhou 7a489da74d Delete fastvideo/configs/models/dits/causal_wanvideo.py 2025-08-26 16:33:25 -07:00
JerryZhou54 cf230dcccd Fix file dir 2025-08-26 23:32:30 +00:00
JerryZhou54 c4521e8953 First version of causalwan 2025-08-26 23:31:46 +00:00
JerryZhou54 026ee8d9f4 checkpoint 2025-08-26 19:49:57 +00:00
41 changed files with 3844 additions and 122 deletions
+3
View File
@@ -64,3 +64,6 @@ docs/source/distillation/examples/
!docs/source/_static/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
sf_output/
@@ -0,0 +1,139 @@
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# Configs
NUM_GPUS=8
# Model paths for DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/SFWan2.1-I2V/validation_better.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
--output_dir "/mnt/sharefs/users/hao.zhang/wl/sf_checkpoints/ode0_SFwan_t2v_finetune_sf_${lr}_c${critic_lr}"
--wandb_run_name "DEBUG${lr}_c${critic_lr}"
--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 832
--num_frames 81
--warp_denoising_step
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus 8 # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 8
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 50
--weight_only_checkpointing_steps 50
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
# Miscellaneous arguments
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 100
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--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.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks False # 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)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port $MASTER_PORT \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,165 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate wei-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# 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 CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=1
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="data/crush-smol-single_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "checkpoints/SFwan_t2v_finetune"
--max_train_steps 500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus 1 # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 10
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 50
--weight_only_checkpointing_steps 50
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
# Miscellaneous arguments
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 100
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--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.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks False # 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)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
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/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,43 @@
#!/bin/bash
# Download the full dataset first
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
# Create a single-example dataset for debugging
SINGLE_EXAMPLE_DIR="data/crush-smol-single"
mkdir -p "$SINGLE_EXAMPLE_DIR/videos"
# Copy the specific video that matches the validation.json style (macaron crushing)
cp "data/crush-smol/videos/7P02AihYkCU-Scene-005.mp4" "$SINGLE_EXAMPLE_DIR/videos/"
# Create a single-line videos.txt
echo "videos/7P02AihYkCU-Scene-005.mp4" > "$SINGLE_EXAMPLE_DIR/videos.txt"
# Create a single-line prompt.txt with the macaron crushing prompt
echo "PIKA_CRUSH A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out." > "$SINGLE_EXAMPLE_DIR/prompt.txt"
# Generate the JSON file and merge.txt for the single example
python scripts/dataset_preparation/prepare_json_file.py --data_folder "$SINGLE_EXAMPLE_DIR" --output "videos2caption.json"
# Create a validation.json that uses the same example for consistency
cat > "$SINGLE_EXAMPLE_DIR/validation.json" << 'EOF'
{
"data": [
{
"caption": "A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 81
}
]
}
EOF
echo "Single example dataset created at $SINGLE_EXAMPLE_DIR"
echo "Contains:"
echo "- 1 video: $(cat $SINGLE_EXAMPLE_DIR/videos.txt)"
echo "- 1 prompt: $(cat $SINGLE_EXAMPLE_DIR/prompt.txt)"
echo "- Validation file created with the same example for consistency"
@@ -0,0 +1,24 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
@@ -0,0 +1,29 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol-single/merge.txt"
OUTPUT_DIR="data/crush-smol-single_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 1 \
--flush_frequency 1 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
# Copy the validation.json to the output directory for consistency
cp "data/crush-smol-single/validation.json" "$OUTPUT_DIR/"
echo "Preprocessing completed. Validation file copied to $OUTPUT_DIR/"
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,31 @@
import os
import time
from fastvideo import VideoGenerator, SamplingParam
OUTPUT_PATH = "video_samples_causal"
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.
model_name = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
generator = VideoGenerator.from_pretrained(
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
text_encoder_cpu_offload=False,
dit_cpu_offload=False,
)
sampling_param = SamplingParam.from_pretrained(model_name)
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, sampling_param=sampling_param)
if __name__ == "__main__":
main()
+7 -1
View File
@@ -92,6 +92,12 @@ class WanVideoArchConfig(DiTArchConfig):
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Causal Wan
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 = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
@@ -103,4 +109,4 @@ class WanVideoArchConfig(DiTArchConfig):
class WanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanVideoArchConfig)
prefix: str = "Wan"
prefix: str = "Wan"
+2
View File
@@ -9,6 +9,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config,
SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
from fastvideo.logger import init_logger
@@ -34,6 +35,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
# Add other specific weight variants
}
+11
View File
@@ -138,3 +138,14 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
@dataclass
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
pass
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
is_causal: bool = True
flow_shift: int = 5
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
+15 -2
View File
@@ -17,6 +17,7 @@ from fastvideo.configs.sample.wan import (
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
SelfForcingWanT2V480PConfig,
)
# isort: on
from fastvideo.logger import init_logger
@@ -28,17 +29,29 @@ logger = init_logger(__name__)
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
# Add other specific weight variants
}
+8
View File
@@ -141,3 +141,11 @@ class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale_2: float = 3.5
num_inference_steps: int = 40
fps: int = 16
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
pass
+97
View File
@@ -591,6 +591,11 @@ class TrainingArgs(FastVideoArgs):
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
# DMD model paths - separate paths for each network
generator_model_path: str = "" # path for generator (student) model
real_score_model_path: str = "" # path for real score (teacher) model
fake_score_model_path: str = "" # path for fake score (critic) model
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
@@ -613,6 +618,7 @@ class TrainingArgs(FastVideoArgs):
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
# optimizer & scheduler
num_train_epochs: int = 0
@@ -644,6 +650,7 @@ class TrainingArgs(FastVideoArgs):
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
@@ -664,16 +671,29 @@ class TrainingArgs(FastVideoArgs):
# distillation args
generator_update_interval: int = 5
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
independent_first_frame: bool = False
enable_gradient_masking: bool = True
gradient_mask_last_n_frames: int = 21
validate_cache_structure: bool = False # Debug flag for cache validation
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
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -775,6 +795,20 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Directory to cache models")
# DMD model paths - separate paths for each network
parser.add_argument(
"--generator-model-path",
type=str,
help="Path to generator (student) model for DMD distillation")
parser.add_argument(
"--real-score-model-path",
type=str,
help="Path to real score (teacher) model for DMD distillation")
parser.add_argument(
"--fake-score-model-path",
type=str,
help="Path to fake score (critic) model for DMD distillation")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
@@ -845,6 +879,10 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument(
"--init-weights-from-safetensors",
type=str,
help="Path to safetensors file for initial weight loading")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
@@ -949,6 +987,10 @@ class TrainingArgs(FastVideoArgs):
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--betas",
type=str,
default=TrainingArgs.betas,
help="Betas for optimizer (format: 'beta1,beta2')")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
@@ -990,6 +1032,13 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument(
"--dfake-gen-update-ratio",
type=int,
default=TrainingArgs.dfake_gen_update_ratio,
help=
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
)
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
@@ -1006,6 +1055,11 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.fake_score_learning_rate,
help="Learning rate for fake score transformer")
parser.add_argument(
"--fake-score-betas",
type=str,
default=TrainingArgs.fake_score_betas,
help="Betas for fake score optimizer (format: 'beta1,beta2')")
parser.add_argument(
"--fake-score-lr-scheduler",
type=str,
@@ -1018,6 +1072,49 @@ class TrainingArgs(FastVideoArgs):
"--simulate-generator-forward",
action=StoreBoolean,
help="Whether to simulate generator forward to match inference")
parser.add_argument(
"--warp-denoising-step",
action=StoreBoolean,
help=
"Whether to warp denoising step according to the scheduler time shift"
)
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=TrainingArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
parser.add_argument(
"--independent-first-frame",
action=StoreBoolean,
help="Whether the first frame is independent in causal generation")
parser.add_argument(
"--enable-gradient-masking",
action=StoreBoolean,
help="Whether to enable frame-level gradient masking")
parser.add_argument(
"--gradient-mask-last-n-frames",
type=int,
default=TrainingArgs.gradient_mask_last_n_frames,
help="Number of last frames to enable gradients for")
parser.add_argument(
"--validate-cache-structure",
action=StoreBoolean,
help="Whether to validate KV cache structure (debug flag)")
parser.add_argument(
"--same-step-across-blocks",
action=StoreBoolean,
help="Whether to use the same exit timestep for all blocks")
parser.add_argument(
"--last-step-only",
action=StoreBoolean,
help="Whether to only use the last timestep for training")
parser.add_argument(
"--context-noise",
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
return parser
+47 -6
View File
@@ -9,6 +9,9 @@ import torch.nn.functional as F
from fastvideo.layers.custom_op import CustomOp
from fastvideo.platforms import current_platform
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@CustomOp.register("rms_norm")
class RMSNorm(CustomOp):
@@ -100,7 +103,13 @@ class ScaleResidual(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
"""Apply gated residual connection."""
return residual + x * gate
# logger.info("x.shape: %s", x.shape)
# if isinstance(gate, torch.Tensor):
# logger.info("gate.shape: %s", gate.shape)
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
return residual + (x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
@@ -172,11 +181,35 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
but before normalization)
"""
# Apply residual connection with gating
residual_output = residual + x * gate
# logger.info("x.shape: %s", x.shape)
if isinstance(gate, int):
# used by cross-attention, should be 1
assert gate == 1
residual_output = residual + x * gate
elif isinstance(gate, torch.Tensor):
# logger.info("gate.shape: %s", gate.shape)
if gate.dim() == 3:
# used by bidirectional self attention
residual_output = residual + x * gate
else:
assert gate.dim() == 4
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
residual_output = residual + (x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
# residual_output = residual + x * gate
else:
raise ValueError(f"Gate type {type(gate)} not supported")
# logger.info("residual_output.shape: %s", residual_output.shape)
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
modulated = normalized * (1.0 + scale) + shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
modulated = (normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
@@ -219,7 +252,15 @@ class LayerNormScaleShift(nn.Module):
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
if scale.dim() == 4:
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
if self.compute_dtype == torch.float32:
return (normalized.float().unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2).to(x.dtype)
else:
return (normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2)
else:
return normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
else:
return normalized * (1.0 + scale) + shift
+10
View File
@@ -29,7 +29,11 @@ import torch
from fastvideo.distributed.parallel_state import get_sp_group
from fastvideo.layers.custom_op import CustomOp
from fastvideo.logger import init_logger
logger = init_logger(__name__)
logger = init_logger(__name__)
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., :x.shape[-1] // 2]
@@ -267,6 +271,7 @@ def get_nd_rotary_pos_embed(
sp_rank: int = 0,
sp_world_size: int = 1,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
@@ -292,6 +297,9 @@ def get_nd_rotary_pos_embed(
full_grid = get_meshgrid_nd(
start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
if start_frame > 0:
full_grid[0] += start_frame
# Shard the grid if using sequence parallelism (sp_world_size > 1)
assert shard_dim < len(
rope_dim_list
@@ -370,6 +378,7 @@ def get_rotary_pos_embed(
interpolation_factor=1.0,
shard_dim: int = 0,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate rotary positional embeddings for the given sizes.
@@ -413,6 +422,7 @@ def get_rotary_pos_embed(
sp_rank=sp_rank,
sp_world_size=sp_world_size,
dtype=dtype,
start_frame=start_frame,
)
return freqs_cos, freqs_sin
+6 -2
View File
@@ -86,12 +86,16 @@ class TimestepEmbedder(nn.Module):
dtype=dtype)
self.freq_dtype = freq_dtype
def forward(self, t: torch.Tensor) -> torch.Tensor:
def forward(self,
t: torch.Tensor,
timestep_seq_len: int | None = None) -> torch.Tensor:
t_freq = timestep_embedding(t,
self.frequency_embedding_size,
self.max_period,
dtype=self.freq_dtype).to(
self.mlp.fc_in.weight.dtype)
if timestep_seq_len is not None:
t_freq = t_freq.unflatten(0, (1, timestep_seq_len))
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb
@@ -172,4 +176,4 @@ def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
return imgs
+710
View File
@@ -0,0 +1,710 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import numpy as np
import torch
import torch.nn as nn
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
from torch.nn.attention.flex_attention import BlockMask
# wan 1.3B model has a weird channel / head configurations and require max-autotune to work with flexattention
# see https://github.com/pytorch/pytorch/issues/133254
# change to default for other models
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
import torch.distributed as dist
import fastvideo.envs as envs
from fastvideo.attention import (DistributedAttention,
LocalAttention)
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
def __init__(self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm=True,
eps=1e-6,
parallel_attention=False) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.parallel_attention = parallel_attention
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Scaled dot product attention
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
seq_lens(Tensor): Shape [B]
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
if cache_start is None:
cache_start = current_start
cos, sin = freqs_cis
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
if kv_cache is None:
# Padding for flex attention
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_roped_query = torch.cat(
[roped_query,
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(
[roped_key, torch.zeros([k.shape[0], padded_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_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
x = 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)
else:
frame_seqlen = q.shape[1]
current_end = current_start + roped_query.shape[1]
sink_tokens = self.sink_size * frame_seqlen
# 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 = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
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["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# 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
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
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
kv_cache["k"] = kv_cache["k"].clone()
kv_cache["v"] = kv_cache["v"].clone()
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
roped_query,
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
return x
class CausalWanTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = CausalWanSelfAttention(
dim,
num_heads,
local_attn_size=local_attn_size,
sink_size=sink_size,
qk_norm=qk_norm,
eps=eps)
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.local_attn_size = local_attn_size
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
print("QK Norm type not supported")
raise Exception
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
) -> torch.Tensor:
# logger.info("temb.shape: %s", temb.shape)
num_frames = temb.shape[1]
# logger.info("first hidden_states.shape: %s", hidden_states.shape)
# logger.info("num_frames: %s", num_frames)
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
frame_seqlen = hidden_states.shape[1] // temb.shape[1]
# logger.info("frame_seqlen: %s", frame_seqlen)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
# logger.info("e.shape: %s", e.shape)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
# logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("scale_msa.shape: %s", scale_msa.shape)
# logger.info("shift_msa.shape: %s", shift_msa.shape)
norm_hidden_states_unflattened = self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
# logger.info("norm_hidden_states_unflattened.shape: %s", norm_hidden_states_unflattened.shape)
# norm_hidden_states = (self.norm1(hidden_states.float()) *
# (1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = (norm_hidden_states_unflattened *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
# logger.info("1 norm_hidden_states.shape: %s", norm_hidden_states.shape)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
attn_output = self.attn1(query, key, value, freqs_cis, block_mask, kv_cache, current_start, cache_start)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
# logger.info("after self_attn_residual_norm norm_hidden_states.shape: %s", norm_hidden_states.shape)
# logger.info("after self_attn_residual_norm hidden_states.shape: %s", hidden_states.shape)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None,
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
# logger.info("after cross_attn_residual_norm norm_hidden_states.shape: %s", norm_hidden_states.shape)
# logger.info("after cross_attn_residual_norm hidden_states.shape: %s", hidden_states.shape)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
# logger.info("after mlp_residual norm_hidden_states.shape: %s", norm_hidden_states.shape)
# logger.info("after mlp_residual hidden_states.shape: %s", hidden_states.shape)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class CausalWanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
param_names_mapping = WanVideoConfig().param_names_mapping
reverse_param_names_mapping = WanVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = WanVideoConfig().lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
self.local_attn_size = config.local_attn_size
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
CausalWanTransformerBlock(inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.local_attn_size,
config.sink_size,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# Debug: Log configuration values
proj_out_dim = config.out_channels * math.prod(config.patch_size)
self.proj_out = nn.Linear(inner_dim, proj_out_dim)
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.independent_first_frame = False
self.__post_init__()
@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
) -> 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
# we do right padding to get to a multiple of 128
padded_length = math.ceil(total_length / 128) * 128 - total_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=0,
end=total_length,
step=frame_seqlen * num_frame_per_block,
device=device
)
for tmp in frame_indices:
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)
else:
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx)
# 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_length + 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_inference(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
kv_cache: dict = None,
crossattn_cache: dict = None,
current_start: int = 0,
cache_start: int = 0,
start_frame: int = 0,
**kwargs) -> torch.Tensor:
r"""
Run the diffusion model with kv caching.
See Algorithm 2 of CausVid paper https://arxiv.org/abs/2412.07772 for details.
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
# logger.info("forward inference hidden_states.shape: %s", hidden_states.shape)
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame # Assume that start_frame is 0 when kv_cache is None
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# logger.info("forward inference flattened and transposed hidden_states.shape: %s", hidden_states.shape)
# logger.info("timestep shape: %s", timestep.shape)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
encoder_hidden_states = encoder_hidden_states.to(
orig_dtype) if current_platform.is_mps(
) else encoder_hidden_states # cast to orig_dtype for MPS
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
for block_index, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
causal_kwargs = {
"kv_cache": kv_cache[block_index],
"current_start": current_start,
"cache_start": cache_start,
"block_mask": self.block_mask
}
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
**causal_kwargs)
else:
causal_kwargs = {
"kv_cache": kv_cache[block_index],
"crossattn_cache": crossattn_cache[block_index],
"current_start": current_start,
"cache_start": cache_start,
"block_mask": self.block_mask
}
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
**causal_kwargs)
# 5. Output norm, projection & unpatchify
# logger.info("===== INFERENCE 5. Output norm, projection & unpatchify")
# logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("temb.shape: %s", temb.shape)
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
# logger.info("WTFWTF train temb.shape: %s", temb.shape)
# logger.info("WTFWTF train self.scale_shift_table.shape: %s", self.scale_shift_table.shape)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
def _forward_train(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
start_frame: int = 0,
**kwargs) -> torch.Tensor:
# logger.info("===== forward train hidden_states.shape: %s", hidden_states.shape)
# logger.info("===== forward train timestep.shape: %s", timestep.shape)
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask 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
)
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# logger.info("forward train flattened and transposed hidden_states.shape: %s", hidden_states.shape)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
# logger.info("forward train timestep_proj.shape: %s", timestep_proj.shape)
# logger.info("forward train timestep.shape: %s", timestep.shape)
# logger.info("forward train temb.shape: %s", temb.shape)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
encoder_hidden_states = encoder_hidden_states.to(
orig_dtype) if current_platform.is_mps(
) else encoder_hidden_states # cast to orig_dtype for MPS
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
else:
for block_index, block in enumerate(self.blocks):
logger.info("===== TRAIN block %d", block_index)
logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("encoder_hidden_states.shape: %s", encoder_hidden_states.shape)
logger.info("timestep_proj.shape: %s", timestep_proj.shape)
# logger.info("freqs_cis.shape: %s", freqs_cis.shape)
# logger.info("block_mask.shape: %s", self.block_mask.shape)
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
# 5. Output norm, projection & unpatchify
# logger.info("===== TRAIN 5. Output norm, projection & unpatchify")
# logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("temb.shape: %s", temb.shape)
# shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
# logger.info("WTFWTF train temb.shape: %s", temb.shape)
# logger.info("WTFWTF train self.scale_shift_table.shape: %s", self.scale_shift_table.shape)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
# logger.info("DEBUG scale.shape: %s", scale.shape)
# logger.info("DEBUG shift.shape: %s", shift.shape)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
# logger.info("DEBUG after proj_out hidden_states.shape: %s", hidden_states.shape)
# logger.info(f"DEBUG reshape dimensions: batch_size={batch_size}, post_patch_num_frames={post_patch_num_frames}")
# logger.info(f"DEBUG reshape dimensions: post_patch_height={post_patch_height}, post_patch_width={post_patch_width}")
# logger.info(f"DEBUG patch dimensions: p_t={p_t}, p_h={p_h}, p_w={p_w}")
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
def forward(
self,
*args,
**kwargs
):
if kwargs.get('kv_cache', None) is not None:
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
+59 -11
View File
@@ -81,8 +81,9 @@ class WanTimeTextImageEmbedding(nn.Module):
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep)
temb = self.time_embedder(timestep, timestep_seq_len)
timestep_proj = self.time_modulation(temb)
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
@@ -145,7 +146,7 @@ class WanSelfAttention(nn.Module):
class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens):
def forward(self, x, context, context_lens, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -156,8 +157,20 @@ class WanT2VCrossAttention(WanSelfAttention):
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
@@ -307,9 +320,24 @@ class WanTransformerBlock(nn.Module):
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
@@ -637,9 +665,21 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten() # batch_size * seq_len
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
# batch_size, 6, inner_dim
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -676,8 +716,15 @@ class WanTransformer3DModel(CachableDiT):
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
if temb.dim() == 3:
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2)
shift = shift.squeeze(2)
scale = scale.squeeze(2)
else:
# batch_size, inner_dim
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -781,3 +828,4 @@ class WanTransformer3DModel(CachableDiT):
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
@@ -430,6 +430,16 @@ class TransformerLoader(ComponentLoader):
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Check if we should use custom initialization weights
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
fastvideo_args.training_mode and
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
if use_custom_weights:
logger.info("Using custom initialization weights from: %s", custom_weights_path)
safetensors_list = [custom_weights_path]
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
+7
View File
@@ -248,6 +248,13 @@ def load_model_from_full_model_state_dict(
sharded_sd = {}
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
full_sd_iterator, param_names_mapping) # type: ignore
print(custom_param_sd.keys())
print("--------------------------------")
print("--------------------------------")
print("--------------------------------")
print("--------------------------------")
print("--------------------------------")
print(meta_sd.keys())
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
+2
View File
@@ -25,12 +25,14 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel")
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
}
_TEXT_ENCODER_MODELS = {
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
noise: torch.Tensor,
timestep: torch.IntTensor,
) -> torch.Tensor:
"""
Args:
clean_latent: the clean latent with shape [B, C, H, W],
where B is batch_size or batch_size * num_frames
noise: the noise with shape [B, C, H, W]
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
Returns:
the corrupted latent with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == clean_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(clean_latent.shape[0])
else:
assert timestep.numel() == clean_latent.shape[0]
else:
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
self.sigmas = self.sigmas.to(noise.device)
timestep = timestep.expand(clean_latent.shape[0])
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
@@ -650,4 +673,4 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
return sample
def __len__(self) -> int:
return self.config.num_train_timesteps
return self.config.num_train_timesteps
+46
View File
@@ -6,6 +6,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):
"""
@@ -137,3 +140,46 @@ def modulate(x: torch.Tensor,
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(
1) # type: ignore[union-attr]
def pred_noise_to_pred_video(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
noise_input_latent = noise_input_latent.float().to(device)
sigmas = scheduler.sigmas.float().to(device)
timesteps = scheduler.timesteps.float().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
@@ -0,0 +1,69 @@
# 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.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
CausalDMDDenosingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
# isort: on
logger = init_logger(__name__)
class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
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")],
tokenizers=[self.get_module("tokenizer")],
))
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", None)))
self.add_stage(stage_name="denoising_stage",
stage=CausalDMDDenosingStage(
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 = WanCausalDMDPipeline
@@ -136,9 +136,11 @@ class ComposedPipelineBase(ABC):
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
logger.info("training args in from_pretrained: %s", fastvideo_args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
for key, value in kwargs.items():
+1 -1
View File
@@ -242,4 +242,4 @@ class TrainingBatch:
@dataclass
class PreprocessBatch(ForwardBatch):
video_loader: list["VideoDecoder"] = field(default_factory=list)
video_file_name: list[str] = field(default_factory=list)
video_file_name: list[str] = field(default_factory=list)
+1
View File
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
}
+2
View File
@@ -7,6 +7,7 @@ complete diffusion pipelines.
"""
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (DenoisingStage,
@@ -30,6 +31,7 @@ __all__ = [
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
@@ -0,0 +1,414 @@
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.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
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 CausalDMDDenosingStage(DenoisingStage):
"""
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
self.scheduler = FlowMatchEulerDiscreteScheduler(shift=8.0)
# KV and cross-attention cache state (initialized on first forward)
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
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
try:
self.local_attn_size = getattr(self.transformer.model,
"local_attn_size",
-1) # type: ignore
except Exception:
self.local_attn_size = -1
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]
patch_ratio = self.transformer.config.arch_config.patch_size[
-1] * self.transformer.config.arch_config.patch_size[-2]
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
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
# "encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
# 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
if self.kv_cache1 is None:
self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device)
else:
assert self.crossattn_cache is not None
# reset cross-attention cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index][
"is_init"] = False # type: ignore
# reset kv cache pointers
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index][
"global_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
self.kv_cache1[block_index][
"local_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
# Optional: cache context features from provided image latents prior to generation
current_start_frame = 0
if getattr(batch, "image_latent", None) is not None:
image_latent = batch.image_latent
assert image_latent is not None
input_frames = image_latent.shape[2]
# timestep zero (or configured context noise) for cache warm-up
t_zero = torch.zeros([latents.shape[0]],
device=latents.device,
dtype=torch.long)
if independent_first_frame and input_frames >= 1:
# warm-up with the very first frame independently
image_first_btchw = image_latent[:, :, :1, :, :].to(
target_dtype).permute(0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
image_first_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += 1
remaining_frames = input_frames - 1
else:
remaining_frames = input_frames
# process remaining input frames in blocks of num_frame_per_block
while remaining_frames > 0:
block = min(self.num_frames_per_block, remaining_frames)
ref_btchw = image_latent[:, :, current_start_frame:
current_start_frame +
block, :, :].to(target_dtype).permute(
0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
ref_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += block
remaining_frames -= block
# Base position offset from any cache warm-up
pos_start_base = current_start_frame
# Determine block sizes
if not independent_first_frame or (independent_first_frame
and batch.image_latent is not None):
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
else:
if (t - 1) % self.num_frames_per_block != 0:
raise ValueError(
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
)
num_blocks = (t - 1) // self.num_frames_per_block
block_sizes = [1] + [self.num_frames_per_block] * num_blocks
start_index = 0
# 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.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.expand(latent_model_input.shape[0])
# t_expand = t_cur * torch.ones((latent_model_input.shape[0], 1), device=latent_model_input.device, dtype=torch.long)
# t_expand = t_expand.repeat(1, self.sliding_window_num_frames)
# 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], 1), device=latent_model_input.device, dtype=torch.long)
pred_noise_btchw = self.transformer(
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
).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]],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
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):
t_expanded_context = t_context * torch.ones((context_bcthw.shape[0], 1), device=context_bcthw.device, dtype=torch.long)
_ = self.transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += current_num_frames
batch.latents = latents
return batch
def _initialize_kv_cache(self, batch_size, dtype, device) -> None:
"""
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
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames
for _ in range(self.num_transformer_blocks):
kv_cache1.append({
"k":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
self.kv_cache1 = kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> None:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
crossattn_cache = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = self.transformer.attention_head_dim
for _ in range(self.num_transformer_blocks):
crossattn_cache.append({
"k":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"is_init":
False,
})
self.crossattn_cache = crossattn_cache
+1 -2
View File
@@ -778,8 +778,7 @@ class DmdDenoisingStage(DenoisingStage):
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
from fastvideo.training.training_utils import (
pred_noise_to_pred_video)
from fastvideo.models.utils import pred_noise_to_pred_video
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
+397 -69
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import json
import os
import time
from abc import abstractmethod
@@ -11,6 +12,7 @@ from typing import Any
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torch.nn.functional as F
import torchvision
from einops import rearrange
@@ -29,16 +31,18 @@ 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)
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases, get_scheduler,
load_distillation_checkpoint, pred_noise_to_pred_video,
save_distillation_checkpoint, shift_timestep)
from fastvideo.utils import is_vsa_available, set_random_seed
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
import wandb # isort: skip
@@ -87,9 +91,27 @@ class DistillationPipeline(TrainingPipeline):
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
# self.transformer is the generator model
self.real_score_transformer = self.get_module("real_score_transformer")
self.fake_score_transformer = self.get_module("fake_score_transformer")
if training_args.real_score_model_path:
logger.info(
f"Loading real score transformer from: {training_args.real_score_model_path}"
)
self.real_score_transformer = self.load_module_from_path(
training_args.real_score_model_path, "transformer",
training_args)
else:
self.real_score_transformer = self.get_module(
"real_score_transformer")
if training_args.fake_score_model_path:
logger.info(
f"Loading fake score transformer from: {training_args.fake_score_model_path}"
)
self.fake_score_transformer = self.load_module_from_path(
training_args.fake_score_model_path, "transformer",
training_args)
else:
self.fake_score_transformer = self.get_module(
"fake_score_transformer")
self.real_score_transformer.requires_grad_(False)
self.real_score_transformer.eval()
@@ -116,10 +138,13 @@ 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=(0.9, 0.999),
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -147,8 +172,19 @@ class DistillationPipeline(TrainingPipeline):
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps",
len(self.denoising_step_list))
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],
dtype=torch.float32))).cuda()
self.denoising_step_list = timesteps[1000 -
self.denoising_step_list]
logger.info("Warping denoising_step_list")
self.denoising_step_list = self.denoising_step_list.to(
get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps: %s",
len(self.denoising_step_list), self.denoising_step_list)
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.min_timestep = int(self.training_args.min_timestep_ratio *
@@ -158,6 +194,82 @@ class DistillationPipeline(TrainingPipeline):
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
self.generator_ema = 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(
f"Initialized generator EMA with decay={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"):
"""
Load a module from a specific path using the same loading logic as the pipeline.
Args:
model_path: Path to the model
module_type: Type of module to load (e.g., "transformer")
training_args: Training arguments
Returns:
The loaded module
"""
logger.info(f"Loading {module_type} from custom path: {model_path}")
# Set flag to prevent custom weight loading for teacher/critic models
training_args._loading_teacher_critic_model = True
try:
from fastvideo.models.loader.component_loader import (
PipelineComponentLoader)
# Download the model if it's a Hugging Face model ID
local_model_path = maybe_download_model(model_path)
logger.info(f"Model downloaded/found at: {local_model_path}")
config = verify_model_config_and_directory(local_model_path)
if module_type not in config:
if hasattr(self, '_extra_config_module_map'
) and module_type in self._extra_config_module_map:
extra_module = self._extra_config_module_map[module_type]
if extra_module in config:
module_type = extra_module
logger.info(f"Using {extra_module} for {module_type}")
else:
raise ValueError(
f"Module {module_type} not found in config at {local_model_path}"
)
else:
raise ValueError(
f"Module {module_type} not found in config at {local_model_path}"
)
module_info = config[module_type]
if module_info is None:
raise ValueError(
f"Module {module_type} has null value in config at {local_model_path}"
)
transformers_or_diffusers, architecture = module_info
component_path = os.path.join(local_model_path, module_type)
module = PipelineComponentLoader.load_module(
module_name=module_type,
component_model_path=component_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=training_args,
)
logger.info(
f"Successfully loaded {module_type} from {component_path}")
return module
finally:
# Always clean up the flag
if hasattr(training_args, '_loading_teacher_critic_model'):
delattr(training_args, '_loading_teacher_critic_model')
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
"""Initialize validation pipeline - must be implemented by subclasses."""
@@ -174,6 +286,110 @@ class DistillationPipeline(TrainingPipeline):
return training_batch
def apply_ema_to_model(self, model):
"""Apply EMA weights to the model for validation or inference."""
if self.generator_ema is not None:
with self.generator_ema.apply_to_model(model):
return model
return model
def get_ema_model_copy(self):
"""Get a copy of the model with EMA weights applied."""
if self.generator_ema is not None:
ema_model = copy.deepcopy(self.transformer)
self.generator_ema.copy_to_unwrapped(ema_model)
return ema_model
return None
def is_ema_ready(self, current_step: int = None):
"""Check if EMA is ready for use (after ema_start_step)."""
if current_step is None:
current_step = getattr(self, 'current_trainstep', 0)
return (self.generator_ema is not None
and current_step >= self.training_args.ema_start_step)
def save_ema_weights(self, output_dir: str, step: int):
"""Save EMA weights separately for inference purposes."""
if self.generator_ema is None:
logger.warning("Cannot save EMA weights: EMA not initialized")
return
if not self.is_ema_ready():
logger.warning(
"Cannot save EMA weights: EMA not ready yet (step < ema_start_step)"
)
return
try:
ema_model = self.get_ema_model_copy()
if ema_model is None:
logger.warning("Failed to create EMA model copy")
return
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
os.makedirs(ema_save_dir, exist_ok=True)
# save as diffusers format
from safetensors.torch import save_file
from fastvideo.training.training_utils import (
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
cpu_state = gather_state_dict_on_cpu_rank0(ema_model, device=None)
if self.global_rank == 0:
weight_path = os.path.join(
ema_save_dir, "diffusion_pytorch_model.safetensors")
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, ema_model.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
config_dict = ema_model.hf_config
if "dtype" in config_dict:
del config_dict["dtype"]
config_path = os.path.join(ema_save_dir, "config.json")
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info(f"EMA weights saved to {weight_path}")
del ema_model
except Exception as e:
logger.error(f"Failed to save EMA weights: {str(e)}")
def get_ema_stats(self):
"""Get EMA statistics for monitoring."""
if self.generator_ema is None:
return {
"ema_enabled": False,
"ema_decay": None,
"ema_start_step": self.training_args.ema_start_step,
"ema_ready": False,
"ema_step": self.current_trainstep,
}
return {
"ema_enabled": True,
"ema_decay": self.training_args.ema_decay,
"ema_start_step": self.training_args.ema_start_step,
"ema_ready": self.is_ema_ready(),
"ema_step": self.current_trainstep,
}
def reset_ema(self):
"""Reset EMA to current model weights."""
if self.generator_ema is not None:
logger.info("Resetting EMA to current model weights")
self.generator_ema.update(self.transformer)
# Force update to current weights by setting decay to 0 temporarily
original_decay = self.generator_ema.decay
self.generator_ema.decay = 0.0
self.generator_ema.update(self.transformer)
self.generator_ema.decay = original_decay
logger.info("EMA reset completed")
else:
logger.warning("Cannot reset EMA: EMA not initialized")
def _build_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
@@ -513,16 +729,17 @@ class DistillationPipeline(TrainingPipeline):
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
}
unconditional_dict = {
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
if getattr(self, "negative_prompt_embeds", None) is not None:
unconditional_dict = {
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
training_batch.unconditional_dict = unconditional_dict
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
training_batch.conditional_dict = conditional_dict
training_batch.unconditional_dict = unconditional_dict
training_batch.raw_latent_shape = training_batch.latents.shape
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
self.video_latent_shape = training_batch.latents.shape
@@ -587,6 +804,10 @@ class DistillationPipeline(TrainingPipeline):
self._clip_model_grad_norm_(batch_gen, self.transformer)
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
avg_dmd_loss = torch.tensor(total_dmd_loss /
gradient_accumulation_steps,
device=self.device)
@@ -637,7 +858,8 @@ class DistillationPipeline(TrainingPipeline):
self.transformer, self.fake_score_transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.fake_score_optimizer, self.train_dataloader, self.lr_scheduler,
self.fake_score_lr_scheduler, self.noise_random_generator)
self.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
if resumed_step > 0:
self.init_steps = resumed_step
@@ -668,6 +890,14 @@ class DistillationPipeline(TrainingPipeline):
sum(p.numel()
for p in self.fake_score_transformer.parameters()) / 1e9)
if self.generator_ema is not None:
logger.info(" Generator EMA enabled with decay: %s",
self.training_args.ema_decay)
logger.info(" Generator EMA start step: %s",
self.training_args.ema_start_step)
else:
logger.info(" Generator EMA disabled")
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
training_args.inference_mode = True
@@ -699,6 +929,18 @@ class DistillationPipeline(TrainingPipeline):
transformer.eval()
# Optionally use EMA model for validation if available and ready
use_ema_for_validation = (self.training_args.use_ema
and self.is_ema_ready(global_step))
if use_ema_for_validation:
logger.info("Using EMA model for validation")
validation_transformer = self.transformer
ema_context = self.generator_ema.apply_to_model(
validation_transformer)
else:
validation_transformer = transformer
ema_context = None
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
@@ -714,50 +956,98 @@ class DistillationPipeline(TrainingPipeline):
step_videos: list[np.ndarray] = []
step_captions: list[str] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
validation_batch,
num_inference_steps)
if ema_context is not None:
with ema_context:
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch,
num_inference_steps)
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
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]
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
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]
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
logger.info(
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
else:
# Use original transformer without EMA
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch,
num_inference_steps)
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
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]
logger.info(
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Log validation results for this step
world_group = get_world_group()
@@ -834,16 +1124,16 @@ class DistillationPipeline(TrainingPipeline):
latents.dtype)
else:
latents += self.vae.shift_factor
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, latents
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, latents
# Process DMD training data if available - use decode_stage instead of self.vae.decode
if 'generator_pred_video' in dmd_latents_vis_dict:
@@ -904,6 +1194,10 @@ class DistillationPipeline(TrainingPipeline):
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", seed)
# Initialize current_trainstep for EMA ready checks
#TODO: check if needed
self.current_trainstep = self.init_steps
# Resume from checkpoint if specified (this will restore random states)
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -947,6 +1241,14 @@ class DistillationPipeline(TrainingPipeline):
self.current_trainstep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
if (step >= self.training_args.ema_start_step) and \
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
self.generator_ema = EMA_FSDP(
self.transformer, decay=self.training_args.ema_decay)
logger.info(
f"Created generator EMA at step {step} with decay={self.training_args.ema_decay}"
)
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
@@ -960,11 +1262,19 @@ class DistillationPipeline(TrainingPipeline):
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"total_loss": f"{total_loss:.4f}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"total_loss":
f"{total_loss:.4f}",
"generator_loss":
f"{generator_loss:.4f}",
"fake_score_loss":
f"{fake_score_loss:.4f}",
"step_time":
f"{step_time:.2f}s",
"grad_norm":
grad_norm,
"ema":
"✓" if (self.generator_ema is not None and self.is_ema_ready())
else "✗",
})
progress_bar.update(1)
@@ -992,6 +1302,15 @@ class DistillationPipeline(TrainingPipeline):
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
if self.generator_ema is not None:
log_data["ema_enabled"] = True
log_data["ema_decay"] = self.training_args.ema_decay
else:
log_data["ema_enabled"] = False
ema_stats = self.get_ema_stats()
log_data.update(ema_stats)
if training_batch.dmd_latent_vis_dict:
dmd_additional_logs = {
"generator_timestep":
@@ -1023,7 +1342,8 @@ class DistillationPipeline(TrainingPipeline):
self.global_rank, self.training_args.output_dir, step,
self.optimizer, self.fake_score_optimizer,
self.train_dataloader, self.lr_scheduler,
self.fake_score_lr_scheduler, self.noise_random_generator)
self.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
if self.transformer:
self.transformer.train()
@@ -1040,7 +1360,11 @@ class DistillationPipeline(TrainingPipeline):
self.global_rank,
self.training_args.output_dir,
f"{step}_weight_only",
only_save_generator_weight=True)
only_save_generator_weight=True,
generator_ema=self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
if self.training_args.log_visualization:
@@ -1060,7 +1384,11 @@ class DistillationPipeline(TrainingPipeline):
self.training_args.output_dir, self.training_args.max_train_steps,
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
self.lr_scheduler, self.fake_score_lr_scheduler,
self.noise_random_generator)
self.noise_random_generator, self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir,
self.training_args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
@@ -0,0 +1,989 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import logging
from typing import Any
from collections import deque
import time
import torch
import torch.nn.functional as F
import wandb
from tqdm.auto import tqdm
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines import TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
EMA_FSDP,
save_distillation_checkpoint,
)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.distributed import get_world_group
import torch.distributed as dist
import numpy as np
from fastvideo.utils import set_random_seed, is_vsa_available
import fastvideo.envs as envs
from einops import rearrange
logger = init_logger(__name__)
vsa_available = is_vsa_available()
class SelfForcingDistillationPipeline(DistillationPipeline):
"""
A self-forcing distillation pipeline that alternates between training
the generator and critic based on the self-forcing methodology.
This implementation follows the self-forcing approach where:
1. Generator and critic are trained in alternating steps
2. Generator loss uses DMD-style loss with the critic as fake score
3. Critic loss trains the fake score model to distinguish real vs fake
"""
def initialize_training_pipeline(self, training_args: TrainingArgs):
"""Initialize the self-forcing training pipeline."""
logger.info("Initializing self-forcing distillation pipeline...")
super().initialize_training_pipeline(training_args)
self.dfake_gen_update_ratio = getattr(training_args, 'dfake_gen_update_ratio', 5)
# Self-forcing specific properties
self.num_frame_per_block = getattr(training_args, 'num_frame_per_block', 3)
self.independent_first_frame = getattr(training_args, 'independent_first_frame', False)
self.same_step_across_blocks = getattr(training_args, 'same_step_across_blocks', False)
self.last_step_only = getattr(training_args, 'last_step_only', False)
self.context_noise = getattr(training_args, 'context_noise', 0)
# Calculate frame sequence length - this will be set properly in _prepare_dit_inputs
self.frame_seq_length = 1560 # TODO: Calculate this dynamically based on patch size
# Cache references (will be initialized per forward pass)
self.kv_cache1 = None
self.crossattn_cache = None
logger.info(f"Self-forcing generator update ratio: {self.dfake_gen_update_ratio}")
def generate_and_sync_list(self, num_blocks, num_denoising_steps, device):
"""Generate and synchronize random exit flags across distributed processes."""
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,
size=(num_blocks,),
device=device
)
if self.last_step_only:
indices = torch.ones_like(indices) * (num_denoising_steps - 1)
else:
indices = torch.empty(num_blocks, dtype=torch.long, device=device)
if dist.is_initialized():
dist.broadcast(indices, src=0) # Broadcast the random indices to all ranks
return indices.tolist()
def generator_loss(self, training_batch: TrainingBatch) -> tuple[torch.Tensor, dict[str, Any]]:
"""
Compute generator loss using DMD-style approach.
The generator tries to fool the critic (fake_score_transformer).
"""
with set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video = self._generator_multi_step_simulation_forward(training_batch)
else:
generator_pred_video = self._generator_forward(training_batch)
with set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=training_batch
)
log_dict = {
"dmdtrain_gradient_norm": torch.tensor(0.0, device=self.device)
}
return dmd_loss, log_dict
def critic_loss(self, training_batch: TrainingBatch) -> tuple[torch.Tensor, dict[str, Any]]:
"""
Compute critic loss using flow matching between noise and generator output.
The critic learns to predict the flow from noise to the generator's output.
"""
updated_batch, flow_matching_loss = self.faker_score_forward(training_batch)
training_batch.fake_score_latent_vis_dict = updated_batch.fake_score_latent_vis_dict
log_dict = {}
return flow_matching_loss, log_dict
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
"""Forward pass through generator with KV cache support for causal generation."""
latents = training_batch.latents
dtype = latents.dtype
batch_size = latents.shape[0]
# Step 1: Sample a timestep from denoising_step_list
index = torch.randint(0,
len(self.denoising_step_list), [1],
device=self.device,
dtype=torch.long)
timestep = self.denoising_step_list[index]
training_batch.dmd_latent_vis_dict["generator_timestep"] = timestep
# Step 2: Initialize KV cache and cross-attention cache for causal generation
kv_cache, crossattn_cache = self._initialize_simulation_caches(batch_size, dtype, self.device)
if getattr(self.training_args, 'validate_cache_structure', False):
self._validate_cache_structure(kv_cache, crossattn_cache, batch_size)
# Step 3: Add noise to latents
noise = torch.randn(self.video_latent_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",
n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
noisy_latent = self.noise_scheduler.add_noise(latents.flatten(0, 1),
noise.flatten(0, 1),
torch.tensor([timestep], device=noise.device))
# Step 4: Build input kwargs with KV cache support
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
# Step 5: Forward pass with KV cache if available
if hasattr(self.transformer, '_forward_inference'):
# Use causal inference forward with KV cache
pred_noise = self.transformer(
hidden_states=training_batch.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch.input_kwargs['encoder_hidden_states'],
timestep=training_batch.input_kwargs['timestep'],
encoder_hidden_states_image=training_batch.input_kwargs.get('encoder_hidden_states_image'),
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=0, # Start from beginning for single-step
cache_start=0
).permute(0, 2, 1, 3, 4)
else:
# Fallback to regular forward
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
# Step 6: Convert noise prediction to video prediction
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noisy_latent.flatten(0, 1),
timestep=torch.tensor([timestep], device=noisy_latent.device),
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
self._reset_simulation_caches(kv_cache, crossattn_cache)
return pred_video
def _generator_multi_step_simulation_forward(
self, training_batch: TrainingBatch, return_sim_steps: bool = False) -> torch.Tensor:
"""Forward pass through student transformer matching inference procedure with KV cache management.
This function is adapted from the reference self-forcing implementation's inference_with_trajectory
and includes gradient masking logic for dynamic frame generation.
"""
latents = training_batch.latents
dtype = latents.dtype
batch_size = latents.shape[0]
initial_latent = getattr(training_batch, 'image_latent', None)
# Dynamic frame generation logic (adapted from _run_generator)
num_training_frames = getattr(self.training_args, 'num_latent_t', 21)
# During training, the number of generated frames should be uniformly sampled from
# [21, self.num_training_frames], but still being a multiple of self.num_frame_per_block
min_num_frames = 20 if self.independent_first_frame else 21
max_num_frames = num_training_frames - 1 if self.independent_first_frame else num_training_frames
assert max_num_frames % self.num_frame_per_block == 0
assert min_num_frames % self.num_frame_per_block == 0
max_num_blocks = max_num_frames // self.num_frame_per_block
min_num_blocks = min_num_frames // self.num_frame_per_block
# Sample number of blocks and sync across processes
num_generated_blocks = torch.randint(min_num_blocks, max_num_blocks + 1, (1,), device=self.device)
if dist.is_initialized():
dist.broadcast(num_generated_blocks, src=0)
num_generated_blocks = num_generated_blocks.item()
num_generated_frames = num_generated_blocks * self.num_frame_per_block
if self.independent_first_frame and initial_latent is None:
num_generated_frames += 1
min_num_frames += 1
# Create noise with dynamic shape
if initial_latent is not None:
noise_shape = [batch_size, num_generated_frames - 1, *self.video_latent_shape[2:]]
else:
noise_shape = [batch_size, num_generated_frames, *self.video_latent_shape[2:]]
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", n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
batch_size, num_frames, num_channels, height, width = noise.shape
# Block size calculation
if not self.independent_first_frame or (self.independent_first_frame and initial_latent is not None):
assert num_frames % self.num_frame_per_block == 0
num_blocks = num_frames // self.num_frame_per_block
else:
assert (num_frames - 1) % self.num_frame_per_block == 0
num_blocks = (num_frames - 1) // self.num_frame_per_block
num_input_frames = initial_latent.shape[1] if initial_latent is not None else 0
num_output_frames = num_frames + num_input_frames
output = torch.zeros([batch_size, num_output_frames, num_channels, height, width],
device=noise.device, dtype=noise.dtype)
# Step 1: Initialize KV cache to all zeros
self.kv_cache1, self.crossattn_cache = self._initialize_simulation_caches(batch_size, dtype, self.device)
# Validate cache structure (can be disabled in production)
if getattr(self.training_args, 'validate_cache_structure', False):
self._validate_cache_structure(self.kv_cache1, self.crossattn_cache, batch_size)
# Step 2: Cache context feature
current_start_frame = 0
if initial_latent is not None:
timestep = torch.ones([batch_size, 1], device=noise.device, dtype=torch.int64) * 0
output[:, :1] = initial_latent
with torch.no_grad():
# Build input kwargs for initial latent
training_batch_temp = self._build_distill_input_kwargs(
initial_latent, timestep * 0, training_batch.conditional_dict, training_batch)
self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
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,
current_start=current_start_frame * self.frame_seq_length
)
current_start_frame += 1
# Step 3: Temporal denoising loop
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)
exit_flags = self.generate_and_sync_list(len(all_num_frames), num_denoising_steps, device=noise.device)
start_gradient_frame_index = max(0, num_output_frames - 21)
for block_index, current_num_frames in enumerate(all_num_frames):
noisy_input = noise[:, current_start_frame - num_input_frames:current_start_frame + current_num_frames - num_input_frames]
# Step 3.1: Spatial denoising loop
for index, current_timestep in enumerate(self.denoising_step_list):
if self.same_step_across_blocks:
exit_flag = (index == exit_flags[0])
else:
exit_flag = (index == exit_flags[block_index])
timestep = torch.ones([batch_size, current_num_frames], device=noise.device, dtype=torch.int64) * current_timestep
if not exit_flag:
with torch.no_grad():
# Build input kwargs
training_batch_temp = self._build_distill_input_kwargs(
noisy_input, timestep, training_batch.conditional_dict, training_batch)
pred_flow = self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
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,
current_start=current_start_frame * self.frame_seq_length
).permute(0, 2, 1, 3, 4)
denoised_pred = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
next_timestep = self.denoising_step_list[index + 1]
noisy_input = self.noise_scheduler.add_noise(
denoised_pred.flatten(0, 1),
torch.randn_like(denoised_pred.flatten(0, 1)),
next_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
).unflatten(0, denoised_pred.shape[:2])
else:
# Final prediction with gradient control
if current_start_frame < start_gradient_frame_index:
with torch.no_grad():
training_batch_temp = self._build_distill_input_kwargs(
noisy_input, timestep, training_batch.conditional_dict, training_batch)
pred_flow = self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
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,
current_start=current_start_frame * self.frame_seq_length
).permute(0, 2, 1, 3, 4)
else:
training_batch_temp = self._build_distill_input_kwargs(
noisy_input, timestep, training_batch.conditional_dict, training_batch)
pred_flow = self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
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,
current_start=current_start_frame * self.frame_seq_length
).permute(0, 2, 1, 3, 4)
denoised_pred = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
break
# Step 3.2: record the model's output
output[:, current_start_frame:current_start_frame + current_num_frames] = denoised_pred
# Step 3.3: rerun with timestep zero to update the cache
context_timestep = torch.ones_like(timestep) * self.context_noise
denoised_pred = self.noise_scheduler.add_noise(
denoised_pred.flatten(0, 1),
torch.randn_like(denoised_pred.flatten(0, 1)),
context_timestep
).unflatten(0, denoised_pred.shape[:2])
with torch.no_grad():
training_batch_temp = self._build_distill_input_kwargs(
denoised_pred, context_timestep, training_batch.conditional_dict, training_batch)
self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
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,
current_start=current_start_frame * self.frame_seq_length
)
# Step 3.4: update the start and end frame indices
current_start_frame += current_num_frames
# Handle last 21 frames logic
pred_image_or_video = output
if num_input_frames > 0:
pred_image_or_video = output[:, num_input_frames:]
# Slice last 21 frames if we generated more
gradient_mask = None
if pred_image_or_video.shape[1] > 21:
with torch.no_grad():
# Reencode to get image latent
latent_to_decode = pred_image_or_video[:, :-20, ...]
# Decode to video
latent_to_decode = latent_to_decode.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
# Apply VAE scaling and shift factors
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent_to_decode = latent_to_decode / self.vae.scaling_factor.to(latent_to_decode.device, latent_to_decode.dtype)
else:
latent_to_decode = latent_to_decode / self.vae.scaling_factor
if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
if isinstance(self.vae.shift_factor, torch.Tensor):
latent_to_decode += self.vae.shift_factor.to(latent_to_decode.device, latent_to_decode.dtype)
else:
latent_to_decode += self.vae.shift_factor
# Decode to pixels
pixels = self.vae.decode(latent_to_decode)
frame = pixels[:, :, -1:, :, :].to(dtype) # Last frame [B, C, 1, H, W]
# Encode frame back to get image latent
image_latent = self.vae.encode(frame).to(dtype)
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
pred_image_or_video_last_21 = torch.cat([image_latent, pred_image_or_video[:, -20:, ...]], dim=1)
else:
pred_image_or_video_last_21 = pred_image_or_video
# Set up gradient mask if we generated more than minimum frames
if num_generated_frames != min_num_frames:
# Currently, we do not use gradient for the first chunk, since it contains image latents
gradient_mask = torch.ones_like(pred_image_or_video_last_21, dtype=torch.bool)
if self.independent_first_frame:
gradient_mask[:, :1] = False
else:
gradient_mask[:, :self.num_frame_per_block] = False
# Apply gradient masking if needed
final_output = pred_image_or_video_last_21.to(dtype)
if gradient_mask is not None:
# Apply gradient masking: detach frames that shouldn't contribute gradients
final_output = torch.where(
gradient_mask,
pred_image_or_video_last_21, # Keep original values where gradient_mask is True
pred_image_or_video_last_21.detach() # Detach where gradient_mask is False
)
# Store visualization data
training_batch.dmd_latent_vis_dict["generator_timestep"] = torch.tensor(
self.denoising_step_list[exit_flags[0]], dtype=torch.float32, device=self.device)
# Store gradient mask information for debugging
if gradient_mask is not None:
training_batch.dmd_latent_vis_dict["gradient_mask"] = gradient_mask.float()
training_batch.dmd_latent_vis_dict["num_generated_frames"] = torch.tensor(
num_generated_frames, dtype=torch.float32, device=self.device)
training_batch.dmd_latent_vis_dict["min_num_frames"] = torch.tensor(
min_num_frames, dtype=torch.float32, device=self.device)
self._reset_simulation_caches(self.kv_cache1, self.crossattn_cache)
return final_output if gradient_mask is not None else pred_image_or_video
def _initialize_simulation_caches(self, batch_size: int, dtype: torch.dtype, device: torch.device):
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
num_transformer_blocks = len(self.transformer.blocks)
# Calculate frame sequence length based on input dimensions and patch size
# From the training batch, we can get the actual latent dimensions
latent_shape = self.video_latent_shape_sp # This is set in _prepare_dit_inputs
batch_size_actual, num_frames, num_channels, height, width = latent_shape
# Get patch size from transformer config
p_t, p_h, p_w = self.transformer.patch_size
post_patch_height = height // p_h
post_patch_width = width // p_w
# Frame sequence length is the spatial sequence length per frame
frame_seq_length = post_patch_height * post_patch_width
# Get local attention size from transformer config
local_attn_size = getattr(self.transformer, 'local_attn_size', -1)
# Get model configuration parameters - handle FSDP wrapping
if hasattr(self.transformer, 'config'):
config = self.transformer.config
num_attention_heads = config.num_attention_heads
attention_head_dim = config.attention_head_dim
text_len = config.text_len
else:
# Fallback to direct attribute access for non-FSDP models
num_attention_heads = getattr(self.transformer, 'num_attention_heads', 40)
attention_head_dim = getattr(self.transformer, 'attention_head_dim', 128)
text_len = getattr(self.transformer, 'text_len', 512)
num_max_frames = getattr(self.training_args, "num_frames", num_frames)
kv_cache_size = num_max_frames * frame_seq_length
kv_cache = []
for _ in range(num_transformer_blocks):
kv_cache.append({
"k": torch.zeros([batch_size, kv_cache_size, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, kv_cache_size, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device)
})
# Initialize cross-attention cache
crossattn_cache = []
for _ in range(num_transformer_blocks):
crossattn_cache.append({
"k": torch.zeros([batch_size, text_len, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, text_len, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"is_init": False
})
return kv_cache, crossattn_cache
def _reset_simulation_caches(self, kv_cache, crossattn_cache):
"""Reset KV cache and cross-attention cache to clean state."""
if kv_cache is not None:
for cache_dict in kv_cache:
cache_dict["global_end_index"].fill_(0)
cache_dict["local_end_index"].fill_(0)
cache_dict["k"].zero_()
cache_dict["v"].zero_()
if crossattn_cache is not None:
for cache_dict in crossattn_cache:
cache_dict["is_init"] = False
cache_dict["k"].zero_()
cache_dict["v"].zero_()
def _validate_cache_structure(self, kv_cache, crossattn_cache, batch_size: int):
"""Validate that cache structures are correctly initialized."""
num_transformer_blocks = len(self.transformer.blocks)
# Get model configuration parameters - handle FSDP wrapping
if hasattr(self.transformer, 'config'):
config = self.transformer.config
num_attention_heads = config.num_attention_heads
attention_head_dim = config.attention_head_dim
text_len = config.text_len
else:
# Fallback to direct attribute access for non-FSDP models
num_attention_heads = getattr(self.transformer, 'num_attention_heads', 40)
attention_head_dim = getattr(self.transformer, 'attention_head_dim', 128)
text_len = getattr(self.transformer, 'text_len', 512)
if kv_cache is not None:
assert len(kv_cache) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(kv_cache)}"
for i, cache_dict in enumerate(kv_cache):
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in kv_cache block {i}"
assert "global_end_index" in cache_dict and "local_end_index" in cache_dict, f"Missing indices in kv_cache block {i}"
assert cache_dict["k"].shape[0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
assert cache_dict["v"].shape[0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
assert cache_dict["k"].shape[2] == num_attention_heads, f"Attention heads mismatch in kv_cache block {i}"
assert cache_dict["k"].shape[3] == attention_head_dim, f"Attention head dim mismatch in kv_cache block {i}"
if crossattn_cache is not None:
assert len(crossattn_cache) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(crossattn_cache)}"
for i, cache_dict in enumerate(crossattn_cache):
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in crossattn_cache block {i}"
assert "is_init" in cache_dict, f"Missing is_init in crossattn_cache block {i}"
assert cache_dict["k"].shape[0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
assert cache_dict["v"].shape[0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
assert cache_dict["k"].shape[1] == text_len, f"Text length mismatch in crossattn_cache block {i}"
assert cache_dict["k"].shape[2] == num_attention_heads, f"Attention heads mismatch in crossattn_cache block {i}"
assert cache_dict["k"].shape[3] == attention_head_dim, f"Attention head dim mismatch in crossattn_cache block {i}"
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
"""
Self-forcing training step that alternates between generator and critic training.
"""
gradient_accumulation_steps = getattr(self.training_args, 'gradient_accumulation_steps', 1)
train_generator = (self.current_trainstep % self.dfake_gen_update_ratio == 0)
batches = []
for _ in range(gradient_accumulation_steps):
batch = self._prepare_distillation(training_batch)
batch = self._get_next_batch(batch)
batch = self._normalize_dit_input(batch)
batch = self._prepare_dit_inputs(batch)
batch = self._build_attention_metadata(batch)
batch.attn_metadata_vsa = copy.deepcopy(batch.attn_metadata)
if batch.attn_metadata is not None:
batch.attn_metadata.VSA_sparsity = 0.0
batches.append(batch)
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
if train_generator:
logger.debug(f"Training generator at step {self.current_trainstep}")
self.optimizer.zero_grad()
total_generator_loss = 0.0
generator_log_dict = {}
for batch in batches:
# Create a new batch with detached tensors
batch_gen = TrainingBatch()
for key, value in batch.__dict__.items():
if isinstance(value, torch.Tensor):
setattr(batch_gen, key, value.detach().clone())
elif isinstance(value, dict):
setattr(batch_gen, key, {k: v.detach().clone() if isinstance(v, torch.Tensor) else copy.deepcopy(v) for k, v in value.items()})
else:
setattr(batch_gen, key, copy.deepcopy(value))
generator_loss, gen_log_dict = self.generator_loss(batch_gen)
with set_forward_context(
current_timestep=batch_gen.timesteps,
attn_metadata=batch_gen.attn_metadata):
(generator_loss / gradient_accumulation_steps).backward()
total_generator_loss += generator_loss.detach().item()
generator_log_dict.update(gen_log_dict)
# Store visualization data from generator training
if hasattr(batch_gen, 'dmd_latent_vis_dict'):
training_batch.dmd_latent_vis_dict.update(batch_gen.dmd_latent_vis_dict)
self._clip_model_grad_norm_(batch_gen, self.transformer)
self.optimizer.step()
self.lr_scheduler.step()
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
avg_generator_loss = torch.tensor(
total_generator_loss / gradient_accumulation_steps,
device=self.device
)
world_group = get_world_group()
world_group.all_reduce(avg_generator_loss, op=torch.distributed.ReduceOp.AVG)
training_batch.generator_loss = avg_generator_loss.item()
else:
training_batch.generator_loss = 0.0
logger.debug(f"Training critic at step {self.current_trainstep}")
self.fake_score_optimizer.zero_grad()
total_critic_loss = 0.0
critic_log_dict = {}
for batch in batches:
# Create a new batch with detached tensors
batch_critic = TrainingBatch()
for key, value in batch.__dict__.items():
if isinstance(value, torch.Tensor):
setattr(batch_critic, key, value.detach().clone())
elif isinstance(value, dict):
setattr(batch_critic, key, {k: v.detach().clone() if isinstance(v, torch.Tensor) else copy.deepcopy(v) for k, v in value.items()})
else:
setattr(batch_critic, key, copy.deepcopy(value))
critic_loss, crit_log_dict = self.critic_loss(batch_critic)
with set_forward_context(
current_timestep=batch_critic.timesteps,
attn_metadata=batch_critic.attn_metadata):
(critic_loss / gradient_accumulation_steps).backward()
total_critic_loss += critic_loss.detach().item()
critic_log_dict.update(crit_log_dict)
# Store visualization data from critic training
if hasattr(batch_critic, 'fake_score_latent_vis_dict'):
training_batch.fake_score_latent_vis_dict.update(batch_critic.fake_score_latent_vis_dict)
self._clip_model_grad_norm_(batch_critic, self.fake_score_transformer)
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
avg_critic_loss = torch.tensor(
total_critic_loss / gradient_accumulation_steps,
device=self.device
)
world_group = get_world_group()
world_group.all_reduce(avg_critic_loss, op=torch.distributed.ReduceOp.AVG)
training_batch.fake_score_loss = avg_critic_loss.item()
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_loss
return training_batch
def _log_training_info(self) -> None:
"""Log self-forcing specific training information."""
super()._log_training_info()
logger.info("Self-forcing specific settings:")
logger.info(" Generator update ratio: %s", self.dfake_gen_update_ratio)
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
wandb_loss_dict = {}
# Debug logging
logger.info(f"Step {step}: Starting visualization")
if hasattr(training_batch, 'dmd_latent_vis_dict'):
logger.info(f"DMD latent keys: {list(training_batch.dmd_latent_vis_dict.keys())}")
if hasattr(training_batch, 'fake_score_latent_vis_dict'):
logger.info(f"Fake score latent keys: {list(training_batch.fake_score_latent_vis_dict.keys())}")
# Process generator predictions if available
if hasattr(training_batch, 'dmd_latent_vis_dict') and training_batch.dmd_latent_vis_dict:
dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
dmd_log_keys = ['generator_pred_video', 'real_score_pred_video', 'faker_score_pred_video']
for latent_key in dmd_log_keys:
if latent_key in dmd_latents_vis_dict:
logger.info(f"Processing DMD latent: {latent_key}")
latents = dmd_latents_vis_dict[latent_key]
if not isinstance(latents, torch.Tensor):
logger.warning(f"Expected tensor for {latent_key}, got {type(latents)}")
continue
latents = latents.detach()
latents = latents.permute(0, 2, 1, 3, 4)
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device, latents.dtype)
else:
latents += self.vae.shift_factor
try:
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(video, fps=24, format="mp4")
logger.info(f"Successfully processed DMD latent: {latent_key}")
except Exception as e:
logger.error(f"Error processing DMD latent {latent_key}: {str(e)}")
del video, latents
# Process critic predictions
if hasattr(training_batch, 'fake_score_latent_vis_dict') and training_batch.fake_score_latent_vis_dict:
fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
fake_score_log_keys = ['generator_pred_video']
for latent_key in fake_score_log_keys:
if latent_key in fake_score_latents_vis_dict:
logger.info(f"Processing critic latent: {latent_key}")
latents = fake_score_latents_vis_dict[latent_key]
if not isinstance(latents, torch.Tensor):
logger.warning(f"Expected tensor for {latent_key}, got {type(latents)}")
continue
latents = latents.detach()
latents = latents.permute(0, 2, 1, 3, 4)
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device, latents.dtype)
else:
latents += self.vae.shift_factor
try:
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(video, fps=24, format="mp4")
logger.info(f"Successfully processed critic latent: {latent_key}")
except Exception as e:
logger.error(f"Error processing critic latent {latent_key}: {str(e)}")
del video, latents
# Log metadata
if hasattr(training_batch, 'dmd_latent_vis_dict') and training_batch.dmd_latent_vis_dict:
if "generator_timestep" in training_batch.dmd_latent_vis_dict:
wandb_loss_dict["generator_timestep"] = training_batch.dmd_latent_vis_dict["generator_timestep"].item()
if "dmd_timestep" in training_batch.dmd_latent_vis_dict:
wandb_loss_dict["dmd_timestep"] = training_batch.dmd_latent_vis_dict["dmd_timestep"].item()
if hasattr(training_batch, 'fake_score_latent_vis_dict') and training_batch.fake_score_latent_vis_dict:
if "fake_score_timestep" in training_batch.fake_score_latent_vis_dict:
wandb_loss_dict["fake_score_timestep"] = training_batch.fake_score_latent_vis_dict["fake_score_timestep"].item()
# Log final dict contents
logger.info(f"Final wandb_loss_dict keys: {list(wandb_loss_dict.keys())}")
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
def train(self) -> None:
"""Main training loop with self-forcing specific logging."""
assert self.training_args.seed is not None, "seed must be set"
seed = self.training_args.seed
# Set the same seed within each SP group to ensure reproducibility
if self.sp_world_size > 1:
# Use the same seed for all processes within the same SP group
sp_group_seed = seed + (self.global_rank // self.sp_world_size)
set_random_seed(sp_group_seed)
logger.info("Rank %s: Using SP group seed %s", self.global_rank,
sp_group_seed)
else:
set_random_seed(seed + self.global_rank)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", seed)
self.current_trainstep = self.init_steps
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
logger.info("Resumed from checkpoint, random states restored")
else:
logger.info("Starting training from scratch")
self.train_loader_iter = iter(self.train_dataloader)
step_times = deque(maxlen=100)
self._log_training_info()
self._log_validation(self.transformer, self.training_args,
self.init_steps)
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
disable=self.local_rank > 0,
)
use_vsa = vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if use_vsa:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
if vsa_decay_interval_steps > 1:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = vsa_sparsity
else:
current_vsa_sparsity = 0.0
training_batch = TrainingBatch()
self.current_trainstep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
if (step >= self.training_args.ema_start_step) and \
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
logger.info(f"Created generator EMA at step {step} with decay={self.training_args.ema_decay}")
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
total_loss = training_batch.total_loss
generator_loss = training_batch.generator_loss
fake_score_loss = training_batch.fake_score_loss
grad_norm = training_batch.grad_norm
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"total_loss": f"{total_loss:.4f}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"ema": "✓" if (self.generator_ema is not None and self.is_ema_ready()) else "✗",
})
progress_bar.update(1)
if self.global_rank == 0:
log_data = {
"train_total_loss": total_loss,
"train_fake_score_loss": fake_score_loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"fake_score_learning_rate": self.fake_score_lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
}
if (step % self.dfake_gen_update_ratio == 0):
log_data["train_generator_loss"] = generator_loss
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
if self.generator_ema is not None:
log_data["ema_enabled"] = True
log_data["ema_decay"] = self.training_args.ema_decay
else:
log_data["ema_enabled"] = False
ema_stats = self.get_ema_stats()
log_data.update(ema_stats)
if training_batch.dmd_latent_vis_dict:
dmd_additional_logs = {
"generator_timestep": training_batch.dmd_latent_vis_dict["generator_timestep"].item(),
"dmd_timestep": training_batch.dmd_latent_vis_dict["dmd_timestep"].item(),
}
log_data.update(dmd_additional_logs)
faker_score_additional_logs = {
"fake_score_timestep": training_batch.fake_score_latent_vis_dict["fake_score_timestep"].item(),
}
log_data.update(faker_score_additional_logs)
wandb.log(log_data, step=step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
if self.training_args.log_visualization:
self.visualize_intermediate_latents(training_batch, self.training_args, step)
if (self.training_args.training_state_checkpointing_steps > 0
and step % self.training_args.training_state_checkpointing_steps == 0):
print("rank", self.global_rank,
"save training state checkpoint at step", step)
save_distillation_checkpoint(
self.transformer, self.fake_score_transformer,
self.global_rank, self.training_args.output_dir, step,
self.optimizer, self.fake_score_optimizer,
self.train_dataloader, self.lr_scheduler,
self.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
if self.transformer:
self.transformer.train()
self.sp_group.barrier()
if (self.training_args.weight_only_checkpointing_steps > 0
and step % self.training_args.weight_only_checkpointing_steps == 0):
print("rank", self.global_rank,
"save weight-only checkpoint at step", step)
save_distillation_checkpoint(self.transformer,
self.fake_score_transformer,
self.global_rank,
self.training_args.output_dir,
f"{step}_weight_only",
only_save_generator_weight=True,
generator_ema=self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
wandb.finish()
print("rank", self.global_rank,
"save final training state checkpoint at step",
self.training_args.max_train_steps)
save_distillation_checkpoint(
self.transformer, self.fake_score_transformer, self.global_rank,
self.training_args.output_dir, self.training_args.max_train_steps,
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
self.lr_scheduler, self.fake_score_lr_scheduler,
self.noise_random_generator, self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir, self.training_args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
+5 -1
View File
@@ -117,10 +117,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
params_to_optimize = self.transformer.parameters()
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=(0.9, 0.999),
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
+170 -22
View File
@@ -202,6 +202,7 @@ def save_distillation_checkpoint(generator_transformer,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None,
only_save_generator_weight=False) -> None:
"""
Save distillation checkpoint with both generator and fake_score models.
@@ -233,6 +234,8 @@ def save_distillation_checkpoint(generator_transformer,
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")
@@ -402,7 +405,8 @@ def load_distillation_checkpoint(generator_transformer,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None) -> int:
noise_generator=None,
generator_ema=None) -> int:
"""
Load distillation checkpoint with both generator and fake_score models.
Returns the step number from which training should resume.
@@ -456,6 +460,18 @@ def load_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
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)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank, str(e))
# Load critic distributed checkpoint
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
"critic")
@@ -826,27 +842,6 @@ def custom_to_hf_state_dict(
return new_state_dict
def pred_noise_to_pred_video(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
"""
timestep = timestep.expand(noise_input_latent.shape[0])
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
noise_input_latent = noise_input_latent.float().to(device)
sigmas = scheduler.sigmas.float().to(device)
timesteps = scheduler.timesteps.float().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
@@ -1299,3 +1294,156 @@ def get_scheduler(
num_warmup_steps=num_warmup_steps,
num_training_steps=num_training_steps,
last_epoch=last_epoch)
class EMA_FSDP:
"""
FSDP2-friendly EMA with two modes:
- mode="local_shard" (default): maintain float32 CPU EMA of local parameter shards on every rank.
Provides a context manager to temporarily swap EMA weights into the live model for teacher forward.
- mode="rank0_full": maintain a consolidated float32 CPU EMA of full parameters on rank 0 only
using gather_state_dict_on_cpu_rank0(). Useful for checkpoint export; not for teacher forward.
Usage (local_shard for CM teacher):
ema = EMA_FSDP(model, decay=0.999, mode="local_shard")
for step in ...:
ema.update(model)
with ema.apply_to_model(model):
with torch.no_grad():
y_teacher = model(...)
Usage (rank0_full for export):
ema = EMA_FSDP(model, decay=0.999, mode="rank0_full")
ema.update(model)
ema.state_dict() # on rank 0
"""
def __init__(self, module, decay: float = 0.999, mode: str = "local_shard"):
self.decay = float(decay)
self.mode = mode
self.shadow: dict[str, torch.Tensor] = {}
self.rank = dist.get_rank() if dist.is_initialized() else 0
if self.mode not in {"local_shard", "rank0_full"}:
raise ValueError(f"Unsupported EMA_FSDP mode: {self.mode}")
self._init_shadow(module)
@staticmethod
def _to_local_tensor(t: torch.Tensor) -> torch.Tensor:
# DTensor-aware to_local fetch; fall back to raw tensor
try:
from torch.distributed.tensor import DTensor # type: ignore
if isinstance(t, DTensor):
return t.to_local()
except Exception:
pass
return t
@torch.no_grad()
def _init_shadow(self, module):
if self.mode == "rank0_full":
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
if self.rank == 0:
self.shadow = {k: v.detach().clone().float().cpu() for k, v in cpu_state.items()}
else:
self.shadow = {}
return
# local_shard: maintain EMA of local shards for requires_grad params
self.shadow = {}
for name, p in module.named_parameters():
if not p.requires_grad:
continue
local = self._to_local_tensor(p.detach())
self.shadow[name] = local.clone().float().cpu()
@torch.no_grad()
def update(self, module):
d = self.decay
if self.mode == "rank0_full":
if self.rank != 0:
return
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
for n, v in cpu_state.items():
v_cpu = v.detach().float().cpu()
if n not in self.shadow:
self.shadow[n] = v_cpu.clone()
else:
self.shadow[n].mul_(d).add_(v_cpu, alpha=1.0 - d)
return
# local_shard: update local shard EMA on every rank
for name, p in module.named_parameters():
if not p.requires_grad:
continue
local = self._to_local_tensor(p.detach())
v_cpu = local.float().cpu()
if name not in self.shadow:
self.shadow[name] = v_cpu.clone()
else:
self.shadow[name].mul_(d).add_(v_cpu, alpha=1.0 - d)
def state_dict(self) -> dict[str, torch.Tensor]:
if self.mode == "rank0_full":
return {k: v.clone() for k, v in self.shadow.items()} if self.rank == 0 else {}
return {k: v.clone() for k, v in self.shadow.items()}
def load_state_dict(self, sd: dict[str, torch.Tensor]):
self.shadow = {k: v.clone() for k, v in sd.items()}
@torch.no_grad()
def copy_to_unwrapped(self, module) -> None:
"""
Copy EMA weights into a non-sharded (unwrapped) module. Intended for export/eval.
For mode="rank0_full", only rank 0 has the full EMA state.
"""
if self.mode == "rank0_full" and self.rank != 0:
return
name_to_param = dict(module.named_parameters())
for n, w in self.shadow.items():
if n in name_to_param:
p = name_to_param[n]
p.data.copy_(w.to(dtype=p.dtype, device=p.device))
class _ApplyEMACtx:
def __init__(self, ema: "EMA_FSDP", module):
self.ema = ema
self.module = module
self.saved: dict[str, torch.Tensor] = {}
def __enter__(self):
if self.ema.mode != "local_shard":
raise RuntimeError("EMA apply_to_model is only supported for mode='local_shard'")
with torch.no_grad():
for name, p in self.module.named_parameters():
if not p.requires_grad:
continue
# Save local shard
p_local = EMA_FSDP._to_local_tensor(p.detach())
if p_local.numel() == 0:
# Nothing to swap on this rank for this param
continue
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
if name in self.ema.shadow:
ema_cpu = self.ema.shadow[name]
if ema_cpu.numel() != p_local.numel():
# Shard shape mismatch (e.g., empty shard here), skip
continue
# Copy EMA shard into local param shard
p_local.copy_(ema_cpu.to(dtype=p_local.dtype, device=p_local.device))
return self.module
def __exit__(self, exc_type, exc, tb):
with torch.no_grad():
for name, p in self.module.named_parameters():
if name in self.saved:
p_local = EMA_FSDP._to_local_tensor(p.detach())
if p_local.numel() == 0:
continue
saved_local = self.saved[name]
if saved_local.numel() != p_local.numel():
continue
p_local.copy_(saved_local)
self.saved.clear()
return False
def apply_to_model(self, module):
return EMA_FSDP._ApplyEMACtx(self, module)
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import WanCausalDMDPipeline
from fastvideo.training.self_forcing_distillation_pipeline import SelfForcingDistillationPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
"""
A self-forcing distillation pipeline for Wan that uses the self-forcing methodology
with DMD for video generation.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
validation_pipeline = WanCausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
self.validation_pipeline = validation_pipeline
def main(args) -> None:
logger.info("Starting Wan self-forcing distillation pipeline...")
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Wan self-forcing distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
main(args)
+1 -1
View File
@@ -34,6 +34,7 @@ from torch.distributed.fsdp import MixedPrecisionPolicy
import fastvideo.envs as envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
T = TypeVar("T")
@@ -614,7 +615,6 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
f"Failed to download or parse model_index.json for {model_name_or_path}: {e}"
) from e
def update_environment_variables(envs: dict[str, str]):
for k, v in envs.items():
if k in os.environ and os.environ[k] != v:
+22
View File
@@ -0,0 +1,22 @@
#!/bin/bash
counter=0
# for pair in "1e-5 1e-5" "1e-5 8e-6" "1e-5 6e-6" "1e-5 4e-6" "1e-5 2e-6" "1e-5 1e-6"; do
# for pair in "1e-5 8e-6" "1e-5 4e-6" "1e-5 2e-6" "8e-6 8e-6" "8e-6 6e-6" "8e-6 2e-6"; do
for pair in "1e-5 8e-6"; do
# for pair in "2e-6 4e-7" "2e-6 6e-7" "4e-6 6e-7" "4e-6 8e-7" "4e-6 1e-6" "6e-6 4e-7" "6e-6 6e-7" "6e-6 8e-7" "6e-6 1e-6"; do
# for pair in "2e-6 4e-7" "2e-6 6e-7" "4e-6 6e-7" "4e-6 8e-7" "4e-6 1e-6" "6e-6 4e-7" "6e-6 6e-7" "6e-6 8e-7" "6e-6 1e-6"; do
port=$((29500 + counter))
read -r lr critic_lr <<< "$pair"
echo "$lr $critic_lr"
echo "sbatch --job-name=sf-${lr}-c${critic_lr} --output=sf_output/sf-${lr}-c${critic_lr}_%j.out --error=sf_output/sf-${lr}-c${critic_lr}_%j.err examples/distill/SFWan2.1-T2V/distill_dmd_t2v_1.3B.slurm $port $lr $critic_lr"
sbatch --job-name=n-${lr}-c${critic_lr} --output=sf_output/sf-${lr}-c${critic_lr}.out --error=sf_output/sf-${lr}-c${critic_lr}.err examples/distill/SFWan2.1-T2V/distill_dmd_t2v_1.3B.slurm $port $lr $critic_lr
((counter++))
done
# port=$((29500 + counter))
# echo "sbatch --job-name=sf-${i} --output=sf_output/sf-${i}_%j.out --error=sf_output/sf-${i}_%j.err distill_sf_1_3B.slurm 7 $port $i"
# sbatch --job-name=sf-${i} --output=sf_output/sf-${i}_%j.out --error=sf_output/sf-${i}_%j.err distill_sf_1_3B.slurm $i $port 5
# ((counter++))
# done
@@ -0,0 +1,161 @@
from huggingface_hub import save_torch_state_dict, load_state_dict_from_file
# from safetensors import safetensors
from safetensors.torch import save_file
import torch
import re
from collections import OrderedDict
_param_names_mapping: dict = {
r"^text_embedding\.0\.(.*)$":
r"condition_embedder.text_embedder.linear_1.\1",
r"^text_embedding\.2\.(.*)$":
r"condition_embedder.text_embedder.linear_2.\1",
r"^time_embedding\.0\.(.*)$":
r"condition_embedder.time_embedder.linear_1.\1",
r"^time_embedding\.2\.(.*)$":
r"condition_embedder.time_embedder.linear_2.\1",
r"^time_projection\.1\.(.*)$":
r"condition_embedder.time_proj.\1",
r"^img_emb\.proj\.0\.(.*)$":
r"condition_embedder.image_embedder.norm1.\1",
r"^img_emb\.proj\.1\.(.*)$":
r"condition_embedder.image_embedder.ff.net.0.proj.\1",
r"^img_emb\.proj\.3\.(.*)$":
r"condition_embedder.image_embedder.ff.net.2.\1",
r"^img_emb\.proj\.4\.(.*)$":
r"condition_embedder.image_embedder.norm2.\1",
r"^head\.modulation":
r"scale_shift_table",
r"^head\.head\.(.*)$":
r"proj_out.\1",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$":
r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$":
r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$":
r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$":
r"blocks.\1.attn1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$":
r"blocks.\1.attn1.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$":
r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$":
r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.k_img\.(.*)$":
r"blocks.\1.attn2.add_k_proj.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$":
r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.v_img\.(.*)$":
r"blocks.\1.attn2.add_v_proj.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$":
r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$":
r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k_img\.(.*)$":
r"blocks.\1.attn2.norm_added_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$":
r"blocks.\1.ffn.net.0.proj.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$":
r"blocks.\1.ffn.net.2.\2",
r"^blocks\.(\d+)\.modulation":
r"blocks.\1.scale_shift_table",
r"^blocks\.(\d+)\.norm3\.(.*)$":
r"blocks.\1.norm2.\2",
}
# The following mapping has an extra 'patch_embedding' field and also contains
# the 'model' prefixes
_self_forcing_to_diffusers_param_names_mapping: dict = {
r"^model.patch_embedding\.(.*)$":
r"patch_embedding.\1",
r"^model.text_embedding\.0\.(.*)$":
r"condition_embedder.text_embedder.linear_1.\1",
r"^model.text_embedding\.2\.(.*)$":
r"condition_embedder.text_embedder.linear_2.\1",
r"^model.time_embedding\.0\.(.*)$":
r"condition_embedder.time_embedder.linear_1.\1",
r"^model.time_embedding\.2\.(.*)$":
r"condition_embedder.time_embedder.linear_2.\1",
r"^model.time_projection\.1\.(.*)$":
r"condition_embedder.time_proj.\1",
r"^model.img_emb\.proj\.0\.(.*)$":
r"condition_embedder.image_embedder.norm1.\1",
r"^model.img_emb\.proj\.1\.(.*)$":
r"condition_embedder.image_embedder.ff.net.0.proj.\1",
r"^model.img_emb\.proj\.3\.(.*)$":
r"condition_embedder.image_embedder.ff.net.2.\1",
r"^model.img_emb\.proj\.4\.(.*)$":
r"condition_embedder.image_embedder.norm2.\1",
r"^model.head\.modulation":
r"scale_shift_table",
r"^model.head\.head\.(.*)$":
r"proj_out.\1",
r"^model.blocks\.(\d+)\.self_attn\.q\.(.*)$":
r"blocks.\1.attn1.to_q.\2",
r"^model.blocks\.(\d+)\.self_attn\.k\.(.*)$":
r"blocks.\1.attn1.to_k.\2",
r"^model.blocks\.(\d+)\.self_attn\.v\.(.*)$":
r"blocks.\1.attn1.to_v.\2",
r"^model.blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\2",
r"^model.blocks\.(\d+)\.self_attn\.norm_q\.(.*)$":
r"blocks.\1.attn1.norm_q.\2",
r"^model.blocks\.(\d+)\.self_attn\.norm_k\.(.*)$":
r"blocks.\1.attn1.norm_k.\2",
r"^model.blocks\.(\d+)\.cross_attn\.q\.(.*)$":
r"blocks.\1.attn2.to_q.\2",
r"^model.blocks\.(\d+)\.cross_attn\.k\.(.*)$":
r"blocks.\1.attn2.to_k.\2",
r"^model.blocks\.(\d+)\.cross_attn\.k_img\.(.*)$":
r"blocks.\1.attn2.add_k_proj.\2",
r"^model.blocks\.(\d+)\.cross_attn\.v\.(.*)$":
r"blocks.\1.attn2.to_v.\2",
r"^model.blocks\.(\d+)\.cross_attn\.v_img\.(.*)$":
r"blocks.\1.attn2.add_v_proj.\2",
r"^model.blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.0.\2",
r"^model.blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$":
r"blocks.\1.attn2.norm_q.\2",
r"^model.blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$":
r"blocks.\1.attn2.norm_k.\2",
r"^model.blocks\.(\d+)\.cross_attn\.norm_k_img\.(.*)$":
r"blocks.\1.attn2.norm_added_k.\2",
r"^model.blocks\.(\d+)\.ffn\.0\.(.*)$":
r"blocks.\1.ffn.net.0.proj.\2",
r"^model.blocks\.(\d+)\.ffn\.2\.(.*)$":
r"blocks.\1.ffn.net.2.\2",
r"^model.blocks\.(\d+)\.modulation":
r"blocks.\1.scale_shift_table",
r"^model.blocks\.(\d+)\.norm3\.(.*)$":
r"blocks.\1.norm2.\2",
}
state_dict = load_state_dict_from_file("checkpoints/self_forcing_dmd.pt")
state_dict = state_dict["generator_ema"]
new_state_dict = OrderedDict()
for k, v in state_dict.items():
new_key = k
for pattern, replacement in _self_forcing_to_diffusers_param_names_mapping.items():
if re.match(pattern, k):
new_key = re.sub(pattern, replacement, k)
break # Stop at the first match
else:
# print(f"No match found for {k}")
raise ValueError(f"No match found for {k}")
new_state_dict[new_key] = v
if "norm_added_k" in new_key:
dummy_key = new_key.replace("norm_added_k", "norm_added_q")
dummy_value = torch.zeros_like(v)
new_state_dict[dummy_key] = dummy_value
del state_dict
save_torch_state_dict(
new_state_dict,
"new2/",
max_shard_size="10GB"
)
+2 -2
View File
@@ -3,7 +3,7 @@ from huggingface_hub import HfApi
api = HfApi()
api.upload_folder(
folder_path="Wan2.2-TI2V-5B-Diffusers",
repo_id="FastVideo/FastWan2.2-TI2V-5B-Diffusers",
folder_path="wow",
repo_id="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
repo_type="model",
)