Compare commits

..
Author SHA1 Message Date
JerryZhou54 4ea813265b Resolve timestep mismatch between dit forward and pred_noise_to_pred_video 2025-09-17 03:51:05 +00:00
SolitaryThinker 7e2c8f8d49 comment out vmoba 2025-09-16 04:27:42 +00:00
SolitaryThinker 421dcbb50b Merge branch 'wei/dit_debug' into will/ode_init 2025-09-16 02:14:14 +00:00
SolitaryThinker b9662bf882 lmdb datasets 2025-09-16 02:06:43 +00:00
JerryZhou54 140fb9f20e Fix test for forward_train 2025-09-15 23:23:33 +00:00
JerryZhou54 2078876b98 Ensure 0 numerical diff for forward_train 2025-09-15 22:49:52 +00:00
JerryZhou54 a953f46bd6 Add test for _forward_train 2025-09-15 22:37:50 +00:00
JerryZhou54 adae957008 Fix numerical diff between causal_wanvideo.py and SF's causal wan 2025-09-14 08:09:57 +00:00
SolitaryThinker fa40553afb t2v to i2v finetune
checkpoint ode

checkpoint

fix t2v to i2v

lint

chekpt

checkpoint

hacked but working

ode_init scripts

WIP fixing time embedding

WIP fixing time embedding

checkpoint

update

fix

revert

revert

revert

update

update

