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