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
William Lin 8741d204a5 [Training] Refactor and improve validation datasets (#539) 2025-06-21 17:58:35 -07:00
46 changed files with 2666 additions and 581 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"
@@ -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": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -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
}
]
}
+17 -20
View File
@@ -1,19 +1,17 @@
import os
# SPDX-License-Identifier: Apache-2.0
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
from fastvideo.v1.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.v1.dataset.preprocessing_datasets import (
VideoCaptionMergedDataset)
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from .parquet_dataset_map_style import build_parquet_map_style_dataloader
__all__ = ["build_parquet_map_style_dataloader"]
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
def getdataset(args, start_idx=0) -> T2V_dataset:
def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
@@ -31,15 +29,14 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
*resize_topcrop,
norm_fun,
])
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
start_idx=start_idx)
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)
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset"
]
+60 -11
View File
@@ -26,15 +26,47 @@ 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()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_i2v_validation = pa.schema([
pa.field("id", 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()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_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()),
# 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()),
@@ -64,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()),
@@ -80,4 +107,26 @@ pyarrow_schema_t2v = pa.schema([
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
])
pyarrow_schema_t2v_validation = pa.schema([
pa.field("id", 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()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
@@ -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(
@@ -0,0 +1,592 @@
# SPDX-License-Identifier: Apache-2.0
import json
import math
import os
import random
from abc import ABC, abstractmethod
from collections import Counter
from dataclasses import dataclass
from os.path import join as opj
from typing import Any, Dict, List, Optional, Union
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from transformers import AutoTokenizer
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclass
class PreprocessBatch:
"""
Batch information for dataset processing stages.
This class holds all the information about a video-caption or image-caption pair
as it moves through the processing pipeline. Fields are populated by different stages.
"""
# Raw metadata
path: str
cap: Union[str, List[str]]
resolution: Optional[Dict] = None
fps: Optional[float] = None
duration: Optional[float] = None
# Processed metadata
num_frames: Optional[int] = None
sample_frame_index: Optional[List[int]] = None
sample_num_frames: Optional[int] = None
# Processed data
pixel_values: Optional[torch.Tensor] = None
text: Optional[str] = None
input_ids: Optional[torch.Tensor] = None
cond_mask: Optional[torch.Tensor] = None
@property
def is_video(self) -> bool:
"""Check if this is a video item."""
return self.path.endswith(".mp4")
@property
def is_image(self) -> bool:
"""Check if this is an image item."""
return self.path.endswith(".jpg")
class DatasetStage(ABC):
"""
Abstract base class for dataset processing stages.
Similar to PipelineStage but designed for dataset preprocessing operations.
"""
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch.
Args:
batch: Dataset batch to process
**kwargs: Additional processing parameters
Returns:
Processed batch
"""
raise NotImplementedError
class DatasetFilterStage(ABC):
"""
Abstract base class for dataset filtering stages.
These stages can filter out items during metadata processing.
"""
@abstractmethod
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if batch should be kept.
Args:
batch: Dataset batch to check
**kwargs: Additional parameters
Returns:
True if batch should be kept, False otherwise
"""
raise NotImplementedError
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch (for non-filtering operations).
Args:
batch: Dataset batch to process
**kwargs: Additional processing parameters
Returns:
Processed batch
"""
raise NotImplementedError
class DataValidationStage(DatasetFilterStage):
"""Stage for validating data items."""
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Validate data item.
Args:
batch: Dataset batch to validate
Returns:
True if valid, False if invalid
"""
# Check for caption
if batch.cap is None:
return False
if batch.is_video:
# Validate video-specific fields
if batch.duration is None or batch.fps is None:
return False
elif not batch.is_image:
return False
return True
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for validation - filtering is handled by should_keep."""
return batch
class ResolutionFilterStage(DatasetFilterStage):
"""Stage for filtering data items based on resolution constraints."""
def __init__(self,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16,
max_height: int = 1024,
max_width: int = 1024):
self.max_h_div_w_ratio = max_h_div_w_ratio
self.min_h_div_w_ratio = min_h_div_w_ratio
self.max_height = max_height
self.max_width = max_width
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if data item passes resolution filtering.
Args:
batch: Dataset batch with resolution information
Returns:
True if passes filter, False otherwise
"""
# Only apply to videos
if not batch.is_video:
return True
if batch.resolution is None:
return False
height = batch.resolution.get("height", None)
width = batch.resolution.get("width", None)
if height is None or width is None:
return False
# Check aspect ratio
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
return self.filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
return batch
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
min_h_div_w_ratio: float) -> bool:
"""Filter based on height/width ratio."""
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class FrameSamplingStage(DatasetFilterStage):
"""Stage for temporal frame sampling and indexing."""
def __init__(self,
num_frames: int,
train_fps: int,
speed_factor: int = 1,
video_length_tolerance_range: float = 5.0,
drop_short_ratio: float = 0.0):
self.num_frames = num_frames
self.train_fps = train_fps
self.speed_factor = speed_factor
self.video_length_tolerance_range = video_length_tolerance_range
self.drop_short_ratio = drop_short_ratio
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if video should be kept based on length constraints.
Args:
batch: Dataset batch
Returns:
True if should be kept, False otherwise
"""
if batch.is_image:
return True
if batch.duration is None or batch.fps is None:
return False
num_frames = math.ceil(batch.fps * batch.duration)
# Check if video is too long
if (num_frames / batch.fps > self.video_length_tolerance_range *
(self.num_frames / self.train_fps * self.speed_factor)):
return False
# Resample frame indices to check length
frame_interval = batch.fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, num_frames,
frame_interval).astype(int)
# Filter short videos
return not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
def process(self,
batch: PreprocessBatch,
temporal_sample_fn=None,
**kwargs) -> PreprocessBatch:
"""
Process frame sampling for video data items.
Args:
batch: Dataset batch
temporal_sample_fn: Function for temporal sampling
Returns:
Updated batch with frame sampling info
"""
if batch.is_image:
# For images, just add sample info
batch.sample_frame_index = [0]
batch.sample_num_frames = 1
return batch
assert batch.duration is not None and batch.fps is not None
batch.num_frames = math.ceil(batch.fps * batch.duration)
# Resample frame indices
frame_interval = batch.fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch.num_frames,
frame_interval).astype(int)
# Temporal crop if too long
if len(frame_indices
) > self.num_frames and temporal_sample_fn is not None:
begin_index, end_index = temporal_sample_fn(len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
batch.sample_frame_index = frame_indices.tolist()
batch.sample_num_frames = len(frame_indices)
return batch
class VideoTransformStage(DatasetStage):
"""Stage for video data transformation."""
def __init__(self, transform) -> None:
self.transform = transform
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform video data.
Args:
batch: Dataset batch with video information
Returns:
Batch with transformed video tensor
"""
if not batch.is_video:
return batch
assert os.path.exists(batch.path), f"file {batch.path} do not exist!"
assert batch.sample_frame_index is not None, "Frame indices must be set before transformation"
torchvision_video, _, metadata = torchvision.io.read_video(
batch.path, output_format="TCHW")
video = torchvision_video[batch.sample_frame_index]
if self.transform is not None:
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
batch.pixel_values = video
return batch
class ImageTransformStage(DatasetStage):
"""Stage for image data transformation."""
def __init__(self, transform, transform_topcrop) -> None:
self.transform = transform
self.transform_topcrop = transform_topcrop
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform image data.
Args:
batch: Dataset batch with image information
Returns:
Batch with transformed image tensor
"""
if not batch.is_image:
return batch
image = Image.open(batch.path).convert("RGB")
image = torch.from_numpy(np.array(image))
image = rearrange(image, "h w c -> c h w").unsqueeze(0)
if self.transform_topcrop is not None:
image = self.transform_topcrop(image)
elif self.transform is not None:
image = self.transform(image)
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
batch.pixel_values = image
return batch
class TextEncodingStage(DatasetStage):
"""Stage for text tokenization and encoding."""
def __init__(self, tokenizer, text_max_length: int, cfg_rate: float = 0.0):
self.tokenizer = tokenizer
self.text_max_length = text_max_length
self.cfg_rate = cfg_rate
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process text data.
Args:
batch: Dataset batch with caption information
Returns:
Batch with encoded text information
"""
text = batch.cap
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg_rate else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
batch.text = text
batch.input_ids = text_tokens_and_mask["input_ids"]
batch.cond_mask = text_tokens_and_mask["attention_mask"]
return batch
class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
torch.distributed.checkpoint.stateful.Stateful):
"""
Merged dataset for video and caption data with stage-based processing.
This dataset processes video and image data through a series of stages:
- Data validation
- Resolution filtering
- Frame sampling
- Transformation
- Text encoding
"""
def __init__(self,
data_merge_path: str,
args,
transform,
temporal_sample,
transform_topcrop,
start_idx: int = 0):
self.data_merge_path = data_merge_path
self.start_idx = start_idx
self.args = args
self.temporal_sample = temporal_sample
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
# Initialize processing stages
self._init_stages(args, transform, transform_topcrop, tokenizer)
# Process metadata
self.processed_batches = self._process_metadata()
def _init_stages(self, args, transform, transform_topcrop,
tokenizer) -> None:
"""Initialize all processing stages."""
self.validation_stage = DataValidationStage()
self.resolution_filter_stage = ResolutionFilterStage(
max_height=args.max_height, max_width=args.max_width)
self.frame_sampling_stage = FrameSamplingStage(
num_frames=args.num_frames,
train_fps=args.train_fps,
speed_factor=args.speed_factor,
video_length_tolerance_range=args.video_length_tolerance_range,
drop_short_ratio=args.drop_short_ratio)
self.video_transform_stage = VideoTransformStage(transform)
self.image_transform_stage = ImageTransformStage(
transform, transform_topcrop)
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.cfg)
def _load_raw_data(self) -> List[Dict]:
"""Load raw data from JSON files."""
all_data = []
# Read folder-annotation pairs
with open(self.data_merge_path) as f:
folder_anno_pairs = [
line.strip().split(",") for line in f if line.strip()
]
# Process each folder-annotation pair
for folder, annotation_file in folder_anno_pairs:
with open(annotation_file) as f:
data_items = json.load(f)
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
all_data.extend(data_items)
return all_data[self.start_idx:]
def _process_metadata(self) -> List[PreprocessBatch]:
"""Process the raw metadata through all filtering stages."""
raw_data = self._load_raw_data()
processed_batches = []
# Initialize counters
filter_counts = {
"validation_failed": 0,
"resolution_failed": 0,
"frame_sampling_failed": 0
}
sample_num_frames: List[int] = []
for item in raw_data:
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):
continue
# Apply frame sampling processing
batch = self.frame_sampling_stage.process(
batch, temporal_sample_fn=self.temporal_sample)
processed_batches.append(batch)
assert batch.sample_num_frames is not None
sample_num_frames.append(batch.sample_num_frames)
self._log_filtering_stats(filter_counts, sample_num_frames,
len(raw_data), len(processed_batches))
return processed_batches
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):
filter_counts["validation_failed"] += 1
return False
if not self.resolution_filter_stage.should_keep(batch):
filter_counts["resolution_failed"] += 1
return False
if not self.frame_sampling_stage.should_keep(batch):
filter_counts["frame_sampling_failed"] += 1
return False
return True
def _log_filtering_stats(self, filter_counts: Dict[str, int],
sample_num_frames: List[int], before_count: int,
after_count: int):
"""Log filtering statistics."""
logger.info(
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
filter_counts['validation_failed'],
filter_counts['resolution_failed'],
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
before_count, after_count)
def __iter__(self):
"""Iterate through processed data items."""
for idx in range(len(self.processed_batches)):
yield self._get_item(idx)
def __len__(self):
return len(self.processed_batches)
def _get_item(self, idx: int) -> Dict:
"""Get a single processed data item."""
batch = self.processed_batches[idx]
# Apply transformation stages
batch = self.video_transform_stage.process(batch)
batch = self.image_transform_stage.process(batch)
batch = self.text_encoding_stage.process(batch)
# Build result dictionary
result = {
"pixel_values": batch.pixel_values,
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
}
# Add video-specific fields
if batch.is_video:
result.update({"fps": batch.fps, "duration": batch.duration})
return result
def state_dict(self) -> Dict[str, Any]:
"""Return state dict for checkpointing."""
return {"processed_batches": self.processed_batches}
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
-352
View File
@@ -1,352 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import json
import math
import os
import random
from collections import Counter
from os.path import join as opj
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from torch.utils.data import Dataset
from fastvideo.utils.dataset_utils import DecordInit
from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances: dict[type, 'SingletonMeta'] = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
class DataSetProg(metaclass=SingletonMeta):
def __init__(self) -> None:
self.cap_list: list[dict] = []
self.elements: list[int] = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements: dict[int, list[int]] = {}
self.n_used_elements: dict[int, int] = {}
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
self.elements = list(range(n_elements))
random.shuffle(self.elements)
print(f"n_elements: {len(self.elements)}", flush=True)
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(
math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info) -> int:
worker_id = 0 if work_info is None else work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
dataset_prog = DataSetProg()
def filter_resolution(h: int,
w: int,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16) -> bool:
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class T2V_dataset(Dataset):
def __init__(self,
args,
transform,
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0) -> None:
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
self.use_image_num = args.use_image_num
self.transform = transform
self.transform_topcrop = transform_topcrop
self.temporal_sample = temporal_sample
self.tokenizer = tokenizer
self.text_max_length = args.text_max_length
self.cfg = args.cfg
self.speed_factor = args.speed_factor
self.max_height = args.max_height
self.max_width = args.max_width
self.drop_short_ratio = args.drop_short_ratio
assert self.speed_factor >= 1
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if "mt5" not in args.text_encoder_name:
self.support_Chinese = False
cap_list = self.get_cap_list()
assert len(cap_list) > 0
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
def set_checkpoint(self, n_used_elements):
for i in range(len(dataset_prog.n_used_elements)):
dataset_prog.n_used_elements[i] = n_used_elements
def __len__(self):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
def get_data(self, idx) -> dict:
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx) -> dict:
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
assert video.dtype == torch.uint8
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
text = dataset_prog.cap_list[idx]["cap"]
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx) -> dict:
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
# for i in image:
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images"
in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps: list[str] = (image_data["cap"] if isinstance(
image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
single_text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
single_text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"] # 1, l
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
return dict(
pixel_values=image,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=image_data["path"],
)
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
cnt_too_short = 0
cnt_no_cap = 0
cnt_no_resolution = 0
cnt_resolution_mismatch = 0
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
if fps is None or duration is None:
continue
# ======resolution mismatch=====
resolution = i.get("resolution", None)
if resolution is None:
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None
or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"][
"width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
if not is_pick:
print("resolution mismatch")
cnt_resolution_mismatch += 1
continue
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
# from IPython import embed; embed()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps * self.speed_factor
): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"],
frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(
len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
new_cap_list.append(i)
i["sample_num_frames"] = 1
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices) -> torch.Tensor:
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data) -> list[dict]:
cap_lists = []
with open(data) as f:
folder_anno = [
i.strip().split(",") for i in f.readlines()
if len(i.strip()) > 0
]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno) as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
return cap_lists
def get_cap_list(self) -> list:
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
+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
+104
View File
@@ -0,0 +1,104 @@
# 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
import torch
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image, load_video
logger = init_logger(__name__)
class ValidationDataset(torch.utils.data.IterableDataset):
def __init__(self, filename: str):
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(
f"File {self.filename.as_posix()} does not exist")
if self.filename.suffix == ".csv":
data = datasets.load_dataset("csv",
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".json":
data = datasets.load_dataset("json",
data_files=self.filename.as_posix(),
split="train",
field="data")
elif self.filename.suffix == ".parquet":
data = datasets.load_dataset("parquet",
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".arrow":
data = datasets.load_dataset("arrow",
data_files=self.filename.as_posix(),
split="train")
else:
_SUPPORTED_FILE_FORMATS = [".csv", ".json", ".parquet", ".arrow"]
raise ValueError(
f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}"
)
self._data = data.to_iterable_dataset()
def __iter__(self):
for sample in self._data:
# For consistency reasons, we mandate that "caption" is always present in the validation dataset.
# However, since the model specifications use "prompt", we create an alias here.
sample["prompt"] = sample["caption"]
# Load image or video if the path is provided
# TODO(aryan): need to handle custom columns here for control conditions
sample["image"] = None
sample["video"] = None
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)
else:
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)
else:
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)
else:
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(control_video_path)
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+7 -3
View File
@@ -388,7 +388,8 @@ class TrainingArgs(FastVideoArgs):
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_dataset_file: str = ""
validation_preprocessed_path: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
@@ -536,9 +537,12 @@ class TrainingArgs(FastVideoArgs):
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
parser.add_argument("--validation-dataset-file",
type=str,
help="Directory containing validation prompts")
help="Path to unprocessed validation dataset")
parser.add_argument("--validation-preprocessed-path",
type=str,
help="Path to processed validation dataset")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
+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):
+83
View File
@@ -1,8 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
import os
import tempfile
from typing import Callable, List, Optional, Tuple, Union
from urllib.parse import unquote, urlparse
import imageio
import numpy as np
import PIL.Image
import PIL.ImageOps
@@ -86,6 +89,7 @@ def normalize(
return 2.0 * images - 1.0
# adapted from diffusers.utils import load_image
def load_image(
image: Union[str, PIL.Image.Image],
convert_method: Optional[Callable[[PIL.Image.Image],
@@ -131,6 +135,85 @@ def load_image(
return image
# adapted from diffusers.utils import load_video
def load_video(
video: str,
convert_method: Optional[Callable[[List[PIL.Image.Image]],
List[PIL.Image.Image]]] = None,
) -> List[PIL.Image.Image]:
"""
Loads `video` to a list of PIL Image.
Args:
video (`str`):
A URL or Path to a video to convert to a list of PIL Image format.
convert_method (Callable[[List[PIL.Image.Image]], List[PIL.Image.Image]], *optional*):
A conversion method to apply to the video after loading it. When set to `None` the images will be converted
to "RGB".
Returns:
`List[PIL.Image.Image]`:
The video as a list of PIL images.
"""
is_url = video.startswith("http://") or video.startswith("https://")
is_file = os.path.isfile(video)
was_tempfile_created = False
if not (is_url or is_file):
raise ValueError(
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {video} is not a valid path."
)
if is_url:
response = requests.get(video, stream=True)
if response.status_code != 200:
raise ValueError(
f"Failed to download video. Status code: {response.status_code}"
)
parsed_url = urlparse(video)
file_name = os.path.basename(unquote(parsed_url.path))
suffix = os.path.splitext(file_name)[1] or ".mp4"
with tempfile.NamedTemporaryFile(suffix=suffix,
delete=False) as temp_file:
video_path = temp_file.name
video_data = response.iter_content(chunk_size=8192)
for chunk in video_data:
temp_file.write(chunk)
video = video_path
pil_images = []
if video.endswith(".gif"):
gif = PIL.Image.open(video)
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
pass
else:
try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError:
raise AttributeError(
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
) from None
with imageio.get_reader(video) as reader:
# Read all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))
if was_tempfile_created:
os.remove(video_path)
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images
def get_default_height_width(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
vae_scale_factor: int,
+12 -1
View File
@@ -11,6 +11,7 @@ import pprint
from dataclasses import asdict, dataclass, field
from typing import Any, Dict, List, Optional, Union
import PIL.Image
import torch
from fastvideo.v1.attention import AttentionMetadata
@@ -37,6 +38,8 @@ class ForwardBatch:
# Image inputs
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
@@ -148,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
@@ -158,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,19 +2,22 @@
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
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import getdataset
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
@@ -46,7 +49,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# 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_validation_text(fastvideo_args, args)
self.preprocess_validation(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def get_extra_features(self, valid_data: Dict[str, Any],
@@ -58,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)
@@ -103,7 +273,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
@@ -114,14 +283,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataset = getdataset(args)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
@@ -215,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 = {}
@@ -233,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)
@@ -285,7 +450,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
def preprocess_validation(self, fastvideo_args: FastVideoArgs, args):
"""Process validation text prompts and save them to parquet files.
This base implementation handles the common validation text processing logic.
@@ -296,22 +461,32 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
validation_dataset = ValidationDataset(args.validation_dataset_file)
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
prompts = [sampling_param.negative_prompt] + prompts
negative_prompt = {
'caption': sampling_param.negative_prompt,
'image_path': None,
'video_path': None,
}
validation_iterable = chain([negative_prompt], validation_dataset)
else:
negative_prompt = None
validation_iterable = validation_dataset
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
pbar = tqdm(enumerate(validation_iterable),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
for idx, sample in pbar:
with torch.inference_mode():
prompt = sample["caption"]
is_negative_prompt = idx == 0
# Text Encoder
batch = ForwardBatch(
data_type="video",
@@ -338,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)
@@ -420,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
@@ -44,7 +44,7 @@ if __name__ == "__main__":
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--validation_dataset_file", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -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)
+19 -15
View File
@@ -12,8 +12,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.models.vision_utils import (get_default_height_width,
load_image, normalize,
numpy_to_pt, pil_to_numpy, resize)
normalize, numpy_to_pt,
pil_to_numpy, resize)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import V # Import validators
@@ -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 = load_image(image_path)
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],
@@ -97,7 +101,7 @@ class EncodingStage(PipelineStage):
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator[0])
latent_condition = self.retrieve_latents(encoder_output, generator)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
@@ -181,7 +185,7 @@ class EncodingStage(PipelineStage):
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage inputs."""
result = VerificationResult()
result.add_check("image_path", batch.image_path, V.string_not_empty)
# 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,
@@ -11,7 +11,6 @@ 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.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
@@ -58,7 +57,7 @@ class ImageEncodingStage(PipelineStage):
if fastvideo_args.use_cpu_offload:
self.image_encoder = self.image_encoder.to(get_torch_device())
image = load_image(batch.image_path)
image = batch.pil_image
image_inputs = self.image_processor(
images=image, return_tensors="pt").to(get_torch_device())
@@ -78,7 +77,7 @@ class ImageEncodingStage(PipelineStage):
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify image encoding stage inputs."""
result = VerificationResult()
result.add_check("image_path", batch.image_path, V.string_not_empty)
result.add_check("pil_image", batch.pil_image, V.not_none)
result.add_check("image_embeds", batch.image_embeds, V.is_list)
return result
@@ -7,6 +7,7 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import (StageValidators,
@@ -91,6 +92,11 @@ class InputValidationStage(PipelineStage):
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# for i2v, get image from image_path
if batch.image_path is not None:
image = load_image(batch.image_path)
batch.pil_image = image
return batch
def verify_input(self, batch: ForwardBatch,
@@ -76,4 +76,36 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
stage=DecodingStage(vae=self.get_module("vae")))
class WanImageToVideoValidationPipeline(ComposedPipelineBase):
"""
I2V Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
"""
_required_config_modules = ["vae", "scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanImageToVideoPipeline
@@ -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()
@@ -80,7 +80,7 @@ def run_preprocessing():
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_prompt_txt", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.txt"),
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.json"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
@@ -100,7 +100,7 @@ def run_training():
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_prompt_dir", LOCAL_VALIDATION_DATA_DIR,
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
@@ -33,7 +33,7 @@ def run_worker():
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/mini_dataset_i2v_VSA/combined_parquet_dataset",
"--validation_prompt_dir", "data/mini_dataset_i2v_VSA/validation_parquet_dataset",
"--validation_preprocessed_path", "data/mini_dataset_i2v_VSA/validation_parquet_dataset",
"--train_batch_size", "1",
"--num_latent_t", "4",
"--num_gpus", "1",
@@ -38,7 +38,7 @@ def run_worker():
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/crush-smol_parq/combined_parquet_dataset",
"--validation_prompt_dir", "data/crush-smol_parq/validation_parquet_dataset",
"--validation_preprocessed_path", "data/crush-smol_parq/validation_parquet_dataset",
"--train_batch_size", "2",
"--num_latent_t", "4",
"--num_gpus", "4",
+117 -70
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,25 +598,25 @@ 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_prompt_dir: %s',
training_args.validation_prompt_dir)
logger.info('fastvideo_args.validation_preprocessed_path: %s',
training_args.validation_preprocessed_path)
validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
training_args.validation_prompt_dir,
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)
+1 -1
View File
@@ -15,7 +15,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=4 \
--num_latent_t 20 \
--sp_size 4 \
+1 -1
View File
@@ -21,7 +21,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_prompt_dir "$VALIDATION_DIR" \
--validation_preprocessed_path "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
@@ -18,7 +18,7 @@ torchrun --nproc_per_node=$GPU_NUM \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -18,7 +18,7 @@ torchrun --nproc_per_node=$GPU_NUM \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 1 \
--flush_frequency 1 \
--video_length_tolerance_range 5 \