Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eb56c7f8c5 | ||
|
|
e63fe531d8 |
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -1,76 +0,0 @@
|
||||
{
|
||||
"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": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"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": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"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": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,13 +0,0 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,151 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29501
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=4
|
||||
|
||||
# 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/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
|
||||
VALIDATION_DATASET_FILE="data/mixkit-64_processed/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -1,24 +0,0 @@
|
||||
#!/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"
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -17,13 +17,12 @@ def main():
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
ti2v_task=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
sampling_param.image_path = "test.jpg"
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
@@ -45,4 +44,4 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -19,15 +19,13 @@ def main():
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.num_frames = 81
|
||||
|
||||
prompts = [
|
||||
"A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
]
|
||||
|
||||
for prompt in prompts:
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
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()
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
|
||||
The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.
|
||||
The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.
|
||||
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
|
||||
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
|
||||
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
|
||||
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.
|
||||
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
|
||||
The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
|
||||
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
|
||||
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
|
||||
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
|
||||
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
|
||||
The video shows a large orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
|
||||
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
|
||||
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
|
||||
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
|
||||
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
|
||||
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
|
||||
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
|
||||
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
|
||||
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
|
||||
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
|
||||
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
|
||||
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
|
||||
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
|
||||
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.
|
||||
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
|
||||
@@ -1,93 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--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 77
|
||||
--warp_denoising_step
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
-25
@@ -1,25 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
|
||||
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 \
|
||||
--flow_shift 5.0 \
|
||||
--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 "ode_trajectory"
|
||||
@@ -1,40 +0,0 @@
|
||||
{
|
||||
"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": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,99 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-extended-t2v-1-3b/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_70k"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "fixed_wan_ode_init_70k_6e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 6000
|
||||
--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 77
|
||||
--warp_denoising_step
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,135 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=1e5B2_16kFV_warp_ode_vidprom
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.out
|
||||
#SBATCH --error=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
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 WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b-81/"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "Dwarp_vidprom_8b16k_test_warp_1e-5"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "Dwarp_vidprom_8b16k_wan_ode_init_1e-5"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--warp_denoising_step
|
||||
--log_visualization
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--dmd_denoising_steps "1000,750,500,250"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,131 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=ode_vidprom2k
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_vidprom2k_output/ode_vidprom2k.out
|
||||
#SBATCH --error=ode_vidprom2k_output/ode_vidprom2k.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
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 WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing/"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_vidprom2k"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "vidprom2k_wan_ode_init_5e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 6001
|
||||
--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 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 8
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,132 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=ode_crush
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_crush_output/ode_crush.out
|
||||
#SBATCH --error=ode_crush_output/ode_crush.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
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 WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
|
||||
NUM_GPUS=2
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_warp_2"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "2warp_fixed_wan_ode_init_5e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
# --warp_denoising_step
|
||||
--max_train_steps 6001
|
||||
--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 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,103 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export MASTER_PORT=29501
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
|
||||
# DATA_DIR="/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_lmdb_sf/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "Dwarp_vidprom_8b16k_test_warp_1e-5"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "Dwarp_vidprom_8b16k_wan_ode_init_1e-5"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--warp_denoising_step
|
||||
--log_visualization
|
||||
--max_train_steps 10
|
||||
--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 77
|
||||
--dmd_denoising_steps "1000,750,500,250"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
# --log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--seed 1024
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,98 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/wl/FastVideo2/data/crush-smol_processed_t2v_1_3b_ode_init_single"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "overfitwan_ode_init_crush_smol"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 2001
|
||||
# --warp_denoising_step
|
||||
--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 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,100 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "debug_ode_init"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "debug_ode_init"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 1000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--warp_denoising_step
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 10
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,25 +0,0 @@
|
||||
#!/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/vidprom_1.txt"
|
||||
OUTPUT_DIR="data/ode_vidprom_1_fv/"
|
||||
|
||||
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 \
|
||||
--flow_shift 5.0 \
|
||||
--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 "ode_trajectory"
|
||||
@@ -1,76 +0,0 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Elon Musk, dressed in a sleek white spacesuit with a reflective visor, walks confidently across the lunar surface. His posture is upright, and he moves steadily with purpose. The moon's rocky terrain and scattered boulders surround him, casting shadows under the dim sunlight. The background shows vast stretches of the moon's barren landscape with craters and dust clouds kicked up by his boots. The scene captures a wide shot, emphasizing the vastness and desolation of the lunar environment. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "In a dynamic action-packed sequence set in the Marvel multiverse, Spider-Man and Venom engage in an intense battle. Spider-Man, in his classic red and blue suit, swings and dodges venomous attacks from the black symbiote-covered Venom. Both characters display a range of acrobatic moves and powerful strikes. The environment is a chaotic urban landscape with crumbling buildings and neon lights, reflecting the multiversal theme. The camera captures the epic fight from various angles, including wide shots to show the scale of destruction and close-ups to highlight their fierce expressions and physical combat. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A warm, family-oriented scene depicting a father getting ready to leave the house to buy milk. The father, a middle-aged man with a kind face and a casual outfit, picks up a jacket from the coat rack. His posture is upright as he bends down slightly to put on his shoes. In the background, there are glimpses of a cozy living room with a family photograph on the wall. The camera focuses closely on the father, capturing his gentle smile and reassuring nod towards the camera before he opens the front door and steps outside. Static medium close-up shot. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Close-up shot of a man with a prosthetic hand that functions as a rocket launcher. He looks at his new hand with a mix of amazement and concern, his facial expression showing a blend of curiosity and apprehension. The prosthetic hand is sleek and metallic, with intricate details that resemble a high-tech weapon. The background is a dimly lit laboratory with various scientific equipment and monitors displaying data. The man stands in a relaxed posture, his other hand resting on his hip, as he inspects his new limb. The scene is rendered in a realistic sci-fi style, emphasizing the futuristic technology and the man's emotional response to his new appendage. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Realistic CCTV footage style, Kim Taehyung from the band BTS is involved in a drug deal, caught on camera. Kim Taehyung appears nervous and cautious, wearing casual clothing typical of a public space. He exchanges items discreetly with another person, who is partially obscured. Both individuals maintain a watchful demeanor, occasionally glancing around to ensure no one is watching them. The lighting is dim, with flickering fluorescent lights casting shadows on their faces. The background shows a typical urban setting with blurred figures moving in the distance. Static camera angle, medium close-up shot focusing on the interaction between Taehyung and the other individual. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Photorealistic studio setup with professional lighting, showcasing detailed cubic dissections of experimental plastic and felt-like materials on a pristine white background. Each cube reveals intricate layers and textures of the materials, emphasizing their unique properties. The scene has a shallow depth of field initially, then slowly pulls out to reveal the full arrangement of cubes, maintaining a wide depth of field throughout the transition. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -6,8 +6,8 @@ export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_old"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=4
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
@@ -52,7 +52,7 @@ dataset_args=(
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 50
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
@@ -4,7 +4,7 @@ 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_old/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
GPU_NUM=2 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATASET_PATH="data/crush-smol/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
@@ -14,7 +14,7 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.preprocess_video_batch_size 8 \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
|
||||
@@ -28,4 +28,4 @@
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -1,94 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_old"
|
||||
VALIDATION_DATASET_FILE="examples/datasets/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_i2v_finetune"
|
||||
--output_dir "checkpoints/wan_t2v_i2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 20
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--t2v_as_i2v_task True
|
||||
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_t2v_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,24 +0,0 @@
|
||||
#!/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_i2v_1_3b/"
|
||||
|
||||
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 2 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--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_ode_trajectory"
|
||||
@@ -3,7 +3,6 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.selector import backend_name_to_enum, get_attn_backend
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
|
||||
@@ -37,7 +36,7 @@ class DistributedAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.bfloat16 if envs.FASTVIDEO_FORCE_ATTN_BF16 else get_compute_dtype()
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
@@ -222,7 +221,7 @@ class LocalAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.bfloat16 if envs.FASTVIDEO_FORCE_ATTN_BF16 else get_compute_dtype()
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
|
||||
@@ -27,7 +27,6 @@ class DiTArchConfig(ArchConfig):
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -4,13 +4,12 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -45,13 +45,10 @@ class PipelineConfig:
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
disable_autocast: bool = False
|
||||
ti2v_task: bool = False
|
||||
t2v_as_i2v_task: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
dit_forward_precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
@@ -90,7 +87,6 @@ class PipelineConfig:
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
@@ -218,24 +214,6 @@ class PipelineConfig:
|
||||
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
|
||||
)
|
||||
|
||||
# TI2V task
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}ti2v-task",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}ti2v_task",
|
||||
default=PipelineConfig.ti2v_task,
|
||||
help="Enable TI2V",
|
||||
)
|
||||
|
||||
# T2V to I2V task
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}t2v-as-i2v-task",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}t2v_as_i2v_task",
|
||||
default=PipelineConfig.t2v_as_i2v_task,
|
||||
help="Enable T2V to I2V task",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
@@ -267,9 +245,7 @@ class PipelineConfig:
|
||||
"""
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
logger.info("WTF model_path: %s", model_path)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
logger.info("pipeline_config_cls: %s", pipeline_config_cls)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
|
||||
|
||||
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
# isort: off
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig)
|
||||
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
|
||||
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -48,9 +48,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -62,7 +60,6 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
@@ -48,8 +48,6 @@ class SamplingParam:
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.data_type = "video" if self.num_frames > 1 else "image"
|
||||
@@ -207,18 +205,6 @@ class SamplingParam:
|
||||
help=
|
||||
"Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-latents",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_latents,
|
||||
help="Whether to return the trajectory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-decoded",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_decoded,
|
||||
help="Whether to return the decoded trajectory",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -37,7 +37,7 @@ def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
|
||||
def gettextdataset(args) -> TextDataset:
|
||||
return TextDataset(data_merge_path=args.data_merge_path,
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.dataset.lmdb_utils import get_array_shape_from_lmdb, retrieve_row_from_lmdb
|
||||
from torch.utils.data import Dataset
|
||||
import numpy as np
|
||||
import torch
|
||||
import lmdb
|
||||
|
||||
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/dataset.py
|
||||
class ODERegressionLMDBDataset(Dataset):
|
||||
def __init__(self, data_path: str, max_pair: int = int(1e8)):
|
||||
print(f"data_path: {data_path}")
|
||||
self.env = lmdb.open(data_path, readonly=True,
|
||||
lock=False, readahead=False, meminit=False)
|
||||
|
||||
self.latents_shape = get_array_shape_from_lmdb(self.env, 'latents')
|
||||
self.max_pair = max_pair
|
||||
|
||||
def __len__(self):
|
||||
return min(self.latents_shape[0], self.max_pair)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""
|
||||
Outputs:
|
||||
- prompts: List of Strings
|
||||
- latents: Tensor of shape (num_denoising_steps, num_frames, num_channels, height, width). It is ordered from pure noise to clean image.
|
||||
"""
|
||||
latents = retrieve_row_from_lmdb(
|
||||
self.env,
|
||||
"latents", np.float16, idx, shape=self.latents_shape[1:]
|
||||
)
|
||||
|
||||
if len(latents.shape) == 4:
|
||||
latents = latents[None, ...]
|
||||
|
||||
prompts = retrieve_row_from_lmdb(
|
||||
self.env,
|
||||
"prompts", str, idx
|
||||
)
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"ode_latent": torch.tensor(latents, dtype=torch.float32)
|
||||
}
|
||||
|
||||
@@ -1,75 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/lmdb.py
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_array_shape_from_lmdb(env, array_name):
|
||||
with env.begin() as txn:
|
||||
image_shape = txn.get(f"{array_name}_shape".encode()).decode()
|
||||
image_shape = tuple(map(int, image_shape.split()))
|
||||
return image_shape
|
||||
|
||||
|
||||
def store_arrays_to_lmdb(env, arrays_dict, start_index=0):
|
||||
"""
|
||||
Store rows of multiple numpy arrays in a single LMDB.
|
||||
Each row is stored separately with a naming convention.
|
||||
"""
|
||||
with env.begin(write=True) as txn:
|
||||
for array_name, array in arrays_dict.items():
|
||||
for i, row in enumerate(array):
|
||||
# Convert row to bytes
|
||||
if isinstance(row, str):
|
||||
row_bytes = row.encode()
|
||||
else:
|
||||
row_bytes = row.tobytes()
|
||||
|
||||
data_key = f'{array_name}_{start_index + i}_data'.encode()
|
||||
|
||||
txn.put(data_key, row_bytes)
|
||||
|
||||
|
||||
def process_data_dict(data_dict, seen_prompts):
|
||||
output_dict = {}
|
||||
|
||||
all_videos = []
|
||||
all_prompts = []
|
||||
for prompt, video in data_dict.items():
|
||||
if prompt in seen_prompts:
|
||||
continue
|
||||
else:
|
||||
seen_prompts.add(prompt)
|
||||
|
||||
video = video.half().numpy()
|
||||
all_videos.append(video)
|
||||
all_prompts.append(prompt)
|
||||
|
||||
if len(all_videos) == 0:
|
||||
return {"latents": np.array([]), "prompts": np.array([])}
|
||||
|
||||
all_videos = np.concatenate(all_videos, axis=0)
|
||||
|
||||
output_dict['latents'] = all_videos
|
||||
output_dict['prompts'] = np.array(all_prompts)
|
||||
|
||||
return output_dict
|
||||
|
||||
|
||||
def retrieve_row_from_lmdb(lmdb_env, array_name, dtype, row_index, shape=None):
|
||||
"""
|
||||
Retrieve a specific row from a specific array in the LMDB.
|
||||
"""
|
||||
data_key = f'{array_name}_{row_index}_data'.encode()
|
||||
|
||||
with lmdb_env.begin() as txn:
|
||||
row_bytes = txn.get(data_key)
|
||||
|
||||
if dtype == str:
|
||||
array = row_bytes.decode()
|
||||
else:
|
||||
array = np.frombuffer(row_bytes, dtype=dtype)
|
||||
|
||||
if shape is not None and len(shape) > 0:
|
||||
array = array.reshape(shape)
|
||||
return array
|
||||
@@ -3,9 +3,6 @@ from typing import Any, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
@@ -344,9 +344,6 @@ class VideoGenerator:
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
"logging_info": logging_info,
|
||||
"trajectory": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
def set_lora_adapter(self,
|
||||
|
||||
@@ -18,7 +18,6 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_PREFIX: str = ""
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_FORCE_ATTN_BF16: bool = False
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_ATTENTION_CONFIG: str | None = None
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
|
||||
@@ -169,10 +168,6 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_TRACE_FUNCTION":
|
||||
lambda: int(os.getenv("FASTVIDEO_TRACE_FUNCTION", "0")),
|
||||
|
||||
# if set, fastvideo will force attention to be computed in bfloat16
|
||||
"FASTVIDEO_FORCE_ATTN_BF16":
|
||||
lambda: bool(int(os.getenv("FASTVIDEO_FORCE_ATTN_BF16", "0"))),
|
||||
|
||||
# Backend for attention computation
|
||||
# Available options:
|
||||
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
|
||||
|
||||
+1
-106
@@ -158,7 +158,6 @@ class FastVideoArgs:
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
# # DMD parameters
|
||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||
@@ -397,12 +396,6 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-transformer-cls-name",
|
||||
type=str,
|
||||
default=FastVideoArgs.override_transformer_cls_name,
|
||||
help="Override transformer cls name",
|
||||
)
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -612,11 +605,6 @@ 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
|
||||
@@ -639,7 +627,6 @@ 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
|
||||
@@ -671,7 +658,6 @@ 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
|
||||
@@ -692,30 +678,16 @@ 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
|
||||
intermediate_latents_visualization: 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":
|
||||
@@ -817,20 +789,6 @@ 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,
|
||||
@@ -901,10 +859,6 @@ 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")
|
||||
@@ -1009,10 +963,6 @@ 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")
|
||||
@@ -1063,13 +1013,6 @@ 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,
|
||||
@@ -1086,11 +1029,6 @@ 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,
|
||||
@@ -1103,49 +1041,6 @@ 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
|
||||
|
||||
@@ -1153,4 +1048,4 @@ class TrainingArgs(FastVideoArgs):
|
||||
def parse_int_list(value: str) -> list[int]:
|
||||
if not value:
|
||||
return []
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
|
||||
@@ -212,9 +212,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
modulated = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
(1.0 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
modulated = normalized * (1 + scale) + shift
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
@@ -267,11 +267,11 @@ class LayerNormScaleShift(nn.Module):
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
output = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
(1.0 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
# scale.shape: [batch_size, 1, inner_dim]
|
||||
# shift.shape: [batch_size, 1, inner_dim]
|
||||
output = normalized * (1 + scale) + shift
|
||||
output = normalized * (1.0 + scale) + shift
|
||||
|
||||
if self.compute_dtype == torch.float32:
|
||||
output = output.to(x.dtype)
|
||||
|
||||
@@ -147,9 +147,6 @@ class CausalWanSelfAttention(nn.Module):
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
# kv_cache["k"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = self.attn(
|
||||
@@ -179,7 +176,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
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)
|
||||
@@ -212,7 +209,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
# Only T2V for now
|
||||
@@ -225,7 +223,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -250,34 +249,29 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
num_frames = temb.shape[1]
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb
|
||||
e = self.scale_shift_table + temb.float()
|
||||
# e.shape: [batch_size, num_frames, 6, inner_dim]
|
||||
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
# assert shift_msa.dtype == torch.float32
|
||||
|
||||
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
|
||||
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
|
||||
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2)
|
||||
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
|
||||
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.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
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))
|
||||
@@ -291,6 +285,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -299,10 +295,13 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
crossattn_cache=crossattn_cache)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -365,7 +364,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -375,7 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.num_frame_per_block = 3
|
||||
self.num_frame_per_block = 1
|
||||
self.independent_first_frame = False
|
||||
|
||||
self.__post_init__()
|
||||
@@ -487,16 +487,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -543,9 +539,14 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
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 torch.stack(output)
|
||||
return output
|
||||
|
||||
def _forward_train(self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -556,8 +557,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
start_frame: int = 0,
|
||||
**kwargs) -> torch.Tensor:
|
||||
|
||||
logger.info("timestep dtype: %s, timestep sum: %s", timestep.dtype, timestep.float().sum().item())
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
@@ -588,8 +587,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
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:
|
||||
@@ -602,12 +601,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -642,9 +637,14 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
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 torch.stack(output)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -652,34 +652,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
**kwargs
|
||||
):
|
||||
if kwargs.get('kv_cache', None) is not None:
|
||||
noise_pred = self._forward_inference(*args, **kwargs)
|
||||
return self._forward_inference(*args, **kwargs)
|
||||
else:
|
||||
noise_pred = self._forward_train(*args, **kwargs)
|
||||
|
||||
return noise_pred
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
return self._forward_train(*args, **kwargs)
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
@@ -39,14 +37,16 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
dtype = encoder_hidden_states_image.dtype
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states).to(dtype)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
@@ -156,12 +156,12 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
|
||||
if crossattn_cache is not None:
|
||||
if not crossattn_cache["is_init"]:
|
||||
crossattn_cache["is_init"] = True
|
||||
k = self.norm_k.forward_native(self.to_k(context)[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)
|
||||
crossattn_cache["k"] = k
|
||||
crossattn_cache["v"] = v
|
||||
@@ -169,16 +169,11 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
k = crossattn_cache["k"]
|
||||
v = crossattn_cache["v"]
|
||||
else:
|
||||
k = self.norm_k.forward_native(self.to_k(context)[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 envs.FASTVIDEO_FORCE_ATTN_BF16:
|
||||
out_dtype = v.dtype
|
||||
# compute attention
|
||||
x = self.attn(q.to(torch.bfloat16), k.to(torch.bfloat16), v.to(torch.bfloat16)).to(out_dtype)
|
||||
else:
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
@@ -218,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
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)
|
||||
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
|
||||
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
|
||||
b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
@@ -252,7 +247,7 @@ class WanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
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)
|
||||
@@ -283,29 +278,29 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
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")
|
||||
@@ -324,11 +319,12 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
|
||||
if temb.dim() == 4:
|
||||
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb
|
||||
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)
|
||||
@@ -339,20 +335,22 @@ class WanTransformerBlock(nn.Module):
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
|
||||
e = self.scale_shift_table + temb
|
||||
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
|
||||
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
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.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
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))
|
||||
@@ -364,12 +362,7 @@ class WanTransformerBlock(nn.Module):
|
||||
is_neox_style=False), _apply_rotary_emb(
|
||||
key, cos, sin, is_neox_style=False)
|
||||
|
||||
if envs.FASTVIDEO_FORCE_ATTN_BF16:
|
||||
out_dtype = value.dtype
|
||||
attn_output, _ = self.attn1(query.to(torch.bfloat16), key.to(torch.bfloat16), value.to(torch.bfloat16))
|
||||
attn_output = attn_output.to(out_dtype)
|
||||
else:
|
||||
attn_output, _ = self.attn1(query, key, value)
|
||||
attn_output, _ = self.attn1(query, key, value)
|
||||
attn_output = attn_output.flatten(2)
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
@@ -377,20 +370,26 @@ class WanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -407,7 +406,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
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)
|
||||
@@ -439,7 +438,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -459,7 +459,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -479,22 +480,23 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb
|
||||
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
|
||||
norm_hidden_states = (self.norm1(hidden_states) *
|
||||
(1 + scale_msa) + shift_msa)
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
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))
|
||||
@@ -519,6 +521,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -526,15 +530,17 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
|
||||
class WanTransformer3DModel(CachableDiT):
|
||||
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = WanVideoConfig()._compile_conditions
|
||||
@@ -592,7 +598,8 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -652,12 +659,10 @@ class WanTransformer3DModel(CachableDiT):
|
||||
rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
|
||||
@@ -667,8 +672,6 @@ class WanTransformer3DModel(CachableDiT):
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
@@ -725,35 +728,14 @@ class WanTransformer3DModel(CachableDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
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 torch.stack(output)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
return output
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
@@ -845,4 +827,5 @@ class WanTransformer3DModel(CachableDiT):
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
return hidden_states + self.previous_residual_odd
|
||||
|
||||
@@ -238,7 +238,6 @@ class TextEncoderLoader(ComponentLoader):
|
||||
1]
|
||||
|
||||
target_device = get_local_torch_device()
|
||||
logger.info("Loading text encoder in %s precision", encoder_precision)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
fastvideo_args, encoder_precision)
|
||||
@@ -416,10 +415,6 @@ class TransformerLoader(ComponentLoader):
|
||||
raise ValueError(
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
logger.info("transformer cls_name: %s", cls_name)
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
cls_name = fastvideo_args.override_transformer_cls_name
|
||||
logger.info("Overriding transformer cls_name to %s", cls_name)
|
||||
|
||||
fastvideo_args.model_paths["transformer"] = model_path
|
||||
|
||||
@@ -435,23 +430,11 @@ 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)
|
||||
|
||||
default_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
param_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_forward_precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s, default_dtype: %s", cls_name,
|
||||
@@ -471,8 +454,7 @@ class TransformerLoader(ComponentLoader):
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
|
||||
fsdp_inference=fastvideo_args.use_fsdp_inference,
|
||||
# TODO(will): make these configurable
|
||||
default_dtype=default_dtype,
|
||||
param_dtype=param_dtype,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
training_mode=fastvideo_args.training_mode)
|
||||
@@ -481,11 +463,9 @@ class TransformerLoader(ComponentLoader):
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
# Need to convert the model to the default_dtype
|
||||
# Otherwise, the model master weights will be in param_dtype, and the gradients will also be in param_dtype
|
||||
# This means the param update will be in lower precision, causing precision loss
|
||||
logger.info("Converting model to dtype: %s", default_dtype)
|
||||
model = model.to(default_dtype)
|
||||
dtypes = set(param.dtype for param in model.parameters())
|
||||
if len(dtypes) > 1:
|
||||
model = model.to(default_dtype)
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -62,7 +62,6 @@ def maybe_load_fsdp_model(
|
||||
device: torch.device,
|
||||
hsdp_replicate_dim: int,
|
||||
hsdp_shard_dim: int,
|
||||
default_dtype: torch.dtype,
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
cpu_offload: bool = False,
|
||||
@@ -88,7 +87,7 @@ def maybe_load_fsdp_model(
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
with set_default_dtype(param_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
# Check if we should use FSDP
|
||||
@@ -126,7 +125,7 @@ def maybe_load_fsdp_model(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
default_dtype,
|
||||
param_dtype,
|
||||
strict=True,
|
||||
cpu_offload=cpu_offload,
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
|
||||
@@ -635,31 +635,8 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
|
||||
"""
|
||||
Args:
|
||||
clean_latent: the clean latent with shape [B, C, H, W],
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise: the noise with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
|
||||
Returns:
|
||||
the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
@@ -22,10 +22,8 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
|
||||
config_name = "scheduler_config.json"
|
||||
|
||||
order = 1
|
||||
@register_to_config
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
@@ -64,15 +62,8 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
elif timestep.ndim == 0:
|
||||
# handles the case where timestep is a scalar, this occurs when we
|
||||
# use this scheduler for ODE trajectory
|
||||
timestep = timestep.unsqueeze(0)
|
||||
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep = timestep.to(model_output.device)
|
||||
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -171,34 +171,12 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
# timestep shape should be [B]
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
pred_noise = pred_noise.double().to(device)
|
||||
noise_input_latent = noise_input_latent.double().to(device)
|
||||
sigmas = scheduler.sigmas.double().to(device)
|
||||
timesteps = scheduler.timesteps.double().to(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 pred_video_to_pred_noise(x0_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor, scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert x0 prediction to flow matching's prediction.
|
||||
x0_pred: the x0 prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = (x_t - x_0) / sigma_t
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = x0_pred.dtype
|
||||
x0_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(x0_pred.device), [x0_pred, xt,
|
||||
scheduler.sigmas,
|
||||
scheduler.timesteps]
|
||||
)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
flow_pred = (xt - x0_pred) / sigma_t
|
||||
return flow_pred.to(original_dtype)
|
||||
|
||||
@@ -28,6 +28,10 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ class ComposedPipelineBase(ABC):
|
||||
_extra_config_module_map: dict[str, str] = {}
|
||||
training_args: TrainingArgs | None = None
|
||||
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
|
||||
modules: dict[str, Any] = {}
|
||||
modules: dict[str, torch.nn.Module] = {}
|
||||
post_init_called: bool = False
|
||||
|
||||
# TODO(will): args should support both inference args and training args
|
||||
@@ -121,25 +121,14 @@ class ComposedPipelineBase(ABC):
|
||||
model_path: str,
|
||||
device: str | None = None,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
pipeline_config: PipelineConfig | None = None,
|
||||
pipeline_config: str | PipelineConfig | None = None,
|
||||
args: argparse.Namespace | None = None,
|
||||
required_config_modules: list[str] | None = None,
|
||||
loaded_modules: dict[str, torch.nn.Module]
|
||||
| None = None,
|
||||
**kwargs) -> "ComposedPipelineBase":
|
||||
"""
|
||||
Load a pipeline from a pretrained model.
|
||||
Few different patterns are supported:
|
||||
- Only provide model_path:
|
||||
- This will load the pipeline in inference mode.
|
||||
- The pipeline will be initialized with the default config.
|
||||
- The pipeline will be initialized with the default modules.
|
||||
- The pipeline will be initialized with the default stages.
|
||||
- The pipeline will be initialized with the default stages.
|
||||
- override the default config using pipeline_config or args or kwargs
|
||||
- override the default modules using loaded_modules
|
||||
- override the pipelineconfig
|
||||
|
||||
Load a pipeline from a pretrained model.
|
||||
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
|
||||
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
|
||||
"""
|
||||
@@ -147,18 +136,9 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
kwargs['model_path'] = model_path
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
|
||||
if pipeline_config is not None:
|
||||
fastvideo_args.pipeline_config = pipeline_config
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
|
||||
fastvideo_args.pipeline_config = pipeline_config
|
||||
else:
|
||||
assert args is not None, "args must be provided for training mode"
|
||||
fastvideo_args = TrainingArgs.from_cli_args(args)
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
|
||||
fastvideo_args.pipeline_config = pipeline_config
|
||||
logger.info("in 2 Overriding transformer cls name to %s", fastvideo_args.override_transformer_cls_name)
|
||||
# TODO(will): fix this so that its not so ugly
|
||||
fastvideo_args.model_path = model_path
|
||||
for key, value in kwargs.items():
|
||||
@@ -169,8 +149,7 @@ class ComposedPipelineBase(ABC):
|
||||
# model is loaded with the correct precision. Subsequently we will
|
||||
# use FSDP2's MixedPrecisionPolicy to set the precision for the
|
||||
# fwd, bwd, and other operations' precision.
|
||||
fastvideo_args.pipeline_config.dit_precision = 'fp32'
|
||||
# assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
|
||||
@@ -148,12 +148,7 @@ class ForwardBatch:
|
||||
modules: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Final output (after pipeline completion)
|
||||
output: torch.Tensor | None = None
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
trajectory_timesteps: list[int] | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_decoded: list[torch.Tensor] | None = None
|
||||
output: Any = None
|
||||
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
@@ -212,10 +207,6 @@ class TrainingBatch:
|
||||
infos: list[dict[str, Any]] | None = None
|
||||
mask_lat_size: torch.Tensor | None = None
|
||||
|
||||
# ODE trajectory supervision
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: torch.Tensor | None = None
|
||||
timesteps: torch.Tensor | None = None
|
||||
@@ -246,7 +237,6 @@ class TrainingBatch:
|
||||
fake_score_loss: float = 0.0
|
||||
|
||||
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
latent_vis_dict: dict[str, torch.Tensor] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -1,443 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
ODE Trajectory Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
|
||||
using the modular pipeline architecture.
|
||||
|
||||
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
ode_text_only_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video, pred_video_to_pred_noise
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
"""ODE Trajectory preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int
|
||||
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for ODE Trajectory pipeline."""
|
||||
return pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
assert fastvideo_args.pipeline_config.flow_shift == 5
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
|
||||
denoising_strength=1.0)
|
||||
|
||||
fastvideo_args.model_loaded["transformer"] = False
|
||||
loader = TransformerLoader()
|
||||
fastvideo_args.pipeline_config.dit_precision = "fp32" # Overwrite the precision to fp32 for transformer
|
||||
fastvideo_args.pipeline_config.dit_forward_precision = "fp32"
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
self.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
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="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=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
pipeline=self,
|
||||
))
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
|
||||
args):
|
||||
"""Preprocess text-only data and generate trajectory information."""
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
# logger.info("transformer weight sum: %s", sum(p.float().sum().item() for p in self.transformer.parameters()))
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, text in enumerate(data["text"]):
|
||||
if text and text.strip(): # Check if text is not empty
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples (text-only)
|
||||
valid_data = {
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
# Add fps and duration if available in data
|
||||
if "fps" in data:
|
||||
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
|
||||
if "duration" in data:
|
||||
valid_data["duration"] = [
|
||||
data["duration"][i] for i in valid_indices
|
||||
]
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
self.prompt_encoding_stage.text_encoders[0] = self.prompt_encoding_stage.text_encoders[0].to(dtype=torch.bfloat16).to(dtype=torch.float32)
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
logger.info("prompt_embeds sum: %s, prompt_embeds shape: %s, prompt_embeds dtype: %s", prompt_embeds.float().sum(), prompt_embeds.shape, prompt_embeds.dtype)
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
sampling_params = SamplingParam.from_pretrained(args.model_path)
|
||||
|
||||
negative_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
|
||||
|
||||
# encode negative prompt for trajectory collection
|
||||
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
|
||||
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
negative_prompt,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
negative_prompt_embed = negative_prompt_embeds_list[0]
|
||||
negative_prompt_attention_mask = negative_prompt_masks_list[
|
||||
0]
|
||||
else:
|
||||
negative_prompt_embed = None
|
||||
negative_prompt_attention_mask = None
|
||||
|
||||
trajectory_latents = []
|
||||
trajectory_timesteps = []
|
||||
trajectory_decoded = []
|
||||
|
||||
for i, (prompt_embed, prompt_attention_mask) in enumerate(
|
||||
zip(prompt_embeds, prompt_attention_masks,
|
||||
strict=False)):
|
||||
prompt_embed = prompt_embed.unsqueeze(0)
|
||||
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
|
||||
|
||||
# Collect the trajectory data (text-to-video generation)
|
||||
batch = ForwardBatch(**shallow_asdict(sampling_params), )
|
||||
batch.prompt_embeds = [prompt_embed]
|
||||
batch.prompt_attention_mask = [prompt_attention_mask]
|
||||
batch.negative_prompt_embeds = [negative_prompt_embed]
|
||||
batch.negative_attention_mask = [
|
||||
negative_prompt_attention_mask
|
||||
]
|
||||
batch.num_inference_steps = 48
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
batch.return_trajectory_decoded = True
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
batch.fps = args.train_fps
|
||||
batch.guidance_scale = 3.0
|
||||
batch.do_classifier_free_guidance = True
|
||||
|
||||
result_batch = self.input_validation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.timestep_preparation_stage(
|
||||
batch, fastvideo_args)
|
||||
# result_batch = self.latent_preparation_stage(
|
||||
# result_batch, fastvideo_args)
|
||||
# result_batch = self.denoising_stage(result_batch,
|
||||
# fastvideo_args)
|
||||
noisy_input = []
|
||||
# latents = result_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
latents = torch.randn(
|
||||
[1, 21, 16, 60, 104], dtype=torch.float32, device=get_local_torch_device()
|
||||
)
|
||||
# logger.info("transformer weight sum: %s", sum(p.float().sum().item() for p in self.transformer.parameters()))
|
||||
logger.info("latents sum: %s, latents shape: %s, latents dtype: %s", latents.float().sum(), latents.shape, latents.dtype)
|
||||
|
||||
logger.info("scheduler timesteps: %s", self.get_module("scheduler").timesteps)
|
||||
for progress_id, t in enumerate(tqdm(self.get_module("scheduler").timesteps)):
|
||||
timestep = t * \
|
||||
torch.ones([1, 21], device=latents.device, dtype=torch.float32)
|
||||
|
||||
noisy_input.append(latents)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
# logger.info("prompt_embed sum: %s, prompt_embed shape: %s, prompt_embed dtype: %s", prompt_embed.float().sum(), prompt_embed.shape, prompt_embed.dtype)
|
||||
# logger.info("timestep: %s", timestep[:, 0])
|
||||
# Run transformer
|
||||
cond_pred_noise_btchw = self.transformer(
|
||||
hidden_states=latents.permute(0, 2, 1, 3, 4),
|
||||
encoder_hidden_states=prompt_embed,
|
||||
timestep=timestep[:, 0]
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# logger.info("cond_pred_noise_btchw sum: %s, cond_pred_noise_btchw shape: %s, cond_pred_noise_btchw dtype: %s", cond_pred_noise_btchw.float().sum(), cond_pred_noise_btchw.shape, cond_pred_noise_btchw.dtype)
|
||||
|
||||
cond_pred_video_btchw = pred_noise_to_pred_video(
|
||||
pred_noise=cond_pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=latents.flatten(0, 1),
|
||||
timestep=timestep.flatten(0, 1),
|
||||
scheduler=self.get_module("scheduler")).unflatten(
|
||||
0, cond_pred_noise_btchw.shape[:2])
|
||||
|
||||
# logger.info("cond_pred_video_btchw sum: %s, cond_pred_video_btchw shape: %s, cond_pred_video_btchw dtype: %s", cond_pred_video_btchw.float().sum(), cond_pred_video_btchw.shape, cond_pred_video_btchw.dtype)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=t,
|
||||
attn_metadata=None,
|
||||
forward_batch=result_batch,
|
||||
):
|
||||
# logger.info("latents sum: %s, latents shape: %s, latents dtype: %s", latents.float().sum(), latents.shape, latents.dtype)
|
||||
# Run transformer
|
||||
uncond_pred_noise_btchw = self.transformer(
|
||||
latents.permute(0, 2, 1, 3, 4),
|
||||
negative_prompt_embed,
|
||||
timestep[:, 0]
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# logger.info("uncond_pred_noise_btchw sum: %s, uncond_pred_noise_btchw shape: %s, uncond_pred_noise_btchw dtype: %s", uncond_pred_noise_btchw.float().sum(), uncond_pred_noise_btchw.shape, uncond_pred_noise_btchw.dtype)
|
||||
|
||||
uncond_pred_video_btchw = pred_noise_to_pred_video(
|
||||
pred_noise=uncond_pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=latents.flatten(0, 1),
|
||||
timestep=timestep.flatten(0, 1),
|
||||
scheduler=self.get_module("scheduler")).unflatten(
|
||||
0, uncond_pred_noise_btchw.shape[:2])
|
||||
|
||||
pred_video_btchw = uncond_pred_video_btchw + batch.guidance_scale * (
|
||||
cond_pred_video_btchw - uncond_pred_video_btchw
|
||||
)
|
||||
|
||||
# logger.info("pred_video_btchw sum: %s, pred_video_btchw shape: %s, pred_video_btchw dtype: %s", pred_video_btchw.float().sum(), pred_video_btchw.shape, pred_video_btchw.dtype)
|
||||
|
||||
pred_noise_btchw = pred_video_to_pred_noise(
|
||||
x0_pred=pred_video_btchw.flatten(0, 1),
|
||||
xt=latents.flatten(0, 1),
|
||||
timestep=timestep.flatten(0, 1),
|
||||
scheduler=self.get_module("scheduler")).unflatten(
|
||||
0, pred_video_btchw.shape[:2])
|
||||
|
||||
# logger.info("pred_noise_btchw sum: %s, pred_noise_btchw shape: %s, pred_noise_btchw dtype: %s", pred_noise_btchw.float().sum(), pred_noise_btchw.shape, pred_noise_btchw.dtype)
|
||||
|
||||
latents = self.get_module("scheduler").step(
|
||||
pred_noise_btchw.flatten(0, 1),
|
||||
self.get_module("scheduler").timesteps[progress_id] * torch.ones(
|
||||
[1, 21], device=latents.device, dtype=torch.long).flatten(0, 1),
|
||||
latents.flatten(0, 1)
|
||||
)[0].unflatten(dim=0, sizes=pred_noise_btchw.shape[:2])
|
||||
|
||||
# logger.info("latents sum: %s, latents shape: %s, latents dtype: %s", latents.float().sum(), latents.shape, latents.dtype)
|
||||
|
||||
noisy_input.append(latents)
|
||||
|
||||
noisy_inputs = torch.stack(noisy_input, dim=1)
|
||||
|
||||
noisy_inputs = noisy_inputs[:, [0, 12, 24, 36, -1]].half()
|
||||
|
||||
logger.info("noisy inputs sum: %s, noisy inputs shape: %s, noisy inputs dtype: %s", noisy_inputs.float().sum(), noisy_inputs.shape, noisy_inputs.dtype)
|
||||
|
||||
result_batch.trajectory_latents = noisy_inputs.permute(0, 1, 3, 2, 4, 5)
|
||||
result_batch.trajectory_timesteps = torch.tensor([self.get_module("scheduler").timesteps[i] for i in [0, 12, 24, 36, -1]])
|
||||
result_batch.latents = latents.permute(0, 2, 1, 3, 4)
|
||||
result_batch = self.decoding_stage(result_batch,
|
||||
fastvideo_args)
|
||||
trajectory_latents.append(
|
||||
result_batch.trajectory_latents.cpu())
|
||||
trajectory_timesteps.append(
|
||||
result_batch.trajectory_timesteps.cpu())
|
||||
trajectory_decoded.append(result_batch.trajectory_decoded)
|
||||
|
||||
trajectory_latents = torch.stack(trajectory_latents, dim=0).squeeze(0)
|
||||
# trajecotry_latents = trajectory_latents[:, [0, 12, 24, 36, -1]]
|
||||
|
||||
# Prepare extra features for text-only processing
|
||||
extra_features = {
|
||||
"trajectory_latents": trajectory_latents,
|
||||
"trajectory_timesteps": trajectory_timesteps
|
||||
}
|
||||
|
||||
if batch.return_trajectory_decoded:
|
||||
for i, decoded_frames in enumerate(trajectory_decoded):
|
||||
for j, decoded_frame in enumerate(decoded_frames):
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data: list[dict[str, Any]] = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
for idx, video_path in save_pbar:
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Get extra features for this sample
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).numpy()
|
||||
else:
|
||||
assert isinstance(value, list)
|
||||
if isinstance(value[idx], torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).float().numpy()
|
||||
else:
|
||||
sample_extra_features[key] = value[idx]
|
||||
|
||||
# Create record for Parquet dataset (text-only ODE schema)
|
||||
record: dict[str, Any] = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
trajectory_latents=sample_extra_features[
|
||||
"trajectory_latents"],
|
||||
trajectory_timesteps=sample_extra_features[
|
||||
"trajectory_timesteps"],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=self.combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
self.num_processed_samples = 0
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
if written:
|
||||
logger.info("Final flush wrote %s samples", written)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for video preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing videos",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.video_data: dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: dict[str, Any] = {} # Store latent tensors
|
||||
self.preprocess_text_and_trajectory(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_ODE_Trajectory
|
||||
@@ -10,8 +10,6 @@ from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
|
||||
PreprocessPipeline_I2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
|
||||
PreprocessPipeline_ODE_Trajectory)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
|
||||
@@ -50,16 +48,13 @@ def main(args) -> None:
|
||||
text_encoder_cpu_offload=False,
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
|
||||
if args.preprocess_task == "t2v":
|
||||
PreprocessPipeline = PreprocessPipeline_T2V
|
||||
elif args.preprocess_task == "i2v":
|
||||
PreprocessPipeline = PreprocessPipeline_I2V
|
||||
elif args.preprocess_task == "text_only":
|
||||
PreprocessPipeline = PreprocessPipeline_Text
|
||||
elif args.preprocess_task == "ode_trajectory":
|
||||
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
|
||||
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
|
||||
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
|
||||
else:
|
||||
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only")
|
||||
@@ -105,11 +100,10 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--flow_shift", type=float, default=None)
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
|
||||
choices=["t2v", "i2v", "text_only"],
|
||||
help="Type of preprocessing task to run")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
|
||||
@@ -78,8 +78,6 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
else:
|
||||
assert False, "warp_denoising_step must be true"
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
logger.info("Using timesteps: %s", timesteps)
|
||||
|
||||
|
||||
@@ -50,63 +50,6 @@ class DecodingStage(PipelineStage):
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
"""
|
||||
Decode latent representations into pixel space using VAE.
|
||||
|
||||
Args:
|
||||
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
|
||||
fastvideo_args: Configuration containing:
|
||||
- disable_autocast: Whether to disable automatic mixed precision (default: False)
|
||||
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
|
||||
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
|
||||
|
||||
Returns:
|
||||
Decoded video tensor with shape (batch, channels, frames, height, width),
|
||||
normalized to [0, 1] range and moved to CPU as float32
|
||||
"""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
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
|
||||
|
||||
# Apply shifting if needed
|
||||
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
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
@@ -116,28 +59,13 @@ class DecodingStage(PipelineStage):
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
This method processes the batch through the VAE decoder, converting latent
|
||||
representations to pixel-space video/images. It also optionally decodes
|
||||
trajectory latents for visualization purposes.
|
||||
|
||||
Args:
|
||||
batch: The current batch containing:
|
||||
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
|
||||
- return_trajectory_decoded (optional): Flag to decode trajectory latents
|
||||
- trajectory_latents (optional): Latents at different timesteps
|
||||
- trajectory_timesteps (optional): Corresponding timesteps
|
||||
fastvideo_args: Configuration containing:
|
||||
- output_type: "latent" to skip decoding, otherwise decode to pixels
|
||||
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
|
||||
- model_loaded: Track VAE loading state
|
||||
- model_paths: Path to VAE model if loading needed
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
Modified batch with:
|
||||
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
|
||||
- trajectory_decoded (if requested): List of decoded frames per timestep
|
||||
The batch with decoded outputs.
|
||||
"""
|
||||
# load vae if not already loaded (used for memory constrained devices)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
loader = VAELoader()
|
||||
@@ -147,29 +75,58 @@ class DecodingStage(PipelineStage):
|
||||
pipeline.add_module("vae", self.vae)
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
if fastvideo_args.output_type == "latent":
|
||||
frames = batch.latents
|
||||
else:
|
||||
frames = self.decode(batch.latents, fastvideo_args)
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
batch.trajectory_decoded = []
|
||||
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
|
||||
for idx in range(batch.trajectory_latents.shape[1]):
|
||||
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
|
||||
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
|
||||
cur_timestep = batch.trajectory_timesteps[idx]
|
||||
logger.info("decoding trajectory latent for timestep: %s",
|
||||
cur_timestep)
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
latents = batch.latents
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if latents is None:
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
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
|
||||
|
||||
# Apply shifting if needed
|
||||
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
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
frames = frames.cpu().float()
|
||||
image = image.cpu().float()
|
||||
|
||||
# Update batch with decoded image
|
||||
batch.output = frames
|
||||
batch.output = image
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
|
||||
@@ -140,12 +140,11 @@ class DenoisingStage(PipelineStage):
|
||||
latents = latents[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.latents = latents
|
||||
if batch.image_latent is not None:
|
||||
if not fastvideo_args.pipeline_config.ti2v_task and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b c (n t) h w -> b c n t h w",
|
||||
n=sp_world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b c (n t) h w -> b c n t h w",
|
||||
n=sp_world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
@@ -205,13 +204,13 @@ class DenoisingStage(PipelineStage):
|
||||
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if batch.boundary_ratio is not None:
|
||||
logger.info("Overriding boundary ratio from %s to %s",
|
||||
boundary_ratio, batch.boundary_ratio)
|
||||
boundary_ratio = batch.boundary_ratio
|
||||
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
|
||||
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if batch.boundary_ratio is not None:
|
||||
logger.info("Overriding boundary ratio from %s to %s",
|
||||
boundary_ratio, batch.boundary_ratio)
|
||||
boundary_ratio = batch.boundary_ratio
|
||||
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
@@ -254,9 +253,6 @@ class DenoisingStage(PipelineStage):
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
trajectory_timesteps: list[int] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -284,27 +280,14 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
if batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
elif batch.image_latent is not None and fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
assert batch.image_latent is not None, "image latents should be provided for T2V to I2V task"
|
||||
if rank_in_sp_group == 0:
|
||||
logger.info("latent_model_input.shape: %s",
|
||||
latent_model_input.shape)
|
||||
latent_model_input = torch.cat([
|
||||
batch.image_latent,
|
||||
latent_model_input[:, :, 1:, :, :],
|
||||
],
|
||||
dim=2).to(target_dtype)
|
||||
logger.info("latent_model_input.shape: %s",
|
||||
latent_model_input.shape)
|
||||
|
||||
assert not torch.isnan(
|
||||
latent_model_input).any(), "latent_model_input contains nan"
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
timestep = torch.stack([t]).to(get_local_torch_device())
|
||||
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
|
||||
@@ -319,13 +302,6 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
if rank_in_sp_group == 0:
|
||||
latent_model_input = torch.cat([
|
||||
batch.image_latent,
|
||||
latent_model_input[:, :, 1:, :, :],
|
||||
],
|
||||
dim=2).to(target_dtype)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
guidance_expand = (
|
||||
@@ -457,11 +433,6 @@ class DenoisingStage(PipelineStage):
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# save trajectory latents if needed
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(t)
|
||||
trajectory_latents.append(latents)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
@@ -469,35 +440,9 @@ class DenoisingStage(PipelineStage):
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
trajectory_tensor: torch.Tensor | None = None
|
||||
if trajectory_latents:
|
||||
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
|
||||
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
|
||||
dim=0)
|
||||
else:
|
||||
trajectory_tensor = None
|
||||
trajectory_timesteps_tensor = None
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
if sp_group:
|
||||
latents = sequence_model_parallel_all_gather(latents, dim=2)
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_tensor = trajectory_tensor.to(
|
||||
get_local_torch_device())
|
||||
trajectory_tensor = sequence_model_parallel_all_gather(
|
||||
trajectory_tensor, dim=3)
|
||||
|
||||
if trajectory_tensor is not None:
|
||||
batch.trajectory_timesteps = torch.tensor(trajectory_timesteps).cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
|
||||
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
latents = torch.cat([
|
||||
batch.image_latent,
|
||||
latents[:, :, 1:, :, :],
|
||||
],
|
||||
dim=2)
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
@@ -105,81 +105,6 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
def __init__(self, vae: ParallelTiledVAE) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def encode_image(self,
|
||||
image: PIL.Image.Image,
|
||||
height: int,
|
||||
width: int,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Encode image into latent space.
|
||||
"""
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=height,
|
||||
width=width).to(get_local_torch_device(), dtype=torch.float32)
|
||||
|
||||
# (B, C, H, W) -> (B, C, 1, H, W)
|
||||
print(f"image.shape: {image.shape}")
|
||||
image = image.unsqueeze(2)
|
||||
print(f"after unsqueeze image.shape: {image.shape}")
|
||||
return self.encode_tensor(image, fastvideo_args, generator)
|
||||
|
||||
def encode_tensor(self,
|
||||
video_condition: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Encode frames into latent space.
|
||||
"""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
video_condition = video_condition.to(device=get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
|
||||
latent_condition = encoder_output.mean
|
||||
else:
|
||||
generator = generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
return latent_condition
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
@@ -232,28 +157,57 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
# (B, C, H, W) -> (B, C, 1, H, W)
|
||||
image = image.unsqueeze(2)
|
||||
|
||||
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
# repeat the image self.vae.temporal_compression_ratio times
|
||||
video_condition = image.repeat(1, 1,
|
||||
self.vae.temporal_compression_ratio,
|
||||
1, 1)
|
||||
# video_condition = image
|
||||
logger.info("video_condition.shape: %s", video_condition.shape)
|
||||
else:
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
|
||||
image.shape[3], image.shape[4])
|
||||
],
|
||||
dim=2)
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
|
||||
image.shape[3], image.shape[4])
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
latent_condition = self.encode_tensor(video_condition, fastvideo_args,
|
||||
batch.generator)
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
|
||||
latent_condition = encoder_output.mean
|
||||
else:
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
|
||||
batch.image_latent = latent_condition
|
||||
elif fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
logger.info("latent_condition.shape: %s", latent_condition.shape)
|
||||
batch.image_latent = latent_condition
|
||||
else:
|
||||
mask_lat_size = torch.ones(1, 1, num_frames, latent_height,
|
||||
|
||||
@@ -35,15 +35,9 @@ class InputValidationStage(PipelineStage):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = batch.seed
|
||||
num_videos_per_prompt = batch.num_videos_per_prompt
|
||||
if isinstance(batch.prompt, list):
|
||||
num_prompts = len(batch.prompt)
|
||||
else:
|
||||
num_prompts = 1
|
||||
|
||||
total_num_videos = num_prompts * num_videos_per_prompt
|
||||
|
||||
assert seed is not None
|
||||
seeds = [seed + i for i in range(total_num_videos)]
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
batch.seeds = seeds
|
||||
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
|
||||
batch.generator = [
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
@@ -94,7 +93,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
# Generate or use provided latents
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape,
|
||||
generator=torch.Generator(device="cuda").manual_seed(1024),
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
else:
|
||||
|
||||
@@ -82,8 +82,8 @@ def rocm_platform_plugin() -> str | None:
|
||||
logger.info("ROCm platform is available")
|
||||
finally:
|
||||
amdsmi.amdsmi_shut_down()
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.info("ROCm platform is unavailable: %s", e)
|
||||
|
||||
return "fastvideo.platforms.rocm.RocmPlatform" if is_rocm else None
|
||||
|
||||
|
||||
@@ -8,9 +8,6 @@ from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.wan.modules.tokenizers import HuggingfaceTokenizer
|
||||
from fastvideo.wan.modules.t5 import umt5_xxl
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -41,22 +38,9 @@ def test_t5_encoder():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision_str = "fp32"
|
||||
precision = PRECISION_TO_TYPE[precision_str]
|
||||
# model1 = UMT5EncoderModel.from_pretrained(TEXT_ENCODER_PATH).to(
|
||||
# precision).to(device).eval()
|
||||
# tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
model1 = umt5_xxl(
|
||||
encoder_only=True,
|
||||
return_tokenizer=False,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).eval().requires_grad_(False)
|
||||
model1.load_state_dict(
|
||||
torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing-clean/wan_models/Wan2.1-T2V-1.3B/models_t5_umt5-xxl-enc-bf16.pth",
|
||||
map_location='cpu', weights_only=False)
|
||||
)
|
||||
|
||||
tokenizer1 = HuggingfaceTokenizer(
|
||||
name="/mnt/weka/home/hao.zhang/wei/Self-Forcing-clean/wan_models/Wan2.1-T2V-1.3B/google/umt5-xxl/", seq_len=512, clean='whitespace')
|
||||
model1 = UMT5EncoderModel.from_pretrained(TEXT_ENCODER_PATH).to(
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH,
|
||||
@@ -65,9 +49,8 @@ def test_t5_encoder():
|
||||
pin_cpu_memory=False)
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, args)
|
||||
model2 = model2.to(dtype=torch.bfloat16).to(precision)
|
||||
model2 = model2.to(precision)
|
||||
model2.eval()
|
||||
tokenizer2 = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
# Sanity check weights between the two models
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
@@ -78,13 +61,8 @@ def test_t5_encoder():
|
||||
logger.info("Model1 has %s parameters", len(params1))
|
||||
logger.info("Model2 has %s parameters", len(params2))
|
||||
|
||||
model1_weight_sum = sum(p.float().sum().item() for p in model1.parameters())
|
||||
model2_weight_sum = sum(p.float().sum().item() for p in model2.parameters())
|
||||
logger.info("Model1 weight sum: %s", model1_weight_sum)
|
||||
logger.info("Model2 weight sum: %s", model2_weight_sum)
|
||||
|
||||
# weight_diffs = []
|
||||
# # check if embed_tokens are the same
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
@@ -92,21 +70,18 @@ def test_t5_encoder():
|
||||
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
# name1 = w.format(idx)
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
# p1 = params1[name1]
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(device)
|
||||
# assert_close(p1, p2, atol=1e-4, rtol=1e-4)
|
||||
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
|
||||
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
# Test with some sample prompts
|
||||
# prompts = [
|
||||
# "Once upon a time", "The quick brown fox jumps over",
|
||||
# "In a galaxy far, far away"
|
||||
# ]
|
||||
prompts = [
|
||||
"A vibrant scene of Kenyan golfers at a lush green golf course on a sunny day. The golfers, dressed in casual yet stylish attire, are teeing off with animated expressions, showcasing their enthusiasm for the game. Rolling hills and pristine greens stretch out behind them, creating a picturesque backdrop. In the foreground, a golf buggy and a caddy stand ready, adding to the serene atmosphere. The camera captures the action from a mid-shot angle, focusing on the golfers' dynamic motions as they swing their clubs."
|
||||
"Once upon a time", "The quick brown fox jumps over",
|
||||
"In a galaxy far, far away"
|
||||
]
|
||||
|
||||
logger.info("Testing T5 encoder with sample prompts")
|
||||
@@ -116,8 +91,7 @@ def test_t5_encoder():
|
||||
logger.info("Testing prompt: %s", prompt)
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens1, mask = tokenizer1(prompt, return_mask=True, add_special_tokens=True)
|
||||
tokens2 = tokenizer2(prompt,
|
||||
tokens = tokenizer(prompt,
|
||||
padding="max_length",
|
||||
max_length=512,
|
||||
truncation=True,
|
||||
@@ -127,23 +101,22 @@ def test_t5_encoder():
|
||||
# filter out padding input_ids
|
||||
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
|
||||
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
|
||||
outputs1 = model1(tokens1.to(device),
|
||||
mask.to(device))
|
||||
outputs1 = model1(input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True).last_hidden_state
|
||||
print("--------------------------------")
|
||||
logger.info("Testing model2")
|
||||
|
||||
# Get outputs from our implementation
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs2 = model2(
|
||||
input_ids=tokens2.input_ids,
|
||||
attention_mask=tokens2.attention_mask,
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
).last_hidden_state
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1[mask == 1]
|
||||
last_hidden_state2 = outputs2[tokens2.attention_mask == 1]
|
||||
logger.info("last_hidden_state1 sum: %s", last_hidden_state1.float().sum())
|
||||
logger.info("last_hidden_state2 sum: %s", last_hidden_state2.float().sum())
|
||||
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
|
||||
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
@@ -28,14 +28,14 @@ HUNYUAN_PARAMS = {
|
||||
"width": 1280,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 6,
|
||||
"guidance_scale": 1,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 17,
|
||||
# "guidance_scale": 1,
|
||||
# "embedded_cfg_scale": 6,
|
||||
# "flow_shift": 17,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
# "vae_sp": True,
|
||||
# "fps": 24,
|
||||
}
|
||||
|
||||
WAN_T2V_PARAMS = {
|
||||
@@ -45,14 +45,14 @@ WAN_T2V_PARAMS = {
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
"guidance_scale": 3,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
# "guidance_scale": 3,
|
||||
# "embedded_cfg_scale": 6,
|
||||
# "flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
# "vae_sp": True,
|
||||
# "fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": ("fp32",)
|
||||
}
|
||||
@@ -64,18 +64,33 @@ WAN_I2V_PARAMS = {
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 6,
|
||||
"guidance_scale": 5.0,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
# "guidance_scale": 5.0,
|
||||
# "embedded_cfg_scale": 6,
|
||||
# "flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
# "vae_sp": True,
|
||||
# "fps": 24,
|
||||
# "neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": ("fp32",)
|
||||
}
|
||||
|
||||
WAN2_2_I2V_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
# "guidance_scale": 5.0,
|
||||
# "embedded_cfg_scale": 6,
|
||||
# "flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
}
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"FastHunyuan-diffusers": HUNYUAN_PARAMS,
|
||||
"Wan2.1-T2V-1.3B-Diffusers": WAN_T2V_PARAMS,
|
||||
@@ -83,6 +98,7 @@ MODEL_TO_PARAMS = {
|
||||
|
||||
I2V_MODEL_TO_PARAMS = {
|
||||
"Wan2.1-I2V-14B-480P-Diffusers": WAN_I2V_PARAMS,
|
||||
"Wan2.2-I2V-A14B-Diffusers": WAN2_2_I2V_PARAMS,
|
||||
}
|
||||
|
||||
TEST_PROMPTS = [
|
||||
@@ -125,7 +141,6 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"flow_shift": BASE_PARAMS["flow_shift"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
}
|
||||
@@ -142,10 +157,7 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"height": BASE_PARAMS["height"],
|
||||
"width": BASE_PARAMS["width"],
|
||||
"num_frames": BASE_PARAMS["num_frames"],
|
||||
"guidance_scale": BASE_PARAMS["guidance_scale"],
|
||||
"embedded_cfg_scale": BASE_PARAMS["embedded_cfg_scale"],
|
||||
"seed": BASE_PARAMS["seed"],
|
||||
"fps": BASE_PARAMS["fps"],
|
||||
}
|
||||
if "neg_prompt" in BASE_PARAMS:
|
||||
generation_kwargs["neg_prompt"] = BASE_PARAMS["neg_prompt"]
|
||||
@@ -225,7 +237,6 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"flow_shift": BASE_PARAMS["flow_shift"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"dit_cpu_offload": True,
|
||||
@@ -242,10 +253,7 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"height": BASE_PARAMS["height"],
|
||||
"width": BASE_PARAMS["width"],
|
||||
"num_frames": BASE_PARAMS["num_frames"],
|
||||
"guidance_scale": BASE_PARAMS["guidance_scale"],
|
||||
"embedded_cfg_scale": BASE_PARAMS["embedded_cfg_scale"],
|
||||
"seed": BASE_PARAMS["seed"],
|
||||
"fps": BASE_PARAMS["fps"],
|
||||
}
|
||||
if "neg_prompt" in BASE_PARAMS:
|
||||
generation_kwargs["neg_prompt"] = BASE_PARAMS["neg_prompt"]
|
||||
|
||||
@@ -1,292 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_ori_causal_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = CausalWanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
|
||||
new_state_dict = {}
|
||||
for k, v in causal_state_dict.items():
|
||||
if k.startswith("model."):
|
||||
new_state_dict[k.replace("model.", "")] = v
|
||||
causal_state_dict = new_state_dict
|
||||
model1.load_state_dict(causal_state_dict)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
12,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
block_sizes = [3 for _ in range(4)]
|
||||
timesteps = [1000, 750, 500, 250]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
output1 = _causal_inference(model1, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
|
||||
logger.info("Finish inference for model1")
|
||||
output2 = _causal_inference(model2, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
logger.info("Output 1 Sum: %s", output1.float().sum().item())
|
||||
logger.info("Output 2 Sum: %s", output2.float().sum().item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
def _causal_inference(transformer, latents, prompt_embeds, block_sizes, timesteps, target_dtype):
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
start_index = 0
|
||||
pos_start_base = 0
|
||||
frame_seq_length = latents.shape[-1] * latents.shape[-2] // (WanVideoConfig().arch_config.patch_size[-1] * WanVideoConfig().arch_config.patch_size[-2])
|
||||
seq_len = frame_seq_length * latents.shape[2]
|
||||
kv_cache1 = _initialize_kv_cache(transformer, batch_size=latents.shape[0],
|
||||
kv_cache_size=frame_seq_length * latents.shape[2],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
crossattn_cache = _initialize_crossattn_cache(
|
||||
transformer,
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=WanVideoConfig().arch_config.text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
for current_num_frames, t_cur in zip(block_sizes, timesteps):
|
||||
# logger.info(f"Current frame idx: {start_index}, Current timestep: {t_cur}")
|
||||
# logger.info(f"k cache sum: {sum(kv_cache['k'].float().sum().item() for kv_cache in kv_cache1)}, v cache sum: {sum(kv_cache['v'].float().sum().item() for kv_cache in kv_cache1)}")
|
||||
# logger.info(f"latents sum: {latents.float().sum().item()}, encoder_hidden_states sum: {prompt_embeds.float().sum().item()}")
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
attn_metadata = None
|
||||
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(current_latents.shape[0], 1),
|
||||
device=current_latents.device,
|
||||
dtype=torch.long)
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
pred_noise_btchw = transformer(
|
||||
x=current_latents,
|
||||
context=prompt_embeds,
|
||||
t=t_expanded_noise,
|
||||
seq_len=seq_len,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length
|
||||
)
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
pred_noise_btchw = transformer(
|
||||
current_latents,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length,
|
||||
start_frame=start_index
|
||||
)
|
||||
|
||||
# Write back and advance
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = pred_noise_btchw.clone()
|
||||
|
||||
# Re-run with context timestep to update KV cache using clean context
|
||||
context_noise = 0
|
||||
t_context = torch.ones([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = pred_noise_btchw.to(target_dtype)
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
_ = transformer(
|
||||
x=context_bcthw,
|
||||
context=prompt_embeds,
|
||||
t=t_expanded_context,
|
||||
seq_len=seq_len,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length
|
||||
)
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
_ = transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length,
|
||||
start_frame=start_index
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
return latents
|
||||
|
||||
def _initialize_kv_cache(transformer, batch_size, kv_cache_size, dtype, device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
num_attention_heads = transformer.num_heads
|
||||
attention_head_dim = transformer.dim // transformer.num_heads
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
num_attention_heads = transformer.num_attention_heads
|
||||
attention_head_dim = transformer.attention_head_dim
|
||||
|
||||
for _ in range(len(transformer.blocks)):
|
||||
kv_cache1.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
return kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(transformer, batch_size, max_text_len, dtype,
|
||||
device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
crossattn_cache = []
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
num_attention_heads = transformer.num_heads
|
||||
attention_head_dim = transformer.dim // transformer.num_heads
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
num_attention_heads = transformer.num_attention_heads
|
||||
attention_head_dim = transformer.attention_head_dim
|
||||
|
||||
for _ in range(len(transformer.blocks)):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
return crossattn_cache
|
||||
@@ -1,144 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.model import WanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
os.environ["FASTVIDEO_FORCE_ATTN_BF16"] = "1"
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_ori_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.float32
|
||||
precision_str = "fp32"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str, dit_forward_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args)
|
||||
|
||||
model1 = WanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing-clean/wan_models/Wan2.1-T2V-1.3B")
|
||||
model1.eval()
|
||||
model1 = model1.to(device).to(precision)
|
||||
model1.requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 120
|
||||
seq_len = math.ceil((104 * 60) /
|
||||
(2 * 2) *
|
||||
21)
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
21,
|
||||
16,
|
||||
60,
|
||||
104,
|
||||
device=device,
|
||||
generator=torch.Generator("cuda").manual_seed(1024),
|
||||
dtype=precision)
|
||||
|
||||
logger.info("Hidden states sum: %s, Hidden states shape: %s, Hidden states dtype: %s", hidden_states.float().sum(), hidden_states.shape, hidden_states.dtype)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
# encoder_hidden_states = torch.randn(batch_size,
|
||||
# text_seq_len + 1,
|
||||
# 4096,
|
||||
# device=device,
|
||||
# dtype=precision)
|
||||
encoder_hidden_states = torch.load("../sf_cond_prompt_embeds.pt").to(device, dtype=precision)
|
||||
|
||||
logger.info("Encoder hidden states sum: %s, Encoder hidden states shape: %s, Encoder hidden states dtype: %s", encoder_hidden_states.float().sum(), encoder_hidden_states.shape, encoder_hidden_states.dtype)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([995.7627], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
# with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
x=hidden_states.permute(0, 2, 1, 3, 4),
|
||||
context=encoder_hidden_states,
|
||||
t=timestep,
|
||||
seq_len=seq_len,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states.permute(0, 2, 1, 3, 4),
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
logger.info("Output 1 sum: %s", output1.float().sum())
|
||||
logger.info("Output 2 sum: %s", output2.float().sum())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -1,144 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_train_ori_causal_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = CausalWanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
|
||||
new_state_dict = {}
|
||||
for k, v in causal_state_dict.items():
|
||||
if k.startswith("model."):
|
||||
new_state_dict[k.replace("model.", "")] = v
|
||||
causal_state_dict = new_state_dict
|
||||
model1.load_state_dict(causal_state_dict)
|
||||
|
||||
model1.num_frame_per_block = 3
|
||||
model2.num_frame_per_block = 3
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
seq_len = math.ceil((160 * 90) /
|
||||
(2 * 2) *
|
||||
21)
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.randint(0, 1000, (batch_size, 21), device=device, dtype=torch.long)
|
||||
logger.info("timestep: %s", timestep)
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
# with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
x=hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
t=timestep,
|
||||
seq_len=seq_len,
|
||||
)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -64,3 +64,5 @@ def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
|
||||
assert len(files2) == 2
|
||||
total = sum(pq.read_table(str(f)).num_rows for f in files2)
|
||||
assert total == 5
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import copy
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import abstractmethod
|
||||
@@ -12,7 +11,6 @@ from typing import Any
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
@@ -38,11 +36,10 @@ from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases, count_trainable,
|
||||
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)
|
||||
from fastvideo.utils import is_vsa_available, set_random_seed
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
@@ -72,11 +69,18 @@ class DistillationPipeline(TrainingPipeline):
|
||||
current_trainstep: int
|
||||
video_latent_shape: tuple[int, ...]
|
||||
video_latent_shape_sp: tuple[int, ...]
|
||||
real_score_transformer: torch.nn.Module
|
||||
fake_score_transformer: torch.nn.Module
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def set_trainable(self) -> None:
|
||||
super().set_trainable()
|
||||
self.modules["real_score_transformer"].requires_grad_(False)
|
||||
self.modules["vae"].requires_grad_(False)
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize the distillation training pipeline with multiple models."""
|
||||
logger.info("Initializing distillation pipeline...")
|
||||
@@ -85,37 +89,14 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=self.timestep_shift)
|
||||
|
||||
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.transformer is the generator model
|
||||
self.real_score_transformer = self.get_module("real_score_transformer")
|
||||
self.fake_score_transformer = self.get_module("fake_score_transformer")
|
||||
self.real_score_transformer.eval()
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
@@ -138,13 +119,10 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if fake_score_lr == 0.0:
|
||||
fake_score_lr = training_args.learning_rate
|
||||
|
||||
betas_str = training_args.fake_score_betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
self.fake_score_optimizer = torch.optim.AdamW(
|
||||
fake_score_params,
|
||||
lr=fake_score_lr,
|
||||
betas=betas,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
@@ -172,19 +150,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
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)
|
||||
logger.info("Distillation generator model to %s denoising steps",
|
||||
len(self.denoising_step_list))
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
|
||||
self.min_timestep = int(self.training_args.min_timestep_ratio *
|
||||
@@ -194,82 +161,6 @@ 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."""
|
||||
@@ -279,117 +170,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
def _prepare_distillation(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Prepare training environment for distillation."""
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
|
||||
return training_batch
|
||||
|
||||
def apply_ema_to_model(self, model):
|
||||
"""Apply EMA weights to the model for validation or inference."""
|
||||
if self.generator_ema is not None:
|
||||
with self.generator_ema.apply_to_model(model):
|
||||
return model
|
||||
return model
|
||||
|
||||
def get_ema_model_copy(self):
|
||||
"""Get a copy of the model with EMA weights applied."""
|
||||
if self.generator_ema is not None:
|
||||
ema_model = copy.deepcopy(self.transformer)
|
||||
self.generator_ema.copy_to_unwrapped(ema_model)
|
||||
return ema_model
|
||||
return None
|
||||
|
||||
def is_ema_ready(self, current_step: int = None):
|
||||
"""Check if EMA is ready for use (after ema_start_step)."""
|
||||
if current_step is None:
|
||||
current_step = getattr(self, 'current_trainstep', 0)
|
||||
return (self.generator_ema is not None
|
||||
and current_step >= self.training_args.ema_start_step)
|
||||
|
||||
def save_ema_weights(self, output_dir: str, step: int):
|
||||
"""Save EMA weights separately for inference purposes."""
|
||||
if self.generator_ema is None:
|
||||
logger.warning("Cannot save EMA weights: EMA not initialized")
|
||||
return
|
||||
|
||||
if not self.is_ema_ready():
|
||||
logger.warning(
|
||||
"Cannot save EMA weights: EMA not ready yet (step < ema_start_step)"
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
ema_model = self.get_ema_model_copy()
|
||||
if ema_model is None:
|
||||
logger.warning("Failed to create EMA model copy")
|
||||
return
|
||||
|
||||
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
|
||||
os.makedirs(ema_save_dir, exist_ok=True)
|
||||
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from fastvideo.training.training_utils import (
|
||||
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(ema_model, device=None)
|
||||
|
||||
if self.global_rank == 0:
|
||||
weight_path = os.path.join(
|
||||
ema_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict = custom_to_hf_state_dict(
|
||||
cpu_state, ema_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, weight_path)
|
||||
|
||||
config_dict = ema_model.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"]
|
||||
config_path = os.path.join(ema_save_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
|
||||
logger.info(f"EMA weights saved to {weight_path}")
|
||||
|
||||
del ema_model
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save EMA weights: {str(e)}")
|
||||
|
||||
def get_ema_stats(self):
|
||||
"""Get EMA statistics for monitoring."""
|
||||
if self.generator_ema is None:
|
||||
return {
|
||||
"ema_enabled": False,
|
||||
"ema_decay": None,
|
||||
"ema_start_step": self.training_args.ema_start_step,
|
||||
"ema_ready": False,
|
||||
"ema_step": self.current_trainstep,
|
||||
}
|
||||
|
||||
return {
|
||||
"ema_enabled": True,
|
||||
"ema_decay": self.training_args.ema_decay,
|
||||
"ema_start_step": self.training_args.ema_start_step,
|
||||
"ema_ready": self.is_ema_ready(),
|
||||
"ema_step": self.current_trainstep,
|
||||
}
|
||||
|
||||
def reset_ema(self):
|
||||
"""Reset EMA to current model weights."""
|
||||
if self.generator_ema is not None:
|
||||
logger.info("Resetting EMA to current model weights")
|
||||
self.generator_ema.update(self.transformer)
|
||||
# Force update to current weights by setting decay to 0 temporarily
|
||||
original_decay = self.generator_ema.decay
|
||||
self.generator_ema.decay = 0.0
|
||||
self.generator_ema.update(self.transformer)
|
||||
self.generator_ema.decay = original_decay
|
||||
logger.info("EMA reset completed")
|
||||
else:
|
||||
logger.warning("Cannot reset EMA: EMA not initialized")
|
||||
|
||||
def _build_distill_input_kwargs(
|
||||
self, noise_input: torch.Tensor, timestep: torch.Tensor,
|
||||
text_dict: dict[str, torch.Tensor] | None,
|
||||
@@ -546,7 +331,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
def _dmd_forward(self, generator_pred_video: torch.Tensor,
|
||||
training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Compute DMD (Diffusion Model Distillation) loss."""
|
||||
original_latent = generator_pred_video
|
||||
with torch.no_grad():
|
||||
timestep = torch.randint(0,
|
||||
self.num_train_timestep, [1],
|
||||
@@ -571,7 +355,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).detach().unflatten(0, (1, generator_pred_video.shape[1]))
|
||||
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
|
||||
|
||||
# fake_score_transformer forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -620,24 +404,24 @@ class DistillationPipeline(TrainingPipeline):
|
||||
pred_real_video_uncond) * self.real_score_guidance_scale
|
||||
|
||||
grad = (faker_score_pred_video - real_score_pred_video) / torch.abs(
|
||||
original_latent - real_score_pred_video).mean()
|
||||
generator_pred_video - real_score_pred_video).mean()
|
||||
grad = torch.nan_to_num(grad)
|
||||
|
||||
dmd_loss = 0.5 * F.mse_loss(
|
||||
original_latent.float(),
|
||||
(original_latent.float() - grad.float()).detach())
|
||||
generator_pred_video.float(),
|
||||
(generator_pred_video.float() - grad.float()).detach())
|
||||
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"training_batch_dmd_fwd_clean_latent":
|
||||
training_batch.latents,
|
||||
"generator_pred_video":
|
||||
original_latent.detach(),
|
||||
generator_pred_video,
|
||||
"real_score_pred_video":
|
||||
real_score_pred_video.detach(),
|
||||
real_score_pred_video,
|
||||
"faker_score_pred_video":
|
||||
faker_score_pred_video.detach(),
|
||||
faker_score_pred_video,
|
||||
"dmd_timestep":
|
||||
timestep.detach(),
|
||||
timestep,
|
||||
})
|
||||
|
||||
return dmd_loss
|
||||
@@ -734,12 +518,12 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
training_batch.conditional_dict = conditional_dict
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
self.video_latent_shape = training_batch.latents.shape
|
||||
@@ -802,15 +586,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
(dmd_loss / gradient_accumulation_steps).backward()
|
||||
total_dmd_loss += dmd_loss.detach().item()
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
for param in self.transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
if self.generator_ema is not None:
|
||||
self.generator_ema.update(self.transformer)
|
||||
|
||||
avg_dmd_loss = torch.tensor(total_dmd_loss /
|
||||
gradient_accumulation_steps,
|
||||
device=self.device)
|
||||
@@ -834,9 +611,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fake_score_latent_vis_dict.update(
|
||||
batch_fake.fake_score_latent_vis_dict)
|
||||
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
|
||||
for param in self.fake_score_transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
self.lr_scheduler.step()
|
||||
@@ -864,8 +638,7 @@ 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.generator_ema)
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator)
|
||||
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
@@ -896,14 +669,6 @@ 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
|
||||
@@ -935,18 +700,6 @@ 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]
|
||||
@@ -962,98 +715,50 @@ class DistillationPipeline(TrainingPipeline):
|
||||
step_videos: list[np.ndarray] = []
|
||||
step_captions: list[str] = []
|
||||
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
# 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()
|
||||
@@ -1130,16 +835,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:
|
||||
@@ -1191,6 +896,14 @@ class DistillationPipeline(TrainingPipeline):
|
||||
else:
|
||||
set_random_seed(seed + self.global_rank)
|
||||
|
||||
# Check trainable params
|
||||
num_trainable_generator = round(
|
||||
count_trainable(self.transformer) / 1e9, 3)
|
||||
num_trainable_critic = round(
|
||||
count_trainable(self.fake_score_transformer) / 1e9, 3)
|
||||
logger.info(
|
||||
"rank: %s: # of trainable params in generator: %sB, # of trainable params in critic: %sB",
|
||||
self.global_rank, num_trainable_generator, num_trainable_critic)
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
@@ -1200,10 +913,6 @@ 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()
|
||||
@@ -1247,14 +956,6 @@ 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)
|
||||
|
||||
@@ -1268,19 +969,11 @@ 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,
|
||||
"ema":
|
||||
"✓" if (self.generator_ema is not None and self.is_ema_ready())
|
||||
else "✗",
|
||||
"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,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
@@ -1308,15 +1001,6 @@ 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":
|
||||
@@ -1348,8 +1032,7 @@ 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.generator_ema)
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator)
|
||||
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
@@ -1366,11 +1049,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
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)
|
||||
only_save_generator_weight=True)
|
||||
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
if self.training_args.log_visualization:
|
||||
@@ -1390,11 +1069,7 @@ 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.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)
|
||||
self.noise_random_generator)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -1,492 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
import wandb
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
Training pipeline for ODE-init using precomputed denoising trajectories.
|
||||
|
||||
Supervision: predict the next latent in the stored trajectory by
|
||||
- feeding current latent at timestep t into the transformer to predict noise
|
||||
- stepping the scheduler with the predicted noise
|
||||
- minimizing MSE to the stored next latent at timestep t_next
|
||||
"""
|
||||
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Match the preprocess/generation scheduler for consistent stepping
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
|
||||
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
# logger.info(f"ARG dmd_denoising_steps: {training_args.pipeline_config.dmd_denoising_steps}")
|
||||
logger.info(
|
||||
f"ARG dmd_denoising_steps: {self.training_args.pipeline_config.dmd_denoising_steps}"
|
||||
)
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
# self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250], dtype=torch.long, device=get_local_torch_device())
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32))).cuda()
|
||||
logger.info(f"timesteps: {timesteps}")
|
||||
self.dmd_denoising_steps = timesteps[1000 -
|
||||
self.dmd_denoising_steps]
|
||||
logger.info(
|
||||
f"warped self.dmd_denoising_steps: {self.dmd_denoising_steps}")
|
||||
# assert False, "warp_denoising_step must be false"
|
||||
else:
|
||||
assert False, "warp_denoising_step must be true"
|
||||
logger.info("not warped")
|
||||
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
|
||||
get_local_torch_device())
|
||||
|
||||
logger.info(f"denoising_step_list: {self.dmd_denoising_steps}")
|
||||
|
||||
logger.info(
|
||||
"Initialized ODE-init training pipeline with %s denoising steps",
|
||||
len(self.dmd_denoising_steps))
|
||||
# Cache for nearest trajectory index per DMD step (computed lazily on first batch)
|
||||
self._cached_closest_idx_per_dmd = None
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
# self.min_timestep = int(self.training_args.min_timestep_ratio *
|
||||
# self.num_train_timestep)
|
||||
# self.max_timestep = int(self.training_args.max_timestep_ratio *
|
||||
# self.num_train_timestep)
|
||||
# self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
# Warm start validation with current transformer
|
||||
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
|
||||
# training_args.model_path,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
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)
|
||||
|
||||
def _get_next_batch(self, training_batch): # type: ignore[override]
|
||||
# batch = next(self.train_loader_iter, None) # type: ignore
|
||||
# if batch is None:
|
||||
# self.current_epoch += 1
|
||||
# logger.info("Starting epoch %s", self.current_epoch)
|
||||
# self.train_loader_iter = iter(self.train_dataloader)
|
||||
# batch = next(self.train_loader_iter)
|
||||
|
||||
# # Required fields from parquet (ODE trajectory schema)
|
||||
# encoder_hidden_states = batch['text_embedding']
|
||||
# encoder_attention_mask = batch['text_attention_mask']
|
||||
# infos = batch['info_list']
|
||||
|
||||
# # Trajectory tensors may include a leading singleton batch dim per row
|
||||
# trajectory_latents = batch['trajectory_latents']
|
||||
# if trajectory_latents.dim() == 7:
|
||||
# # [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
|
||||
# trajectory_latents = trajectory_latents[:, 0]
|
||||
# elif trajectory_latents.dim() == 6:
|
||||
# # already [B, S, C, T, H, W]
|
||||
# pass
|
||||
# else:
|
||||
# raise ValueError(
|
||||
# f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
|
||||
# )
|
||||
|
||||
# trajectory_timesteps = batch['trajectory_timesteps']
|
||||
# if trajectory_timesteps.dim() == 3:
|
||||
# # [B, 1, S] -> [B, S]
|
||||
# trajectory_timesteps = trajectory_timesteps[:, 0]
|
||||
# elif trajectory_timesteps.dim() == 2:
|
||||
# # [B, S]
|
||||
# pass
|
||||
# else:
|
||||
# raise ValueError(
|
||||
# f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
|
||||
# )
|
||||
# # [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
|
||||
# trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
|
||||
|
||||
# # Move to device
|
||||
device = get_local_torch_device()
|
||||
# training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
# device, dtype=torch.bfloat16)
|
||||
# training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
# device, dtype=torch.bfloat16)
|
||||
# training_batch.infos = infos
|
||||
|
||||
# return training_batch, trajectory_latents.to(
|
||||
# device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
|
||||
## TEMP
|
||||
path = "/mnt/weka/home/hao.zhang/wei/FastVideo/data/ode_vidprom_1_fv/00000.pt"
|
||||
b = torch.load(path)
|
||||
for k, v in b.items():
|
||||
logger.info(f"b[{k}]: {type(v)}")
|
||||
if isinstance(v, torch.Tensor):
|
||||
logger.info(f"b[{k}]: {v.shape}")
|
||||
else:
|
||||
logger.info(f"b[{k}]: {v}")
|
||||
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
|
||||
0).to(device, dtype=torch.bfloat16)
|
||||
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
|
||||
logger.info(f"trajectory_latents: {trajectory_latents.shape}")
|
||||
logger.info(
|
||||
f"encoder_hidden_states: {training_batch.encoder_hidden_states.shape}"
|
||||
)
|
||||
|
||||
return training_batch, trajectory_latents.to(device,
|
||||
dtype=torch.bfloat16), None
|
||||
|
||||
def _get_timestep(self,
|
||||
min_timestep: int,
|
||||
max_timestep: int,
|
||||
batch_size: int,
|
||||
num_frame: int,
|
||||
num_frame_per_block: int,
|
||||
uniform_timestep: bool = False) -> torch.Tensor:
|
||||
if uniform_timestep:
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, 1],
|
||||
device=self.device,
|
||||
dtype=torch.long).repeat(1, num_frame)
|
||||
return timestep
|
||||
else:
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, num_frame],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
# logger.info(f"individual timestep: {timestep}")
|
||||
# make the noise level the same within every block
|
||||
timestep = timestep.reshape(timestep.shape[0], -1,
|
||||
num_frame_per_block)
|
||||
timestep[:, :, 1:] = timestep[:, :, 0:1]
|
||||
timestep = timestep.reshape(timestep.shape[0], -1)
|
||||
return timestep
|
||||
|
||||
def _step_predict_next_latent(
|
||||
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
|
||||
torch.Tensor]]:
|
||||
latent_vis_dict = {}
|
||||
device = get_local_torch_device()
|
||||
target_latent = traj_latents[:, -1]
|
||||
|
||||
# logger.info(f"traj_latents: {traj_latents.shape}")
|
||||
# logger.info(f"traj_timesteps: {traj_timesteps.shape}")
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
B, S, num_frames, num_channels, height, width = traj_latents.shape
|
||||
|
||||
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
|
||||
if self._cached_closest_idx_per_dmd is None:
|
||||
# Use the first sample's trajectory timesteps; assumed identical across batches
|
||||
# s_steps = traj_timesteps[0].to(torch.long) # [S]
|
||||
# dmd = cast(torch.Tensor, self.dmd_denoising_steps).to(s_steps.device) # [K]
|
||||
# distances_ks: [K, S] = |s_steps - dmd|
|
||||
# distances_ks = (s_steps.unsqueeze(0) - dmd.unsqueeze(1)).abs()
|
||||
# self._cached_closest_idx_per_dmd = distances_ks.argmin(dim=1).to(torch.long).cpu() # [K]
|
||||
self._cached_closest_idx_per_dmd = torch.tensor(
|
||||
[0, 12, 24, 36], dtype=torch.long).cpu()
|
||||
logger.info(
|
||||
f"self._cached_closest_idx_per_dmd: {self._cached_closest_idx_per_dmd}"
|
||||
)
|
||||
logger.info(
|
||||
f"corresponding timesteps: {self.noise_scheduler.timesteps[self._cached_closest_idx_per_dmd]}"
|
||||
)
|
||||
|
||||
# logger.info(f"traj_latents: {traj_latents.shape}")
|
||||
# Select the K indexes from traj_latents using self._cached_closest_idx_per_dmd
|
||||
# traj_latents: [B, S, C, T, H, W], self._cached_closest_idx_per_dmd: [K]
|
||||
# Output: [B, K, C, T, H, W]
|
||||
# relevant_traj_latents = torch.index_select(
|
||||
# traj_latents,
|
||||
# dim=1,
|
||||
# index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
|
||||
relevant_traj_latents = traj_latents
|
||||
logger.info(f"relevant_traj_latents sum: {relevant_traj_latents.float().sum().item()}, relevant_traj_latents: {relevant_traj_latents.shape}")
|
||||
# assert relevant_traj_latents.shape[0] == 1
|
||||
|
||||
indexes = self._get_timestep( # [B, num_frames]
|
||||
0,
|
||||
len(self.dmd_denoising_steps),
|
||||
B,
|
||||
num_frames,
|
||||
3,
|
||||
uniform_timestep=False)
|
||||
logger.info(f"indexes: {indexes.shape}")
|
||||
logger.info(f"indexes: {indexes}")
|
||||
# noisy_input = relevant_traj_latents[indexes]
|
||||
noisy_input = torch.gather(
|
||||
relevant_traj_latents,
|
||||
dim=1,
|
||||
index=indexes.reshape(B, 1, num_frames, 1, 1,
|
||||
1).expand(-1, -1, -1, num_channels, height,
|
||||
width).to(self.device)).squeeze(1)
|
||||
# noisy_input = noisy_input.unsqueeze(0)
|
||||
|
||||
# # Sample a single DMD step for the whole batch and fetch its cached nearest S-index
|
||||
# K = len(self.dmd_denoising_steps)
|
||||
# dmd_idx = torch.randint(0, K, (1,), device=device)
|
||||
# logger.info(f"dmd_idx: {dmd_idx}")
|
||||
# assert self._cached_closest_idx_per_dmd is not None
|
||||
# nearest_s_idx = int(self._cached_closest_idx_per_dmd[int(dmd_idx.item())])
|
||||
# nearest_idx = torch.full((B,), nearest_s_idx, device=device, dtype=torch.long)
|
||||
|
||||
# batch_indices = torch.arange(B, device=device)
|
||||
# noisy_input = traj_latents[batch_indices, nearest_idx] # [B, C, T, H, W]
|
||||
# target_latent = traj_latents[batch_indices, -1] # [B, C, T, H, W]
|
||||
# t = traj_timesteps[batch_indices, nearest_idx] # [B]
|
||||
|
||||
# Scale model input as in inference for consistency with stored trajectories
|
||||
# noisy_input = self.modules["scheduler"].scale_model_input(noisy_input, t)
|
||||
# logger.info(f"indexes: {indexes.shape}")
|
||||
# logger.info(f"indexes: {indexes}")
|
||||
timestep = self.dmd_denoising_steps[indexes]
|
||||
# logger.info(f"timestep: {timestep.shape}")
|
||||
# logger.info(f"timestep: {timestep}")
|
||||
|
||||
# Prepare inputs for transformer
|
||||
latent_vis_dict["noisy_input"] = noisy_input.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
|
||||
4).detach().clone().cpu()
|
||||
|
||||
logger.info("noisy input sum: %s, noisy input shape: %s, noisy input dtype: %s", noisy_input.float().sum().item(), noisy_input.shape, noisy_input.dtype)
|
||||
logger.info("encoder hidden states sum: %s, encoder hidden states shape: %s, encoder hidden states dtype: %s", encoder_hidden_states.float().sum().item(), encoder_hidden_states.shape, encoder_hidden_states.dtype)
|
||||
logger.info("timestep sum: %s, timestep shape: %s, timestep dtype: %s", timestep.float().sum().item(), timestep.shape, timestep.dtype)
|
||||
logger.info("model dtype set: %s. model weight sum: %s", set(p.dtype for p in self.transformer.parameters()), sum(p.float().sum().item() for p in self.transformer.parameters()))
|
||||
|
||||
# model_dtype = next(self.transformer.parameters()).dtype
|
||||
model_dtype = torch.bfloat16
|
||||
input_kwargs = {
|
||||
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timestep,
|
||||
"return_dict": False,
|
||||
"scheduler": self.modules["scheduler"]
|
||||
}
|
||||
# Predict noise and step the scheduler to obtain next latent
|
||||
with set_forward_context(current_timestep=timestep,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
pred_video = self.transformer(**input_kwargs)
|
||||
# logger.info(f"noise_pred: {noise_pred.shape}")
|
||||
# logger.info("noise pred sum: %s, noise pred shape: %s, noise pred dtype: %s", noise_pred.float().sum().item(), noise_pred.shape, noise_pred.dtype)
|
||||
# if isinstance(noise_pred, (tuple, list)):
|
||||
# noise_pred = noise_pred[0]
|
||||
|
||||
# from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
# pred_video = pred_noise_to_pred_video(
|
||||
# pred_noise=noise_pred.flatten(0, 1),
|
||||
# noise_input_latent=noisy_input.flatten(0, 1),
|
||||
# timestep=timestep.to(dtype=model_dtype).flatten(0, 1),
|
||||
# scheduler=self.modules["scheduler"]).unflatten(
|
||||
# 0, noise_pred.shape[:2])
|
||||
logger.info("pred video sum: %s, pred video shape: %s, pred video dtype: %s", pred_video.float().sum().item(), pred_video.shape, pred_video.dtype)
|
||||
|
||||
latent_vis_dict["pred_video"] = pred_video.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
|
||||
# noisy_input = pred_noise_to_pred_video(noise_pred, noisy_input, t, self.modules["scheduler"])
|
||||
# next_latent_pred = self.modules["scheduler"].step(
|
||||
# noise_pred, t, current_latents, return_dict=False)[0]
|
||||
return pred_video, target_latent, timestep, latent_vis_dict
|
||||
|
||||
def train_one_step(self, training_batch): # type: ignore[override]
|
||||
self.transformer.train()
|
||||
self.optimizer.zero_grad()
|
||||
training_batch.total_loss = 0.0
|
||||
args = cast(TrainingArgs, self.training_args)
|
||||
|
||||
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
|
||||
|
||||
for _ in range(args.gradient_accumulation_steps):
|
||||
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
|
||||
training_batch)
|
||||
text_embeds = training_batch.encoder_hidden_states
|
||||
text_attention_mask = training_batch.encoder_attention_mask
|
||||
assert traj_latents.shape[0] == 1
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
B, S = traj_latents.shape[0], traj_latents.shape[1]
|
||||
if S < 2:
|
||||
raise ValueError("Trajectory must contain at least 2 steps")
|
||||
|
||||
# Sample per-sample current step i in [0, S-2]
|
||||
|
||||
# idx = torch.randint(low=0, high=S - 1, size=(B, ),
|
||||
# device=traj_latents.device)
|
||||
|
||||
# Gather current latents and next latents
|
||||
# batch_indices = torch.arange(B, device=traj_latents.device)
|
||||
# current_latents = traj_latents[batch_indices, idx] # [B, C, T,H,W]
|
||||
# current_latent = traj_timesteps[:, -1, :, :, :, :]
|
||||
# target_latents = traj_latents[:, -1, :, :, :, :]
|
||||
|
||||
# Corresponding timesteps t (long) -> cast per sample
|
||||
# t = traj_timesteps[:, -1, :, :, :, :]
|
||||
# if t.dtype != torch.long:
|
||||
# t = t.long()
|
||||
|
||||
# Forward to predict next latent by stepping scheduler with predicted noise
|
||||
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
|
||||
traj_latents, training_batch.current_timestep, text_embeds, text_attention_mask)
|
||||
|
||||
training_batch.latent_vis_dict.update(latent_vis_dict)
|
||||
|
||||
mask = t != 0
|
||||
|
||||
# Compute loss
|
||||
loss = F.mse_loss(noise_pred[mask],
|
||||
target_latent[mask],
|
||||
reduction="mean")
|
||||
loss = loss / args.gradient_accumulation_steps
|
||||
|
||||
with set_forward_context(current_timestep=t,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
loss.backward()
|
||||
|
||||
logger.info("Loss sum: %s, Loss dtype: %s", loss.float().sum().item(), loss.dtype)
|
||||
assert self.transformer.blocks[0].to_q.weight.dtype == torch.float32
|
||||
logger.info("blocks[0].to_q param sum before backprop: %s", self.transformer.blocks[0].to_q.weight.float().sum().item())
|
||||
logger.info("blocks[0].to_q param grad sum: %s", self.transformer.blocks[0].to_q.weight.grad.float().sum().item())
|
||||
logger.info("Transformer grad dtype: %s", set(p.grad.dtype for p in self.transformer.parameters()))
|
||||
logger.info("Transformer param dtype before backprop: %s", set(p.dtype for p in self.transformer.parameters()))
|
||||
avg_loss = loss.detach().clone()
|
||||
training_batch.total_loss += avg_loss.item()
|
||||
|
||||
# Clip grad and step optimizers
|
||||
# grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
# [p for p in self.transformer.parameters() if p.requires_grad],
|
||||
# args.max_grad_norm if args.max_grad_norm is not None else 0.0)
|
||||
grad_norm = None
|
||||
|
||||
self.optimizer.step()
|
||||
logger.info("Transformer param dtype after backprop: %s", set(p.dtype for p in self.transformer.parameters()))
|
||||
logger.info("blocks[0].to_q param param after backprop sum: %s", self.transformer.blocks[0].to_q.weight.float().sum().item())
|
||||
|
||||
assert False
|
||||
# self.lr_scheduler.step()
|
||||
|
||||
if grad_norm is None:
|
||||
grad_value = 0.0
|
||||
else:
|
||||
try:
|
||||
if isinstance(grad_norm, torch.Tensor):
|
||||
grad_value = float(grad_norm.detach().float().item())
|
||||
else:
|
||||
grad_value = float(grad_norm)
|
||||
except Exception:
|
||||
grad_value = 0.0
|
||||
training_batch.grad_norm = grad_value
|
||||
return training_batch
|
||||
|
||||
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 = {}
|
||||
latents_vis_dict = training_batch.latent_vis_dict
|
||||
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
|
||||
for latent_key in latent_log_keys:
|
||||
assert latent_key in latents_vis_dict and latents_vis_dict[
|
||||
latent_key] is not None
|
||||
latent = latents_vis_dict[latent_key]
|
||||
pixel_latent = self.validation_pipeline.decoding_stage.decode(
|
||||
latent, training_args)
|
||||
|
||||
video = pixel_latent.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=16, format="mp4") # change to 16 for Wan2.1
|
||||
# Clean up references
|
||||
del video, pixel_latent, latent
|
||||
|
||||
# Log to wandb
|
||||
if self.global_rank == 0:
|
||||
wandb.log(wandb_loss_dict, step=step)
|
||||
|
||||
# dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
|
||||
# fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
|
||||
# fake_score_log_keys = ['generator_pred_video']
|
||||
# dmd_log_keys = ['faker_score_pred_video', 'real_score_pred_video']
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting ODE-init training pipeline...")
|
||||
logger.info(f"ARG dmd_denoising_steps: {args.dmd_denoising_steps}")
|
||||
pipeline = ODEInitTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("ODE-init training pipeline done")
|
||||
|
||||
|
||||
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()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -22,7 +22,7 @@ from tqdm.auto import tqdm
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
# from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
@@ -39,26 +39,20 @@ from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_scheduler, get_sigmas,
|
||||
load_checkpoint, normalize_dit_input, save_checkpoint,
|
||||
compute_density_for_timestep_sampling, count_trainable, get_scheduler,
|
||||
get_sigmas, load_checkpoint, normalize_dit_input, save_checkpoint,
|
||||
shard_latents_across_sp)
|
||||
# from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
# set_random_seed, shallow_asdict)
|
||||
from fastvideo.utils import (is_vsa_available,
|
||||
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
set_random_seed, shallow_asdict)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
# vmoba_available = is_vmoba_available()
|
||||
vmoba_available = is_vmoba_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _get_trainable_params(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class TrainingPipeline(LoRAPipeline, ABC):
|
||||
"""
|
||||
A pipeline for training a model. All training pipelines should inherit from this class.
|
||||
@@ -118,18 +112,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
# Set grads for proper modules based on the training mode (Distill, LoRA, etc.)
|
||||
self.set_trainable()
|
||||
params_to_optimize = self.transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
# Parse betas from string format "beta1,beta2"
|
||||
betas_str = training_args.betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=training_args.learning_rate,
|
||||
betas=betas,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
@@ -281,20 +272,20 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=current_vsa_sparsity,
|
||||
device=get_local_torch_device())
|
||||
# elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
# moba_params = self.training_args.moba_config.copy()
|
||||
# moba_params.update({
|
||||
# "current_timestep":
|
||||
# training_batch.timesteps,
|
||||
# "raw_latent_shape":
|
||||
# training_batch.raw_latent_shape[2:5],
|
||||
# "patch_size":
|
||||
# self.training_args.pipeline_config.dit_config.patch_size,
|
||||
# "device":
|
||||
# get_local_torch_device(),
|
||||
# })
|
||||
# training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
|
||||
# ).build(**moba_params)
|
||||
elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
moba_params = self.training_args.moba_config.copy()
|
||||
moba_params.update({
|
||||
"current_timestep":
|
||||
training_batch.timesteps,
|
||||
"raw_latent_shape":
|
||||
training_batch.raw_latent_shape[2:5],
|
||||
"patch_size":
|
||||
self.training_args.pipeline_config.dit_config.patch_size,
|
||||
"device":
|
||||
get_local_torch_device(),
|
||||
})
|
||||
training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
|
||||
).build(**moba_params)
|
||||
else:
|
||||
training_batch.attn_metadata = None
|
||||
|
||||
@@ -319,8 +310,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
def _transformer_forward_and_compute_loss(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
# if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
assert training_batch.attn_metadata is not None
|
||||
else:
|
||||
assert training_batch.attn_metadata is None
|
||||
@@ -441,7 +431,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
local_main_process_only=False)
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
num_trainable_params = _get_trainable_params(self.transformer)
|
||||
num_trainable_params = count_trainable(self.transformer)
|
||||
logger.info("Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
@@ -486,9 +476,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
current_decay_times = min(step // vsa_decay_interval_steps,
|
||||
vsa_sparsity // vsa_decay_rate)
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
# elif vmoba_available:
|
||||
# # TODO: add vmoba sparsity scheduling here
|
||||
# current_vsa_sparsity = 0.0
|
||||
elif vmoba_available:
|
||||
# TODO: add vmoba sparsity scheduling here
|
||||
current_vsa_sparsity = 0.0
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
@@ -530,14 +520,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
if self.training_args.log_visualization:
|
||||
self.visualize_intermediate_latents(training_batch,
|
||||
self.training_args,
|
||||
step)
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
trainable_params = round(
|
||||
_get_trainable_params(self.transformer) / 1e9, 3)
|
||||
count_trainable(self.transformer) / 1e9, 3)
|
||||
logger.info(
|
||||
"GPU memory usage after validation: %s MB, trainable params: %sB",
|
||||
gpu_memory_usage, trainable_params)
|
||||
@@ -573,7 +559,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
logger.info(" Total optimization steps = %s",
|
||||
self.training_args.max_train_steps)
|
||||
logger.info(" Total training parameters per FSDP shard = %s B",
|
||||
round(_get_trainable_params(self.transformer) / 1e9, 3))
|
||||
round(count_trainable(self.transformer) / 1e9, 3))
|
||||
# print dtype
|
||||
logger.info(" Master weight dtype: %s",
|
||||
self.transformer.parameters().__next__().dtype)
|
||||
@@ -641,7 +627,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
validation_dataloader = DataLoader(validation_dataset,
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
transformer.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
@@ -735,8 +720,3 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
raise NotImplementedError("Visualize intermediate latents is not implemented for training pipeline")
|
||||
|
||||
@@ -202,7 +202,6 @@ 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.
|
||||
@@ -234,8 +233,6 @@ 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")
|
||||
@@ -349,14 +346,10 @@ def load_checkpoint(transformer,
|
||||
"""
|
||||
if not os.path.exists(checkpoint_path):
|
||||
logger.warning("Checkpoint path %s does not exist", checkpoint_path)
|
||||
assert False
|
||||
return 0
|
||||
|
||||
# Extract step number from checkpoint path
|
||||
try:
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
except:
|
||||
step = 1
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
|
||||
if rank == 0:
|
||||
logger.info("Loading checkpoint from step %s", step)
|
||||
@@ -409,8 +402,7 @@ def load_distillation_checkpoint(generator_transformer,
|
||||
dataloader=None,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None,
|
||||
generator_ema=None) -> int:
|
||||
noise_generator=None) -> int:
|
||||
"""
|
||||
Load distillation checkpoint with both generator and fake_score models.
|
||||
Returns the step number from which training should resume.
|
||||
@@ -464,18 +456,6 @@ 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")
|
||||
@@ -1300,154 +1280,5 @@ def get_scheduler(
|
||||
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)
|
||||
def count_trainable(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
@@ -1,72 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import WanCausalDMDPipeline
|
||||
from fastvideo.training.self_forcing_distillation_pipeline import SelfForcingDistillationPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
|
||||
"""
|
||||
A self-forcing distillation pipeline for Wan that uses the self-forcing methodology
|
||||
with DMD for video generation.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
]
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
validation_pipeline = WanCausalDMDPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={"transformer": self.get_module("transformer")},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting Wan self-forcing distillation pipeline...")
|
||||
|
||||
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Wan self-forcing distillation pipeline completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,211 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_i2v_pipeline import (
|
||||
WanImageToVideoPipeline)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanT2VI2VTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Wan.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_t2v
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.dit_cpu_offload = True
|
||||
# args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
|
||||
pipeline_config = PipelineConfig.from_pretrained(
|
||||
training_args.model_path)
|
||||
pipeline_config.vae_config.load_encoder = True
|
||||
self.validation_pipeline = WanImageToVideoPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
inference_mode=True,
|
||||
pipeline_config=pipeline_config,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
required_config_modules=[
|
||||
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
|
||||
],
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
dit_cpu_offload=True,
|
||||
)
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
latents = batch['vae_latent']
|
||||
latents = latents[:, :, :self.training_args.num_latent_t]
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
# clip_features = batch['clip_feature']
|
||||
# image_latents = batch['first_frame_latent']
|
||||
# image_latents = image_latents[:, :, :self.training_args.num_latent_t]
|
||||
# pil_image = batch['pil_image']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
# training_batch.preprocessed_image = pil_image.to(
|
||||
# get_local_torch_device())
|
||||
# training_batch.image_embeds = clip_features.to(get_local_torch_device())
|
||||
# training_batch.image_latents = image_latents.to(
|
||||
# get_local_torch_device())
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
|
||||
|
||||
# First, call parent method to prepare noise, timesteps, etc. for video latents
|
||||
training_batch = super()._prepare_dit_inputs(training_batch)
|
||||
|
||||
latents = training_batch.latents
|
||||
logger.info("latents.shape: %s", latents.shape)
|
||||
first_frame_latent = latents[:, :, 0, :, :]
|
||||
|
||||
logger.info("first_frame_latent.shape: %s", first_frame_latent.shape)
|
||||
logger.info("training_batch.noisy_model_input.shape: %s",
|
||||
training_batch.noisy_model_input.shape)
|
||||
|
||||
training_batch.noisy_model_input = torch.cat([
|
||||
first_frame_latent.unsqueeze(2),
|
||||
training_batch.noisy_model_input[:, :, 1:, :, :]
|
||||
],
|
||||
dim=2)
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_input_kwargs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
|
||||
# Image Embeds for conditioning
|
||||
# image_embeds = training_batch.image_embeds
|
||||
# assert torch.isnan(image_embeds).sum() == 0
|
||||
# image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
# dtype=torch.bfloat16)
|
||||
# encoder_hidden_states_image = image_embeds
|
||||
|
||||
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states":
|
||||
training_batch.noisy_model_input,
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states,
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
"encoder_attention_mask":
|
||||
training_batch.encoder_attention_mask,
|
||||
# "encoder_hidden_states_image":
|
||||
# encoder_hidden_states_image,
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
return training_batch
|
||||
|
||||
def _prepare_validation_batch(self, sampling_param: SamplingParam,
|
||||
training_args: TrainingArgs,
|
||||
validation_batch: dict[str, Any],
|
||||
num_inference_steps: int) -> ForwardBatch:
|
||||
sampling_param.prompt = validation_batch['prompt']
|
||||
sampling_param.height = training_args.num_height
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.image_path = validation_batch['video_path']
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
assert self.seed is not None
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
VSA_sparsity=training_args.VSA_sparsity,
|
||||
)
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = WanT2VI2VTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
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()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
@@ -2,10 +2,6 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
|
||||
import argparse
|
||||
from einops import rearrange
|
||||
import torchvision
|
||||
import numpy as np
|
||||
import imageio
|
||||
import ctypes
|
||||
import hashlib
|
||||
import importlib
|
||||
@@ -27,14 +23,10 @@ from typing import Any, TypeVar, cast
|
||||
|
||||
import cloudpickle
|
||||
import filelock
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
import yaml
|
||||
from diffusers.loaders.lora_base import (
|
||||
_best_guess_weight_name) # watch out for potetential removal from diffusers
|
||||
from einops import rearrange
|
||||
from huggingface_hub import snapshot_download
|
||||
from remote_pdb import RemotePdb
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy
|
||||
@@ -894,17 +886,3 @@ def best_output_size(w, h, dw, dh, expected_area):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
|
||||
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor],
|
||||
output_path: str, fps: int):
|
||||
# Process outputs
|
||||
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
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))
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
Code in this folder is modified from https://github.com/Wan-Video/Wan2.1
|
||||
Apache-2.0 License
|
||||
@@ -1,3 +0,0 @@
|
||||
from . import configs, distributed, modules
|
||||
from .image2video import WanI2V
|
||||
from .text2video import WanT2V
|
||||
@@ -1,42 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
from .wan_t2v_14B import t2v_14B
|
||||
from .wan_t2v_1_3B import t2v_1_3B
|
||||
from .wan_i2v_14B import i2v_14B
|
||||
import copy
|
||||
import os
|
||||
|
||||
os.environ['TOKENIZERS_PARALLELISM'] = 'false'
|
||||
|
||||
|
||||
# the config of t2i_14B is the same as t2v_14B
|
||||
t2i_14B = copy.deepcopy(t2v_14B)
|
||||
t2i_14B.__name__ = 'Config: Wan T2I 14B'
|
||||
|
||||
WAN_CONFIGS = {
|
||||
't2v-14B': t2v_14B,
|
||||
't2v-1.3B': t2v_1_3B,
|
||||
'i2v-14B': i2v_14B,
|
||||
't2i-14B': t2i_14B,
|
||||
}
|
||||
|
||||
SIZE_CONFIGS = {
|
||||
'720*1280': (720, 1280),
|
||||
'1280*720': (1280, 720),
|
||||
'480*832': (480, 832),
|
||||
'832*480': (832, 480),
|
||||
'1024*1024': (1024, 1024),
|
||||
}
|
||||
|
||||
MAX_AREA_CONFIGS = {
|
||||
'720*1280': 720 * 1280,
|
||||
'1280*720': 1280 * 720,
|
||||
'480*832': 480 * 832,
|
||||
'832*480': 832 * 480,
|
||||
}
|
||||
|
||||
SUPPORTED_SIZES = {
|
||||
't2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
|
||||
't2v-1.3B': ('480*832', '832*480'),
|
||||
'i2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
|
||||
't2i-14B': tuple(SIZE_CONFIGS.keys()),
|
||||
}
|
||||
@@ -1,19 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
from easydict import EasyDict
|
||||
|
||||
# ------------------------ Wan shared config ------------------------#
|
||||
wan_shared_cfg = EasyDict()
|
||||
|
||||
# t5
|
||||
wan_shared_cfg.t5_model = 'umt5_xxl'
|
||||
wan_shared_cfg.t5_dtype = torch.bfloat16
|
||||
wan_shared_cfg.text_len = 512
|
||||
|
||||
# transformer
|
||||
wan_shared_cfg.param_dtype = torch.bfloat16
|
||||
|
||||
# inference
|
||||
wan_shared_cfg.num_train_timesteps = 1000
|
||||
wan_shared_cfg.sample_fps = 16
|
||||
wan_shared_cfg.sample_neg_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
|
||||
@@ -1,35 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
from easydict import EasyDict
|
||||
|
||||
from .shared_config import wan_shared_cfg
|
||||
|
||||
# ------------------------ Wan I2V 14B ------------------------#
|
||||
|
||||
i2v_14B = EasyDict(__name__='Config: Wan I2V 14B')
|
||||
i2v_14B.update(wan_shared_cfg)
|
||||
|
||||
i2v_14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
|
||||
i2v_14B.t5_tokenizer = 'google/umt5-xxl'
|
||||
|
||||
# clip
|
||||
i2v_14B.clip_model = 'clip_xlm_roberta_vit_h_14'
|
||||
i2v_14B.clip_dtype = torch.float16
|
||||
i2v_14B.clip_checkpoint = 'models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth'
|
||||
i2v_14B.clip_tokenizer = 'xlm-roberta-large'
|
||||
|
||||
# vae
|
||||
i2v_14B.vae_checkpoint = 'Wan2.1_VAE.pth'
|
||||
i2v_14B.vae_stride = (4, 8, 8)
|
||||
|
||||
# transformer
|
||||
i2v_14B.patch_size = (1, 2, 2)
|
||||
i2v_14B.dim = 5120
|
||||
i2v_14B.ffn_dim = 13824
|
||||
i2v_14B.freq_dim = 256
|
||||
i2v_14B.num_heads = 40
|
||||
i2v_14B.num_layers = 40
|
||||
i2v_14B.window_size = (-1, -1)
|
||||
i2v_14B.qk_norm = True
|
||||
i2v_14B.cross_attn_norm = True
|
||||
i2v_14B.eps = 1e-6
|
||||
@@ -1,29 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
from easydict import EasyDict
|
||||
|
||||
from .shared_config import wan_shared_cfg
|
||||
|
||||
# ------------------------ Wan T2V 14B ------------------------#
|
||||
|
||||
t2v_14B = EasyDict(__name__='Config: Wan T2V 14B')
|
||||
t2v_14B.update(wan_shared_cfg)
|
||||
|
||||
# t5
|
||||
t2v_14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
|
||||
t2v_14B.t5_tokenizer = 'google/umt5-xxl'
|
||||
|
||||
# vae
|
||||
t2v_14B.vae_checkpoint = 'Wan2.1_VAE.pth'
|
||||
t2v_14B.vae_stride = (4, 8, 8)
|
||||
|
||||
# transformer
|
||||
t2v_14B.patch_size = (1, 2, 2)
|
||||
t2v_14B.dim = 5120
|
||||
t2v_14B.ffn_dim = 13824
|
||||
t2v_14B.freq_dim = 256
|
||||
t2v_14B.num_heads = 40
|
||||
t2v_14B.num_layers = 40
|
||||
t2v_14B.window_size = (-1, -1)
|
||||
t2v_14B.qk_norm = True
|
||||
t2v_14B.cross_attn_norm = True
|
||||
t2v_14B.eps = 1e-6
|
||||
@@ -1,29 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
from easydict import EasyDict
|
||||
|
||||
from .shared_config import wan_shared_cfg
|
||||
|
||||
# ------------------------ Wan T2V 1.3B ------------------------#
|
||||
|
||||
t2v_1_3B = EasyDict(__name__='Config: Wan T2V 1.3B')
|
||||
t2v_1_3B.update(wan_shared_cfg)
|
||||
|
||||
# t5
|
||||
t2v_1_3B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
|
||||
t2v_1_3B.t5_tokenizer = 'google/umt5-xxl'
|
||||
|
||||
# vae
|
||||
t2v_1_3B.vae_checkpoint = 'Wan2.1_VAE.pth'
|
||||
t2v_1_3B.vae_stride = (4, 8, 8)
|
||||
|
||||
# transformer
|
||||
t2v_1_3B.patch_size = (1, 2, 2)
|
||||
t2v_1_3B.dim = 1536
|
||||
t2v_1_3B.ffn_dim = 8960
|
||||
t2v_1_3B.freq_dim = 256
|
||||
t2v_1_3B.num_heads = 12
|
||||
t2v_1_3B.num_layers = 30
|
||||
t2v_1_3B.window_size = (-1, -1)
|
||||
t2v_1_3B.qk_norm = True
|
||||
t2v_1_3B.cross_attn_norm = True
|
||||
t2v_1_3B.eps = 1e-6
|
||||
@@ -1,33 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
|
||||
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
|
||||
|
||||
|
||||
def shard_model(
|
||||
model,
|
||||
device_id,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
process_group=None,
|
||||
sharding_strategy=ShardingStrategy.FULL_SHARD,
|
||||
sync_module_states=True,
|
||||
):
|
||||
model = FSDP(
|
||||
module=model,
|
||||
process_group=process_group,
|
||||
sharding_strategy=sharding_strategy,
|
||||
auto_wrap_policy=partial(
|
||||
lambda_auto_wrap_policy, lambda_fn=lambda m: m in model.blocks),
|
||||
mixed_precision=MixedPrecision(
|
||||
param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
buffer_dtype=buffer_dtype),
|
||||
device_id=device_id,
|
||||
use_orig_params=True,
|
||||
sync_module_states=sync_module_states)
|
||||
return model
|
||||
@@ -1,192 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
from xfuser.core.distributed import (get_sequence_parallel_rank,
|
||||
get_sequence_parallel_world_size,
|
||||
get_sp_group)
|
||||
from xfuser.core.long_ctx_attention import xFuserLongContextAttention
|
||||
|
||||
from ..modules.model import sinusoidal_embedding_1d
|
||||
|
||||
|
||||
def pad_freqs(original_tensor, target_len):
|
||||
seq_len, s1, s2 = original_tensor.shape
|
||||
pad_size = target_len - seq_len
|
||||
padding_tensor = torch.ones(
|
||||
pad_size,
|
||||
s1,
|
||||
s2,
|
||||
dtype=original_tensor.dtype,
|
||||
device=original_tensor.device)
|
||||
padded_tensor = torch.cat([original_tensor, padding_tensor], dim=0)
|
||||
return padded_tensor
|
||||
|
||||
|
||||
@amp.autocast(enabled=False)
|
||||
def rope_apply(x, grid_sizes, freqs):
|
||||
"""
|
||||
x: [B, L, N, C].
|
||||
grid_sizes: [B, 3].
|
||||
freqs: [M, C // 2].
|
||||
"""
|
||||
s, n, c = x.size(1), x.size(2), x.size(3) // 2
|
||||
# split freqs
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
# loop over samples
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
|
||||
# precompute multipliers
|
||||
x_i = torch.view_as_complex(x[i, :s].to(torch.float64).reshape(
|
||||
s, n, -1, 2))
|
||||
freqs_i = torch.cat([
|
||||
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||
],
|
||||
dim=-1).reshape(seq_len, 1, -1)
|
||||
|
||||
# apply rotary embedding
|
||||
sp_size = get_sequence_parallel_world_size()
|
||||
sp_rank = get_sequence_parallel_rank()
|
||||
freqs_i = pad_freqs(freqs_i, s * sp_size)
|
||||
s_per_rank = s
|
||||
freqs_i_rank = freqs_i[(sp_rank * s_per_rank):((sp_rank + 1) *
|
||||
s_per_rank), :, :]
|
||||
x_i = torch.view_as_real(x_i * freqs_i_rank).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, s:]])
|
||||
|
||||
# append to collection
|
||||
output.append(x_i)
|
||||
return torch.stack(output).float()
|
||||
|
||||
|
||||
def usp_dit_forward(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
context,
|
||||
seq_len,
|
||||
clip_fea=None,
|
||||
y=None,
|
||||
):
|
||||
"""
|
||||
x: A list of videos each with shape [C, T, H, W].
|
||||
t: [B].
|
||||
context: A list of text embeddings each with shape [L, C].
|
||||
"""
|
||||
if self.model_type == 'i2v':
|
||||
assert clip_fea is not None and y is not None
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
if self.freqs.device != device:
|
||||
self.freqs = self.freqs.to(device)
|
||||
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
||||
|
||||
# embeddings
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1)
|
||||
for u in x
|
||||
])
|
||||
|
||||
# time embeddings
|
||||
with amp.autocast(dtype=torch.float32):
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t).float())
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
||||
assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
|
||||
# context
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]))
|
||||
|
||||
if clip_fea is not None:
|
||||
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
context = torch.concat([context_clip, context], dim=1)
|
||||
|
||||
# arguments
|
||||
kwargs = dict(
|
||||
e=e0,
|
||||
seq_lens=seq_lens,
|
||||
grid_sizes=grid_sizes,
|
||||
freqs=self.freqs,
|
||||
context=context,
|
||||
context_lens=context_lens)
|
||||
|
||||
# Context Parallel
|
||||
x = torch.chunk(
|
||||
x, get_sequence_parallel_world_size(),
|
||||
dim=1)[get_sequence_parallel_rank()]
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x, **kwargs)
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
# Context Parallel
|
||||
x = get_sp_group().all_gather(x, dim=1)
|
||||
|
||||
# unpatchify
|
||||
x = self.unpatchify(x, grid_sizes)
|
||||
return [u.float() for u in x]
|
||||
|
||||
|
||||
def usp_attn_forward(self,
|
||||
x,
|
||||
seq_lens,
|
||||
grid_sizes,
|
||||
freqs,
|
||||
dtype=torch.bfloat16):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
q = rope_apply(q, grid_sizes, freqs)
|
||||
k = rope_apply(k, grid_sizes, freqs)
|
||||
|
||||
# TODO: We should use unpaded q,k,v for attention.
|
||||
# k_lens = seq_lens // get_sequence_parallel_world_size()
|
||||
# if k_lens is not None:
|
||||
# q = torch.cat([u[:l] for u, l in zip(q, k_lens)]).unsqueeze(0)
|
||||
# k = torch.cat([u[:l] for u, l in zip(k, k_lens)]).unsqueeze(0)
|
||||
# v = torch.cat([u[:l] for u, l in zip(v, k_lens)]).unsqueeze(0)
|
||||
|
||||
x = xFuserLongContextAttention()(
|
||||
None,
|
||||
query=half(q),
|
||||
key=half(k),
|
||||
value=half(v),
|
||||
window_size=self.window_size)
|
||||
|
||||
# TODO: padding after attention.
|
||||
# x = torch.cat([x, x.new_zeros(b, s - x.size(1), n, d)], dim=1)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
@@ -1,347 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import gc
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import types
|
||||
from contextlib import contextmanager
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
import torch.distributed as dist
|
||||
import torchvision.transforms.functional as TF
|
||||
from tqdm import tqdm
|
||||
|
||||
from .distributed.fsdp import shard_model
|
||||
from .modules.clip import CLIPModel
|
||||
from .modules.model import WanModel
|
||||
from .modules.t5 import T5EncoderModel
|
||||
from .modules.vae import WanVAE
|
||||
from .utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
|
||||
get_sampling_sigmas, retrieve_timesteps)
|
||||
from .utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
|
||||
|
||||
class WanI2V:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
checkpoint_dir,
|
||||
device_id=0,
|
||||
rank=0,
|
||||
t5_fsdp=False,
|
||||
dit_fsdp=False,
|
||||
use_usp=False,
|
||||
t5_cpu=False,
|
||||
init_on_cpu=True,
|
||||
):
|
||||
r"""
|
||||
Initializes the image-to-video generation model components.
|
||||
|
||||
Args:
|
||||
config (EasyDict):
|
||||
Object containing model parameters initialized from config.py
|
||||
checkpoint_dir (`str`):
|
||||
Path to directory containing model checkpoints
|
||||
device_id (`int`, *optional*, defaults to 0):
|
||||
Id of target GPU device
|
||||
rank (`int`, *optional*, defaults to 0):
|
||||
Process rank for distributed training
|
||||
t5_fsdp (`bool`, *optional*, defaults to False):
|
||||
Enable FSDP sharding for T5 model
|
||||
dit_fsdp (`bool`, *optional*, defaults to False):
|
||||
Enable FSDP sharding for DiT model
|
||||
use_usp (`bool`, *optional*, defaults to False):
|
||||
Enable distribution strategy of USP.
|
||||
t5_cpu (`bool`, *optional*, defaults to False):
|
||||
Whether to place T5 model on CPU. Only works without t5_fsdp.
|
||||
init_on_cpu (`bool`, *optional*, defaults to True):
|
||||
Enable initializing Transformer Model on CPU. Only works without FSDP or USP.
|
||||
"""
|
||||
self.device = torch.device(f"cuda:{device_id}")
|
||||
self.config = config
|
||||
self.rank = rank
|
||||
self.use_usp = use_usp
|
||||
self.t5_cpu = t5_cpu
|
||||
|
||||
self.num_train_timesteps = config.num_train_timesteps
|
||||
self.param_dtype = config.param_dtype
|
||||
|
||||
shard_fn = partial(shard_model, device_id=device_id)
|
||||
self.text_encoder = T5EncoderModel(
|
||||
text_len=config.text_len,
|
||||
dtype=config.t5_dtype,
|
||||
device=torch.device('cpu'),
|
||||
checkpoint_path=os.path.join(checkpoint_dir, config.t5_checkpoint),
|
||||
tokenizer_path=os.path.join(checkpoint_dir, config.t5_tokenizer),
|
||||
shard_fn=shard_fn if t5_fsdp else None,
|
||||
)
|
||||
|
||||
self.vae_stride = config.vae_stride
|
||||
self.patch_size = config.patch_size
|
||||
self.vae = WanVAE(
|
||||
vae_pth=os.path.join(checkpoint_dir, config.vae_checkpoint),
|
||||
device=self.device)
|
||||
|
||||
self.clip = CLIPModel(
|
||||
dtype=config.clip_dtype,
|
||||
device=self.device,
|
||||
checkpoint_path=os.path.join(checkpoint_dir,
|
||||
config.clip_checkpoint),
|
||||
tokenizer_path=os.path.join(checkpoint_dir, config.clip_tokenizer))
|
||||
|
||||
logging.info(f"Creating WanModel from {checkpoint_dir}")
|
||||
self.model = WanModel.from_pretrained(checkpoint_dir)
|
||||
self.model.eval().requires_grad_(False)
|
||||
|
||||
if t5_fsdp or dit_fsdp or use_usp:
|
||||
init_on_cpu = False
|
||||
|
||||
if use_usp:
|
||||
from xfuser.core.distributed import \
|
||||
get_sequence_parallel_world_size
|
||||
|
||||
from .distributed.xdit_context_parallel import (usp_attn_forward,
|
||||
usp_dit_forward)
|
||||
for block in self.model.blocks:
|
||||
block.self_attn.forward = types.MethodType(
|
||||
usp_attn_forward, block.self_attn)
|
||||
self.model.forward = types.MethodType(usp_dit_forward, self.model)
|
||||
self.sp_size = get_sequence_parallel_world_size()
|
||||
else:
|
||||
self.sp_size = 1
|
||||
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
if dit_fsdp:
|
||||
self.model = shard_fn(self.model)
|
||||
else:
|
||||
if not init_on_cpu:
|
||||
self.model.to(self.device)
|
||||
|
||||
self.sample_neg_prompt = config.sample_neg_prompt
|
||||
|
||||
def generate(self,
|
||||
input_prompt,
|
||||
img,
|
||||
max_area=720 * 1280,
|
||||
frame_num=81,
|
||||
shift=5.0,
|
||||
sample_solver='unipc',
|
||||
sampling_steps=40,
|
||||
guide_scale=5.0,
|
||||
n_prompt="",
|
||||
seed=-1,
|
||||
offload_model=True):
|
||||
r"""
|
||||
Generates video frames from input image and text prompt using diffusion process.
|
||||
|
||||
Args:
|
||||
input_prompt (`str`):
|
||||
Text prompt for content generation.
|
||||
img (PIL.Image.Image):
|
||||
Input image tensor. Shape: [3, H, W]
|
||||
max_area (`int`, *optional*, defaults to 720*1280):
|
||||
Maximum pixel area for latent space calculation. Controls video resolution scaling
|
||||
frame_num (`int`, *optional*, defaults to 81):
|
||||
How many frames to sample from a video. The number should be 4n+1
|
||||
shift (`float`, *optional*, defaults to 5.0):
|
||||
Noise schedule shift parameter. Affects temporal dynamics
|
||||
[NOTE]: If you want to generate a 480p video, it is recommended to set the shift value to 3.0.
|
||||
sample_solver (`str`, *optional*, defaults to 'unipc'):
|
||||
Solver used to sample the video.
|
||||
sampling_steps (`int`, *optional*, defaults to 40):
|
||||
Number of diffusion sampling steps. Higher values improve quality but slow generation
|
||||
guide_scale (`float`, *optional*, defaults 5.0):
|
||||
Classifier-free guidance scale. Controls prompt adherence vs. creativity
|
||||
n_prompt (`str`, *optional*, defaults to ""):
|
||||
Negative prompt for content exclusion. If not given, use `config.sample_neg_prompt`
|
||||
seed (`int`, *optional*, defaults to -1):
|
||||
Random seed for noise generation. If -1, use random seed
|
||||
offload_model (`bool`, *optional*, defaults to True):
|
||||
If True, offloads models to CPU during generation to save VRAM
|
||||
|
||||
Returns:
|
||||
torch.Tensor:
|
||||
Generated video frames tensor. Dimensions: (C, N H, W) where:
|
||||
- C: Color channels (3 for RGB)
|
||||
- N: Number of frames (81)
|
||||
- H: Frame height (from max_area)
|
||||
- W: Frame width from max_area)
|
||||
"""
|
||||
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(self.device)
|
||||
|
||||
F = frame_num
|
||||
h, w = img.shape[1:]
|
||||
aspect_ratio = h / w
|
||||
lat_h = round(
|
||||
np.sqrt(max_area * aspect_ratio) // self.vae_stride[1] //
|
||||
self.patch_size[1] * self.patch_size[1])
|
||||
lat_w = round(
|
||||
np.sqrt(max_area / aspect_ratio) // self.vae_stride[2] //
|
||||
self.patch_size[2] * self.patch_size[2])
|
||||
h = lat_h * self.vae_stride[1]
|
||||
w = lat_w * self.vae_stride[2]
|
||||
|
||||
max_seq_len = ((F - 1) // self.vae_stride[0] + 1) * lat_h * lat_w // (
|
||||
self.patch_size[1] * self.patch_size[2])
|
||||
max_seq_len = int(math.ceil(max_seq_len / self.sp_size)) * self.sp_size
|
||||
|
||||
seed = seed if seed >= 0 else random.randint(0, sys.maxsize)
|
||||
seed_g = torch.Generator(device=self.device)
|
||||
seed_g.manual_seed(seed)
|
||||
noise = torch.randn(
|
||||
16,
|
||||
21,
|
||||
lat_h,
|
||||
lat_w,
|
||||
dtype=torch.float32,
|
||||
generator=seed_g,
|
||||
device=self.device)
|
||||
|
||||
msk = torch.ones(1, 81, lat_h, lat_w, device=self.device)
|
||||
msk[:, 1:] = 0
|
||||
msk = torch.concat([
|
||||
torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]
|
||||
],
|
||||
dim=1)
|
||||
msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w)
|
||||
msk = msk.transpose(1, 2)[0]
|
||||
|
||||
if n_prompt == "":
|
||||
n_prompt = self.sample_neg_prompt
|
||||
|
||||
# preprocess
|
||||
if not self.t5_cpu:
|
||||
self.text_encoder.model.to(self.device)
|
||||
context = self.text_encoder([input_prompt], self.device)
|
||||
context_null = self.text_encoder([n_prompt], self.device)
|
||||
if offload_model:
|
||||
self.text_encoder.model.cpu()
|
||||
else:
|
||||
context = self.text_encoder([input_prompt], torch.device('cpu'))
|
||||
context_null = self.text_encoder([n_prompt], torch.device('cpu'))
|
||||
context = [t.to(self.device) for t in context]
|
||||
context_null = [t.to(self.device) for t in context_null]
|
||||
|
||||
self.clip.model.to(self.device)
|
||||
clip_context = self.clip.visual([img[:, None, :, :]])
|
||||
if offload_model:
|
||||
self.clip.model.cpu()
|
||||
|
||||
y = self.vae.encode([
|
||||
torch.concat([
|
||||
torch.nn.functional.interpolate(
|
||||
img[None].cpu(), size=(h, w), mode='bicubic').transpose(
|
||||
0, 1),
|
||||
torch.zeros(3, 80, h, w)
|
||||
],
|
||||
dim=1).to(self.device)
|
||||
])[0]
|
||||
y = torch.concat([msk, y])
|
||||
|
||||
@contextmanager
|
||||
def noop_no_sync():
|
||||
yield
|
||||
|
||||
no_sync = getattr(self.model, 'no_sync', noop_no_sync)
|
||||
|
||||
# evaluation mode
|
||||
with amp.autocast(dtype=self.param_dtype), torch.no_grad(), no_sync():
|
||||
|
||||
if sample_solver == 'unipc':
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(
|
||||
num_train_timesteps=self.num_train_timesteps,
|
||||
shift=1,
|
||||
use_dynamic_shifting=False)
|
||||
sample_scheduler.set_timesteps(
|
||||
sampling_steps, device=self.device, shift=shift)
|
||||
timesteps = sample_scheduler.timesteps
|
||||
elif sample_solver == 'dpm++':
|
||||
sample_scheduler = FlowDPMSolverMultistepScheduler(
|
||||
num_train_timesteps=self.num_train_timesteps,
|
||||
shift=1,
|
||||
use_dynamic_shifting=False)
|
||||
sampling_sigmas = get_sampling_sigmas(sampling_steps, shift)
|
||||
timesteps, _ = retrieve_timesteps(
|
||||
sample_scheduler,
|
||||
device=self.device,
|
||||
sigmas=sampling_sigmas)
|
||||
else:
|
||||
raise NotImplementedError("Unsupported solver.")
|
||||
|
||||
# sample videos
|
||||
latent = noise
|
||||
|
||||
arg_c = {
|
||||
'context': [context[0]],
|
||||
'clip_fea': clip_context,
|
||||
'seq_len': max_seq_len,
|
||||
'y': [y],
|
||||
}
|
||||
|
||||
arg_null = {
|
||||
'context': context_null,
|
||||
'clip_fea': clip_context,
|
||||
'seq_len': max_seq_len,
|
||||
'y': [y],
|
||||
}
|
||||
|
||||
if offload_model:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
self.model.to(self.device)
|
||||
for _, t in enumerate(tqdm(timesteps)):
|
||||
latent_model_input = [latent.to(self.device)]
|
||||
timestep = [t]
|
||||
|
||||
timestep = torch.stack(timestep).to(self.device)
|
||||
|
||||
noise_pred_cond = self.model(
|
||||
latent_model_input, t=timestep, **arg_c)[0].to(
|
||||
torch.device('cpu') if offload_model else self.device)
|
||||
if offload_model:
|
||||
torch.cuda.empty_cache()
|
||||
noise_pred_uncond = self.model(
|
||||
latent_model_input, t=timestep, **arg_null)[0].to(
|
||||
torch.device('cpu') if offload_model else self.device)
|
||||
if offload_model:
|
||||
torch.cuda.empty_cache()
|
||||
noise_pred = noise_pred_uncond + guide_scale * (
|
||||
noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
latent = latent.to(
|
||||
torch.device('cpu') if offload_model else self.device)
|
||||
|
||||
temp_x0 = sample_scheduler.step(
|
||||
noise_pred.unsqueeze(0),
|
||||
t,
|
||||
latent.unsqueeze(0),
|
||||
return_dict=False,
|
||||
generator=seed_g)[0]
|
||||
latent = temp_x0.squeeze(0)
|
||||
|
||||
x0 = [latent.to(self.device)]
|
||||
del latent_model_input, timestep
|
||||
|
||||
if offload_model:
|
||||
self.model.cpu()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if self.rank == 0:
|
||||
videos = self.vae.decode(x0)
|
||||
|
||||
del noise, latent
|
||||
del sample_scheduler
|
||||
if offload_model:
|
||||
gc.collect()
|
||||
torch.cuda.synchronize()
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
|
||||
return videos[0] if self.rank == 0 else None
|
||||
@@ -1,16 +0,0 @@
|
||||
from .attention import flash_attention
|
||||
from .model import WanModel
|
||||
from .t5 import T5Decoder, T5Encoder, T5EncoderModel, T5Model
|
||||
from .tokenizers import HuggingfaceTokenizer
|
||||
from .vae import WanVAE
|
||||
|
||||
__all__ = [
|
||||
'WanVAE',
|
||||
'WanModel',
|
||||
'T5Model',
|
||||
'T5Encoder',
|
||||
'T5Decoder',
|
||||
'T5EncoderModel',
|
||||
'HuggingfaceTokenizer',
|
||||
'flash_attention',
|
||||
]
|
||||
@@ -1,185 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
|
||||
try:
|
||||
import flash_attn_interface
|
||||
|
||||
def is_hopper_gpu():
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
device_name = torch.cuda.get_device_name(0).lower()
|
||||
return "h100" in device_name or "hopper" in device_name
|
||||
FLASH_ATTN_3_AVAILABLE = is_hopper_gpu()
|
||||
except ModuleNotFoundError:
|
||||
FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
try:
|
||||
import flash_attn
|
||||
FLASH_ATTN_2_AVAILABLE = True
|
||||
except ModuleNotFoundError:
|
||||
FLASH_ATTN_2_AVAILABLE = False
|
||||
|
||||
# FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
import warnings
|
||||
|
||||
__all__ = [
|
||||
'flash_attention',
|
||||
'attention',
|
||||
]
|
||||
|
||||
|
||||
def flash_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_lens=None,
|
||||
k_lens=None,
|
||||
dropout_p=0.,
|
||||
softmax_scale=None,
|
||||
q_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
dtype=torch.bfloat16,
|
||||
version=None,
|
||||
):
|
||||
"""
|
||||
q: [B, Lq, Nq, C1].
|
||||
k: [B, Lk, Nk, C1].
|
||||
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
|
||||
q_lens: [B].
|
||||
k_lens: [B].
|
||||
dropout_p: float. Dropout probability.
|
||||
softmax_scale: float. The scaling of QK^T before applying softmax.
|
||||
causal: bool. Whether to apply causal attention mask.
|
||||
window_size: (left right). If not (-1, -1), apply sliding window local attention.
|
||||
deterministic: bool. If True, slightly slower and uses more memory.
|
||||
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
|
||||
"""
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
assert dtype in half_dtypes
|
||||
assert q.device.type == 'cuda' and q.size(-1) <= 256
|
||||
|
||||
# params
|
||||
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q = half(q.flatten(0, 1))
|
||||
q_lens = torch.tensor(
|
||||
[lq] * b, dtype=torch.int32).to(
|
||||
device=q.device, non_blocking=True)
|
||||
else:
|
||||
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if k_lens is None:
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
k_lens = torch.tensor(
|
||||
[lk] * b, dtype=torch.int32).to(
|
||||
device=k.device, non_blocking=True)
|
||||
else:
|
||||
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
|
||||
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
|
||||
|
||||
q = q.to(v.dtype)
|
||||
k = k.to(v.dtype)
|
||||
|
||||
if q_scale is not None:
|
||||
q = q * q_scale
|
||||
|
||||
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
||||
warnings.warn(
|
||||
'Flash attention 3 is not available, use flash attention 2 instead.'
|
||||
)
|
||||
|
||||
# apply attention
|
||||
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
||||
# Note: dropout_p, window_size are not supported in FA3 now.
|
||||
x = flash_attn_interface.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic)[0].unflatten(0, (b, lq))
|
||||
else:
|
||||
assert FLASH_ATTN_2_AVAILABLE
|
||||
x = flash_attn.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic).unflatten(0, (b, lq))
|
||||
|
||||
# output
|
||||
return x.type(out_dtype)
|
||||
|
||||
|
||||
def attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_lens=None,
|
||||
k_lens=None,
|
||||
dropout_p=0.,
|
||||
softmax_scale=None,
|
||||
q_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
dtype=torch.bfloat16,
|
||||
fa_version=None,
|
||||
):
|
||||
if FLASH_ATTN_2_AVAILABLE or FLASH_ATTN_3_AVAILABLE:
|
||||
return flash_attention(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
q_lens=q_lens,
|
||||
k_lens=k_lens,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
q_scale=q_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic,
|
||||
dtype=dtype,
|
||||
version=fa_version,
|
||||
)
|
||||
else:
|
||||
if q_lens is not None or k_lens is not None:
|
||||
warnings.warn(
|
||||
'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
|
||||
)
|
||||
attn_mask = None
|
||||
|
||||
q = q.transpose(1, 2).to(dtype)
|
||||
k = k.transpose(1, 2).to(dtype)
|
||||
v = v.transpose(1, 2).to(dtype)
|
||||
|
||||
out = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
|
||||
|
||||
out = out.transpose(1, 2).contiguous()
|
||||
return out
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,542 +0,0 @@
|
||||
# Modified from ``https://github.com/openai/CLIP'' and ``https://github.com/mlfoundations/open_clip''
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import logging
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as T
|
||||
|
||||
from .attention import flash_attention
|
||||
from .tokenizers import HuggingfaceTokenizer
|
||||
from .xlm_roberta import XLMRoberta
|
||||
|
||||
__all__ = [
|
||||
'XLMRobertaCLIP',
|
||||
'clip_xlm_roberta_vit_h_14',
|
||||
'CLIPModel',
|
||||
]
|
||||
|
||||
|
||||
def pos_interpolate(pos, seq_len):
|
||||
if pos.size(1) == seq_len:
|
||||
return pos
|
||||
else:
|
||||
src_grid = int(math.sqrt(pos.size(1)))
|
||||
tar_grid = int(math.sqrt(seq_len))
|
||||
n = pos.size(1) - src_grid * src_grid
|
||||
return torch.cat([
|
||||
pos[:, :n],
|
||||
F.interpolate(
|
||||
pos[:, n:].float().reshape(1, src_grid, src_grid, -1).permute(
|
||||
0, 3, 1, 2),
|
||||
size=(tar_grid, tar_grid),
|
||||
mode='bicubic',
|
||||
align_corners=False).flatten(2).transpose(1, 2)
|
||||
],
|
||||
dim=1)
|
||||
|
||||
|
||||
class QuickGELU(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
|
||||
def forward(self, x):
|
||||
return super().forward(x.float()).type_as(x)
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
causal=False,
|
||||
attn_dropout=0.0,
|
||||
proj_dropout=0.0):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.causal = causal
|
||||
self.attn_dropout = attn_dropout
|
||||
self.proj_dropout = proj_dropout
|
||||
|
||||
# layers
|
||||
self.to_qkv = nn.Linear(dim, dim * 3)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
x: [B, L, C].
|
||||
"""
|
||||
b, s, c, n, d = *x.size(), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q, k, v = self.to_qkv(x).view(b, s, 3, n, d).unbind(2)
|
||||
|
||||
# compute attention
|
||||
p = self.attn_dropout if self.training else 0.0
|
||||
x = flash_attention(q, k, v, dropout_p=p, causal=self.causal, version=2)
|
||||
x = x.reshape(b, s, c)
|
||||
|
||||
# output
|
||||
x = self.proj(x)
|
||||
x = F.dropout(x, self.proj_dropout, self.training)
|
||||
return x
|
||||
|
||||
|
||||
class SwiGLU(nn.Module):
|
||||
|
||||
def __init__(self, dim, mid_dim):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.mid_dim = mid_dim
|
||||
|
||||
# layers
|
||||
self.fc1 = nn.Linear(dim, mid_dim)
|
||||
self.fc2 = nn.Linear(dim, mid_dim)
|
||||
self.fc3 = nn.Linear(mid_dim, dim)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.silu(self.fc1(x)) * self.fc2(x)
|
||||
x = self.fc3(x)
|
||||
return x
|
||||
|
||||
|
||||
class AttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
mlp_ratio,
|
||||
num_heads,
|
||||
post_norm=False,
|
||||
causal=False,
|
||||
activation='quick_gelu',
|
||||
attn_dropout=0.0,
|
||||
proj_dropout=0.0,
|
||||
norm_eps=1e-5):
|
||||
assert activation in ['quick_gelu', 'gelu', 'swi_glu']
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.num_heads = num_heads
|
||||
self.post_norm = post_norm
|
||||
self.causal = causal
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
# layers
|
||||
self.norm1 = LayerNorm(dim, eps=norm_eps)
|
||||
self.attn = SelfAttention(dim, num_heads, causal, attn_dropout,
|
||||
proj_dropout)
|
||||
self.norm2 = LayerNorm(dim, eps=norm_eps)
|
||||
if activation == 'swi_glu':
|
||||
self.mlp = SwiGLU(dim, int(dim * mlp_ratio))
|
||||
else:
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(dim, int(dim * mlp_ratio)),
|
||||
QuickGELU() if activation == 'quick_gelu' else nn.GELU(),
|
||||
nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout))
|
||||
|
||||
def forward(self, x):
|
||||
if self.post_norm:
|
||||
x = x + self.norm1(self.attn(x))
|
||||
x = x + self.norm2(self.mlp(x))
|
||||
else:
|
||||
x = x + self.attn(self.norm1(x))
|
||||
x = x + self.mlp(self.norm2(x))
|
||||
return x
|
||||
|
||||
|
||||
class AttentionPool(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
mlp_ratio,
|
||||
num_heads,
|
||||
activation='gelu',
|
||||
proj_dropout=0.0,
|
||||
norm_eps=1e-5):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.proj_dropout = proj_dropout
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
# layers
|
||||
gain = 1.0 / math.sqrt(dim)
|
||||
self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))
|
||||
self.to_q = nn.Linear(dim, dim)
|
||||
self.to_kv = nn.Linear(dim, dim * 2)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.norm = LayerNorm(dim, eps=norm_eps)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(dim, int(dim * mlp_ratio)),
|
||||
QuickGELU() if activation == 'quick_gelu' else nn.GELU(),
|
||||
nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout))
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
x: [B, L, C].
|
||||
"""
|
||||
b, s, c, n, d = *x.size(), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.to_q(self.cls_embedding).view(1, 1, n, d).expand(b, -1, -1, -1)
|
||||
k, v = self.to_kv(x).view(b, s, 2, n, d).unbind(2)
|
||||
|
||||
# compute attention
|
||||
x = flash_attention(q, k, v, version=2)
|
||||
x = x.reshape(b, 1, c)
|
||||
|
||||
# output
|
||||
x = self.proj(x)
|
||||
x = F.dropout(x, self.proj_dropout, self.training)
|
||||
|
||||
# mlp
|
||||
x = x + self.mlp(self.norm(x))
|
||||
return x[:, 0]
|
||||
|
||||
|
||||
class VisionTransformer(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
image_size=224,
|
||||
patch_size=16,
|
||||
dim=768,
|
||||
mlp_ratio=4,
|
||||
out_dim=512,
|
||||
num_heads=12,
|
||||
num_layers=12,
|
||||
pool_type='token',
|
||||
pre_norm=True,
|
||||
post_norm=False,
|
||||
activation='quick_gelu',
|
||||
attn_dropout=0.0,
|
||||
proj_dropout=0.0,
|
||||
embedding_dropout=0.0,
|
||||
norm_eps=1e-5):
|
||||
if image_size % patch_size != 0:
|
||||
print(
|
||||
'[WARNING] image_size is not divisible by patch_size',
|
||||
flush=True)
|
||||
assert pool_type in ('token', 'token_fc', 'attn_pool')
|
||||
out_dim = out_dim or dim
|
||||
super().__init__()
|
||||
self.image_size = image_size
|
||||
self.patch_size = patch_size
|
||||
self.num_patches = (image_size // patch_size)**2
|
||||
self.dim = dim
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.out_dim = out_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.pool_type = pool_type
|
||||
self.post_norm = post_norm
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
# embeddings
|
||||
gain = 1.0 / math.sqrt(dim)
|
||||
self.patch_embedding = nn.Conv2d(
|
||||
3,
|
||||
dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=not pre_norm)
|
||||
if pool_type in ('token', 'token_fc'):
|
||||
self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))
|
||||
self.pos_embedding = nn.Parameter(gain * torch.randn(
|
||||
1, self.num_patches +
|
||||
(1 if pool_type in ('token', 'token_fc') else 0), dim))
|
||||
self.dropout = nn.Dropout(embedding_dropout)
|
||||
|
||||
# transformer
|
||||
self.pre_norm = LayerNorm(dim, eps=norm_eps) if pre_norm else None
|
||||
self.transformer = nn.Sequential(*[
|
||||
AttentionBlock(dim, mlp_ratio, num_heads, post_norm, False,
|
||||
activation, attn_dropout, proj_dropout, norm_eps)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
self.post_norm = LayerNorm(dim, eps=norm_eps)
|
||||
|
||||
# head
|
||||
if pool_type == 'token':
|
||||
self.head = nn.Parameter(gain * torch.randn(dim, out_dim))
|
||||
elif pool_type == 'token_fc':
|
||||
self.head = nn.Linear(dim, out_dim)
|
||||
elif pool_type == 'attn_pool':
|
||||
self.head = AttentionPool(dim, mlp_ratio, num_heads, activation,
|
||||
proj_dropout, norm_eps)
|
||||
|
||||
def forward(self, x, interpolation=False, use_31_block=False):
|
||||
b = x.size(0)
|
||||
|
||||
# embeddings
|
||||
x = self.patch_embedding(x).flatten(2).permute(0, 2, 1)
|
||||
if self.pool_type in ('token', 'token_fc'):
|
||||
x = torch.cat([self.cls_embedding.expand(b, -1, -1), x], dim=1)
|
||||
if interpolation:
|
||||
e = pos_interpolate(self.pos_embedding, x.size(1))
|
||||
else:
|
||||
e = self.pos_embedding
|
||||
x = self.dropout(x + e)
|
||||
if self.pre_norm is not None:
|
||||
x = self.pre_norm(x)
|
||||
|
||||
# transformer
|
||||
if use_31_block:
|
||||
x = self.transformer[:-1](x)
|
||||
return x
|
||||
else:
|
||||
x = self.transformer(x)
|
||||
return x
|
||||
|
||||
|
||||
class XLMRobertaWithHead(XLMRoberta):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.out_dim = kwargs.pop('out_dim')
|
||||
super().__init__(**kwargs)
|
||||
|
||||
# head
|
||||
mid_dim = (self.dim + self.out_dim) // 2
|
||||
self.head = nn.Sequential(
|
||||
nn.Linear(self.dim, mid_dim, bias=False), nn.GELU(),
|
||||
nn.Linear(mid_dim, self.out_dim, bias=False))
|
||||
|
||||
def forward(self, ids):
|
||||
# xlm-roberta
|
||||
x = super().forward(ids)
|
||||
|
||||
# average pooling
|
||||
mask = ids.ne(self.pad_id).unsqueeze(-1).to(x)
|
||||
x = (x * mask).sum(dim=1) / mask.sum(dim=1)
|
||||
|
||||
# head
|
||||
x = self.head(x)
|
||||
return x
|
||||
|
||||
|
||||
class XLMRobertaCLIP(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
embed_dim=1024,
|
||||
image_size=224,
|
||||
patch_size=14,
|
||||
vision_dim=1280,
|
||||
vision_mlp_ratio=4,
|
||||
vision_heads=16,
|
||||
vision_layers=32,
|
||||
vision_pool='token',
|
||||
vision_pre_norm=True,
|
||||
vision_post_norm=False,
|
||||
activation='gelu',
|
||||
vocab_size=250002,
|
||||
max_text_len=514,
|
||||
type_size=1,
|
||||
pad_id=1,
|
||||
text_dim=1024,
|
||||
text_heads=16,
|
||||
text_layers=24,
|
||||
text_post_norm=True,
|
||||
text_dropout=0.1,
|
||||
attn_dropout=0.0,
|
||||
proj_dropout=0.0,
|
||||
embedding_dropout=0.0,
|
||||
norm_eps=1e-5):
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.image_size = image_size
|
||||
self.patch_size = patch_size
|
||||
self.vision_dim = vision_dim
|
||||
self.vision_mlp_ratio = vision_mlp_ratio
|
||||
self.vision_heads = vision_heads
|
||||
self.vision_layers = vision_layers
|
||||
self.vision_pre_norm = vision_pre_norm
|
||||
self.vision_post_norm = vision_post_norm
|
||||
self.activation = activation
|
||||
self.vocab_size = vocab_size
|
||||
self.max_text_len = max_text_len
|
||||
self.type_size = type_size
|
||||
self.pad_id = pad_id
|
||||
self.text_dim = text_dim
|
||||
self.text_heads = text_heads
|
||||
self.text_layers = text_layers
|
||||
self.text_post_norm = text_post_norm
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
# models
|
||||
self.visual = VisionTransformer(
|
||||
image_size=image_size,
|
||||
patch_size=patch_size,
|
||||
dim=vision_dim,
|
||||
mlp_ratio=vision_mlp_ratio,
|
||||
out_dim=embed_dim,
|
||||
num_heads=vision_heads,
|
||||
num_layers=vision_layers,
|
||||
pool_type=vision_pool,
|
||||
pre_norm=vision_pre_norm,
|
||||
post_norm=vision_post_norm,
|
||||
activation=activation,
|
||||
attn_dropout=attn_dropout,
|
||||
proj_dropout=proj_dropout,
|
||||
embedding_dropout=embedding_dropout,
|
||||
norm_eps=norm_eps)
|
||||
self.textual = XLMRobertaWithHead(
|
||||
vocab_size=vocab_size,
|
||||
max_seq_len=max_text_len,
|
||||
type_size=type_size,
|
||||
pad_id=pad_id,
|
||||
dim=text_dim,
|
||||
out_dim=embed_dim,
|
||||
num_heads=text_heads,
|
||||
num_layers=text_layers,
|
||||
post_norm=text_post_norm,
|
||||
dropout=text_dropout)
|
||||
self.log_scale = nn.Parameter(math.log(1 / 0.07) * torch.ones([]))
|
||||
|
||||
def forward(self, imgs, txt_ids):
|
||||
"""
|
||||
imgs: [B, 3, H, W] of torch.float32.
|
||||
- mean: [0.48145466, 0.4578275, 0.40821073]
|
||||
- std: [0.26862954, 0.26130258, 0.27577711]
|
||||
txt_ids: [B, L] of torch.long.
|
||||
Encoded by data.CLIPTokenizer.
|
||||
"""
|
||||
xi = self.visual(imgs)
|
||||
xt = self.textual(txt_ids)
|
||||
return xi, xt
|
||||
|
||||
def param_groups(self):
|
||||
groups = [{
|
||||
'params': [
|
||||
p for n, p in self.named_parameters()
|
||||
if 'norm' in n or n.endswith('bias')
|
||||
],
|
||||
'weight_decay': 0.0
|
||||
}, {
|
||||
'params': [
|
||||
p for n, p in self.named_parameters()
|
||||
if not ('norm' in n or n.endswith('bias'))
|
||||
]
|
||||
}]
|
||||
return groups
|
||||
|
||||
|
||||
def _clip(pretrained=False,
|
||||
pretrained_name=None,
|
||||
model_cls=XLMRobertaCLIP,
|
||||
return_transforms=False,
|
||||
return_tokenizer=False,
|
||||
tokenizer_padding='eos',
|
||||
dtype=torch.float32,
|
||||
device='cpu',
|
||||
**kwargs):
|
||||
# init a model on device
|
||||
with torch.device(device):
|
||||
model = model_cls(**kwargs)
|
||||
|
||||
# set device
|
||||
model = model.to(dtype=dtype, device=device)
|
||||
output = (model,)
|
||||
|
||||
# init transforms
|
||||
if return_transforms:
|
||||
# mean and std
|
||||
if 'siglip' in pretrained_name.lower():
|
||||
mean, std = [0.5, 0.5, 0.5], [0.5, 0.5, 0.5]
|
||||
else:
|
||||
mean = [0.48145466, 0.4578275, 0.40821073]
|
||||
std = [0.26862954, 0.26130258, 0.27577711]
|
||||
|
||||
# transforms
|
||||
transforms = T.Compose([
|
||||
T.Resize((model.image_size, model.image_size),
|
||||
interpolation=T.InterpolationMode.BICUBIC),
|
||||
T.ToTensor(),
|
||||
T.Normalize(mean=mean, std=std)
|
||||
])
|
||||
output += (transforms,)
|
||||
return output[0] if len(output) == 1 else output
|
||||
|
||||
|
||||
def clip_xlm_roberta_vit_h_14(
|
||||
pretrained=False,
|
||||
pretrained_name='open-clip-xlm-roberta-large-vit-huge-14',
|
||||
**kwargs):
|
||||
cfg = dict(
|
||||
embed_dim=1024,
|
||||
image_size=224,
|
||||
patch_size=14,
|
||||
vision_dim=1280,
|
||||
vision_mlp_ratio=4,
|
||||
vision_heads=16,
|
||||
vision_layers=32,
|
||||
vision_pool='token',
|
||||
activation='gelu',
|
||||
vocab_size=250002,
|
||||
max_text_len=514,
|
||||
type_size=1,
|
||||
pad_id=1,
|
||||
text_dim=1024,
|
||||
text_heads=16,
|
||||
text_layers=24,
|
||||
text_post_norm=True,
|
||||
text_dropout=0.1,
|
||||
attn_dropout=0.0,
|
||||
proj_dropout=0.0,
|
||||
embedding_dropout=0.0)
|
||||
cfg.update(**kwargs)
|
||||
return _clip(pretrained, pretrained_name, XLMRobertaCLIP, **cfg)
|
||||
|
||||
|
||||
class CLIPModel:
|
||||
|
||||
def __init__(self, dtype, device, checkpoint_path, tokenizer_path):
|
||||
self.dtype = dtype
|
||||
self.device = device
|
||||
self.checkpoint_path = checkpoint_path
|
||||
self.tokenizer_path = tokenizer_path
|
||||
|
||||
# init model
|
||||
self.model, self.transforms = clip_xlm_roberta_vit_h_14(
|
||||
pretrained=False,
|
||||
return_transforms=True,
|
||||
return_tokenizer=False,
|
||||
dtype=dtype,
|
||||
device=device)
|
||||
self.model = self.model.eval().requires_grad_(False)
|
||||
logging.info(f'loading {checkpoint_path}')
|
||||
self.model.load_state_dict(
|
||||
torch.load(checkpoint_path, map_location='cpu'))
|
||||
|
||||
# init tokenizer
|
||||
self.tokenizer = HuggingfaceTokenizer(
|
||||
name=tokenizer_path,
|
||||
seq_len=self.model.max_text_len - 2,
|
||||
clean='whitespace')
|
||||
|
||||
def visual(self, videos):
|
||||
# preprocess
|
||||
size = (self.model.image_size,) * 2
|
||||
videos = torch.cat([
|
||||
F.interpolate(
|
||||
u.transpose(0, 1),
|
||||
size=size,
|
||||
mode='bicubic',
|
||||
align_corners=False) for u in videos
|
||||
])
|
||||
videos = self.transforms.transforms[-1](videos.mul_(0.5).add_(0.5))
|
||||
|
||||
# forward
|
||||
with torch.cuda.amp.autocast(dtype=self.dtype):
|
||||
out = self.model.visual(videos, use_31_block=True)
|
||||
return out
|
||||
@@ -1,934 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from einops import repeat
|
||||
|
||||
from .attention import flash_attention
|
||||
|
||||
__all__ = ['WanModel']
|
||||
|
||||
|
||||
def sinusoidal_embedding_1d(dim, position):
|
||||
# preprocess
|
||||
assert dim % 2 == 0
|
||||
half = dim // 2
|
||||
position = position.type(torch.float64)
|
||||
|
||||
# calculation
|
||||
sinusoid = torch.outer(
|
||||
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
|
||||
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
return x
|
||||
|
||||
|
||||
# @amp.autocast(enabled=False)
|
||||
def rope_params(max_seq_len, dim, theta=10000):
|
||||
assert dim % 2 == 0
|
||||
freqs = torch.outer(
|
||||
torch.arange(max_seq_len),
|
||||
1.0 / torch.pow(theta,
|
||||
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
|
||||
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
||||
return freqs
|
||||
|
||||
|
||||
# @amp.autocast(enabled=False)
|
||||
def rope_apply(x, grid_sizes, freqs):
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
|
||||
# split freqs
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
# loop over samples
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
|
||||
# precompute multipliers
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
|
||||
seq_len, n, -1, 2))
|
||||
freqs_i = torch.cat([
|
||||
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||
],
|
||||
dim=-1).reshape(seq_len, 1, -1)
|
||||
|
||||
# apply rotary embedding
|
||||
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]])
|
||||
|
||||
# append to collection
|
||||
output.append(x_i)
|
||||
return torch.stack(output).type_as(x)
|
||||
|
||||
|
||||
class WanRMSNorm(nn.Module):
|
||||
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, C]
|
||||
"""
|
||||
return self._norm(x.float()).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class WanLayerNorm(nn.LayerNorm):
|
||||
|
||||
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||
|
||||
def forward(self, x):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, C]
|
||||
"""
|
||||
return super().forward(x).type_as(x)
|
||||
|
||||
|
||||
class WanSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim)
|
||||
self.k = nn.Linear(dim, dim)
|
||||
self.v = nn.Linear(dim, dim)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs):
|
||||
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]
|
||||
"""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
print(f"query sum: {torch.sum(q.float()).item()}")
|
||||
|
||||
q = rope_apply(q, grid_sizes, freqs)
|
||||
|
||||
print(f"query after rotary embeddings sum: {torch.sum(q.float()).item()}")
|
||||
|
||||
x = flash_attention(
|
||||
q=q,
|
||||
k=rope_apply(k, grid_sizes, freqs),
|
||||
v=v,
|
||||
k_lens=seq_lens,
|
||||
window_size=self.window_size)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
|
||||
print(f"attn_output sum: {torch.sum(x.float()).item()}")
|
||||
return x
|
||||
|
||||
|
||||
class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, context_lens, crossattn_cache=None):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
crossattn_cache (List[dict], *optional*): Contains the cached key and value tensors for context embedding.
|
||||
"""
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.q(x)).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.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).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.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
x = flash_attention(q, k, v, k_lens=context_lens)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
|
||||
|
||||
class WanGanCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, crossattn_cache=None):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
crossattn_cache (List[dict], *optional*): Contains the cached key and value tensors for context embedding.
|
||||
"""
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
qq = self.norm_q(self.q(context)).view(b, 1, -1, d)
|
||||
|
||||
kk = self.norm_k(self.k(x)).view(b, -1, n, d)
|
||||
vv = self.v(x).view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
x = flash_attention(qq, kk, vv)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
|
||||
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6):
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps)
|
||||
|
||||
self.k_img = nn.Linear(dim, dim)
|
||||
self.v_img = nn.Linear(dim, dim)
|
||||
# self.alpha = nn.Parameter(torch.zeros((1, )))
|
||||
self.norm_k_img = WanRMSNorm(
|
||||
dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, context, context_lens):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
"""
|
||||
context_img = context[:, :257]
|
||||
context = context[:, 257:]
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
k_img = self.norm_k_img(self.k_img(context_img)).view(b, -1, n, d)
|
||||
v_img = self.v_img(context_img).view(b, -1, n, d)
|
||||
img_x = flash_attention(q, k_img, v_img, k_lens=None)
|
||||
# compute attention
|
||||
x = flash_attention(q, k, v, k_lens=context_lens)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
img_x = img_x.flatten(2)
|
||||
x = x + img_x
|
||||
x = self.o(x)
|
||||
return x
|
||||
|
||||
|
||||
WAN_CROSSATTENTION_CLASSES = {
|
||||
't2v_cross_attn': WanT2VCrossAttention,
|
||||
'i2v_cross_attn': WanI2VCrossAttention,
|
||||
}
|
||||
|
||||
|
||||
class WanAttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
cross_attn_type,
|
||||
dim,
|
||||
ffn_dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
cross_attn_norm=False,
|
||||
eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(dim, eps)
|
||||
self.self_attn = WanSelfAttention(dim, num_heads, window_size, qk_norm,
|
||||
eps)
|
||||
self.norm3 = WanLayerNorm(
|
||||
dim, eps,
|
||||
elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](dim,
|
||||
num_heads,
|
||||
(-1, -1),
|
||||
qk_norm,
|
||||
eps)
|
||||
self.norm2 = WanLayerNorm(dim, eps)
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn_dim, dim))
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
e,
|
||||
seq_lens,
|
||||
grid_sizes,
|
||||
freqs,
|
||||
context,
|
||||
context_lens,
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, C]
|
||||
e(Tensor): Shape [B, 6, C]
|
||||
seq_lens(Tensor): Shape [B], length of each sequence in batch
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
# assert e.dtype == torch.float32
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
e = (self.modulation + e).chunk(6, dim=1)
|
||||
# assert e[0].dtype == torch.float32
|
||||
|
||||
|
||||
norm_x = self.norm1(x) * (1 + e[1]) + e[0]
|
||||
print(f"norm_hidden_states sum: {torch.sum(norm_x.float()).item()}")
|
||||
# self-attention
|
||||
y = self.self_attn(
|
||||
norm_x, seq_lens, grid_sizes,
|
||||
freqs)
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
x = x + y * e[2]
|
||||
|
||||
# cross-attention & ffn function
|
||||
def cross_attn_ffn(x, context, context_lens, e):
|
||||
x = x + self.cross_attn(self.norm3(x), context, context_lens)
|
||||
y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3])
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
x = x + y * e[5]
|
||||
return x
|
||||
|
||||
x = cross_attn_ffn(x, context, context_lens, e)
|
||||
return x
|
||||
|
||||
|
||||
class GanAttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim=1536,
|
||||
ffn_dim=8192,
|
||||
num_heads=12,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
# self.norm1 = WanLayerNorm(dim, eps)
|
||||
# self.self_attn = WanSelfAttention(dim, num_heads, window_size, qk_norm,
|
||||
# eps)
|
||||
self.norm3 = WanLayerNorm(
|
||||
dim, eps,
|
||||
elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
|
||||
self.norm2 = WanLayerNorm(dim, eps)
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn_dim, dim))
|
||||
|
||||
self.cross_attn = WanGanCrossAttention(dim, num_heads,
|
||||
(-1, -1),
|
||||
qk_norm,
|
||||
eps)
|
||||
|
||||
# modulation
|
||||
# self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
context,
|
||||
# seq_lens,
|
||||
# grid_sizes,
|
||||
# freqs,
|
||||
# context,
|
||||
# context_lens,
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, C]
|
||||
e(Tensor): Shape [B, 6, C]
|
||||
seq_lens(Tensor): Shape [B], length of each sequence in batch
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
# assert e.dtype == torch.float32
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
# e = (self.modulation + e).chunk(6, dim=1)
|
||||
# assert e[0].dtype == torch.float32
|
||||
|
||||
# # self-attention
|
||||
# y = self.self_attn(
|
||||
# self.norm1(x) * (1 + e[1]) + e[0], seq_lens, grid_sizes,
|
||||
# freqs)
|
||||
# # with amp.autocast(dtype=torch.float32):
|
||||
# x = x + y * e[2]
|
||||
|
||||
# cross-attention & ffn function
|
||||
def cross_attn_ffn(x, context):
|
||||
token = context + self.cross_attn(self.norm3(x), context)
|
||||
y = self.ffn(self.norm2(token)) + token # * (1 + e[4]) + e[3])
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
# x = x + y * e[5]
|
||||
return y
|
||||
|
||||
x = cross_attn_ffn(x, context)
|
||||
return x
|
||||
|
||||
|
||||
class Head(nn.Module):
|
||||
|
||||
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.out_dim = out_dim
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
out_dim = math.prod(patch_size) * out_dim
|
||||
self.norm = WanLayerNorm(dim, eps)
|
||||
self.head = nn.Linear(dim, out_dim)
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
||||
|
||||
def forward(self, x, e):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
e(Tensor): Shape [B, C]
|
||||
"""
|
||||
# assert e.dtype == torch.float32
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
e = (self.modulation + e.unsqueeze(1)).chunk(2, dim=1)
|
||||
x = (self.head(self.norm(x) * (1 + e[1]) + e[0]))
|
||||
return x
|
||||
|
||||
|
||||
class MLPProj(torch.nn.Module):
|
||||
|
||||
def __init__(self, in_dim, out_dim):
|
||||
super().__init__()
|
||||
|
||||
self.proj = torch.nn.Sequential(
|
||||
torch.nn.LayerNorm(in_dim), torch.nn.Linear(in_dim, in_dim),
|
||||
torch.nn.GELU(), torch.nn.Linear(in_dim, out_dim),
|
||||
torch.nn.LayerNorm(out_dim))
|
||||
|
||||
def forward(self, image_embeds):
|
||||
clip_extra_context_tokens = self.proj(image_embeds)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
|
||||
class RegisterTokens(nn.Module):
|
||||
def __init__(self, num_registers: int, dim: int):
|
||||
super().__init__()
|
||||
self.register_tokens = nn.Parameter(torch.randn(num_registers, dim) * 0.02)
|
||||
self.rms_norm = WanRMSNorm(dim, eps=1e-6)
|
||||
|
||||
def forward(self):
|
||||
return self.rms_norm(self.register_tokens)
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.normal_(self.register_tokens, std=0.02)
|
||||
|
||||
|
||||
class WanModel(ModelMixin, ConfigMixin):
|
||||
r"""
|
||||
Wan diffusion backbone supporting both text-to-video and image-to-video.
|
||||
"""
|
||||
|
||||
ignore_for_config = [
|
||||
'patch_size', 'cross_attn_norm', 'qk_norm', 'text_dim', 'window_size'
|
||||
]
|
||||
_no_split_modules = ['WanAttentionBlock']
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(self,
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
ffn_dim=8192,
|
||||
freq_dim=256,
|
||||
text_dim=4096,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
|
||||
Args:
|
||||
model_type (`str`, *optional*, defaults to 't2v'):
|
||||
Model variant - 't2v' (text-to-video) or 'i2v' (image-to-video)
|
||||
patch_size (`tuple`, *optional*, defaults to (1, 2, 2)):
|
||||
3D patch dimensions for video embedding (t_patch, h_patch, w_patch)
|
||||
text_len (`int`, *optional*, defaults to 512):
|
||||
Fixed length for text embeddings
|
||||
in_dim (`int`, *optional*, defaults to 16):
|
||||
Input video channels (C_in)
|
||||
dim (`int`, *optional*, defaults to 2048):
|
||||
Hidden dimension of the transformer
|
||||
ffn_dim (`int`, *optional*, defaults to 8192):
|
||||
Intermediate dimension in feed-forward network
|
||||
freq_dim (`int`, *optional*, defaults to 256):
|
||||
Dimension for sinusoidal time embeddings
|
||||
text_dim (`int`, *optional*, defaults to 4096):
|
||||
Input dimension for text embeddings
|
||||
out_dim (`int`, *optional*, defaults to 16):
|
||||
Output video channels (C_out)
|
||||
num_heads (`int`, *optional*, defaults to 16):
|
||||
Number of attention heads
|
||||
num_layers (`int`, *optional*, defaults to 32):
|
||||
Number of transformer blocks
|
||||
window_size (`tuple`, *optional*, defaults to (-1, -1)):
|
||||
Window size for local attention (-1 indicates global attention)
|
||||
qk_norm (`bool`, *optional*, defaults to True):
|
||||
Enable query/key normalization
|
||||
cross_attn_norm (`bool`, *optional*, defaults to False):
|
||||
Enable cross-attention normalization
|
||||
eps (`float`, *optional*, defaults to 1e-6):
|
||||
Epsilon value for normalization layers
|
||||
"""
|
||||
|
||||
super().__init__()
|
||||
|
||||
assert model_type in ['t2v', 'i2v']
|
||||
self.model_type = model_type
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
self.in_dim = in_dim
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.freq_dim = freq_dim
|
||||
self.text_dim = text_dim
|
||||
self.out_dim = out_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
self.local_attn_size = 21
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.text_embedding = nn.Sequential(
|
||||
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(dim, dim))
|
||||
|
||||
self.time_embedding = nn.Sequential(
|
||||
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
|
||||
# blocks
|
||||
cross_attn_type = 't2v_cross_attn' if model_type == 't2v' else 'i2v_cross_attn'
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock(cross_attn_type, dim, ffn_dim, num_heads,
|
||||
window_size, qk_norm, cross_attn_norm, eps)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
# head
|
||||
self.head = Head(dim, out_dim, patch_size, eps)
|
||||
|
||||
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
d = dim // num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
],
|
||||
dim=1)
|
||||
|
||||
if model_type == 'i2v':
|
||||
self.img_emb = MLPProj(1280, dim)
|
||||
|
||||
# initialize weights
|
||||
self.init_weights()
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
self.gradient_checkpointing = value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
*args,
|
||||
**kwargs
|
||||
):
|
||||
# if kwargs.get('classify_mode', False) is True:
|
||||
# kwargs.pop('classify_mode')
|
||||
# return self._forward_classify(*args, **kwargs)
|
||||
# else:
|
||||
return self._forward(*args, **kwargs)
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
context,
|
||||
seq_len,
|
||||
classify_mode=False,
|
||||
concat_time_embeddings=False,
|
||||
register_tokens=None,
|
||||
cls_pred_branch=None,
|
||||
gan_ca_blocks=None,
|
||||
clip_fea=None,
|
||||
y=None,
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of input video tensors, each with shape [C_in, F, H, W]
|
||||
t (Tensor):
|
||||
Diffusion timesteps tensor of shape [B]
|
||||
context (List[Tensor]):
|
||||
List of text embeddings each with shape [L, C]
|
||||
seq_len (`int`):
|
||||
Maximum sequence length for positional encoding
|
||||
clip_fea (Tensor, *optional*):
|
||||
CLIP image features for image-to-video mode
|
||||
y (List[Tensor], *optional*):
|
||||
Conditional video inputs for image-to-video mode, same shape as x
|
||||
|
||||
Returns:
|
||||
List[Tensor]:
|
||||
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
if self.model_type == 'i2v':
|
||||
assert clip_fea is not None and y is not None
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
if self.freqs.device != device:
|
||||
self.freqs = self.freqs.to(device)
|
||||
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
||||
|
||||
# embeddings
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
dim=1) for u in x
|
||||
])
|
||||
|
||||
# time embeddings
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t).type_as(x))
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
||||
# assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
|
||||
# context
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat(
|
||||
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]))
|
||||
|
||||
if clip_fea is not None:
|
||||
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
context = torch.concat([context_clip, context], dim=1)
|
||||
|
||||
# arguments
|
||||
kwargs = dict(
|
||||
e=e0,
|
||||
seq_lens=seq_lens,
|
||||
grid_sizes=grid_sizes,
|
||||
freqs=self.freqs,
|
||||
context=context,
|
||||
context_lens=context_lens)
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs, **kwargs):
|
||||
return module(*inputs, **kwargs)
|
||||
return custom_forward
|
||||
|
||||
# TODO: Tune the number of blocks for feature extraction
|
||||
final_x = None
|
||||
if classify_mode:
|
||||
assert register_tokens is not None
|
||||
assert gan_ca_blocks is not None
|
||||
assert cls_pred_branch is not None
|
||||
|
||||
final_x = []
|
||||
registers = repeat(register_tokens(), "n d -> b n d", b=x.shape[0])
|
||||
# x = torch.cat([registers, x], dim=1)
|
||||
|
||||
gan_idx = 0
|
||||
for ii, block in enumerate(self.blocks):
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
x = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
x, **kwargs,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
x = block(x, **kwargs)
|
||||
|
||||
if classify_mode and ii in [13, 21, 29]:
|
||||
gan_token = registers[:, gan_idx: gan_idx + 1]
|
||||
final_x.append(gan_ca_blocks[gan_idx](x, gan_token))
|
||||
gan_idx += 1
|
||||
|
||||
if classify_mode:
|
||||
final_x = torch.cat(final_x, dim=1)
|
||||
if concat_time_embeddings:
|
||||
final_x = cls_pred_branch(torch.cat([final_x, 10 * e[:, None, :]], dim=1).view(final_x.shape[0], -1))
|
||||
else:
|
||||
final_x = cls_pred_branch(final_x.view(final_x.shape[0], -1))
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
# unpatchify
|
||||
x = self.unpatchify(x, grid_sizes)
|
||||
|
||||
if classify_mode:
|
||||
return torch.stack(x), final_x
|
||||
|
||||
return torch.stack(x)
|
||||
|
||||
def _forward_classify(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
context,
|
||||
seq_len,
|
||||
register_tokens,
|
||||
cls_pred_branch,
|
||||
clip_fea=None,
|
||||
y=None,
|
||||
):
|
||||
r"""
|
||||
Feature extraction through the diffusion model
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of input video tensors, each with shape [C_in, F, H, W]
|
||||
t (Tensor):
|
||||
Diffusion timesteps tensor of shape [B]
|
||||
context (List[Tensor]):
|
||||
List of text embeddings each with shape [L, C]
|
||||
seq_len (`int`):
|
||||
Maximum sequence length for positional encoding
|
||||
clip_fea (Tensor, *optional*):
|
||||
CLIP image features for image-to-video mode
|
||||
y (List[Tensor], *optional*):
|
||||
Conditional video inputs for image-to-video mode, same shape as x
|
||||
|
||||
Returns:
|
||||
List[Tensor]:
|
||||
List of video features with original input shapes [C_block, F, H / 8, W / 8]
|
||||
"""
|
||||
if self.model_type == 'i2v':
|
||||
assert clip_fea is not None and y is not None
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
if self.freqs.device != device:
|
||||
self.freqs = self.freqs.to(device)
|
||||
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
||||
|
||||
# embeddings
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
dim=1) for u in x
|
||||
])
|
||||
|
||||
# time embeddings
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t).type_as(x))
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
||||
# assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
|
||||
# context
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat(
|
||||
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]))
|
||||
|
||||
if clip_fea is not None:
|
||||
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
context = torch.concat([context_clip, context], dim=1)
|
||||
|
||||
# arguments
|
||||
kwargs = dict(
|
||||
e=e0,
|
||||
seq_lens=seq_lens,
|
||||
grid_sizes=grid_sizes,
|
||||
freqs=self.freqs,
|
||||
context=context,
|
||||
context_lens=context_lens)
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs, **kwargs):
|
||||
return module(*inputs, **kwargs)
|
||||
return custom_forward
|
||||
|
||||
# TODO: Tune the number of blocks for feature extraction
|
||||
for block in self.blocks[:16]:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
x = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
x, **kwargs,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
x = block(x, **kwargs)
|
||||
|
||||
# unpatchify
|
||||
x = self.unpatchify(x, grid_sizes, c=self.dim // 4)
|
||||
return torch.stack(x)
|
||||
|
||||
def unpatchify(self, x, grid_sizes, c=None):
|
||||
r"""
|
||||
Reconstruct video tensors from patch embeddings.
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
shape [B, 3] (3 dimensions correspond to F_patches, H_patches, W_patches)
|
||||
|
||||
Returns:
|
||||
List[Tensor]:
|
||||
Reconstructed video tensors with shape [C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_dim if c is None else c
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = torch.einsum('fhwpqrc->cfphqwr', u)
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
|
||||
def init_weights(self):
|
||||
r"""
|
||||
Initialize model parameters using Xavier initialization.
|
||||
"""
|
||||
|
||||
# basic init
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
# init embeddings
|
||||
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
|
||||
for m in self.text_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=.02)
|
||||
for m in self.time_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=.02)
|
||||
|
||||
# init output layer
|
||||
nn.init.zeros_(self.head.head.weight)
|
||||
@@ -1,513 +0,0 @@
|
||||
# Modified from transformers.models.t5.modeling_t5
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import logging
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .tokenizers import HuggingfaceTokenizer
|
||||
|
||||
__all__ = [
|
||||
'T5Model',
|
||||
'T5Encoder',
|
||||
'T5Decoder',
|
||||
'T5EncoderModel',
|
||||
]
|
||||
|
||||
|
||||
def fp16_clamp(x):
|
||||
if x.dtype == torch.float16 and torch.isinf(x).any():
|
||||
clamp = torch.finfo(x.dtype).max - 1000
|
||||
x = torch.clamp(x, min=-clamp, max=clamp)
|
||||
return x
|
||||
|
||||
|
||||
def init_weights(m):
|
||||
if isinstance(m, T5LayerNorm):
|
||||
nn.init.ones_(m.weight)
|
||||
elif isinstance(m, T5Model):
|
||||
nn.init.normal_(m.token_embedding.weight, std=1.0)
|
||||
elif isinstance(m, T5FeedForward):
|
||||
nn.init.normal_(m.gate[0].weight, std=m.dim**-0.5)
|
||||
nn.init.normal_(m.fc1.weight, std=m.dim**-0.5)
|
||||
nn.init.normal_(m.fc2.weight, std=m.dim_ffn**-0.5)
|
||||
elif isinstance(m, T5Attention):
|
||||
nn.init.normal_(m.q.weight, std=(m.dim * m.dim_attn)**-0.5)
|
||||
nn.init.normal_(m.k.weight, std=m.dim**-0.5)
|
||||
nn.init.normal_(m.v.weight, std=m.dim**-0.5)
|
||||
nn.init.normal_(m.o.weight, std=(m.num_heads * m.dim_attn)**-0.5)
|
||||
elif isinstance(m, T5RelativeEmbedding):
|
||||
nn.init.normal_(
|
||||
m.embedding.weight, std=(2 * m.num_buckets * m.num_heads)**-0.5)
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return 0.5 * x * (1.0 + torch.tanh(
|
||||
math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
|
||||
|
||||
|
||||
class T5LayerNorm(nn.Module):
|
||||
|
||||
def __init__(self, dim, eps=1e-6):
|
||||
super(T5LayerNorm, self).__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) +
|
||||
self.eps)
|
||||
if self.weight.dtype in [torch.float16, torch.bfloat16]:
|
||||
x = x.type_as(self.weight)
|
||||
return self.weight * x
|
||||
|
||||
|
||||
class T5Attention(nn.Module):
|
||||
|
||||
def __init__(self, dim, dim_attn, num_heads, dropout=0.1):
|
||||
assert dim_attn % num_heads == 0
|
||||
super(T5Attention, self).__init__()
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim_attn // num_heads
|
||||
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim_attn, bias=False)
|
||||
self.k = nn.Linear(dim, dim_attn, bias=False)
|
||||
self.v = nn.Linear(dim, dim_attn, bias=False)
|
||||
self.o = nn.Linear(dim_attn, dim, bias=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def forward(self, x, context=None, mask=None, pos_bias=None):
|
||||
"""
|
||||
x: [B, L1, C].
|
||||
context: [B, L2, C] or None.
|
||||
mask: [B, L2] or [B, L1, L2] or None.
|
||||
"""
|
||||
# check inputs
|
||||
context = x if context is None else context
|
||||
b, n, c = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.q(x).view(b, -1, n, c)
|
||||
k = self.k(context).view(b, -1, n, c)
|
||||
v = self.v(context).view(b, -1, n, c)
|
||||
|
||||
# attention bias
|
||||
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
|
||||
if pos_bias is not None:
|
||||
attn_bias += pos_bias
|
||||
if mask is not None:
|
||||
assert mask.ndim in [2, 3]
|
||||
mask = mask.view(b, 1, 1,
|
||||
-1) if mask.ndim == 2 else mask.unsqueeze(1)
|
||||
attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
|
||||
|
||||
# compute attention (T5 does not use scaling)
|
||||
attn = torch.einsum('binc,bjnc->bnij', q, k) + attn_bias
|
||||
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
|
||||
x = torch.einsum('bnij,bjnc->binc', attn, v)
|
||||
|
||||
# output
|
||||
x = x.reshape(b, -1, n * c)
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
|
||||
class T5FeedForward(nn.Module):
|
||||
|
||||
def __init__(self, dim, dim_ffn, dropout=0.1):
|
||||
super(T5FeedForward, self).__init__()
|
||||
self.dim = dim
|
||||
self.dim_ffn = dim_ffn
|
||||
|
||||
# layers
|
||||
self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
|
||||
self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
|
||||
self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x) * self.gate(x)
|
||||
x = self.dropout(x)
|
||||
x = self.fc2(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
|
||||
class T5SelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
dim_attn,
|
||||
dim_ffn,
|
||||
num_heads,
|
||||
num_buckets,
|
||||
shared_pos=True,
|
||||
dropout=0.1):
|
||||
super(T5SelfAttention, self).__init__()
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.dim_ffn = dim_ffn
|
||||
self.num_heads = num_heads
|
||||
self.num_buckets = num_buckets
|
||||
self.shared_pos = shared_pos
|
||||
|
||||
# layers
|
||||
self.norm1 = T5LayerNorm(dim)
|
||||
self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
||||
self.norm2 = T5LayerNorm(dim)
|
||||
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
|
||||
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
|
||||
num_buckets, num_heads, bidirectional=True)
|
||||
|
||||
def forward(self, x, mask=None, pos_bias=None):
|
||||
e = pos_bias if self.shared_pos else self.pos_embedding(
|
||||
x.size(1), x.size(1))
|
||||
x = fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
|
||||
x = fp16_clamp(x + self.ffn(self.norm2(x)))
|
||||
return x
|
||||
|
||||
|
||||
class T5CrossAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
dim_attn,
|
||||
dim_ffn,
|
||||
num_heads,
|
||||
num_buckets,
|
||||
shared_pos=True,
|
||||
dropout=0.1):
|
||||
super(T5CrossAttention, self).__init__()
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.dim_ffn = dim_ffn
|
||||
self.num_heads = num_heads
|
||||
self.num_buckets = num_buckets
|
||||
self.shared_pos = shared_pos
|
||||
|
||||
# layers
|
||||
self.norm1 = T5LayerNorm(dim)
|
||||
self.self_attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
||||
self.norm2 = T5LayerNorm(dim)
|
||||
self.cross_attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
||||
self.norm3 = T5LayerNorm(dim)
|
||||
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
|
||||
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
|
||||
num_buckets, num_heads, bidirectional=False)
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
mask=None,
|
||||
encoder_states=None,
|
||||
encoder_mask=None,
|
||||
pos_bias=None):
|
||||
e = pos_bias if self.shared_pos else self.pos_embedding(
|
||||
x.size(1), x.size(1))
|
||||
x = fp16_clamp(x + self.self_attn(self.norm1(x), mask=mask, pos_bias=e))
|
||||
x = fp16_clamp(x + self.cross_attn(
|
||||
self.norm2(x), context=encoder_states, mask=encoder_mask))
|
||||
x = fp16_clamp(x + self.ffn(self.norm3(x)))
|
||||
return x
|
||||
|
||||
|
||||
class T5RelativeEmbedding(nn.Module):
|
||||
|
||||
def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):
|
||||
super(T5RelativeEmbedding, self).__init__()
|
||||
self.num_buckets = num_buckets
|
||||
self.num_heads = num_heads
|
||||
self.bidirectional = bidirectional
|
||||
self.max_dist = max_dist
|
||||
|
||||
# layers
|
||||
self.embedding = nn.Embedding(num_buckets, num_heads)
|
||||
|
||||
def forward(self, lq, lk):
|
||||
device = self.embedding.weight.device
|
||||
# rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
|
||||
# torch.arange(lq).unsqueeze(1).to(device)
|
||||
rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \
|
||||
torch.arange(lq, device=device).unsqueeze(1)
|
||||
rel_pos = self._relative_position_bucket(rel_pos)
|
||||
rel_pos_embeds = self.embedding(rel_pos)
|
||||
rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(
|
||||
0) # [1, N, Lq, Lk]
|
||||
return rel_pos_embeds.contiguous()
|
||||
|
||||
def _relative_position_bucket(self, rel_pos):
|
||||
# preprocess
|
||||
if self.bidirectional:
|
||||
num_buckets = self.num_buckets // 2
|
||||
rel_buckets = (rel_pos > 0).long() * num_buckets
|
||||
rel_pos = torch.abs(rel_pos)
|
||||
else:
|
||||
num_buckets = self.num_buckets
|
||||
rel_buckets = 0
|
||||
rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
|
||||
|
||||
# embeddings for small and large positions
|
||||
max_exact = num_buckets // 2
|
||||
rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /
|
||||
math.log(self.max_dist / max_exact) *
|
||||
(num_buckets - max_exact)).long()
|
||||
rel_pos_large = torch.min(
|
||||
rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
|
||||
rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
|
||||
return rel_buckets
|
||||
|
||||
|
||||
class T5Encoder(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
vocab,
|
||||
dim,
|
||||
dim_attn,
|
||||
dim_ffn,
|
||||
num_heads,
|
||||
num_layers,
|
||||
num_buckets,
|
||||
shared_pos=True,
|
||||
dropout=0.1):
|
||||
super(T5Encoder, self).__init__()
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.dim_ffn = dim_ffn
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.num_buckets = num_buckets
|
||||
self.shared_pos = shared_pos
|
||||
|
||||
# layers
|
||||
self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
|
||||
else nn.Embedding(vocab, dim)
|
||||
self.pos_embedding = T5RelativeEmbedding(
|
||||
num_buckets, num_heads, bidirectional=True) if shared_pos else None
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.blocks = nn.ModuleList([
|
||||
T5SelfAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
|
||||
shared_pos, dropout) for _ in range(num_layers)
|
||||
])
|
||||
self.norm = T5LayerNorm(dim)
|
||||
|
||||
# initialize weights
|
||||
self.apply(init_weights)
|
||||
|
||||
def forward(self, ids, mask=None):
|
||||
x = self.token_embedding(ids)
|
||||
x = self.dropout(x)
|
||||
e = self.pos_embedding(x.size(1),
|
||||
x.size(1)) if self.shared_pos else None
|
||||
for block in self.blocks:
|
||||
x = block(x, mask, pos_bias=e)
|
||||
x = self.norm(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
|
||||
class T5Decoder(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
vocab,
|
||||
dim,
|
||||
dim_attn,
|
||||
dim_ffn,
|
||||
num_heads,
|
||||
num_layers,
|
||||
num_buckets,
|
||||
shared_pos=True,
|
||||
dropout=0.1):
|
||||
super(T5Decoder, self).__init__()
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.dim_ffn = dim_ffn
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.num_buckets = num_buckets
|
||||
self.shared_pos = shared_pos
|
||||
|
||||
# layers
|
||||
self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
|
||||
else nn.Embedding(vocab, dim)
|
||||
self.pos_embedding = T5RelativeEmbedding(
|
||||
num_buckets, num_heads, bidirectional=False) if shared_pos else None
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.blocks = nn.ModuleList([
|
||||
T5CrossAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
|
||||
shared_pos, dropout) for _ in range(num_layers)
|
||||
])
|
||||
self.norm = T5LayerNorm(dim)
|
||||
|
||||
# initialize weights
|
||||
self.apply(init_weights)
|
||||
|
||||
def forward(self, ids, mask=None, encoder_states=None, encoder_mask=None):
|
||||
b, s = ids.size()
|
||||
|
||||
# causal mask
|
||||
if mask is None:
|
||||
mask = torch.tril(torch.ones(1, s, s).to(ids.device))
|
||||
elif mask.ndim == 2:
|
||||
mask = torch.tril(mask.unsqueeze(1).expand(-1, s, -1))
|
||||
|
||||
# layers
|
||||
x = self.token_embedding(ids)
|
||||
x = self.dropout(x)
|
||||
e = self.pos_embedding(x.size(1),
|
||||
x.size(1)) if self.shared_pos else None
|
||||
for block in self.blocks:
|
||||
x = block(x, mask, encoder_states, encoder_mask, pos_bias=e)
|
||||
x = self.norm(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
|
||||
class T5Model(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
vocab_size,
|
||||
dim,
|
||||
dim_attn,
|
||||
dim_ffn,
|
||||
num_heads,
|
||||
encoder_layers,
|
||||
decoder_layers,
|
||||
num_buckets,
|
||||
shared_pos=True,
|
||||
dropout=0.1):
|
||||
super(T5Model, self).__init__()
|
||||
self.vocab_size = vocab_size
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.dim_ffn = dim_ffn
|
||||
self.num_heads = num_heads
|
||||
self.encoder_layers = encoder_layers
|
||||
self.decoder_layers = decoder_layers
|
||||
self.num_buckets = num_buckets
|
||||
|
||||
# layers
|
||||
self.token_embedding = nn.Embedding(vocab_size, dim)
|
||||
self.encoder = T5Encoder(self.token_embedding, dim, dim_attn, dim_ffn,
|
||||
num_heads, encoder_layers, num_buckets,
|
||||
shared_pos, dropout)
|
||||
self.decoder = T5Decoder(self.token_embedding, dim, dim_attn, dim_ffn,
|
||||
num_heads, decoder_layers, num_buckets,
|
||||
shared_pos, dropout)
|
||||
self.head = nn.Linear(dim, vocab_size, bias=False)
|
||||
|
||||
# initialize weights
|
||||
self.apply(init_weights)
|
||||
|
||||
def forward(self, encoder_ids, encoder_mask, decoder_ids, decoder_mask):
|
||||
x = self.encoder(encoder_ids, encoder_mask)
|
||||
x = self.decoder(decoder_ids, decoder_mask, x, encoder_mask)
|
||||
x = self.head(x)
|
||||
return x
|
||||
|
||||
|
||||
def _t5(name,
|
||||
encoder_only=False,
|
||||
decoder_only=False,
|
||||
return_tokenizer=False,
|
||||
tokenizer_kwargs={},
|
||||
dtype=torch.float32,
|
||||
device='cpu',
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert not (encoder_only and decoder_only)
|
||||
|
||||
# params
|
||||
if encoder_only:
|
||||
model_cls = T5Encoder
|
||||
kwargs['vocab'] = kwargs.pop('vocab_size')
|
||||
kwargs['num_layers'] = kwargs.pop('encoder_layers')
|
||||
_ = kwargs.pop('decoder_layers')
|
||||
elif decoder_only:
|
||||
model_cls = T5Decoder
|
||||
kwargs['vocab'] = kwargs.pop('vocab_size')
|
||||
kwargs['num_layers'] = kwargs.pop('decoder_layers')
|
||||
_ = kwargs.pop('encoder_layers')
|
||||
else:
|
||||
model_cls = T5Model
|
||||
|
||||
# init model
|
||||
with torch.device(device):
|
||||
model = model_cls(**kwargs)
|
||||
|
||||
# set device
|
||||
model = model.to(dtype=dtype, device=device)
|
||||
|
||||
# init tokenizer
|
||||
if return_tokenizer:
|
||||
from .tokenizers import HuggingfaceTokenizer
|
||||
tokenizer = HuggingfaceTokenizer(f'google/{name}', **tokenizer_kwargs)
|
||||
return model, tokenizer
|
||||
else:
|
||||
return model
|
||||
|
||||
|
||||
def umt5_xxl(**kwargs):
|
||||
cfg = dict(
|
||||
vocab_size=256384,
|
||||
dim=4096,
|
||||
dim_attn=4096,
|
||||
dim_ffn=10240,
|
||||
num_heads=64,
|
||||
encoder_layers=24,
|
||||
decoder_layers=24,
|
||||
num_buckets=32,
|
||||
shared_pos=False,
|
||||
dropout=0.1)
|
||||
cfg.update(**kwargs)
|
||||
return _t5('umt5-xxl', **cfg)
|
||||
|
||||
|
||||
class T5EncoderModel:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_len,
|
||||
dtype=torch.bfloat16,
|
||||
device=torch.cuda.current_device(),
|
||||
checkpoint_path=None,
|
||||
tokenizer_path=None,
|
||||
shard_fn=None,
|
||||
):
|
||||
self.text_len = text_len
|
||||
self.dtype = dtype
|
||||
self.device = device
|
||||
self.checkpoint_path = checkpoint_path
|
||||
self.tokenizer_path = tokenizer_path
|
||||
|
||||
# init model
|
||||
model = umt5_xxl(
|
||||
encoder_only=True,
|
||||
return_tokenizer=False,
|
||||
dtype=dtype,
|
||||
device=device).eval().requires_grad_(False)
|
||||
logging.info(f'loading {checkpoint_path}')
|
||||
model.load_state_dict(torch.load(checkpoint_path, map_location='cpu'))
|
||||
self.model = model
|
||||
if shard_fn is not None:
|
||||
self.model = shard_fn(self.model, sync_module_states=False)
|
||||
else:
|
||||
self.model.to(self.device)
|
||||
# init tokenizer
|
||||
self.tokenizer = HuggingfaceTokenizer(
|
||||
name=tokenizer_path, seq_len=text_len, clean='whitespace')
|
||||
|
||||
def __call__(self, texts, device):
|
||||
ids, mask = self.tokenizer(
|
||||
texts, return_mask=True, add_special_tokens=True)
|
||||
ids = ids.to(device)
|
||||
mask = mask.to(device)
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
context = self.model(ids, mask)
|
||||
return [u[:v] for u, v in zip(context, seq_lens)]
|
||||
@@ -1,82 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import html
|
||||
import string
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
__all__ = ['HuggingfaceTokenizer']
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
def canonicalize(text, keep_punctuation_exact_string=None):
|
||||
text = text.replace('_', ' ')
|
||||
if keep_punctuation_exact_string:
|
||||
text = keep_punctuation_exact_string.join(
|
||||
part.translate(str.maketrans('', '', string.punctuation))
|
||||
for part in text.split(keep_punctuation_exact_string))
|
||||
else:
|
||||
text = text.translate(str.maketrans('', '', string.punctuation))
|
||||
text = text.lower()
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
return text.strip()
|
||||
|
||||
|
||||
class HuggingfaceTokenizer:
|
||||
|
||||
def __init__(self, name, seq_len=None, clean=None, **kwargs):
|
||||
assert clean in (None, 'whitespace', 'lower', 'canonicalize')
|
||||
self.name = name
|
||||
self.seq_len = seq_len
|
||||
self.clean = clean
|
||||
|
||||
# init tokenizer
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
|
||||
self.vocab_size = self.tokenizer.vocab_size
|
||||
|
||||
def __call__(self, sequence, **kwargs):
|
||||
return_mask = kwargs.pop('return_mask', False)
|
||||
|
||||
# arguments
|
||||
_kwargs = {'return_tensors': 'pt'}
|
||||
if self.seq_len is not None:
|
||||
_kwargs.update({
|
||||
'padding': 'max_length',
|
||||
'truncation': True,
|
||||
'max_length': self.seq_len
|
||||
})
|
||||
_kwargs.update(**kwargs)
|
||||
|
||||
# tokenization
|
||||
if isinstance(sequence, str):
|
||||
sequence = [sequence]
|
||||
if self.clean:
|
||||
sequence = [self._clean(u) for u in sequence]
|
||||
ids = self.tokenizer(sequence, **_kwargs)
|
||||
|
||||
# output
|
||||
if return_mask:
|
||||
return ids.input_ids, ids.attention_mask
|
||||
else:
|
||||
return ids.input_ids
|
||||
|
||||
def _clean(self, text):
|
||||
if self.clean == 'whitespace':
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
elif self.clean == 'lower':
|
||||
text = whitespace_clean(basic_clean(text)).lower()
|
||||
elif self.clean == 'canonicalize':
|
||||
text = canonicalize(basic_clean(text))
|
||||
return text
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user