Compare commits

...
9 Commits
23 changed files with 3403 additions and 286 deletions
+1
View File
@@ -64,3 +64,4 @@ docs/source/distillation/examples/
!docs/source/_static/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
@@ -0,0 +1,151 @@
#!/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
# 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=29503
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# 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/test-text-preprocessing/Node_0_GPU_1_File_1/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 "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
--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 # 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 $NUM_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_GPUS
)
# 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 "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 500
--weight_only_checkpointing_steps 500
--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'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
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,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,151 @@
#!/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
# 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=29503
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=8
# Model paths for Self-Forcing DMD distillation with Wan2.2:
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
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/test-text-preprocessing/Node_0_GPU_1_File_1/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 SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 16
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--num_frames 61 # 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 4 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 16
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS # 64
--sp_size 4
--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 "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 500
--weight_only_checkpointing_steps 500
--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'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
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[@]}"
+4 -10
View File
@@ -9,9 +9,9 @@ def main():
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
num_gpus=4,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
@@ -25,9 +25,7 @@ def main():
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
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."
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
@@ -35,11 +33,7 @@ def main():
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
"The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
+98 -1
View File
@@ -605,6 +605,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
@@ -627,6 +632,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
@@ -658,6 +664,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
@@ -678,16 +685,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":
@@ -789,6 +809,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,
@@ -859,6 +893,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")
@@ -963,6 +1001,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")
@@ -1013,6 +1055,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,
@@ -1029,6 +1078,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,
@@ -1041,6 +1095,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
@@ -1048,4 +1145,4 @@ class TrainingArgs(FastVideoArgs):
def parse_int_list(value: str) -> list[int]:
if not value:
return []
return [int(x.strip()) for x in value.split(",")]
return [int(x.strip()) for x in value.split(",")]
+5 -5
View File
@@ -212,9 +212,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
frame_seqlen = normalized.shape[1] // num_frames
modulated = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
(1 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1.0 + scale) + shift
modulated = normalized * (1 + scale) + shift
return modulated, residual_output
@@ -267,13 +267,13 @@ class LayerNormScaleShift(nn.Module):
frame_seqlen = normalized.shape[1] // num_frames
output = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
(1 + scale) + shift).flatten(1, 2)
else:
# scale.shape: [batch_size, 1, inner_dim]
# shift.shape: [batch_size, 1, inner_dim]
output = normalized * (1.0 + scale) + shift
output = normalized * (1 + scale) + shift
if self.compute_dtype == torch.float32:
output = output.to(x.dtype)
return output
return output
+63 -38
View File
@@ -147,6 +147,9 @@ class CausalWanSelfAttention(nn.Module):
# 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"].detach()
# kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
@@ -176,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(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)
@@ -209,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -223,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -249,29 +250,34 @@ class CausalWanTransformerBlock(nn.Module):
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
# e.shape: [batch_size, num_frames, 6, inner_dim]
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
assert shift_msa.dtype == torch.float32
# assert shift_msa.dtype == torch.float32
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), 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)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -285,8 +291,6 @@ class CausalWanTransformerBlock(nn.Module):
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)
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,
@@ -295,13 +299,10 @@ class CausalWanTransformerBlock(nn.Module):
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)
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)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -364,8 +365,7 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -375,7 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.num_frame_per_block = 3
self.independent_first_frame = False
self.__post_init__()
@@ -487,12 +487,16 @@ class CausalWanTransformer3DModel(BaseDiT):
)
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
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
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)
@@ -539,14 +543,9 @@ class CausalWanTransformer3DModel(BaseDiT):
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)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -587,8 +586,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
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
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
@@ -601,8 +600,12 @@ class CausalWanTransformer3DModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
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)
@@ -637,14 +640,9 @@ class CausalWanTransformer3DModel(BaseDiT):
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)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def forward(
self,
@@ -655,3 +653,30 @@ class CausalWanTransformer3DModel(BaseDiT):
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
+80 -73
View File
@@ -1,3 +1,5 @@
import torch
import torch.nn as nn
# SPDX-License-Identifier: Apache-2.0
import math
@@ -37,16 +39,14 @@ class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = FP32LayerNorm(in_features)
self.norm1 = nn.LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = FP32LayerNorm(out_features)
self.norm2 = nn.LayerNorm(out_features)
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
dtype = encoder_hidden_states_image.dtype
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states).to(dtype)
hidden_states = self.norm2(hidden_states)
return hidden_states
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
@@ -156,12 +156,12 @@ class WanT2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
q = self.norm_q.forward_native(self.to_q(x)[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)
k = self.norm_k.forward_native(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
@@ -169,7 +169,7 @@ class WanT2VCrossAttention(WanSelfAttention):
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
@@ -213,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# 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)
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
@@ -247,7 +247,7 @@ class WanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(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)
@@ -278,29 +278,29 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
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)
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -319,12 +319,11 @@ class WanTransformerBlock(nn.Module):
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
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()
self.scale_shift_table.unsqueeze(0) + temb
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
@@ -335,22 +334,20 @@ class WanTransformerBlock(nn.Module):
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()
e = self.scale_shift_table + temb
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
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
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)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -370,26 +367,20 @@ class WanTransformerBlock(nn.Module):
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)
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,
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
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)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
@@ -406,7 +397,7 @@ class WanTransformerBlock_VSA(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(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)
@@ -438,8 +429,7 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -459,8 +449,7 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -480,23 +469,22 @@ class WanTransformerBlock_VSA(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()
e = self.scale_shift_table + temb
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
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states) *
(1 + scale_msa) + shift_msa)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -521,8 +509,6 @@ class WanTransformerBlock_VSA(nn.Module):
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)
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,
@@ -530,17 +516,15 @@ class WanTransformerBlock_VSA(nn.Module):
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
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)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformer3DModel(CachableDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
@@ -598,8 +582,7 @@ class WanTransformer3DModel(CachableDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -659,10 +642,12 @@ class WanTransformer3DModel(CachableDiT):
rope_theta=10000)
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
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
@@ -672,6 +657,8 @@ class WanTransformer3DModel(CachableDiT):
else:
ts_seq_len = None
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
@@ -728,14 +715,35 @@ class WanTransformer3DModel(CachableDiT):
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)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
@@ -827,5 +835,4 @@ class WanTransformer3DModel(CachableDiT):
if self.is_even:
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
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)
@@ -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)
@@ -22,8 +22,10 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
config_name = "scheduler_config.json"
order = 1
@register_to_config
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
@@ -28,10 +28,6 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"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."""
@@ -0,0 +1,292 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.causal_model import CausalWanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_ori_causal_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = CausalWanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
new_state_dict = {}
for k, v in causal_state_dict.items():
if k.startswith("model."):
new_state_dict[k.replace("model.", "")] = v
causal_state_dict = new_state_dict
model1.load_state_dict(causal_state_dict)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
12,
160,
90,
device=device,
dtype=precision)
block_sizes = [3 for _ in range(4)]
timesteps = [1000, 750, 500, 250]
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
output1 = _causal_inference(model1, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
logger.info("Finish inference for model1")
output2 = _causal_inference(model2, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
logger.info("Output 1 Sum: %s", output1.float().sum().item())
logger.info("Output 2 Sum: %s", output2.float().sum().item())
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
def _causal_inference(transformer, latents, prompt_embeds, block_sizes, timesteps, target_dtype):
forward_batch = ForwardBatch(
data_type="dummy",
)
start_index = 0
pos_start_base = 0
frame_seq_length = latents.shape[-1] * latents.shape[-2] // (WanVideoConfig().arch_config.patch_size[-1] * WanVideoConfig().arch_config.patch_size[-2])
seq_len = frame_seq_length * latents.shape[2]
kv_cache1 = _initialize_kv_cache(transformer, batch_size=latents.shape[0],
kv_cache_size=frame_seq_length * latents.shape[2],
dtype=target_dtype,
device=latents.device)
crossattn_cache = _initialize_crossattn_cache(
transformer,
batch_size=latents.shape[0],
max_text_len=WanVideoConfig().arch_config.text_len,
dtype=target_dtype,
device=latents.device)
for current_num_frames, t_cur in zip(block_sizes, timesteps):
# logger.info(f"Current frame idx: {start_index}, Current timestep: {t_cur}")
# logger.info(f"k cache sum: {sum(kv_cache['k'].float().sum().item() for kv_cache in kv_cache1)}, v cache sum: {sum(kv_cache['v'].float().sum().item() for kv_cache in kv_cache1)}")
# logger.info(f"latents sum: {latents.float().sum().item()}, encoder_hidden_states sum: {prompt_embeds.float().sum().item()}")
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
attn_metadata = None
with set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=forward_batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(current_latents.shape[0], 1),
device=current_latents.device,
dtype=torch.long)
if isinstance(transformer, CausalWanModel):
pred_noise_btchw = transformer(
x=current_latents,
context=prompt_embeds,
t=t_expanded_noise,
seq_len=seq_len,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length
)
elif isinstance(transformer, CausalWanTransformer3DModel):
pred_noise_btchw = transformer(
current_latents,
prompt_embeds,
t_expanded_noise,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length,
start_frame=start_index
)
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = pred_noise_btchw.clone()
# Re-run with context timestep to update KV cache using clean context
context_noise = 0
t_context = torch.ones([latents.shape[0]],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = pred_noise_btchw.to(target_dtype)
with set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=forward_batch):
t_expanded_context = t_context.unsqueeze(1)
if isinstance(transformer, CausalWanModel):
_ = transformer(
x=context_bcthw,
context=prompt_embeds,
t=t_expanded_context,
seq_len=seq_len,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length
)
elif isinstance(transformer, CausalWanTransformer3DModel):
_ = transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length,
start_frame=start_index
)
start_index += current_num_frames
return latents
def _initialize_kv_cache(transformer, batch_size, kv_cache_size, dtype, device) -> None:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
if isinstance(transformer, CausalWanModel):
num_attention_heads = transformer.num_heads
attention_head_dim = transformer.dim // transformer.num_heads
elif isinstance(transformer, CausalWanTransformer3DModel):
num_attention_heads = transformer.num_attention_heads
attention_head_dim = transformer.attention_head_dim
for _ in range(len(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),
})
return kv_cache1
def _initialize_crossattn_cache(transformer, batch_size, max_text_len, dtype,
device) -> None:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
crossattn_cache = []
if isinstance(transformer, CausalWanModel):
num_attention_heads = transformer.num_heads
attention_head_dim = transformer.dim // transformer.num_heads
elif isinstance(transformer, CausalWanTransformer3DModel):
num_attention_heads = transformer.num_attention_heads
attention_head_dim = transformer.attention_head_dim
for _ in range(len(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,
})
return crossattn_cache
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.model import WanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_ori_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = WanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
seq_len = math.ceil((160 * 90) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states,
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
@@ -0,0 +1,144 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.causal_model import CausalWanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_train_ori_causal_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = CausalWanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
new_state_dict = {}
for k, v in causal_state_dict.items():
if k.startswith("model."):
new_state_dict[k.replace("model.", "")] = v
causal_state_dict = new_state_dict
model1.load_state_dict(causal_state_dict)
model1.num_frame_per_block = 3
model2.num_frame_per_block = 3
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
seq_len = math.ceil((160 * 90) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.randint(0, 1000, (batch_size, 21), device=device, dtype=torch.long)
logger.info("timestep: %s", timestep)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states,
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
+492 -99
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
@@ -36,10 +38,11 @@ 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, count_trainable,
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, set_random_seed
shift_timestep, compute_density_for_timestep_sampling, get_sigmas)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
import wandb # isort: skip
@@ -69,18 +72,11 @@ class DistillationPipeline(TrainingPipeline):
current_trainstep: int
video_latent_shape: tuple[int, ...]
video_latent_shape_sp: tuple[int, ...]
real_score_transformer: torch.nn.Module
fake_score_transformer: torch.nn.Module
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_trainable(self) -> None:
super().set_trainable()
self.modules["real_score_transformer"].requires_grad_(False)
self.modules["vae"].requires_grad_(False)
def initialize_training_pipeline(self, training_args: TrainingArgs):
"""Initialize the distillation training pipeline with multiple models."""
logger.info("Initializing distillation pipeline...")
@@ -89,14 +85,40 @@ class DistillationPipeline(TrainingPipeline):
self.noise_scheduler = self.get_module("scheduler")
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
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")
self.transformer_2 = self.get_module("transformer_2", None)
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
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()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
@@ -108,6 +130,39 @@ class DistillationPipeline(TrainingPipeline):
self.real_score_transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
# Initialize optimizers
fake_score_params = list(
@@ -119,10 +174,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,
)
@@ -150,8 +208,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 *
@@ -161,6 +230,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."""
@@ -170,11 +315,120 @@ class DistillationPipeline(TrainingPipeline):
def _prepare_distillation(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Prepare training environment for distillation."""
self.transformer.requires_grad_(True)
self.transformer.train()
if self.transformer_2 is not None:
self.transformer_2.requires_grad_(True)
self.transformer_2.train()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
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,
@@ -221,7 +475,9 @@ class DistillationPipeline(TrainingPipeline):
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
pred_noise = current_model(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
@@ -263,6 +519,7 @@ class DistillationPipeline(TrainingPipeline):
max_target_idx = len(self.denoising_step_list) - 1
noise_latents = []
noise_latent_index = target_timestep_idx_int - 1
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
if max_target_idx > 0:
# Run student model for all steps before the target timestep
with torch.no_grad():
@@ -274,7 +531,7 @@ class DistillationPipeline(TrainingPipeline):
training_batch_temp = self._build_distill_input_kwargs(
current_noise_latents, current_timestep_tensor,
training_batch.conditional_dict, training_batch)
pred_flow = self.transformer(
pred_flow = current_model(
**training_batch_temp.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_clean = pred_noise_to_pred_video(
@@ -317,7 +574,7 @@ class DistillationPipeline(TrainingPipeline):
training_batch = self._build_distill_input_kwargs(
noisy_input, target_timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
pred_noise = current_model(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
@@ -331,6 +588,7 @@ class DistillationPipeline(TrainingPipeline):
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
original_latent = generator_pred_video
with torch.no_grad():
timestep = torch.randint(0,
self.num_train_timestep, [1],
@@ -355,7 +613,7 @@ class DistillationPipeline(TrainingPipeline):
noisy_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
timestep).detach().unflatten(0, (1, generator_pred_video.shape[1]))
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
@@ -404,24 +662,24 @@ class DistillationPipeline(TrainingPipeline):
pred_real_video_uncond) * self.real_score_guidance_scale
grad = (faker_score_pred_video - real_score_pred_video) / torch.abs(
generator_pred_video - real_score_pred_video).mean()
original_latent - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
generator_pred_video.float(),
(generator_pred_video.float() - grad.float()).detach())
original_latent.float(),
(original_latent.float() - grad.float()).detach())
training_batch.dmd_latent_vis_dict.update({
"training_batch_dmd_fwd_clean_latent":
training_batch.latents,
"generator_pred_video":
generator_pred_video,
original_latent.detach(),
"real_score_pred_video":
real_score_pred_video,
real_score_pred_video.detach(),
"faker_score_pred_video":
faker_score_pred_video,
faker_score_pred_video.detach(),
"dmd_timestep":
timestep,
timestep.detach(),
})
return dmd_loss
@@ -514,16 +772,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
@@ -557,6 +816,9 @@ class DistillationPipeline(TrainingPipeline):
batches.append(batch)
self.optimizer.zero_grad()
# TODO: confirm this
if self.transformer_2 is not None:
self.optimizer_2.zero_grad()
total_dmd_loss = 0.0
dmd_latent_vis_dict = {}
fake_score_latent_vis_dict = {}
@@ -585,9 +847,32 @@ class DistillationPipeline(TrainingPipeline):
attn_metadata=batch_gen.attn_metadata_vsa):
(dmd_loss / gradient_accumulation_steps).backward()
total_dmd_loss += dmd_loss.detach().item()
self._clip_model_grad_norm_(batch_gen, self.transformer)
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self._clip_model_grad_norm_(batch_gen, self.transformer_2)
for param in self.transformer_2.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer_2.step()
self.optimizer_2.zero_grad(set_to_none=True)
else:
self._clip_model_grad_norm_(batch_gen, self.transformer)
for param in self.transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.generator_ema is not None:
# TODO: support EMA for transformer_2?
if self.train_transformer_2 and self.transformer_2 is not None:
# Note: EMA currently only supports the main transformer
# Could be extended to support transformer_2 in the future
pass
else:
self.generator_ema.update(self.transformer)
avg_dmd_loss = torch.tensor(total_dmd_loss /
gradient_accumulation_steps,
device=self.device)
@@ -611,9 +896,18 @@ class DistillationPipeline(TrainingPipeline):
fake_score_latent_vis_dict.update(
batch_fake.fake_score_latent_vis_dict)
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
for param in self.fake_score_transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
self.lr_scheduler.step()
# Step the appropriate scheduler
if self.train_transformer_2 and self.transformer_2 is not None:
self.lr_scheduler_2.step()
else:
self.lr_scheduler.step()
self.fake_score_optimizer.zero_grad(set_to_none=True)
avg_fake_score_loss = torch.tensor(total_fake_score_loss /
gradient_accumulation_steps,
@@ -638,7 +932,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
@@ -669,6 +964,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
@@ -700,6 +1003,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]
@@ -715,50 +1030,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()
@@ -835,16 +1198,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:
@@ -896,14 +1259,6 @@ class DistillationPipeline(TrainingPipeline):
else:
set_random_seed(seed + self.global_rank)
# Check trainable params
num_trainable_generator = round(
count_trainable(self.transformer) / 1e9, 3)
num_trainable_critic = round(
count_trainable(self.fake_score_transformer) / 1e9, 3)
logger.info(
"rank: %s: # of trainable params in generator: %sB, # of trainable params in critic: %sB",
self.global_rank, num_trainable_generator, num_trainable_critic)
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
@@ -913,6 +1268,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()
@@ -956,6 +1315,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)
@@ -969,11 +1336,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)
@@ -1001,6 +1376,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":
@@ -1032,7 +1416,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()
@@ -1049,7 +1434,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:
@@ -1069,7 +1458,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()
cleanup_dist_env_and_memory()
File diff suppressed because it is too large Load Diff
+144 -51
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import dataclasses
import math
import os
@@ -22,7 +23,7 @@ from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
# from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
@@ -39,20 +40,26 @@ from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, count_trainable, get_scheduler,
get_sigmas, load_checkpoint, normalize_dit_input, save_checkpoint,
compute_density_for_timestep_sampling, get_scheduler, get_sigmas,
load_checkpoint, normalize_dit_input, save_checkpoint,
shard_latents_across_sp)
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
# from fastvideo.utils import (is_vmoba_available, is_vsa_available,
# set_random_seed, shallow_asdict)
from fastvideo.utils import (is_vsa_available,
set_random_seed, shallow_asdict)
import wandb # isort: skip
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
# vmoba_available = is_vmoba_available()
logger = init_logger(__name__)
def _get_trainable_params(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
class TrainingPipeline(LoRAPipeline, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
@@ -63,6 +70,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
train_transformer_2: bool = False
def __init__(
self,
@@ -98,6 +106,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
self.transformer_2 = self.get_module("transformer_2", None)
self.seed = training_args.seed
self.set_schemas()
@@ -110,17 +119,25 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
# Set grads for proper modules based on the training mode (Distill, LoRA, etc.)
self.set_trainable()
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,
)
@@ -138,6 +155,30 @@ class TrainingPipeline(LoRAPipeline, ABC):
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
if self.transformer_2 is not None:
# Ensure transformer_2 has trainable parameters before creating optimizer
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
@@ -152,7 +193,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
seed=self.seed)
self.noise_scheduler = noise_scheduler
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
@@ -178,9 +219,25 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
self.transformer.train()
self.optimizer.zero_grad()
if self.transformer_2 is not None:
self.transformer_2.train()
self.optimizer_2.zero_grad()
training_batch.total_loss = 0.0
return training_batch
def _enable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
"""Enable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = True
model.train()
optimizer.zero_grad()
def _disable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
"""Disable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = False
optimizer.zero_grad(set_to_none=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
@@ -224,17 +281,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=latents.device)
timesteps = self._sample_timesteps(batch_size, latents.device)
# Enable training for the model that will be trained next and disable the other
if self.train_transformer_2:
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self._enable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
@@ -257,6 +314,38 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _sample_timesteps(self, batch_size, device):
# Determine which model to train based on the boundary timestep
if (self.transformer_2 is not None and self.boundary_timestep is not None and
torch.rand(1, generator=self.noise_random_generator).item() <= self.training_args.boundary_ratio):
self.train_transformer_2 = True
else:
self.train_transformer_2 = False
# Broadcast the decision to all processes
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
dist.broadcast(decision, src=0)
self.train_transformer_2 = decision.item() == 1.0
# Sample u from the appropriate range
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
boundary_ratio = self.training_args.boundary_ratio
if self.train_transformer_2:
u = (1 - boundary_ratio) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
else:
u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
latents_shape = training_batch.raw_latent_shape
@@ -272,20 +361,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
patch_size=patch_size,
VSA_sparsity=current_vsa_sparsity,
device=get_local_torch_device())
elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
moba_params = self.training_args.moba_config.copy()
moba_params.update({
"current_timestep":
training_batch.timesteps,
"raw_latent_shape":
training_batch.raw_latent_shape[2:5],
"patch_size":
self.training_args.pipeline_config.dit_config.patch_size,
"device":
get_local_torch_device(),
})
training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
).build(**moba_params)
else:
training_batch.attn_metadata = None
@@ -310,7 +385,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
@@ -321,11 +396,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
# [1000.0],
# device=training_batch.noisy_model_input.device,
# dtype=torch.bfloat16)
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
with set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
model_pred = current_model(**input_kwargs)
if self.training_args.precondition_outputs:
assert training_batch.sigmas is not None
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
@@ -356,7 +432,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
model_parts = [self.transformer]
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
@@ -401,9 +482,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch = self._clip_grad_norm(training_batch)
self.optimizer.step()
self.lr_scheduler.step()
# Only step the optimizer and scheduler for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
return training_batch
@@ -431,10 +517,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
local_main_process_only=False)
if not self.post_init_called:
self.post_init()
num_trainable_params = count_trainable(self.transformer)
num_trainable_params = _get_trainable_params(self.transformer)
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
if getattr(self, "transformer_2", None) is not None:
num_trainable_params = _get_trainable_params(self.transformer_2)
logger.info("Transformer 2: Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
@@ -455,7 +546,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self._log_training_info()
self._log_validation(self.transformer, self.training_args,
self._log_validation(self.training_args,
self.init_steps)
# Train!
@@ -476,9 +567,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
elif vmoba_available:
# TODO: add vmoba sparsity scheduling here
pass
else:
current_vsa_sparsity = 0.0
@@ -520,10 +608,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
self._log_validation(self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
count_trainable(self.transformer) / 1e9, 3)
_get_trainable_params(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
@@ -559,7 +647,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
logger.info(" Total optimization steps = %s",
self.training_args.max_train_steps)
logger.info(" Total training parameters per FSDP shard = %s B",
round(count_trainable(self.transformer) / 1e9, 3))
round(_get_trainable_params(self.transformer) / 1e9, 3))
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
@@ -601,12 +689,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
return batch
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
def _log_validation(self, training_args, global_step) -> None:
"""
Generate a validation video and log it to wandb to check the quality during training.
"""
training_args.inference_mode = True
training_args.dit_cpu_offload = True
training_args.dit_cpu_offload = False
if not training_args.log_validation:
return
if self.validation_pipeline is None:
@@ -627,7 +715,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
transformer.eval()
self.transformer.eval()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
@@ -719,4 +810,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
self.transformer.train()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.train()
+168 -3
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")
@@ -1280,5 +1296,154 @@ def get_scheduler(
last_epoch=last_epoch)
def count_trainable(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
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,72 @@
# 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 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)
+57
View File
@@ -0,0 +1,57 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='2f25ad37933894dbf0966c838c0b8494987f9f2f'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/
DATA_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset
# VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
VALIDATION_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16\
--sp_size 4 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp-shard-dim 8 \
--train_sp_batch_size 1 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps 1 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 1000 \
--validation_steps 30 \
--validation_sampling_steps "40" \
--log_validation True \
--checkpoints_total_limit 3 \
--ema_start_step 0 \
--training_cfg_rate 0.1 \
--seed 1024 \
--output_dir "outputs_train_test/wan_finetune_v1" \
--tracker_project_name VSA_finetune \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 5 \
--validation_guidance_scale "5.0" \
--num_euler_timesteps 50 \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--weight_decay 0.01 \
--max_grad_norm 1.0