Compare commits

..
Author SHA1 Message Date
SolitaryThinker b1fb0fc130 update script 2025-06-26 00:19:03 -07:00
SolitaryThinker 8ad0752d16 fix i2v 2025-06-26 00:19:03 -07:00
SolitaryThinker ba45929acc checkpoint 2025-06-26 00:19:02 -07:00
SolitaryThinker f893fc9e4a udate 2025-06-26 00:19:02 -07:00
SolitaryThinker a366c98194 update 2025-06-26 00:19:01 -07:00
SolitaryThinker 79b5d82d98 fix sp 2025-06-26 00:19:01 -07:00
SolitaryThinker 522b224c44 update i2v script 2025-06-26 00:19:00 -07:00
SolitaryThinker 29ae201575 improve script format 2025-06-26 00:19:00 -07:00
SolitaryThinker 4af9d06794 remove print and enable first val 2025-06-26 00:19:00 -07:00
SolitaryThinker f3f12bf2d0 update scripts 2025-06-26 00:18:59 -07:00
SolitaryThinker 528063a43a i2v working 2025-06-26 00:18:59 -07:00
SolitaryThinker 5fc25d43e6 t2v working again 2025-06-26 00:18:58 -07:00
SolitaryThinker c6b601ed33 t2v example 2025-06-26 00:18:57 -07:00
SolitaryThinker c620f8549e f 2025-06-26 00:18:57 -07:00
SolitaryThinker 1212ed991f update 2025-06-26 00:18:56 -07:00
SolitaryThinker 7419829ebc update 2025-06-26 00:18:56 -07:00
SolitaryThinker 82128e5b25 slrm 2025-06-26 00:18:56 -07:00
SolitaryThinker 9a160c3230 update path 2025-06-26 00:18:55 -07:00
SolitaryThinker e4eeae93e6 exmaple scripts 2025-06-26 00:18:55 -07:00
SolitaryThinker c966ffc2fc cleanup 2025-06-26 00:18:55 -07:00
SolitaryThinker db7037a39f fix pil image 2025-06-26 00:18:54 -07:00
SolitaryThinker ef7788d2a5 i2v preprocess 2025-06-26 00:18:53 -07:00
SolitaryThinker 3e55580906 checkpoint 2025-06-26 00:18:53 -07:00
SolitaryThinker 72e21cfe37 checkpoint 2025-06-26 00:18:52 -07:00
81 changed files with 536 additions and 927 deletions
+21 -93
View File
@@ -1,138 +1,66 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
- wait
- block: "Start Build"
blocked_state: "running"
prompt: "Approve build?"
- label: "Trigger Tests"
command: |
echo "Current working directory: $(pwd)"
echo "Current branch:"
git branch --show-current
echo "Full diff:"
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
plugins:
- monorepo-diff#v1.4.0:
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
watch:
- path:
- "fastvideo/v1/models/encoders/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/encoders/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=encoder
agents:
queue: "default"
- path:
- "fastvideo/v1/models/vaes/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/vaes/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=vae
agents:
queue: "default"
- path:
- "fastvideo/v1/models/dits/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/transformers/**"
- "fastvideo/v1/layers/**"
- "fastvideo/v1/attention/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
- path: "fastvideo/v1/**/*.py"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- TEST_TYPE=training
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- TEST_TYPE=training_vsa
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- TEST_TYPE=inference_sta
agents:
queue: "default"
- path:
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- TEST_TYPE=precision_vsa
agents:
queue: "default"
+4 -30
View File
@@ -31,10 +31,6 @@ log "Setting up Modal authentication from Buildkite secrets..."
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
log "Retrieved Modal credentials from Buildkite secrets"
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
@@ -58,44 +54,22 @@ if [ -z "${TEST_TYPE:-}" ]; then
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
;;
"vae")
log "Running VAE tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
;;
"transformer")
log "Running transformer tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
;;
"training_vsa")
log "Running training VSA tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
;;
"inference_sta")
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"precision_sta")
log "Running precision STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
;;
"precision_vsa")
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
-40
View File
@@ -1,40 +0,0 @@
#!/bin/bash
set -uo pipefail
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
log "=== Starting pre-commit checks ==="
cd "$(dirname "$0")/../.."
PROJECT_ROOT=$(pwd)
log "Project root: $PROJECT_ROOT"
if ! python3 -m pre_commit --version &> /dev/null; then
log "pre-commit not found, installing..."
python3 -m pip install --user pre-commit==4.0.1
if ! python3 -m pre_commit --version &> /dev/null; then
log "Error: Failed to install pre-commit."
exit 1
fi
fi
log "Pre-commit version: $(python3 -m pre_commit --version)"
log "Installing/updating pre-commit hooks..."
python3 -m pre_commit install --install-hooks
log "Running pre-commit checks on all files..."
python3 -m pre_commit run --all-files
PRE_COMMIT_EXIT_CODE=$?
if [ $PRE_COMMIT_EXIT_CODE -eq 0 ]; then
log "Pre-commit checks completed successfully"
else
log "Error: Pre-commit checks failed with exit code: $PRE_COMMIT_EXIT_CODE"
fi
log "=== Pre-commit checks completed with exit code: $PRE_COMMIT_EXIT_CODE ==="
exit $PRE_COMMIT_EXIT_CODE
+2 -2
View File
@@ -264,7 +264,7 @@ jobs:
with:
job_id: "training-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 2
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
@@ -284,7 +284,7 @@ jobs:
with:
job_id: "inference-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 2
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
-14
View File
@@ -15,11 +15,6 @@ With FastVideo's optimizations, you can achieve more than 3x inference improveme
<img src=assets/perf.png width="90%"/>
</div>
## NEWS
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
## Key Features
FastVideo has the following features:
@@ -133,15 +128,6 @@ We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support thro
If you use FastVideo for your research, please cite our paper:
```bibtex
@misc{zhang2025vsafastervideodiffusion,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
year={2025},
eprint={2505.13389},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2505.13389},
}
@misc{zhang2025fastvideogenerationsliding,
title={Fast Video Generation with Sliding Tile Attention},
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
@@ -1,10 +0,0 @@
This directory contain e2e examples scripts for finetuning Wan2.1 I2V.
Execute the following commands from `FastVideo/` to run training:
- Download crush-smol dataset:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
- Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
- Edit the following file and run finetuning:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,10 +1,7 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
NUM_GPUS=8
@@ -15,7 +12,7 @@ NUM_GPUS=8
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 2000
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
@@ -36,8 +33,8 @@ parallel_args=(
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--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
@@ -59,7 +56,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--checkpointing_steps 6000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -69,7 +66,7 @@ miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--cfg 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
@@ -1,5 +1,5 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
@@ -9,8 +9,8 @@
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
@@ -30,9 +30,7 @@ nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
@@ -40,91 +38,60 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
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
# Training arguments
training_args=(
--tracker_project_name wan_i2v_finetune
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"
--max_train_steps=2000
--train_batch_size=2
--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 $SLURM_JOB_NUM_NODES
--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 10
)
# 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=1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
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 \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
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=2\
--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 10\
--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
@@ -1,5 +1,4 @@
#!/bin/bash
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
@@ -22,5 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "i2v"
@@ -1,10 +0,0 @@
This directory contain e2e examples scripts for finetuning Wan2.1 T2v.
Execute the following commands from `FastVideo/` to run training:
- Download crush-smol dataset:
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
- Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
- Edit the following file and run finetuning:
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,5 +1,3 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
@@ -27,11 +25,11 @@ training_args=(
# 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
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim $NUM_GPUS \
)
# Model arguments
@@ -69,7 +67,7 @@ miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--cfg 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
@@ -7,10 +7,10 @@
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=t2v_output/t2v_%j.out
#SBATCH --error=t2v_output/t2v_%j.err
#SBATCH --output=1n_t2v/1n_t2v_%j.out
#SBATCH --error=1n_t2v/1n_t2v_%j.err
#SBATCH --exclusive
set -e -x
@@ -30,98 +30,69 @@ nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
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]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_finetune
--output_dir="outputs/wan_t2v_finetune"
--max_train_steps=1000
--train_batch_size=4
--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 4
--tp_size 4
--hsdp_replicate_dim 2
--hsdp_shard_dim 4
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 10
)
# 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=5e-5
--mixed_precision="bf16"
--checkpointing_steps=500
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
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 \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
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=4\
--num_latent_t 8 \
--num_gpus $NUM_GPUS \
--sp_size 4 \
--tp_size 4 \
--hsdp_replicate_dim 2 \
--hsdp_shard_dim 4 \
--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 50\
--validation_sampling_steps "50" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_t2v_finetune_2n"\
--tracker_project_name wan_t2v_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
}
]
}
@@ -1,5 +1,4 @@
#!/bin/bash
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
@@ -22,5 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
+1 -1
View File
@@ -52,7 +52,7 @@ class PipelineConfig:
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp32"
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
+1 -1
View File
@@ -50,7 +50,7 @@ class WanT2V480PConfig(PipelineConfig):
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp32"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
@@ -4,7 +4,7 @@
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
@@ -4,7 +4,7 @@
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
@@ -8,7 +8,6 @@ import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dist_cp
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.v1.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.v1.distributed import get_world_rank
@@ -68,18 +67,14 @@ def main() -> None:
# Create DataLoader with proper settings
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader with %d batches", len(dataloader))
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
@@ -105,18 +100,14 @@ def main() -> None:
# Recreate dataloader and load state
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
args.path, args.batch_size, args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
@@ -125,16 +116,11 @@ def main() -> None:
break
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
args.path, args.batch_size, args.num_data_workers)
# Second pass - verify latent sums match
second_pass_sums = []
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
for i, (latents, embeddings, masks) in enumerate(dataloader):
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
@@ -158,9 +144,8 @@ def main() -> None:
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
@@ -185,6 +185,9 @@ class LatentsParquetMapStyleDataset(Dataset):
Note:
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", "clip_feature",
"first_frame_latent", "pil_image"]
def __init__(
self,
@@ -201,6 +204,10 @@ class LatentsParquetMapStyleDataset(Dataset):
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"
)
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
path)
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
@@ -236,8 +243,7 @@ class LatentsParquetMapStyleDataset(Dataset):
batch = collate_rows_from_parquet_schema([row_dict],
self.parquet_schema,
self.text_padding_length,
cfg_rate=0.0)
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']
@@ -259,10 +265,11 @@ class LatentsParquetMapStyleDataset(Dataset):
for idx in indices
]
batch = collate_rows_from_parquet_schema(rows,
self.parquet_schema,
self.text_padding_length,
cfg_rate=self.cfg_rate)
# 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):
@@ -484,7 +484,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate)
cfg_rate=args.cfg)
def _load_raw_data(self) -> List[Dict]:
"""Load raw data from JSON files."""
@@ -493,13 +493,10 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
folder_anno_pairs = [
line.strip().split(",") for line in f if line.strip()
]
assert len(
folder_anno_pairs) == 1, "Only support one folder-annotation pair"
assert len(folder_anno_pairs[0]
) == 2, "Folder-annotation pair should have two elements"
assert len(folder_anno_pairs) == 1, "Only support one folder-annotation pair"
assert len(folder_anno_pairs[0]) == 2, "Folder-annotation pair should have two elements"
folder, annotation_file = folder_anno_pairs[0]
data_items: List[Dict] = []
with open(annotation_file) as f:
data_items = json.load(f)
+94 -67
View File
@@ -1,9 +1,12 @@
import random
from typing import Any, Dict, List, cast
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:
"""
@@ -21,7 +24,7 @@ def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
return t[:padding_length], torch.ones(padding_length)
def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
"""
Get the latents and prompts from a row dictionary.
"""
@@ -39,45 +42,70 @@ def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> 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
if key == 'text_embedding' and random.random() < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
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
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,
cfg_rate=0.0
) -> 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, cfg_rate)
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"])
@@ -88,14 +116,17 @@ 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,
cfg_rate=0.0) -> Dict[str, Any]:
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.
@@ -108,10 +139,10 @@ def collate_rows_from_parquet_schema(rows,
Dict containing batched tensors and metadata
"""
if not rows:
return cast(Dict[str, Any], {})
return {}
# Initialize containers for different data types
batch_data: Dict[str, Any] = {}
batch_data = {}
# Get tensor and metadata field names from schema (fields ending with '_bytes')
tensor_fields = []
@@ -128,7 +159,7 @@ def collate_rows_from_parquet_schema(rows,
# Only add actual metadata fields, not the shape/dtype helper fields
metadata_fields.append(field)
# Process each tensor field
# Process each tensor field efficiently
for tensor_name in tensor_fields:
tensor_list = []
@@ -138,6 +169,9 @@ def collate_rows_from_parquet_schema(rows,
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]
@@ -145,12 +179,11 @@ def collate_rows_from_parquet_schema(rows,
tensor = torch.zeros(0, dtype=torch.bfloat16)
else:
# Convert bytes to tensor using float32 as default
if tensor_name == 'text_embedding' and random.random(
) < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
else:
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
# 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
@@ -163,44 +196,38 @@ def collate_rows_from_parquet_schema(rows,
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_name == 'text_embedding':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
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))
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 all tensors to preserve batch consistency
# Don't filter out None or empty tensors as this breaks batch sizing
try:
batch_data[tensor_name] = torch.stack(tensor_list)
except ValueError as e:
shapes = [
t.shape
if t is not None and hasattr(t, 'shape') else 'None/Invalid'
for t in tensor_list
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
]
raise ValueError(
f"Failed to stack tensors for field '{tensor_name}'. "
f"Tensor shapes: {shapes}. "
f"All tensors in a batch must have compatible shapes. "
f"Original error: {e}") from e
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 into info_list
# Process metadata fields efficiently into info_list
info_list = []
for row in rows:
info = {}
@@ -19,6 +19,7 @@ class ValidationDataset(torch.utils.data.IterableDataset):
self.filename = pathlib.Path(filename)
# get directory of filename
# TODO(will)
self.dir = os.path.abspath(self.filename.parent)
if not self.filename.exists():
+6 -8
View File
@@ -384,7 +384,7 @@ class TrainingArgs(FastVideoArgs):
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
training_cfg_rate: float = 0.0
cfg: float = 0.0
precondition_outputs: bool = False
# validation & logs
@@ -413,7 +413,7 @@ class TrainingArgs(FastVideoArgs):
lr_scheduler: str = "constant"
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
enable_gradient_checkpointing_type: Optional[str] = None
gradient_checkpointing: bool = False
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
@@ -528,7 +528,7 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--training-cfg-rate",
parser.add_argument("--cfg",
type=float,
help="Classifier-free guidance scale")
parser.add_argument(
@@ -612,11 +612,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--max-grad-norm",
type=float,
help="Maximum gradient norm")
parser.add_argument("--enable-gradient-checkpointing-type",
type=str,
choices=["full", "ops", "block_skip"],
default=None,
help="Gradient checkpointing type")
parser.add_argument("--gradient-checkpointing",
action=StoreBoolean,
help="Whether to use gradient checkpointing")
parser.add_argument("--selective-checkpointing",
type=float,
help="Selective checkpointing threshold")
-6
View File
@@ -38,12 +38,6 @@ class RMSNorm(CustomOp):
if self.has_weight:
self.weight = nn.Parameter(self.weight)
# if we do fully_shard(model.layer_norm), and we call layer_form.forward_native(input) instead of layer_norm(input),
# we need to call model.layer_norm.register_fsdp_forward_method(model, "forward_native") to make sure fsdp2 hooks are triggered
# for mixed precision and cpu offloading
# the even better way might be fully_shard(model.layer_norm, mp_policy=, cpu_offloading=), and call model.layer_norm(input). everything should work out of the box
# because fsdp2 hooks will be triggered with model.layer_norm.__call__
def forward_native(
self,
x: torch.Tensor,
+6 -7
View File
@@ -71,16 +71,15 @@ class BaseLayerWithLoRA(nn.Module):
f"cuda:{torch.cuda.current_device()}").full_tensor()
data += (self.slice_lora_b_weights(self.lora_B)
@ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh,
placements=placements).to(current_device))
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placements).to(current_device)
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.to(
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}")
data += \
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight = nn.Parameter(data.to(current_device))
self.base_layer.weight.data = data.to(current_device)
self.merged = True
@torch.no_grad()
@@ -107,8 +106,8 @@ class BaseLayerWithLoRA(nn.Module):
f"cuda:{torch.cuda.current_device()}").full_tensor()
data -= self.slice_lora_b_weights(
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh, placements=placement).to(device))
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placement).to(device)
else:
self.base_layer.weight.data -= \
self.slice_lora_b_weights(self.lora_B) @\
+3
View File
@@ -25,9 +25,12 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class WanImageEmbedding(torch.nn.Module):
@@ -149,9 +149,11 @@ class TrainingBatch:
# Dataloader batch outputs
latents: Optional[torch.Tensor] = None
# original_latents: Optional[torch.Tensor] = None
encoder_hidden_states: Optional[torch.Tensor] = None
encoder_attention_mask: Optional[torch.Tensor] = 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
@@ -104,7 +104,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if strict:
raise ValueError(
f"Failed to convert tensor {tensor_name} to bytes: {e}"
) from e
)
record[field] = b'' # Empty bytes for missing data
else:
record[field] = b'' # Empty bytes for missing data
@@ -139,8 +139,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to int: {e}"
) from e
f"Failed to convert field {field} to int: {e}")
record[field] = 0
else:
record[field] = 0
@@ -158,7 +157,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if strict:
raise ValueError(
f"Failed to convert field {field} to float: {e}"
) from e
)
record[field] = 0.0
else:
record[field] = 0.0
@@ -212,8 +211,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Log unfilled fields as warning if not in strict mode
if unfilled_fields:
logger.warning(
"Some fields were not filled and got default values: %s",
unfilled_fields)
f"Some fields were not filled and got default values: {unfilled_fields}"
)
return record
@@ -222,6 +221,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
# text_attention_mask: np.ndarray,
valid_data: Dict[str, Any],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
@@ -380,6 +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)
# Get extra features for this sample if needed
sample_extra_features = {}
@@ -396,6 +398,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
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=sample_extra_features)
@@ -540,13 +543,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
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,
valid_data=valid_data,
idx=0,
extra_features=sample_extra_features)
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)
@@ -58,6 +58,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
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,
@@ -83,6 +84,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
# 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 = {}
@@ -185,16 +187,19 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: Dict[str, Any],
# 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,
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"]
@@ -84,7 +84,7 @@ if __name__ == "__main__":
type=str,
default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--training_cfg_rate", type=float, default=0.0)
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
+1 -3
View File
@@ -68,10 +68,8 @@ class EncodingStage(PipelineStage):
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
else:
# assumes image is loaded from parquet file and used for validation
image = image.transpose(1, 2)
logger.info("image: %s", image.shape)
image = image.transpose(1, 2)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
@@ -168,5 +168,5 @@ def test_clip_encoder():
f"Pooler outputs differ significantly: mean diff = {mean_diff_pooler.item()}"
assert max_diff_hidden < 1e-1, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
assert max_diff_pooler < 2e-2, \
assert max_diff_pooler < 1e-2, \
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
@@ -5,7 +5,7 @@ from pathlib import Path
import pytest
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
NUM_GPUS_PER_NODE = "1"
# Set environment variables
os.environ["FASTVIDEO_ATTENTION_CONFIG"] = "assets/mask_strategy_wan.json"
@@ -17,9 +17,9 @@ def test_inference():
cmd = [
"fastvideo", "generate",
"--model-path", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
"--sp-size", "2",
"--tp-size", "2",
"--num-gpus", "2",
"--sp-size", "1",
"--tp-size", "1",
"--num-gpus", "1",
"--height", "768",
"--width", "1280",
"--num-frames", "69",
+64 -65
View File
@@ -4,56 +4,31 @@ app = modal.App()
import os
image_version = os.getenv("IMAGE_VERSION")
image_version = os.getenv("IMAGE_VERSION", "latest")
image_tag = f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{image_version}"
print(f"Using image: {image_tag}")
image = (
modal.Image.from_registry(image_tag, add_python="3.12")
.run_commands("rm -rf /FastVideo")
.apt_install("cmake", "pkg-config", "build-essential", "curl", "libssl-dev")
.run_commands("curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable")
.run_commands("echo 'source ~/.cargo/env' >> ~/.bashrc")
.env({
"PATH": "/root/.cargo/bin:$PATH",
"BUILDKITE_REPO": os.environ.get("BUILDKITE_REPO", ""),
"BUILDKITE_COMMIT": os.environ.get("BUILDKITE_COMMIT", ""),
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
"IMAGE_VERSION": os.environ.get("IMAGE_VERSION", ""),
})
.env({"PATH": "/root/.cargo/bin:$PATH"})
.run_commands("/bin/bash -c 'source $HOME/.local/bin/env && source /opt/venv/bin/activate && cd /FastVideo && uv pip install -e .[test]'")
)
def run_test(pytest_command: str):
"""Helper function to run a test suite with custom pytest command"""
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_encoder_tests():
"""Run encoder tests on L40S GPU"""
import subprocess
import sys
import os
git_repo = os.environ.get("BUILDKITE_REPO")
git_commit = os.environ.get("BUILDKITE_COMMIT")
pr_number = os.environ.get("BUILDKITE_PULL_REQUEST")
os.chdir("/FastVideo")
print(f"Cloning repository: {git_repo}")
print(f"Target commit: {git_commit}")
if pr_number:
print(f"PR number: {pr_number}")
# For PRs (including forks), use GitHub's PR refs to get the correct commit
if pr_number and pr_number != "false":
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
print(f"Using PR ref for checkout: {checkout_command}")
else:
checkout_command = f"git checkout {git_commit}"
print(f"Using direct commit checkout: {checkout_command}")
command = f"""
source $HOME/.local/bin/env &&
source /opt/venv/bin/activate &&
git clone {git_repo} /FastVideo &&
cd /FastVideo &&
{checkout_command} &&
uv pip install -e .[test] &&
{pytest_command}
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/encoders -s
"""
result = subprocess.run([
@@ -62,38 +37,62 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_encoder_tests():
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
@app.function(gpu="L40S:1", image=image, timeout=900)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_vae_tests():
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
"""Run VAE tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/vaes -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=900)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_transformer_tests():
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
"""Run transformer tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/transformers -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:2", image=image, timeout=1800)
@app.function(gpu="L40S:2", image=image, timeout=3600)
def run_ssim_tests():
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/Vanilla -srP")
@app.function(gpu="H100:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests_VSA():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/VSA -srP")
@app.function(gpu="H100:2", image=image, timeout=900)
def run_inference_tests_STA():
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
@app.function(gpu="H100:1", image=image, timeout=900)
def run_precision_tests_STA():
run_test("python csrc/attn/tests/test_sta.py")
@app.function(gpu="H100:1", image=image, timeout=900)
def run_precision_tests_VSA():
run_test("python csrc/attn/tests/test_block_sparse.py")
"""Run SSIM tests on 2x L40S GPUs"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/ssim -vs
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@@ -122,7 +122,7 @@ def run_training():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.1",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
"--num_height", "480",
@@ -31,9 +31,12 @@ 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, exist_ok=True)
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
@@ -119,7 +122,7 @@ def run_training():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_finetune_overfit_ci",
"--num_height", "480",
+6 -7
View File
@@ -1,14 +1,13 @@
The reference videos in the `*_reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
The reference videos in the `reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
`A40_reference_videos` are generated on A40s and so on.
run `bash update_reference_videos.sh` from inside the `fastvideo/v1/tests/ssim/` directory after running `test_inference_similarity.py` to update reference videos. Note: make sure to update the path to the corresponding device.
all reference videos are were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
`reference_videos/FastHunyuan-diffusers/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
`reference_videos/FastHunyuan-diffusers/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
`reference_videos/Wan2.1-T2V-1.3B-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
`reference_videos/Wan2.1-I2V-14B-480P-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
## Generation Details
2 x NVIDIA L40S GPUs
2 x NVIDIA A40 GPUs
## Generation Parameters
@@ -2,7 +2,6 @@
import json
import os
import torch
import pytest
from fastvideo import VideoGenerator
@@ -12,14 +11,6 @@ from fastvideo.v1.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = '_reference_videos'
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
# Base parameters from the shell script
HUNYUAN_PARAMS = {
"num_gpus": 2,
@@ -197,8 +188,8 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
@@ -297,8 +288,8 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
@@ -1,63 +0,0 @@
#!/bin/bash
# Script to update reference videos using videos from generated_videos directory
# Both directories should exist in the same directory as this script
set -e # Exit on any error
# Define directory paths
GENERATED_DIR="generated_videos"
REFERENCE_DIR="set_me_to_correct_path"
# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
NC='\033[0m' # No Color
echo -e "${YELLOW}Starting reference video update...${NC}"
# Check if generated_videos directory exists
if [ ! -d "$GENERATED_DIR" ]; then
echo -e "${RED}Error: $GENERATED_DIR directory not found!${NC}"
exit 1
fi
# Check if reference_videos directory exists
if [ ! -d "$REFERENCE_DIR" ]; then
echo -e "${RED}Error: $REFERENCE_DIR directory not found!${NC}"
exit 1
fi
# Function to copy videos recursively
copy_videos() {
local src_dir="$1"
local dst_dir="$2"
# Find all video files in the source directory
find "$src_dir" -type f \( -name "*.mp4" -o -name "*.avi" -o -name "*.mov" -o -name "*.mkv" -o -name "*.webm" -o -name "*.flv" \) | while read -r video_file; do
# Get relative path from source directory
relative_path="${video_file#$src_dir/}"
# Construct destination path
dst_file="$dst_dir/$relative_path"
# Create destination directory if it doesn't exist
dst_file_dir=$(dirname "$dst_file")
mkdir -p "$dst_file_dir"
# Copy the video file
echo -e "${GREEN}Copying: $relative_path${NC}"
cp "$video_file" "$dst_file"
done
}
# Perform the copy operation
echo -e "${YELLOW}Copying videos from $GENERATED_DIR to $REFERENCE_DIR...${NC}"
copy_videos "$GENERATED_DIR" "$REFERENCE_DIR"
echo -e "${GREEN}Reference videos updated successfully!${NC}"
# Show summary
video_count=$(find "$GENERATED_DIR" -type f \( -name "*.mp4" -o -name "*.avi" -o -name "*.mov" -o -name "*.mkv" -o -name "*.webm" -o -name "*.flv" \) | wc -l)
echo -e "${YELLOW}Total videos processed: $video_count${NC}"
@@ -1 +1 @@
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
{"grad_norm":0.478515625,"_runtime":95.727033597,"_wandb":{"runtime":95},"_step":5,"validation_videos_50_steps":{"videos":[{"_type":"video-file","sha256":"42a1c311521a9d460db788713be1cbf2db767494e02619b43be5bf3eed8381d8","size":158632,"path":"media/videos/validation_videos_50_steps_0_42a1c311521a9d460db7.mp4"},{"path":"media/videos/validation_videos_50_steps_0_818505095b4b5e8b7f51.mp4","_type":"video-file","sha256":"818505095b4b5e8b7f511012d45f04d151ce3344bc058fc0f3225a414a851e4a","size":147825},{"sha256":"fc334ba9ed5e66c8527ee3b408e3be2d76167fef03588bf2840f4a0792f2fe34","size":136933,"path":"media/videos/validation_videos_50_steps_0_fc334ba9ed5e66c8527e.mp4","_type":"video-file"},{"size":201797,"path":"media/videos/validation_videos_50_steps_0_ccd98f6f907635d266a7.mp4","_type":"video-file","sha256":"ccd98f6f907635d266a74783688e7ecf1dac752d79d72d69eab9ef0e3f7413eb"},{"_type":"video-file","sha256":"ca79f40a0aed38f676f12779b349ce40e9e3fb7f36c578f49a20854c70508fb4","size":147114,"path":"media/videos/validation_videos_50_steps_0_ca79f40a0aed38f676f1.mp4"},{"size":175104,"path":"media/videos/validation_videos_50_steps_0_32c9b33ff920c17e5881.mp4","_type":"video-file","sha256":"32c9b33ff920c17e588133d7a27aa400ff3dc529b01ed4f16ac4d6bb2afa0f00"},{"sha256":"2cf520bfb93401c914e93c87ef791c2f12a4e043b95dfdc98115c930e11dfe67","size":139655,"path":"media/videos/validation_videos_50_steps_0_2cf520bfb93401c914e9.mp4","_type":"video-file"},{"_type":"video-file","sha256":"1d73aba17ce582c7aef4af4d64079e3e9d3df205634eff453446bdaf2340b214","size":149028,"path":"media/videos/validation_videos_50_steps_0_1d73aba17ce582c7aef4.mp4"}],"captions":false,"_type":"videos","count":8},"train_loss":0.08922439813613892,"_timestamp":1.750202051751466e+09,"avg_step_time":0.7536672964692116,"step_time":0.4742048177868128,"learning_rate":1e-05,"vsa_sparsity":0.05}
@@ -15,7 +15,7 @@ wandb_name = "test_training_loss_VSA"
reference_wandb_summary_file = "fastvideo/v1/tests/training/VSA/reference_wandb_summary_VSA.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
NUM_GPUS_PER_NODE = "1"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
@@ -31,18 +31,19 @@ def run_worker():
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/mini_dataset_i2v_VSA/combined_parquet_dataset",
"--validation_preprocessed_path", "data/mini_dataset_i2v_VSA/validation_parquet_dataset",
"--train_batch_size", "1",
"--num_latent_t", "4",
"--num_gpus", "2",
"--sp_size", "2",
"--tp_size", "2",
"--num_gpus", "1",
"--sp_size", "1",
"--tp_size", "1",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "2",
"--hsdp_shard_dim", "1",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "4",
"--gradient_accumulation_steps", "2",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "5",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
@@ -53,7 +54,7 @@ def run_worker():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test_VSA",
"--tracker_project_name", "wan_finetune_ci_VSA",
"--wandb_run_name", wandb_name,
@@ -110,7 +111,7 @@ def test_distributed_training():
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 1.0,
'step_time': 0.5,
'train_loss': 0.001
}
@@ -1 +0,0 @@
{"step_time":5.501357046999999,"grad_norm":0.384765625,"train_loss":0.07890288904309273,"avg_step_time":5.831571423200001}
@@ -1 +1 @@
{"_timestamp":1.7496170016478686e+09,"validation_videos_8_steps":{"_type":"videos","count":5,"videos":[{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"},{"sha256":"cd72a3d513eca6b41b03b80e6fa044ce7219c35e969d2ca20b9cb48c91e585c6","size":477837,"path":"media/videos/validation_videos_8_steps_0_cd72a3d513eca6b41b03.mp4","_type":"video-file"},{"_type":"video-file","sha256":"43d47c211a69bf0be3544738e76e7d8fa58bb108c00cb21dc34eaad3c6ce7cc3","size":409419,"path":"media/videos/validation_videos_8_steps_0_43d47c211a69bf0be354.mp4"},{"_type":"video-file","sha256":"ea674ec9e200bc97563c9d87d9dc07110c3f42c0ab7277dd237ae96dd8f90a10","size":333966,"path":"media/videos/validation_videos_8_steps_0_ea674ec9e200bc97563c.mp4"},{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"}],"captions":false},"step_time":2.5065076276659966,"_wandb":{"runtime":53},"learning_rate":1e-06,"_step":5,"_runtime":53.172758961,"grad_norm":0.408203125,"train_loss":0.07883700542151928,"avg_step_time":2.8116052336990833}
{"_timestamp":1.7496170016478686e+09,"validation_videos_8_steps":{"_type":"videos","count":5,"videos":[{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"},{"sha256":"cd72a3d513eca6b41b03b80e6fa044ce7219c35e969d2ca20b9cb48c91e585c6","size":477837,"path":"media/videos/validation_videos_8_steps_0_cd72a3d513eca6b41b03.mp4","_type":"video-file"},{"_type":"video-file","sha256":"43d47c211a69bf0be3544738e76e7d8fa58bb108c00cb21dc34eaad3c6ce7cc3","size":409419,"path":"media/videos/validation_videos_8_steps_0_43d47c211a69bf0be354.mp4"},{"_type":"video-file","sha256":"ea674ec9e200bc97563c9d87d9dc07110c3f42c0ab7277dd237ae96dd8f90a10","size":333966,"path":"media/videos/validation_videos_8_steps_0_ea674ec9e200bc97563c.mp4"},{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"}],"captions":false},"step_time":2.5065076276659966,"_wandb":{"runtime":53},"learning_rate":1e-06,"_step":5,"_runtime":53.172758961,"grad_norm":5.65625,"train_loss":0.3915919363498688,"avg_step_time":2.8116052336990833}
@@ -18,8 +18,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
wandb_name = "test_training_loss"
a40_reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/a40_reference_wandb_summary.json"
l40s_reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/l40s_reference_wandb_summary.json"
reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/reference_wandb_summary.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "4"
@@ -37,8 +36,9 @@ def run_worker():
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--data_path", "data/crush-smol_processed_t2v/combined_parquet_dataset",
"--validation_preprocessed_path", "data/crush-smol_processed_t2v/validation_parquet_dataset",
"--cache_dir", "/home/.cache",
"--data_path", "data/crush-smol_parq/combined_parquet_dataset",
"--validation_preprocessed_path", "data/crush-smol_parq/validation_parquet_dataset",
"--train_batch_size", "2",
"--num_latent_t", "4",
"--num_gpus", "4",
@@ -59,7 +59,7 @@ def run_worker():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test",
"--tracker_project_name", "wan_finetune_ci",
"--wandb_run_name", wandb_name,
@@ -83,12 +83,12 @@ def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "online"
data_dir = Path("data/crush-smol_processed_t2v")
data_dir = Path("data/crush-smol_parq")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="wlsaidhi/crush-smol_processed_t2v",
repo_id="PY007/crush-smol",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
@@ -109,21 +109,13 @@ def test_distributed_training():
summary_file = 'wandb/latest-run/files/wandb-summary.json'
device_name = torch.cuda.get_device_name()
if "A40" in device_name:
reference_wandb_summary_file = a40_reference_wandb_summary_file
elif "L40S" in device_name:
reference_wandb_summary_file = l40s_reference_wandb_summary_file
else:
raise ValueError(f"Unknown device: {device_name}")
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 6.0,
'grad_norm': 0.3,
'step_time': 6.0,
'avg_step_time': 1.0,
'grad_norm': 0.2,
'step_time': 0.5,
'train_loss': 0.0025
}
@@ -92,7 +92,7 @@ def test_hunyuanvideo_distributed():
# Move to GPU based on local rank (0 or 1 for 2 GPUs)
device = torch.device(f"cuda:0")
model = model
model = model.to(device)
batch_size = 1
seq_len = 3
@@ -62,7 +62,7 @@ def test_hunyuanvideo_distributed():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=True,
use_cpu_offload=False,
pipeline_config=PipelineConfig(dit_config=HunyuanVideoConfig(), dit_precision=precision_str))
args.device = torch.device(f"cuda:{LOCAL_RANK}")
@@ -34,12 +34,12 @@ def test_wan_transformer():
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=True,
use_cpu_offload=False,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(dtype=precision)
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
model1 = WanTransformer3DModel.from_pretrained(
TRANSFORMER_PATH, device=device,
+2 -14
View File
@@ -28,12 +28,8 @@ MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
VAE_PATH = os.path.join(MODEL_PATH, "vae")
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
# Latent generated on commit d71a4ebffc2034922fc379568b6a6aa722f3744c with 1 x A40
# torch 2.7.1
A40_REFERENCE_LATENT = -106.22467041015625
# Latent generated on commit 2b54068960c41d42221e8b8719a374b499855029 with 1 x L40S
L40S_REFERENCE_LATENT = -158.32318115234375
# Latent generated on commit 250f0b916cebb18a1c15c4aae1a0b480604d066a with 1 x A40
REFERENCE_LATENT = -105.51324462890625
@pytest.mark.usefixtures("distributed_setup")
@@ -70,14 +66,6 @@ def test_hunyuan_vae():
latent = model.encode(input_tensor).mean.double().sum().item()
# Check if latents are similar
device_name = torch.cuda.get_device_name()
if "A40" in device_name:
REFERENCE_LATENT = A40_REFERENCE_LATENT
elif "L40S" in device_name:
REFERENCE_LATENT = L40S_REFERENCE_LATENT
else:
raise ValueError(f"Unknown device: {device_name}")
diff_encoded_latents = abs(REFERENCE_LATENT - latent)
logger.info(
f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}"
@@ -1,91 +0,0 @@
import collections
from enum import Enum
from typing import Optional
import torch
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
checkpoint_wrapper)
TRANSFORMER_BLOCK_NAMES = [
"blocks",
"double_blocks",
"single_blocks",
"transformer_blocks",
"temporal_transformer_blocks",
"transformer_double_blocks",
"transformer_single_blocks",
]
class CheckpointType(str, Enum):
FULL = "full"
OPS = "ops"
BLOCK_SKIP = "block_skip"
_SELECTIVE_ACTIVATION_CHECKPOINTING_OPS = {
torch.ops.aten.mm.default,
torch.ops.aten._scaled_dot_product_efficient_attention.default,
torch.ops.aten._scaled_dot_product_flash_attention.default,
torch.ops._c10d_functional.reduce_scatter_tensor.default,
}
def apply_activation_checkpointing(
module: torch.nn.Module,
checkpointing_type: str = CheckpointType.FULL,
n_layer: int = 1) -> torch.nn.Module:
if checkpointing_type == CheckpointType.FULL:
module = _apply_activation_checkpointing_blocks(module)
elif checkpointing_type == CheckpointType.OPS:
module = _apply_activation_checkpointing_ops(
module, _SELECTIVE_ACTIVATION_CHECKPOINTING_OPS)
elif checkpointing_type == CheckpointType.BLOCK_SKIP:
module = _apply_activation_checkpointing_blocks(module, n_layer)
else:
raise ValueError(
f"Checkpointing type '{checkpointing_type}' not supported. Supported types are {CheckpointType.__members__.keys()}"
)
return module
def _apply_activation_checkpointing_blocks(
module: torch.nn.Module,
n_layer: Optional[int] = None) -> torch.nn.Module:
for transformer_block_name in TRANSFORMER_BLOCK_NAMES:
blocks: torch.nn.Module = getattr(module, transformer_block_name, None)
if blocks is None:
continue
for index, (layer_id, block) in enumerate(blocks.named_children()):
if n_layer is None or index % n_layer == 0:
block = checkpoint_wrapper(block, preserve_rng_state=False)
blocks.register_module(layer_id, block)
return module
def _apply_activation_checkpointing_ops(module: torch.nn.Module,
ops) -> torch.nn.Module:
from torch.utils.checkpoint import (CheckpointPolicy,
create_selective_checkpoint_contexts)
def _get_custom_policy(meta: dict[str, int]) -> CheckpointPolicy:
def _custom_policy(ctx, func, *args, **kwargs):
mode = "recompute" if ctx.is_recompute else "forward"
mm_count_key = f"{mode}_mm_count"
if func == torch.ops.aten.mm.default:
meta[mm_count_key] += 1
# Saves output of all compute ops, except every second mm
to_save = func in ops and not (func == torch.ops.aten.mm.default
and meta[mm_count_key] % 2 == 0)
return CheckpointPolicy.MUST_SAVE if to_save else CheckpointPolicy.PREFER_RECOMPUTE
return _custom_policy
def selective_checkpointing_context_fn():
meta: dict[str, int] = collections.defaultdict(int)
return create_selective_checkpoint_contexts(_get_custom_policy(meta))
return checkpoint_wrapper(module,
context_fn=selective_checkpointing_context_fn,
preserve_rng_state=False)
+47 -40
View File
@@ -25,14 +25,13 @@ 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)
get_torch_device, get_world_group,
sequence_model_parallel_all_gather)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
from fastvideo.v1.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
@@ -61,7 +60,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self) -> None:
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_t2v
self.validation_dataset_schema = pyarrow_schema_t2v_validation
@@ -81,14 +80,11 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
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()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
@@ -121,7 +117,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
@@ -163,7 +158,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
@@ -176,11 +170,18 @@ 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']
latents = latents[:, :, :self.training_args.num_latent_t]
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(),
@@ -189,6 +190,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
@@ -252,8 +254,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
current_vsa_sparsity = training_batch.current_vsa_sparsity
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
dit_seq_shape = [
latents.shape[2] * self.sp_world_size // patch_size[0],
latents.shape[2] // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
@@ -293,10 +296,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.noisy_model_input 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
@@ -316,18 +317,19 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
# if self.training_args.sp_size > 1:
# model_pred = sequence_model_parallel_all_gather(model_pred, dim=2)
loss.backward()
avg_loss = loss.detach().clone()
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
@@ -366,12 +368,20 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
for _ in range(self.training_args.gradient_accumulation_steps):
training_batch = self._get_next_batch(training_batch)
training_batch.latents = training_batch.latents[:, :, :self.
training_args.
num_latent_t]
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
# Create noisy model input
# Save original latents
# training_batch.original_latents = training_batch.latents.clone()
training_batch = self._prepare_dit_inputs(training_batch)
# Shard latents across sp groups
# # Shard latents across sp groups
training_batch.latents = shard_latents_across_sp(
training_batch.latents,
num_latent_t=self.training_args.num_latent_t)
@@ -379,11 +389,10 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
training_batch.noisy_model_input = shard_latents_across_sp(
training_batch.noisy_model_input,
num_latent_t=self.training_args.num_latent_t)
# shard noise to match latents
# CRITICAL FIX: Also shard noise to match latents
training_batch.noise = shard_latents_across_sp(
training_batch.noise,
num_latent_t=self.training_args.num_latent_t)
training_batch = self._build_attention_metadata(training_batch)
training_batch = self._build_input_kwargs(training_batch)
training_batch = self._transformer_forward_and_compute_loss(
@@ -554,8 +563,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
assert len(validation_batch['info_list']
) == 1, "Only batch size 1 is supported for validation"
# 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']
@@ -612,6 +621,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Set deterministic seed for validation
set_random_seed(self.seed)
logger.info("Using validation seed: %s", self.seed)
# Prepare validation prompts
@@ -622,13 +632,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
batch_size=1,
parquet_schema=self.validation_dataset_schema,
num_data_workers=0,
cfg_rate=0.0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None)
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, negative_prompt = validation_dataset.get_validation_negative_prompt(
)
logger.info("Using negative_prompt: %s", negative_prompt)
logger.info("negative_prompt: %s", negative_prompt)
transformer.eval()
@@ -639,18 +649,15 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
step_videos: List[np.ndarray] = []
step_captions: List[str] = []
step_captions: List[str | None] = []
# 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)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
step_captions.extend([None]) # TODO(peiyuan): add caption
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
dtype=torch.bfloat16):
+5 -5
View File
@@ -162,8 +162,8 @@ def save_checkpoint(transformer,
weight_path,
local_main_process_only=False)
# Convert fastvideo custom format to diffusers format and save
diffusers_state_dict = convert_custom_format_to_diffusers_format(
# Convert training format to diffusers format and save
diffusers_state_dict = convert_training_to_diffusers_format(
cpu_state, transformer)
save_file(diffusers_state_dict, weight_path)
@@ -487,10 +487,10 @@ def _has_foreach_support(tensors: List[torch.Tensor],
t is None or type(t) in [torch.Tensor] for t in tensors)
def convert_custom_format_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
def convert_training_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
"""
Convert fastvideo custom format state dict to diffusers format using reverse_param_names_mapping.
Convert training format state dict to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in training format
@@ -18,6 +18,7 @@ from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
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()
@@ -63,7 +64,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
@@ -75,14 +76,21 @@ class WanI2VTrainingPipeline(TrainingPipeline):
# 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']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
@@ -94,8 +102,13 @@ class WanI2VTrainingPipeline(TrainingPipeline):
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
assert training_batch.image_latents is not None
# if self.training_args and hasattr(self.training_args, 'num_latent_t'):
training_batch.image_latents = training_batch.image_latents[:, :, :self.training_args.num_latent_t]
return training_batch
def _prepare_dit_inputs(self,
@@ -110,14 +123,12 @@ class WanI2VTrainingPipeline(TrainingPipeline):
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
image_latents = training_batch.image_latents.to(get_torch_device(), dtype=torch.bfloat16)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
return training_batch
def _build_input_kwargs(self,
@@ -159,9 +170,12 @@ class WanI2VTrainingPipeline(TrainingPipeline):
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']
pil_image = validation_batch['pil_image']
infos = validation_batch['info_list']
prompt = infos[0]['prompt']
@@ -169,6 +183,21 @@ class WanI2VTrainingPipeline(TrainingPipeline):
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
clip_features = clip_features.to(get_torch_device())
# 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,
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.1.1"
version = "0.1.0"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
@@ -19,7 +19,7 @@ dependencies = [
# Machine Learning & Transformers
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
"timm==1.0.11", "peft>=0.15.0", "diffusers>=0.33.1", "bitsandbytes",
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
"torch==2.7.1", "torchvision",
# Acceleration & Optimization
+2 -3
View File
@@ -36,7 +36,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--training_cfg_rate 0.0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_finetune"\
--tracker_project_name wan_finetune \
--num_height 480 \
@@ -48,5 +48,4 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--weight_decay 0.01 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0 \
--enable_gradient_checkpointing_type "full"
--max_grad_norm 1.0
+2 -3
View File
@@ -42,7 +42,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--cfg 0.0 \
--output_dir "$DATA_DIR/outputs/wan_finetune" \
--tracker_project_name VSA_finetune \
--num_height 448 \
@@ -58,6 +58,5 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--VSA_decay_sparsity 0.9 \
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 30 \
--VSA_val_sparsity 0.9 \
--enable_gradient_checkpointing_type "full"
--VSA_val_sparsity 0.9
# --resume_from_checkpoint "$CHECKPOINT_PATH"
-27
View File
@@ -1,27 +0,0 @@
#!/bin/bash
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# change model path to local dir if you want to inference using your checkpoint
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--num-gpus $num_gpus \
--height 448 \
--width 832 \
--num-frames 77 \
--num-inference-steps 50 \
--fps 16 \
--guidance-scale 6.0 \
--flow-shift 8.0 \
--VSA-sparsity 0.9 \
--prompt "A beautiful woman in a red dress walking down a street" \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output-path outputs_video_1.3B_VSA/sparsity_0.9/