Compare commits

...
Author SHA1 Message Date
SolitaryThinker 48f9690c47 load_video 2025-06-24 10:46:20 -07:00
SolitaryThinker 32df3bcca3 update i2v script 2025-06-24 03:20:09 -07:00
SolitaryThinker b8c81191b6 improve script format 2025-06-24 03:01:55 -07:00
SolitaryThinker e506c4074f remove print and enable first val 2025-06-24 01:55:05 -07:00
SolitaryThinker 1b3914da84 update scripts 2025-06-22 21:09:30 -07:00
SolitaryThinker f471fd3f02 i2v working 2025-06-22 20:59:37 -07:00
SolitaryThinker 2e6d5c5304 t2v working again 2025-06-22 19:45:44 -07:00
SolitaryThinker 5db34184c7 t2v example 2025-06-22 21:53:05 +00:00
SolitaryThinker 37252bf62c f 2025-06-22 11:58:18 +00:00
SolitaryThinker e9263f7d2b update 2025-06-22 04:20:34 -07:00
SolitaryThinker 93afb86c20 update 2025-06-22 04:19:10 -07:00
SolitaryThinker 0d944ba9c1 slrm 2025-06-22 04:09:09 -07:00
SolitaryThinker 3d78604281 update path 2025-06-22 03:54:39 -07:00
SolitaryThinker 0694b0c5eb exmaple scripts 2025-06-22 03:44:57 -07:00
SolitaryThinker b6c5644d40 cleanup 2025-06-22 03:07:37 -07:00
SolitaryThinker a9089fa358 fix pil image 2025-06-22 02:29:51 -07:00
SolitaryThinker 6079a98fd7 i2v preprocess 2025-06-21 19:18:39 -07:00
SolitaryThinker 65c0fcb633 checkpoint 2025-06-21 18:02:43 -07:00
SolitaryThinker 82e3641264 checkpoint 2025-06-21 18:00:29 -07:00
30 changed files with 1749 additions and 245 deletions
@@ -0,0 +1,88 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 8
--tp_size 8
--hsdp_replicate_dim 1
--hsdp_shard_dim 8
)
# Model arguments
model_args=(
--model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
--pretrained_model_name_or_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--cfg 0.0
--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/v1/training/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,97 @@
#!/bin/bash
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/crush-smol_processed_i2v/combined_parquet_dataset
VALIDATION_DIR=data/crush-smol_processed_i2v/validation_parquet_dataset
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/v1/training/wan_i2v_training_pipeline.py\
--model_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 16 \
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
--hsdp_shard_dim $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
--gradient_accumulation_steps=2\
--max_train_steps=10000 \
--learning_rate=5e-5\
--mixed_precision="bf16"\
--checkpointing_steps=11000 \
--validation_steps 100\
--validation_sampling_steps "40" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
--tracker_project_name wan_i2v_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 1e-4 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -3,8 +3,8 @@
{
"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": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 50,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
@@ -12,8 +12,8 @@
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 50,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
@@ -21,8 +21,8 @@
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 50,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
@@ -0,0 +1,88 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_finetune"
--output_dir "outputs/wan_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim $NUM_GPUS \
)
# 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_preprocessed_path $VALIDATION_DIR
--validation_steps 100
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--cfg 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,98 @@
#!/bin/bash
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/v1/training/wan_training_pipeline.py\
--model_path $MODEL_PATH \
--inference_mode False\
--pretrained_model_name_or_path $MODEL_PATH \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 8 \
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
--hsdp_shard_dim $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
--gradient_accumulation_steps=1\
--max_train_steps=10000 \
--learning_rate=5e-5\
--mixed_precision="bf16"\
--checkpointing_steps=11000 \
--validation_steps 100\
--validation_sampling_steps "40" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
--tracker_project_name wan_i2v_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 1e-4 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -0,0 +1,13 @@
{
"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": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "t2v"
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
+5 -8
View File
@@ -29,14 +29,11 @@ def getdataset(args) -> VideoCaptionMergedDataset:
*resize_topcrop,
norm_fun,
])
if args.dataset == "t2v":
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop)
raise NotImplementedError(args.dataset)
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop)
__all__ = [
+7 -34
View File
@@ -26,15 +26,17 @@ pyarrow_schema_i2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -50,13 +52,6 @@ pyarrow_schema_i2v = pa.schema([
pyarrow_schema_i2v_validation = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
@@ -64,11 +59,6 @@ pyarrow_schema_i2v_validation = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
@@ -106,11 +96,6 @@ pyarrow_schema_t2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -126,13 +111,6 @@ pyarrow_schema_t2v = pa.schema([
pyarrow_schema_t2v_validation = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
@@ -140,11 +118,6 @@ pyarrow_schema_t2v_validation = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -4,6 +4,7 @@ import random
from typing import Dict, List, Tuple
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
import tqdm
@@ -70,10 +71,12 @@ class LatentsParquetIterStyleDataset(IterableDataset):
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32):
read_batch_size: int = 32,
parquet_schema: pa.Schema = None):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.parquet_schema = parquet_schema
self.cfg_rate = cfg_rate
self.text_padding_length = text_padding_length
self.seed = seed
@@ -3,6 +3,7 @@ import os
import pickle
from typing import Any, Dict, List, Tuple
import pyarrow as pa
import pyarrow.parquet as pq
# Torch in general
import torch
@@ -11,7 +12,7 @@ import tqdm
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
@@ -185,12 +186,14 @@ class LatentsParquetMapStyleDataset(Dataset):
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), "text_embedding"]
keys = [("vae_latent", "latent"), "text_embedding", "clip_feature",
"first_frame_latent", "pil_image"]
def __init__(
self,
path: str,
batch_size: int,
parquet_schema: pa.Schema,
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
@@ -200,6 +203,7 @@ class LatentsParquetMapStyleDataset(Dataset):
super().__init__()
self.path = path
self.cfg_rate = cfg_rate
self.parquet_schema = parquet_schema
if cfg_rate > 0.0:
raise ValueError(
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
@@ -209,15 +213,6 @@ class LatentsParquetMapStyleDataset(Dataset):
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
self.batch = batch_size
self.text_padding_length = text_padding_length
self._cols = [
"vae_latent_bytes",
"vae_latent_shape",
"text_embedding_bytes",
"text_embedding_shape",
"text_embedding_dtype",
"height",
"width",
]
self.sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=sum(self.lengths),
@@ -232,7 +227,7 @@ class LatentsParquetMapStyleDataset(Dataset):
len(self.parquet_files), sum(self.lengths))
def get_validation_negative_prompt(
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
self) -> tuple[torch.Tensor, torch.Tensor, str]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
@@ -246,19 +241,22 @@ class LatentsParquetMapStyleDataset(Dataset):
row_dict = read_row_from_parquet_file([file_path], row_idx,
[self.lengths[0]])
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
[row_dict], self.text_padding_length, self.keys)
all_latents, all_embs, all_masks, caption_text = all_latents_list[
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
# add batch dimension
if len(all_embs.shape) == 2:
all_embs = all_embs.unsqueeze(0)
if len(all_masks.shape) == 1:
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
return all_latents, all_embs, all_masks, caption_text
batch = collate_rows_from_parquet_schema([row_dict],
self.parquet_schema,
self.text_padding_length)
negative_prompt = batch['info_list'][0]['prompt']
negative_prompt_embedding = batch['text_embedding']
negative_prompt_attention_mask = batch['text_attention_mask']
if len(negative_prompt_embedding.shape) == 2:
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
if len(negative_prompt_attention_mask.shape) == 1:
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
0).unsqueeze(0)
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
# PyTorch calls this ONLY because the batch_sampler yields a list
def __getitems__(self, indices: List[int]):
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
"""
Batch fetch using read_row_from_parquet_file for each index.
"""
@@ -267,9 +265,12 @@ class LatentsParquetMapStyleDataset(Dataset):
for idx in indices
]
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
rows, self.text_padding_length, self.keys)
return all_latents, all_embs, all_masks, caption_text
# all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos = collate_latents_embs_masks(
# rows, self.text_padding_length, self.keys)
# return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
batch = collate_rows_from_parquet_schema(rows, self.parquet_schema,
self.text_padding_length)
return batch
def __len__(self):
return sum(self.lengths)
@@ -286,6 +287,7 @@ def build_parquet_map_style_dataloader(
path,
batch_size,
num_data_workers,
parquet_schema,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
@@ -298,6 +300,7 @@ def build_parquet_map_style_dataloader(
drop_last=drop_last,
drop_first_row=drop_first_row,
text_padding_length=text_padding_length,
parquet_schema=parquet_schema,
seed=seed)
loader = StatefulDataLoader(
+21 -21
View File
@@ -22,7 +22,7 @@ logger = init_logger(__name__)
@dataclass
class DatasetBatch:
class PreprocessBatch:
"""
Batch information for dataset processing stages.
@@ -66,7 +66,7 @@ class DatasetStage(ABC):
"""
@abstractmethod
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch.
@@ -88,7 +88,7 @@ class DatasetFilterStage(ABC):
"""
@abstractmethod
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if batch should be kept.
@@ -102,7 +102,7 @@ class DatasetFilterStage(ABC):
raise NotImplementedError
@abstractmethod
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch (for non-filtering operations).
@@ -119,7 +119,7 @@ class DatasetFilterStage(ABC):
class DataValidationStage(DatasetFilterStage):
"""Stage for validating data items."""
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Validate data item.
@@ -142,7 +142,7 @@ class DataValidationStage(DatasetFilterStage):
return True
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for validation - filtering is handled by should_keep."""
return batch
@@ -160,7 +160,7 @@ class ResolutionFilterStage(DatasetFilterStage):
self.max_height = max_height
self.max_width = max_width
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if data item passes resolution filtering.
@@ -193,7 +193,7 @@ class ResolutionFilterStage(DatasetFilterStage):
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
return batch
@@ -218,7 +218,7 @@ class FrameSamplingStage(DatasetFilterStage):
self.video_length_tolerance_range = video_length_tolerance_range
self.drop_short_ratio = drop_short_ratio
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if video should be kept based on length constraints.
@@ -252,9 +252,9 @@ class FrameSamplingStage(DatasetFilterStage):
and random.random() < self.drop_short_ratio)
def process(self,
batch: DatasetBatch,
batch: PreprocessBatch,
temporal_sample_fn=None,
**kwargs) -> DatasetBatch:
**kwargs) -> PreprocessBatch:
"""
Process frame sampling for video data items.
@@ -298,7 +298,7 @@ class VideoTransformStage(DatasetStage):
def __init__(self, transform) -> None:
self.transform = transform
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform video data.
@@ -339,7 +339,7 @@ class ImageTransformStage(DatasetStage):
self.transform = transform
self.transform_topcrop = transform_topcrop
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform image data.
@@ -375,7 +375,7 @@ class TextEncodingStage(DatasetStage):
self.text_max_length = text_max_length
self.cfg_rate = cfg_rate
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process text data.
@@ -486,7 +486,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
return all_data[self.start_idx:]
def _process_metadata(self) -> List[DatasetBatch]:
def _process_metadata(self) -> List[PreprocessBatch]:
"""Process the raw metadata through all filtering stages."""
raw_data = self._load_raw_data()
processed_batches = []
@@ -500,11 +500,11 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
sample_num_frames: List[int] = []
for item in raw_data:
batch = DatasetBatch(path=item["path"],
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
batch = PreprocessBatch(path=item["path"],
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
@@ -522,7 +522,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
len(raw_data), len(processed_batches))
return processed_batches
def _apply_filter_stages(self, batch: DatasetBatch,
def _apply_filter_stages(self, batch: PreprocessBatch,
filter_counts: Dict[str, int]) -> bool:
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
if not self.validation_stage.should_keep(batch):
+174 -12
View File
@@ -3,6 +3,10 @@ from typing import Any, Dict, List
import numpy as np
import torch
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
@@ -38,39 +42,70 @@ def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
if shape is None or bytes is None:
raise ValueError(f"Key {key} not found in row_dict")
else:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
try:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
except KeyError:
continue
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
if len(bytes) == 0:
return_dict[key] = torch.zeros(0, dtype=torch.bfloat16)
else:
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
return return_dict
def collate_latents_embs_masks(
batch_to_process, text_padding_length,
keys) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
batch_to_process, text_padding_length, keys
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str], Dict[str, Any],
List[Dict[str, Any]]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
all_clip_features = []
all_first_frame_latents = []
all_pil_images = []
all_infos = []
caption_text = []
# Process each row individually
for i, row in enumerate(batch_to_process):
# Get info from row
info_keys = [
"caption", "file_name", "media_type", "width", "height",
"num_frames", "duration_sec", "fps"
]
info = {}
for key in info_keys:
if key in row:
info[key] = row[key]
else:
info[key] = ""
info["prompt"] = info["caption"]
# Get tensors from row
data = get_torch_tensors_from_row_dict(row, keys)
latents, emb = data["vae_latent"], data["text_embedding"]
clip_feature = data.get("clip_feature", None)
first_frame_latent = data.get("first_frame_latent", None)
pil_image = data.get("pil_image", None)
padded_emb, mask = pad(emb, text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
all_clip_features.append(clip_feature)
all_first_frame_latents.append(first_frame_latent)
all_pil_images.append(pil_image)
all_infos.append(info)
# TODO(py): remove this once we fix preprocess
try:
caption_text.append(row["prompt"])
@@ -81,5 +116,132 @@ def collate_latents_embs_masks(
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
all_extra_latents = {
"clip_feature": torch.stack(all_clip_features),
"first_frame_latent": torch.stack(all_first_frame_latents),
"pil_image": all_pil_images,
}
return all_latents, all_embs, all_masks, caption_text
return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
def collate_rows_from_parquet_schema(rows, parquet_schema,
text_padding_length) -> Dict[str, Any]:
"""
Collate rows from parquet files based on the provided schema.
Dynamically processes tensor fields based on schema and returns batched data.
Args:
rows: List of row dictionaries from parquet files
parquet_schema: PyArrow schema defining the structure of the data
Returns:
Dict containing batched tensors and metadata
"""
if not rows:
return {}
# Initialize containers for different data types
batch_data = {}
# Get tensor and metadata field names from schema (fields ending with '_bytes')
tensor_fields = []
metadata_fields = []
for field in parquet_schema.names:
if field.endswith('_bytes'):
shape_field = field.replace('_bytes', '_shape')
dtype_field = field.replace('_bytes', '_dtype')
tensor_name = field.replace('_bytes', '')
tensor_fields.append(tensor_name)
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
elif not field.endswith('_shape') and not field.endswith('_dtype'):
# Only add actual metadata fields, not the shape/dtype helper fields
metadata_fields.append(field)
# Process each tensor field efficiently
for tensor_name in tensor_fields:
tensor_list = []
for row in rows:
# Get tensor data from row using the existing helper function pattern
shape_key = f"{tensor_name}_shape"
bytes_key = f"{tensor_name}_bytes"
if shape_key in row and bytes_key in row:
# logger.info("row: %s", row)
# logger.info("shape_key: %s", shape_key)
# logger.info("bytes_key: %s", bytes_key)
shape = row[shape_key]
bytes_data = row[bytes_key]
if len(bytes_data) == 0:
tensor = torch.zeros(0, dtype=torch.bfloat16)
else:
# Convert bytes to tensor using float32 as default
# logger.info("len(bytes_data): %s", len(bytes_data))
# logger.info("shape: %s", shape)
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
# assert B == 1, "Batch size must be 1"
# tensor = tensor.squeeze(0)
tensor_list.append(tensor)
else:
# Handle missing tensor data
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_list:
if tensor_name == 'text_embedding':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
for tensor in tensor_list:
if tensor.numel() > 0:
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
else:
# Handle empty embeddings - assume default embedding dimension
padded_tensors.append(
torch.zeros(text_padding_length,
768,
dtype=torch.bfloat16))
attention_masks.append(torch.zeros(text_padding_length))
batch_data[tensor_name] = torch.stack(padded_tensors)
batch_data['text_attention_mask'] = torch.stack(attention_masks)
else:
# Stack other tensors directly, handling None values
valid_tensors = [
t for t in tensor_list if t is not None and t.numel() > 0
]
if valid_tensors:
batch_data[tensor_name] = torch.stack(valid_tensors)
elif tensor_list: # All tensors are empty but exist
batch_data[tensor_name] = torch.stack(tensor_list)
# Process metadata fields efficiently into info_list
info_list = []
for row in rows:
info = {}
for field in metadata_fields:
info[field] = row.get(field, "")
# Add prompt field for backward compatibility
info["prompt"] = info.get("caption", "")
info_list.append(info)
batch_data['info_list'] = info_list
# Add caption_text for backward compatibility
if info_list and 'caption' in info_list[0]:
batch_data['caption_text'] = [info['caption'] for info in info_list]
return batch_data
+15 -11
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
import os
import pathlib
import datasets
@@ -17,6 +18,9 @@ class ValidationDataset(torch.utils.data.IterableDataset):
super().__init__()
self.filename = pathlib.Path(filename)
# get directory of filename
# TODO(will)
self.dir = os.path.abspath(self.filename.parent)
if not self.filename.exists():
raise FileNotFoundError(
@@ -60,41 +64,41 @@ class ValidationDataset(torch.utils.data.IterableDataset):
if sample.get("image_path", None) is not None:
image_path = sample["image_path"]
image_path = os.path.join(self.dir, image_path)
if not pathlib.Path(image_path).is_file(
) and not image_path.startswith("http"):
logger.warning("Image file %s does not exist.",
image_path.as_posix())
logger.warning("Image file %s does not exist.", image_path)
else:
sample["image"] = load_image(sample["image_path"])
sample["image"] = load_image(image_path)
if sample.get("video_path", None) is not None:
video_path = sample["video_path"]
video_path = os.path.join(self.dir, video_path)
if not pathlib.Path(video_path).is_file(
) and not video_path.startswith("http"):
logger.warning("Video file %s does not exist.",
video_path.as_posix())
logger.warning("Video file %s does not exist.", video_path)
else:
sample["video"] = load_video(sample["video_path"])
sample["video"] = load_video(video_path)
if sample.get("control_image_path", None) is not None:
control_image_path = sample["control_image_path"]
control_image_path = os.path.join(self.dir, control_image_path)
if not pathlib.Path(control_image_path).is_file(
) and not control_image_path.startswith("http"):
logger.warning("Control Image file %s does not exist.",
control_image_path.as_posix())
control_image_path)
else:
sample["control_image"] = load_image(
sample["control_image_path"])
sample["control_image"] = load_image(control_image_path)
if sample.get("control_video_path", None) is not None:
control_video_path = sample["control_video_path"]
control_video_path = os.path.join(self.dir, control_video_path)
if not pathlib.Path(control_video_path).is_file(
) and not control_video_path.startswith("http"):
logger.warning("Control Video file %s does not exist.",
control_video_path)
else:
sample["control_video"] = load_video(
sample["control_video_path"])
sample["control_video"] = load_video(control_video_path)
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+3
View File
@@ -25,9 +25,12 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class WanImageEmbedding(torch.nn.Module):
+10 -1
View File
@@ -39,6 +39,7 @@ class ForwardBatch:
image_path: Optional[str] = None
image_embeds: List[torch.Tensor] = field(default_factory=list)
pil_image: Optional[PIL.Image.Image] = None
preprocessed_image: Optional[torch.Tensor] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
@@ -150,7 +151,12 @@ class TrainingBatch:
latents: Optional[torch.Tensor] = None
encoder_hidden_states: Optional[torch.Tensor] = None
encoder_attention_mask: Optional[torch.Tensor] = None
info: Optional[Dict[str, Any]] = None
# i2v
# extra_latents: Optional[Dict[str, Any]] = None
preprocessed_image: Optional[torch.Tensor] = None
image_embeds: Optional[torch.Tensor] = None
image_latents: Optional[torch.Tensor] = None
infos: Optional[List[Dict[str, Any]]] = None
# Transformer inputs
noisy_model_input: Optional[torch.Tensor] = None
@@ -160,6 +166,9 @@ class TrainingBatch:
attn_metadata: Optional[AttentionMetadata] = None
# input kwargs
input_kwargs: Optional[Dict[str, Any]] = None
# Training loss
loss: torch.Tensor | None = None
@@ -2,11 +2,13 @@
import gc
import multiprocessing
import os
from collections import defaultdict
from concurrent.futures import ProcessPoolExecutor
from itertools import chain
from typing import Any, Dict, List, Optional
import numpy as np
import PIL.Image
import pyarrow as pa
import pyarrow.parquet as pq
import torch
@@ -15,12 +17,13 @@ from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import ValidationDataset, getdataset
from fastvideo.v1.dataset.preprocessing_datasets import PreprocessBatch
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import EncodingStage, TextEncodingStage
from fastvideo.v1.pipelines.stages import TextEncodingStage
logger = init_logger(__name__)
@@ -36,9 +39,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae"), ))
@torch.no_grad()
def forward(
self,
@@ -61,39 +61,206 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"""Get the schema fields for the pipeline type. Override in subclasses."""
raise NotImplementedError
def create_record_for_schema(self,
preprocess_batch: PreprocessBatch,
schema: pa.Schema,
strict: bool = False) -> Dict[str, Any]:
"""Create a record for the Parquet dataset using a generic schema-based approach.
Args:
preprocess_batch: The batch containing the data to extract
schema: PyArrow schema defining the expected fields
strict: If True, raises an exception when required fields are missing or unfilled
Returns:
Dictionary record matching the schema
Raises:
ValueError: If strict=True and required fields are missing or unfilled
"""
record = {}
unfilled_fields = []
for field in schema.names:
field_filled = False
if field.endswith('_bytes'):
# Handle binary tensor data - convert numpy array or tensor to bytes
tensor_name = field.replace('_bytes', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None:
try:
if hasattr(tensor_data, 'numpy'): # torch tensor
record[field] = tensor_data.cpu().numpy().tobytes()
field_filled = True
elif hasattr(tensor_data, 'tobytes'): # numpy array
record[field] = tensor_data.tobytes()
field_filled = True
else:
raise ValueError(
f"Unsupported tensor type for field {field}: {type(tensor_data)}"
)
except Exception as e:
if strict:
raise ValueError(
f"Failed to convert tensor {tensor_name} to bytes: {e}"
)
record[field] = b'' # Empty bytes for missing data
else:
record[field] = b'' # Empty bytes for missing data
elif field.endswith('_shape'):
# Handle tensor shape info
tensor_name = field.replace('_shape', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None and hasattr(tensor_data, 'shape'):
record[field] = list(tensor_data.shape)
field_filled = True
else:
record[field] = []
elif field.endswith('_dtype'):
# Handle tensor dtype info
tensor_name = field.replace('_dtype', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None and hasattr(tensor_data, 'dtype'):
record[field] = str(tensor_data.dtype)
field_filled = True
else:
record[field] = 'unknown'
elif field in ['width', 'height', 'num_frames']:
# Handle integer metadata fields
value = getattr(preprocess_batch, field, None)
if value is not None:
try:
record[field] = int(value)
field_filled = True
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to int: {e}")
record[field] = 0
else:
record[field] = 0
elif field in ['duration_sec', 'fps']:
# Handle float metadata fields
# Map schema field names to batch attribute names
attr_name = 'duration' if field == 'duration_sec' else field
value = getattr(preprocess_batch, attr_name, None)
if value is not None:
try:
record[field] = float(value)
field_filled = True
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to float: {e}"
)
record[field] = 0.0
else:
record[field] = 0.0
else:
# Handle string fields (id, file_name, caption, media_type, etc.)
# Map common schema field names to batch attribute names
attr_name = field
if field == 'caption':
attr_name = 'text'
elif field == 'file_name':
attr_name = 'path'
elif field == 'id':
# Generate ID from path if available
path_value = getattr(preprocess_batch, 'path', None)
if path_value:
import os
record[field] = os.path.basename(path_value).split(
'.')[0]
field_filled = True
else:
record[field] = ""
continue
elif field == 'media_type':
# Determine media type from path
path_value = getattr(preprocess_batch, 'path', None)
if path_value:
record[field] = 'video' if path_value.endswith(
'.mp4') else 'image'
field_filled = True
else:
record[field] = ""
continue
value = getattr(preprocess_batch, attr_name, None)
if value is not None:
record[field] = str(value)
field_filled = True
else:
record[field] = ""
# Track unfilled fields
if not field_filled:
unfilled_fields.append(field)
# Handle strict mode
if strict and unfilled_fields:
raise ValueError(
f"Required fields were not filled: {unfilled_fields}")
# Log unfilled fields as warning if not in strict mode
if unfilled_fields:
logger.warning(
f"Some fields were not filled and got default values: {unfilled_fields}"
)
return record
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
# text_attention_mask: np.ndarray,
valid_data: Dict[str, Any],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset."""
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx] if valid_data else "",
"media_type": "video",
"id":
video_name,
"vae_latent_bytes":
vae_latent.tobytes(),
"vae_latent_shape":
list(vae_latent.shape),
"vae_latent_dtype":
str(vae_latent.dtype),
"text_embedding_bytes":
text_embedding.tobytes(),
"text_embedding_shape":
list(text_embedding.shape),
"text_embedding_dtype":
str(text_embedding.dtype),
"file_name":
video_name,
"caption":
valid_data["text"][idx] if len(valid_data["text"]) > 0 else "",
"media_type":
"video",
"width":
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
valid_data["pixel_values"][idx].shape[-2]
if len(valid_data["pixel_values"]) > 0 else 0,
"height":
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
valid_data["pixel_values"][idx].shape[-1]
if len(valid_data["pixel_values"]) > 0 else 0,
"num_frames":
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
"duration_sec":
float(valid_data["duration"][idx]) if valid_data else 0.0,
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
float(valid_data["duration"][idx])
if len(valid_data["duration"]) > 0 else 0.0,
"fps":
float(valid_data["fps"][idx])
if len(valid_data["fps"]) > 0 else 0.0,
}
if extra_features:
record.update(extra_features)
@@ -213,8 +380,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
# ).astype(np.uint8)
# Get extra features for this sample if needed
sample_extra_features = {}
@@ -231,7 +398,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
@@ -318,7 +485,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
for idx, sample in pbar:
with torch.inference_mode():
prompt = sample["caption"]
# is_negative_prompt = idx == 0
is_negative_prompt = idx == 0
# Text Encoder
batch = ForwardBatch(
@@ -346,15 +513,44 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"Shape after removing padding - Embeddings: %s, Mask: %s",
text_embedding.shape, text_attention_mask.shape)
extra_features = {}
if not is_negative_prompt:
height = sample["height"]
width = sample["width"]
if "image_path" in sample and "video_path" in sample:
raise ValueError(
"Only one of image_path or video_path should be provided"
)
if "image" in sample:
extra_features = self.preprocess_image(
sample["image"], height, width, fastvideo_args)
if "video" in sample:
extra_features = self.preprocess_video(
sample["video"], height, width, fastvideo_args)
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value.cpu().numpy()
else:
sample_extra_features[key] = value
valid_data = defaultdict(list)
valid_data["text"] = [prompt]
# Create record for Parquet dataset
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=None,
idx=0,
extra_features=None)
record = self.create_record(
video_name=file_name,
vae_latent=np.array([], dtype=np.float32),
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=0,
extra_features=sample_extra_features)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
@@ -428,6 +624,15 @@ class BasePreprocessPipeline(ComposedPipelineBase):
del table
gc.collect() # Force garbage collection
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return {}
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return {}
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
@@ -8,6 +8,7 @@ using the modular pipeline architecture.
from typing import Any, Dict, List, Optional
import numpy as np
import PIL
import torch
from PIL import Image
@@ -15,8 +16,13 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.models.vision_utils import (get_default_height_width,
normalize, numpy_to_pt,
pil_to_numpy, resize)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.v1.pipelines.stages import ImageEncodingStage, TextEncodingStage
class PreprocessPipeline_I2V(BasePreprocessPipeline):
@@ -26,18 +32,75 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
assert hasattr(
self,
"image_encoding_stage"), "Image encoding stage must be created"
batch = ForwardBatch(
data_type="video",
pil_image=image,
)
result_batch = self.image_encoding_stage(batch, fastvideo_args)
clip_features = result_batch.image_embeds[0]
# image = self.pil_to_tensor(image)
image = self.preprocess(
image,
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
height=height,
width=width)
return {
"clip_feature": clip_features[0],
"pil_image": image,
}
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return self.preprocess_image(video[0], height, width, fastvideo_args)
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for I2V pipeline."""
return [f.name for f in pyarrow_schema_i2v]
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("image_encoder").to(get_torch_device())
# self.get_module("image_processor").to(get_torch_device())
self.get_module("vae").to(get_torch_device())
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)
batch_size, _, 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
processed_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))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
@@ -53,25 +116,90 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
return {"clip_feature": clip_features}
features["clip_feature"] = clip_features
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0],
processed_img.shape[1], num_frames - 1,
height, width)
],
dim=2)
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor.to(latent_condition.device,
latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, list(range(1, num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
dim=2,
repeats=self.get_module("vae").temporal_compression_ratio)
mask_lat_size = torch.concat(
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
mask_lat_size = mask_lat_size.view(
batch_size, -1,
self.get_module("vae").temporal_compression_ratio, latent_height,
latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
features["first_frame_latent"] = image_latent
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
# text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = 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,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
record = super().create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
@@ -87,7 +215,69 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
return record # type: ignore
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes":
first_frame_latent.tobytes(),
"first_frame_latent_shape":
list(first_frame_latent.shape),
"first_frame_latent_dtype":
str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_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 pil_to_tensor(self, image: PIL.Image.Image) -> torch.Tensor:
image = image
image = np.array(image).astype(np.float32)
image = torch.from_numpy(image)
return image
def preprocess(self,
image: PIL.Image.Image,
vae_scale_factor: int,
height: int,
width: int,
resize_mode: str = "default") -> torch.Tensor:
image = [image]
height, width = get_default_height_width(image[0], vae_scale_factor,
height, width)
image = [
resize(i, height, width, resize_mode=resize_mode) for i in image
]
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
do_normalize = True
if image.min() < 0:
do_normalize = False
if do_normalize:
image = normalize(image)
return image
EntryClass = PreprocessPipeline_I2V
@@ -79,7 +79,6 @@ 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("--dataset", default="t2v")
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
+16 -12
View File
@@ -51,22 +51,26 @@ class EncodingStage(PipelineStage):
"""
self.vae = self.vae.to(get_torch_device())
image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
raise ValueError("Image Path must be provided")
assert batch.height is not None
assert batch.width is not None
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
image = batch.pil_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
image = batch.preprocessed_image
# TODO(will)
if image is None:
assert batch.pil_image is not None
image = batch.pil_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
else:
image = image.transpose(1, 2)
logger.info("image: %s", image.shape)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
@@ -181,7 +185,7 @@ class EncodingStage(PipelineStage):
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage inputs."""
result = VerificationResult()
result.add_check("pil_image", batch.pil_image, V.not_none)
# result.add_check("pil_image", batch.pil_image)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("generator", batch.generator,
@@ -0,0 +1 @@
{"step_time":2.245914653001819,"_wandb":{"runtime":1434},"learning_rate":1e-05,"grad_norm":0.57421875,"avg_step_time":1.1814782944297622,"train_loss":0.07932619750499725,"vsa_sparsity":0,"_timestamp":1.750578625921253e+09,"validation_videos_40_steps":{"count":1,"videos":[{"size":420969,"path":"media/videos/validation_videos_40_steps_900_581ff5eae2909d3a7b36.mp4","_type":"video-file","sha256":"581ff5eae2909d3a7b362dcb24d060c006c09e4d4deb44b82f4aa697f6789ba7"}],"captions":false,"_type":"videos"},"_runtime":1434.62395329,"_step":901}
@@ -0,0 +1,177 @@
import os
from pathlib import Path
from huggingface_hub import snapshot_download
import shutil
import subprocess
import sys
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
NUM_NODES = "1"
MODEL_PATH = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
# preprocessing
DATA_DIR = "data"
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_preprocessed_data_i2v"))
# training
NUM_GPUS_PER_NODE_TRAINING = "8"
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_i2v_training_pipeline.py"
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
LOCAL_VALIDATION_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "validation_parquet_dataset")
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
# if data_dir.exists():
# print(f"Removing existing data directory at {data_dir}")
# shutil.rmtree(data_dir)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
# result = snapshot_download(
# repo_id="wlsaidhi/cats-overfit-merged",
# local_dir=str(LOCAL_RAW_DATA_DIR),
# repo_type="dataset",
# resume_download=True,
# token=os.environ.get("HF_TOKEN"), # In case authentication is needed
# )
print(f"Download completed successfully. Files downloaded to: {result}")
# Verify the download
if not LOCAL_RAW_DATA_DIR.exists():
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
# List downloaded files
print("Downloaded files:")
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
if file.is_file():
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
except Exception as e:
print(f"Error during download: {str(e)}")
raise
def run_preprocessing():
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
PREPROCESSING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge_1_sample.txt"),
"--preprocess_video_batch_size", "1",
"--max_height", "480",
"--max_width", "832",
"--num_frames", "77",
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_i2v_prompt_1_sample.json"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
"--preprocess_task", "i2v",
]
process = subprocess.run(cmd, check=True)
def run_training():
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
TRAINING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
"--tp_size", NUM_GPUS_PER_NODE_TRAINING,
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "40",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
]
print(f"Running training with command: {cmd}")
process = subprocess.run(cmd, check=True)
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
# download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
print(f"reference_video_file: {reference_video_file}")
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
print(f"final_validation_video_file: {final_validation_video_file}")
# Ensure both files exist
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
# Compute SSIM
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
reference_video_file,
final_validation_video_file,
use_ms_ssim=True # Using MS-SSIM for better quality assessment
)
print("\n===== SSIM Results for Step 900 Validation =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
if __name__ == "__main__":
test_e2e_overfit_single_sample()
+114 -67
View File
@@ -22,6 +22,8 @@ from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadata)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_t2v, pyarrow_schema_t2v_validation)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -50,14 +52,17 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
_required_config_modules = ["scheduler", "transformer"]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
Dict[str, Any]]]
train_loader_iter: Iterator[Dict[str, Any]]
current_epoch: int = 0
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_t2v
self.validation_dataset_schema = pyarrow_schema_t2v_validation
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.training_args = training_args
@@ -70,7 +75,12 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
assert training_args.seed is not None
self.seed = training_args.seed
assert self.transformer is not None
self.set_schemas()
# self.train_dataset_schema = pyarrow_schema_t2v
# self.validation_dataset_schema = pyarrow_schema_t2v_validation
self.transformer.requires_grad_(True)
self.transformer.train()
@@ -104,12 +114,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=training_args.seed)
seed=self.seed)
self.noise_scheduler = noise_scheduler
@@ -158,14 +169,28 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents, encoder_hidden_states, encoder_attention_mask, infos = batch
# latents, encoder_hidden_states, encoder_attention_mask, caption_text, extra_latents, infos = batch
# for key, value in batch.items():
# if isinstance(value, torch.Tensor):
# logger.info("key: %s, shape: %s", key, value.shape)
# else:
# logger.info("key: %s, value: %s", key, value)
# print("--------------------------------")
# logger.info("batch: %s", batch)
latents = batch['vae_latent']
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.info = infos
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
@@ -243,25 +268,15 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
return training_batch
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
# assert training_batch.attn_metadata is not None
assert training_batch.latents is not None
assert training_batch.noise is not None
assert training_batch.sigmas is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
input_kwargs = {
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
@@ -274,6 +289,23 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
"return_dict":
False,
}
return training_batch
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.noise is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
assert training_batch.input_kwargs is not None
input_kwargs = training_batch.input_kwargs
# if 'hunyuan' in self.training_args.model_type:
# input_kwargs["guidance"] = torch.tensor(
# [1000.0],
@@ -340,6 +372,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
training_batch = self._normalize_dit_input(training_batch)
training_batch = self._prepare_dit_inputs(training_batch)
training_batch = self._build_attention_metadata(training_batch)
training_batch = self._build_input_kwargs(training_batch)
training_batch = self._transformer_forward_and_compute_loss(
training_batch)
@@ -372,14 +405,10 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
assert self.training_args is not None
# Set random seeds for deterministic training
assert self.training_args.seed is not None, "seed must be set"
seed = self.training_args.seed
set_random_seed(seed)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
set_random_seed(self.seed)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
@@ -505,6 +534,54 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
def _prepare_validation_inputs(
self, sampling_param: SamplingParam, training_args: TrainingArgs,
validation_batch: Dict[str, Any], num_inference_steps: int,
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
# logger.info("validation_batch: %s", validation_batch)
# latents, embeddings, masks, caption_text, extra_latents, infos = validation_batch
prompt = validation_batch['info_list'][0]['prompt']
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
prompt_embeds = prompt_embeds.to(get_torch_device())
prompt_attention_mask = prompt_attention_mask.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
prompt=prompt,
data_type="video",
latents=None,
seed=self.seed, # Use deterministic seed
generator=torch.Generator(device="cpu").manual_seed(self.seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
assert training_args is not None
@@ -521,11 +598,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
# Set deterministic seed for validation
validation_seed = training_args.seed if training_args.seed is not None else 42
torch.manual_seed(validation_seed)
torch.cuda.manual_seed_all(validation_seed)
set_random_seed(self.seed)
logger.info("Using validation seed: %s", validation_seed)
logger.info("Using validation seed: %s", self.seed)
# Prepare validation prompts
logger.info('fastvideo_args.validation_preprocessed_path: %s',
@@ -533,13 +608,15 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
training_args.validation_preprocessed_path,
batch_size=1,
parquet_schema=self.validation_dataset_schema,
num_data_workers=0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None,
cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
negative_prompt_embeds, negative_prompt_attention_mask, negative_prompt = validation_dataset.get_validation_negative_prompt(
)
logger.info("negative_prompt: %s", negative_prompt)
transformer.eval()
@@ -552,43 +629,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
for _, embeddings, masks, infos in validation_dataloader:
# for _, embeddings, masks, caption_text, extra_latents, infos in validation_dataloader:
for validation_batch in validation_dataloader:
batch = self._prepare_validation_inputs(
sampling_param, training_args, validation_batch,
num_inference_steps, negative_prompt_embeds,
negative_prompt_attention_mask)
step_captions.extend([None]) # TODO(peiyuan): add caption
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
dtype=torch.bfloat16):
@@ -0,0 +1,264 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any, Dict
import torch
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
TrainingBatch)
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
WanImageToVideoValidationPipeline)
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import shard_latents_across_sp
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanI2VTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_i2v
self.validation_dataset_schema = pyarrow_schema_i2v_validation
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus)
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, caption_text, extra_latents, infos = batch
# for key, value in batch.items():
# if isinstance(value, torch.Tensor):
# logger.info("key: %s, shape: %s", key, value.shape)
# else:
# logger.info("key: %s, value: %s", key, value)
# print("--------------------------------")
# logger.info("batch: %s", batch)
latents = batch['vae_latent']
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
pil_image = batch['pil_image']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.preprocessed_image = pil_image.to(get_torch_device())
training_batch.image_embeds = clip_features.to(get_torch_device())
training_batch.image_latents = image_latents.to(get_torch_device())
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
assert training_batch.preprocessed_image is not None
assert training_batch.image_embeds is not None
assert training_batch.image_latents is not None
# assert training_batch.extra_latents is not None
# extra_latents = training_batch.extra_latents
# if extra_latents:
# image_embeds, image_latents = extra_latents[
# "clip_feature"], extra_latents["first_frame_latent"]
# image_
# Image Embeds
image_embeds = training_batch.image_embeds
image_latents = training_batch.image_latents
preprocessed_image = training_batch.preprocessed_image
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
# Image Latents
assert torch.isnan(image_latents).sum() == 0
image_latents = image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
image_latents = shard_latents_across_sp(
image_latents, num_latent_t=self.training_args.num_latent_t)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
"encoder_hidden_states_image":
encoder_hidden_states_image,
"return_dict":
False,
}
return training_batch
def _prepare_validation_inputs(
self, sampling_param: SamplingParam, training_args: TrainingArgs,
validation_batch: Dict[str, Any], num_inference_steps: int,
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
# latents, embeddings, masks, caption_text, extra_latents, infos = validation_batch
# latents = validation_batch['vae_latent']
embeddings = validation_batch['text_embedding']
masks = validation_batch['text_attention_mask']
clip_features = validation_batch['clip_feature']
# extra_latents = validation_batch['extra_latents']
preprocessed_image = validation_batch['pil_image']
infos = validation_batch['info_list']
prompt = infos[0]['prompt']
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
clip_features = clip_features.to(get_torch_device())
if 'colorful candies' in prompt:
logger.info("colorful candies")
from fastvideo.v1.models.vision_utils import load_video
# video_path = 'validation_dataset/yYcK4nANZz4-Scene-030.mp4'
video_path = '/mnt/user_storage/fv/FastVideo/examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4'
video = load_video(video_path)
pil_image = video[0]
preprocessed_image = None
else:
pil_image = None
# clip_features = extra_latents.get("clip_feature")
# first_frame_latent = extra_latents.get("first_frame_latent")
# pil_image = extra_latents.get("pil_image")
# if clip_features is not None and clip_features.numel() > 0:
# clip_features = clip_features.to(get_torch_device())
# if first_frame_latent is not None and first_frame_latent.numel() > 0:
# first_frame_latent = first_frame_latent.to(get_torch_device())
# if pil_image is not None and pil_image[0] is not None and pil_image[
# 0].numel() > 0:
# pil_image = pil_image[0].to(get_torch_device())
# else:
# clip_features = None
# first_frame_latent = None
# pil_image = None
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
prompt=prompt,
data_type="video",
latents=None,
seed=self.seed, # Use deterministic seed
generator=torch.Generator(device="cpu").manual_seed(self.seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
image_embeds=[clip_features],
preprocessed_image=preprocessed_image,
pil_image=pil_image,
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanI2VTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)