visualize
2025-09-14 02:23:42 +00:00
RandNMR73 93ebd15a0d text preprocessing ready 2025-09-10 11:12:48 +00:00
JerryZhou54 1110474065 checkpoint 2025-09-10 08:57:54 +00:00
JerryZhou54 80baffd540 Enable timestep warping & using SelfForcing scheduler 2025-09-09 23:30:55 +00:00
JerryZhou54 918180048e Stop backprop through kv_cache 2025-09-09 10:02:03 +00:00
RandNMR73 b7dbd7cb9e new branch 2025-09-09 10:02:00 +00:00
RandNMR73 71159b6416 inference works after changes added 2025-09-09 10:01:31 +00:00
40 changed files with 687 additions and 1100 deletions
+1 -3
View File
@@ -20,7 +20,5 @@ setup(
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[
"flash-attn >= 2.7.1",
]
install_requires=[]
)
+2 -10
View File
@@ -6,16 +6,8 @@ import time
import os
import torch
from typing import Tuple
try:
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
from functools import lru_cache
from einops import rearrange
+8 -10
View File
@@ -26,22 +26,20 @@ def main():
sampling_param.image_path = "test.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
"A girl is packing a suitcase when stuff suddently starts flying around the room."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
# prompt2 = (
# "A majestic lion strides across the golden savanna, its powerful frame "
# "glistening under the warm afternoon sun. The tall grass ripples gently in "
# "the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
# "embodying the raw energy of the wild. Low angle, steady tracking shot, "
# "cinematic.")
# video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__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[@]}"
@@ -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,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,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[@]}"
@@ -3,8 +3,8 @@
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/"
DATA_MERGE_PATH="data/crush-smol_single/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_single/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
@@ -15,7 +15,6 @@ torchrun --nproc_per_node=$GPU_NUM \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
@@ -1,5 +1,14 @@
{
"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,
@@ -19,52 +28,7 @@
"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. ",
"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,
@@ -28,4 +28,4 @@
"num_frames": 77
}
]
}
}
+4 -4
View File
@@ -5,6 +5,7 @@ from dataclasses import dataclass
import torch
from einops import rearrange
from flash_attn.bert_padding import pad_input
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
@@ -133,8 +134,6 @@ class VMOBAAttentionImpl(AttentionImpl):
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
from flash_attn.bert_padding import pad_input
self.pad_input = pad_input
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
@@ -170,6 +169,7 @@ class VMOBAAttentionImpl(AttentionImpl):
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
# torch.distributed.breakpoint()
query, chunk_size = process_moba_input(query,
attn_metadata.patch_resolution,
moba_chunk_size)
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = process_moba_output(hidden_states,
attn_metadata.patch_resolution,
moba_chunk_size)
+2 -3
View File
@@ -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,
-1
View File
@@ -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:
+2 -3
View File
@@ -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"
]
-5
View File
@@ -51,7 +51,6 @@ class PipelineConfig:
# 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)
@@ -88,10 +87,6 @@ class PipelineConfig:
# DMD parameters
dmd_denoising_steps: list[int] | None = field(default=None)
# Wan2.2 TI2V parameters
ti2v_task: bool = False
boundary_ratio: float | None = None
# Compilation
# enable_torch_compile: bool = False
+4 -5
View File
@@ -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, SelfForcingWanT2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -48,7 +48,6 @@ 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,8 +61,8 @@ 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,
"stepvideo": StepVideoT2VConfig,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
+1 -1
View File
@@ -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,
-5
View File
@@ -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
+5 -6
View File
@@ -1141,11 +1141,10 @@ class TrainingArgs(FastVideoArgs):
"--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")
parser.add_argument("--context-noise",
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
return parser
@@ -1153,4 +1152,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(",")]
+2 -5
View File
@@ -556,8 +556,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]
@@ -652,11 +650,10 @@ 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 self._forward_train(*args, **kwargs)
return noise_pred
def unpatchify(self, x, grid_sizes):
r"""
+3 -13
View File
@@ -172,13 +172,8 @@ class WanT2VCrossAttention(WanSelfAttention):
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
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)
@@ -364,12 +359,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)
+4 -10
View File
@@ -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)
@@ -450,8 +449,6 @@ class TransformerLoader(ComponentLoader):
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 +468,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 +477,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
+2 -3
View File
@@ -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,
@@ -64,15 +64,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)
+3 -22
View File
@@ -171,6 +171,9 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
# Convert to double following Self-Forcing
# https://github.com/guandeh17/Self-Forcing/blob/main/utils/wan_wrapper.py#L184
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
@@ -180,25 +183,3 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
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)
@@ -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
@@ -21,35 +21,147 @@ 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.dataset import getdataset
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory
from fastvideo.distributed import get_local_torch_device
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,
ImageVAEEncodingStage,
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 FlowMatchScheduler:
order = 1
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):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.inverse_timesteps = inverse_timesteps
self.extra_one_step = extra_one_step
self.reverse_sigmas = reverse_sigmas
self.set_timesteps(num_inference_steps)
def set_timesteps(self,
num_inference_steps=100,
denoising_strength=1.0,
training=False,
device=None):
sigma_start = self.sigma_min + \
(self.sigma_max - self.sigma_min) * denoising_strength
if self.extra_one_step:
self.sigmas = torch.linspace(sigma_start, self.sigma_min,
num_inference_steps + 1)[:-1]
else:
self.sigmas = torch.linspace(sigma_start, self.sigma_min,
num_inference_steps)
if self.inverse_timesteps:
self.sigmas = torch.flip(self.sigmas, dims=[0])
self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas)
if self.reverse_sigmas:
self.sigmas = 1 - self.sigmas
self.timesteps = self.sigmas * self.num_train_timesteps
if training:
x = self.timesteps
y = torch.exp(
-2 * ((x - num_inference_steps / 2) / num_inference_steps)**2)
y_shifted = y - y.min()
bsmntw_weighing = y_shifted * \
(num_inference_steps / y_shifted.sum())
self.linear_timesteps_weights = bsmntw_weighing
def step(self,
model_output,
timestep,
sample,
to_final=False,
return_dict=False,
**kwargs):
assert return_dict is False
assert kwargs == {}
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
logger.info('step timestep: %s', timestep)
logger.info('step timestep: %s', timestep.shape)
# timestep is [num_frames]
# timestep_id = torch.argmin(
# (self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
# assert timestep.ndim == 1
# assert timestep.shape[0] == 1
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
sigma_ = 1 if (self.inverse_timesteps or self.reverse_sigmas) else 0
else:
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
prev_sample = sample + model_output * (sigma_ - sigma)
return (prev_sample, )
def scale_model_input(self, sample: torch.Tensor, *args,
**kwargs) -> torch.Tensor:
"""
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
current timestep.
Args:
sample (`torch.Tensor`):
The input sample.
Returns:
`torch.Tensor`:
A scaled input sample.
"""
return sample
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B, C, H, W]
- noise: the noise with shape [B, C, H, W]
- timestep: the timestep with shape [B]
Output: the corrupted latent with shape [B, C, H, W]
"""
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.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)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
def training_weight(self, timestep):
timestep_id = torch.argmin(
(self.timesteps - timestep.to(self.timesteps.device)).abs())
weights = self.linear_timesteps_weights[timestep_id]
return weights
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
@@ -58,31 +170,28 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
]
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 get_schema_fields(self):
"""Get the schema fields for ODE Trajectory pipeline."""
return [f.name for f in pyarrow_schema_ode_trajectory]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
fastvideo_args.pipeline_config.flow_shift = 5
logger.info('WTF flow_shift: %s',
fastvideo_args.pipeline_config.flow_shift)
assert fastvideo_args.pipeline_config.flow_shift == 5
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
# self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
# shift=fastvideo_args.pipeline_config.flow_shift)
self.modules["scheduler"] = FlowMatchScheduler(
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
logger.info('WTF scheduler timesteps: %s',
self.modules["scheduler"].timesteps)
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
@@ -91,6 +200,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="vae_encoding_stage",
stage=ImageVAEEncodingStage(
vae=self.get_module("vae"), ))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
@@ -101,49 +213,56 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
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."""
def preprocess_video_and_text_and_trajectory(self,
fastvideo_args: FastVideoArgs,
args):
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
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
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)
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][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
]
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
get_local_torch_device())).mean
# Get extra features if needed
extra_features = self.get_extra_features(
valid_data, fastvideo_args)
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)
logger.info(f"===== batch_captions: {batch_captions}")
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
@@ -152,25 +271,43 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
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)
# # Get sequence lengths from attention masks (number of 1s)
# seq_lens = prompt_attention_mask.sum(dim=1)
negative_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
# non_padded_embeds = []
# non_padded_masks = []
# # Process each item in the batch
# for i in range(prompt_embeds.size(0)):
# seq_len = seq_lens[i].item()
# # Slice the embeddings and masks to keep only non-padding parts
# non_padded_embeds.append(prompt_embeds[i, :seq_len])
# non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
# prompt_embeds = non_padded_embeds
# prompt_attention_masks = non_padded_masks
# prompt_embeds = prompt_embeds
# logger.info(f"===== prompt_embeds: {prompt_embeds[0].shape}")
# logger.info(f"===== prompt_attention_masks: {prompt_attention_masks[0].shape}")
sampling_params = SamplingParam.from_pretrained(args.model_path)
# 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,
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0]
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0]
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
@@ -178,232 +315,306 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
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), )
logger.info("what")
logger.info(f"===== prompt_embed: {prompt_embed.shape}")
logger.info(
f"===== prompt_attention_mask: {prompt_attention_mask.shape}"
)
# Collect the trajectory data
batch = ForwardBatch(
**shallow_asdict(sampling_params),
# data_type="video",
# seed=args.seed,
# prompt=batch_captions[i],
# prompt_embeds=[prompt_embed],
# prompt_attention_mask=[prompt_attention_mask],
# height=args.max_height,
# width=args.max_width,
# num_frames=81,
# fps=args.train_fps,
# return_trajectory_latents=True,
# guidance_scale=3.0,
# do_classifier_free_guidance=True,
)
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.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.num_inference_steps = 48
# batch.num_frames = 81
batch.fps = args.train_fps
batch.guidance_scale = 3.0
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
# fastvideo_args.pipeline_config.ti2v_task = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
# result_batch = self.prompt_encoding_stage(result_batch, fastvideo_args)
# result_batch = self.vae_encoding_stage(result_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.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
# trajectory_latents = result_batch.trajectory_latents
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]]
extra_features["trajectory_latents"] = trajectory_latents
extra_features["trajectory_timesteps"] = trajectory_timesteps
logger.info(
f"===== trajectory_latents: {trajectory_latents[0].shape}")
logger.info(
f"===== trajectory_latents len: {len(trajectory_latents)}")
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
logger.info(
f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
# 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,
if batch.return_trajectory_decoded:
logger.info("===== SAVING TRAJECTORY DECODED")
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
logger.info(
f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}"
)
self.dataset_writer.append_table(table)
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# assert False
# Prepare batch data for Parquet dataset
batch_data = []
logger.info("Collected batch with %s samples", len(table))
# 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:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
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
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
# 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)
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
logger.info(f"===== key: {key}")
if isinstance(value, torch.Tensor):
logger.info(f"===== value: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().numpy(
)
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
logger.info(
f"===== value in list: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
else:
logger.info("===== value in list: not tensor")
sample_extra_features[key] = value[idx]
# logger.info(f"===== value: not tensor")
# sample_extra_features[key] = value[idx]
# Create record for Parquet dataset
record = self.create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("vae").to(get_local_torch_device())
# generator = torch.Generator(device=get_local_torch_device(), seed=42)
generator = torch.Generator("cpu").manual_seed(42)
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
unprocessed_images = []
pil_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
pil_images.append(frame_pil)
# processed_img = self.get_module("image_processor")(
# images=frame_pil, return_tensors="pt")
unprocessed_images.append(frame_pil)
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in unprocessed_images:
latent = self.vae_encoding_stage.encode_image(
frame, height, width, fastvideo_args, generator)
video_conditions.append(latent)
features["image_condition_latents"] = video_conditions
features["pil_images"] = pil_images
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "image_condition_latents" in extra_features:
image_condition_latents = extra_features["image_condition_latents"]
record.update({
"image_condition_latents_bytes":
image_condition_latents.tobytes(),
"image_condition_latents_shape":
list(image_condition_latents.shape),
"image_condition_latents_dtype":
str(image_condition_latents.dtype),
})
else:
record.update({
"image_condition_latents_bytes": b"",
"image_condition_latents_shape": [],
"image_condition_latents_dtype": "",
})
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes":
trajectory_latents.tobytes(),
"trajectory_latents_shape":
list(trajectory_latents.shape),
"trajectory_latents_dtype":
str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes":
trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape":
list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype":
str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
@@ -417,7 +628,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
train_dataset = gettextdataset(args)
train_dataset = getdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
@@ -437,7 +648,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
# 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)
self.preprocess_video_and_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_ODE_Trajectory
@@ -56,10 +56,6 @@ def main(args) -> None:
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 +101,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)
+9 -34
View File
@@ -53,20 +53,7 @@ class DecodingStage(PipelineStage):
@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
"""
"""Decode latents into pixel space."""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
@@ -116,26 +103,12 @@ 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
@@ -155,13 +128,15 @@ class DecodingStage(PipelineStage):
# decode trajectory latents if needed
if batch.return_trajectory_decoded:
batch.trajectory_decoded = []
logger.info(f"batch.trajectory_latents.shape: {batch.trajectory_latents.shape}")
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]
# bathc.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
logger.info(f"cur_latent.shape: {cur_latent.shape}")
cur_timestep = batch.trajectory_timesteps[idx]
logger.info("decoding trajectory latent for timestep: %s",
cur_timestep)
logger.info(
f"decoding trajectory latent for timestep: {cur_timestep}")
decoded_frames = self.decode(cur_latent, fastvideo_args)
batch.trajectory_decoded.append(decoded_frames.cpu().float())
+10 -11
View File
@@ -205,14 +205,14 @@ 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_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_timestep is not None:
logger.info("Overriding boundary timestep from %s to %s",
boundary_timestep, batch.boundary_timestep)
boundary_timestep = batch.boundary_timestep
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
boundary_timestep *= self.scheduler.num_train_timesteps
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
@@ -460,6 +460,7 @@ class DenoisingStage(PipelineStage):
# save trajectory latents if needed
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
# trajectory_latents.append(latents.cpu())
trajectory_latents.append(latents)
# Update progress bar
@@ -473,16 +474,14 @@ class DenoisingStage(PipelineStage):
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:
# logger.info("before stack trajectory_latents.shape: %s", trajectory_latents[0].shape)
logger.info("after stack trajectory_latents.shape: %s", trajectory_tensor.shape)
trajectory_tensor = trajectory_tensor.to(
get_local_torch_device())
trajectory_tensor = sequence_model_parallel_all_gather(
@@ -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:
+20 -47
View File
@@ -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}"
@@ -22,7 +22,6 @@ 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(
@@ -33,21 +32,19 @@ 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"
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, dit_forward_precision=precision_str))
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args)
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
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)
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
@@ -78,35 +75,29 @@ def test_ori_wan_transformer():
# Create identical inputs for both models
batch_size = 1
text_seq_len = 120
seq_len = math.ceil((104 * 60) /
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,
21,
16,
60,
104,
21,
160,
90,
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)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([995.7627], device=device, dtype=precision)
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(
data_type="dummy",
@@ -114,19 +105,19 @@ def test_ori_wan_transformer():
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states.permute(0, 2, 1, 3, 4),
x=hidden_states,
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),
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep).permute(0, 2, 1, 3, 4)
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}"
@@ -137,8 +128,6 @@ def test_ori_wan_transformer():
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()}"
@@ -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
+77 -126
View File
@@ -3,13 +3,12 @@ import sys
from copy import deepcopy
from typing import cast
import numpy as np
import torch
import torch.nn.functional as F
import numpy as np
import wandb
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
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
@@ -18,8 +17,8 @@ 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.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
@@ -122,75 +121,55 @@ class ODEInitTrainingPipeline(TrainingPipeline):
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)
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']
# 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 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)
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
# 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
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
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
def _get_timestep(self,
min_timestep: int,
@@ -222,8 +201,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
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]]:
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str, torch.Tensor]]:
latent_vis_dict = {}
device = get_local_torch_device()
target_latent = traj_latents[:, -1]
@@ -255,12 +233,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# 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}")
relevant_traj_latents = torch.index_select(
traj_latents,
dim=1,
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
# assert relevant_traj_latents.shape[0] == 1
indexes = self._get_timestep( # [B, num_frames]
@@ -303,46 +279,34 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# 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()
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
model_dtype = next(self.transformer.parameters()).dtype
input_kwargs = {
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep,
"timestep": timestep.to(device, dtype=model_dtype),
"encoder_attention_mask": encoder_attention_mask,
"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)
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
# 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]
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()
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])
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(
@@ -387,8 +351,8 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
traj_latents, training_batch.current_timestep, text_embeds, text_attention_mask)
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
training_batch.latent_vis_dict.update(latent_vis_dict)
mask = t != 0
@@ -403,28 +367,16 @@ class ODEInitTrainingPipeline(TrainingPipeline):
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
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)
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()
self.lr_scheduler.step()
if grad_norm is None:
grad_value = 0.0
@@ -446,11 +398,9 @@ class ODEInitTrainingPipeline(TrainingPipeline):
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
assert latent_key in latents_vis_dict and latents_vis_dict[latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
pixel_latent = self.validation_pipeline.decoding_stage.decode(latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
@@ -464,6 +414,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
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']
+6 -5
View File
@@ -44,8 +44,7 @@ from fastvideo.training.training_utils import (
shard_latents_across_sp)
# from fastvideo.utils import (is_vmoba_available, is_vsa_available,
# set_random_seed, shallow_asdict)
from fastvideo.utils import (is_vsa_available,
set_random_seed, shallow_asdict)
from fastvideo.utils import is_vsa_available, set_random_seed, shallow_asdict
import wandb # isort: skip
@@ -125,7 +124,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# 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,
@@ -735,8 +734,10 @@ 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")
raise NotImplementedError(
"Visualize intermediate latents is not implemented for training pipeline"
)
+23 -8
View File
@@ -470,11 +470,13 @@ def load_distillation_checkpoint(generator_transformer,
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)
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))
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",
@@ -1321,6 +1323,7 @@ class EMA_FSDP:
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
@@ -1346,7 +1349,10 @@ class EMA_FSDP:
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()}
self.shadow = {
k: v.detach().clone().float().cpu()
for k, v in cpu_state.items()
}
else:
self.shadow = {}
return
@@ -1387,7 +1393,10 @@ class EMA_FSDP:
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()
} 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]):
@@ -1408,6 +1417,7 @@ class EMA_FSDP:
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
@@ -1415,7 +1425,9 @@ class EMA_FSDP:
def __enter__(self):
if self.ema.mode != "local_shard":
raise RuntimeError("EMA apply_to_model is only supported for 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:
@@ -1425,14 +1437,17 @@ class EMA_FSDP:
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)
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))
p_local.copy_(
ema_cpu.to(dtype=p_local.dtype,
device=p_local.device))
return self.module
def __exit__(self, exc_type, exc, tb):
@@ -1450,4 +1465,4 @@ class EMA_FSDP:
return False
def apply_to_model(self, module):
return EMA_FSDP._ApplyEMACtx(self, module)
return EMA_FSDP._ApplyEMACtx(self, module)
+1 -6
View File
@@ -27,14 +27,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
@@ -896,8 +892,7 @@ def best_output_size(w, h, dw, dh, expected_area):
return ow2, oh2
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor],
output_path: str, fps: int):
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 = []