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