Compare commits

..
Author SHA1 Message Date
SolitaryThinker 480868bef9 update 2025-06-24 15:30:55 -07:00
83 changed files with 465 additions and 2219 deletions
+13 -95
View File
@@ -2,26 +2,25 @@ env:
IMAGE_VERSION: "py3.12-latest"
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 30m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
@@ -32,10 +31,8 @@ steps:
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 30m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
@@ -46,12 +43,10 @@ steps:
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 30m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
@@ -60,8 +55,7 @@ steps:
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
- path: "fastvideo/v1/**/*.py"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
@@ -70,79 +64,3 @@ steps:
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 30m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 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
+9 -8
View File
@@ -212,7 +212,8 @@ jobs:
ssim-test:
needs: change-filter
if: >-
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
strategy:
fail-fast: false
matrix:
@@ -238,7 +239,7 @@ jobs:
training-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
@@ -258,13 +259,13 @@ jobs:
training-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
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"
@@ -278,13 +279,13 @@ jobs:
inference-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
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"
@@ -297,7 +298,7 @@ jobs:
precision-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
@@ -316,7 +317,7 @@ jobs:
precision-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
@@ -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,91 +0,0 @@
#!/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
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 8
--tp_size 8
--hsdp_replicate_dim 1
--hsdp_shard_dim 8
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "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
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,130 +0,0 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
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 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[@]}"
@@ -1,25 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -3,8 +3,8 @@
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
@@ -12,8 +12,8 @@
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
@@ -21,8 +21,8 @@
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 40,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
@@ -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,90 +0,0 @@
#!/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-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_finetune"
--output_dir "outputs/wan_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 8
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path $DATA_DIR
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path $VALIDATION_DIR
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--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
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,127 +0,0 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=t2v_output/t2v_%j.out
#SBATCH --error=t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
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/"
# 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
)
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[@]}"
@@ -1,31 +0,0 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
+8 -5
View File
@@ -29,11 +29,14 @@ def getdataset(args) -> VideoCaptionMergedDataset:
*resize_topcrop,
norm_fun,
])
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop)
if args.dataset == "t2v":
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop)
raise NotImplementedError(args.dataset)
__all__ = [
@@ -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
+34 -7
View File
@@ -26,17 +26,15 @@ pyarrow_schema_i2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -52,6 +50,13 @@ pyarrow_schema_i2v = pa.schema([
pyarrow_schema_i2v_validation = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
@@ -59,6 +64,11 @@ pyarrow_schema_i2v_validation = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
@@ -96,6 +106,11 @@ pyarrow_schema_t2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -111,6 +126,13 @@ pyarrow_schema_t2v = pa.schema([
pyarrow_schema_t2v_validation = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
@@ -118,6 +140,11 @@ pyarrow_schema_t2v_validation = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -4,7 +4,6 @@ import random
from typing import Dict, List, Tuple
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
import tqdm
@@ -71,12 +70,10 @@ class LatentsParquetIterStyleDataset(IterableDataset):
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32,
parquet_schema: pa.Schema = None):
read_batch_size: int = 32):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.parquet_schema = parquet_schema
self.cfg_rate = cfg_rate
self.text_padding_length = text_padding_length
self.seed = seed
@@ -3,7 +3,6 @@ import os
import pickle
from typing import Any, Dict, List, Tuple
import pyarrow as pa
import pyarrow.parquet as pq
# Torch in general
import torch
@@ -12,7 +11,7 @@ import tqdm
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
@@ -185,12 +184,13 @@ 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"]
def __init__(
self,
path: str,
batch_size: int,
parquet_schema: pa.Schema,
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
@@ -200,12 +200,24 @@ class LatentsParquetMapStyleDataset(Dataset):
super().__init__()
self.path = path
self.cfg_rate = cfg_rate
self.parquet_schema = parquet_schema
if cfg_rate > 0.0:
raise ValueError(
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
)
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
path)
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
self.batch = batch_size
self.text_padding_length = text_padding_length
self._cols = [
"vae_latent_bytes",
"vae_latent_shape",
"text_embedding_bytes",
"text_embedding_shape",
"text_embedding_dtype",
"height",
"width",
]
self.sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=sum(self.lengths),
@@ -220,7 +232,7 @@ class LatentsParquetMapStyleDataset(Dataset):
len(self.parquet_files), sum(self.lengths))
def get_validation_negative_prompt(
self) -> tuple[torch.Tensor, torch.Tensor, str]:
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
@@ -234,23 +246,19 @@ class LatentsParquetMapStyleDataset(Dataset):
row_dict = read_row_from_parquet_file([file_path], row_idx,
[self.lengths[0]])
batch = collate_rows_from_parquet_schema([row_dict],
self.parquet_schema,
self.text_padding_length,
cfg_rate=0.0)
negative_prompt = batch['info_list'][0]['prompt']
negative_prompt_embedding = batch['text_embedding']
negative_prompt_attention_mask = batch['text_attention_mask']
if len(negative_prompt_embedding.shape) == 2:
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
if len(negative_prompt_attention_mask.shape) == 1:
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
0).unsqueeze(0)
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
[row_dict], self.text_padding_length, self.keys)
all_latents, all_embs, all_masks, caption_text = all_latents_list[
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
# add batch dimension
if len(all_embs.shape) == 2:
all_embs = all_embs.unsqueeze(0)
if len(all_masks.shape) == 1:
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
return all_latents, all_embs, all_masks, caption_text
# PyTorch calls this ONLY because the batch_sampler yields a list
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
def __getitems__(self, indices: List[int]):
"""
Batch fetch using read_row_from_parquet_file for each index.
"""
@@ -259,11 +267,9 @@ 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)
return batch
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
rows, self.text_padding_length, self.keys)
return all_latents, all_embs, all_masks, caption_text
def __len__(self):
return sum(self.lengths)
@@ -280,7 +286,6 @@ def build_parquet_map_style_dataloader(
path,
batch_size,
num_data_workers,
parquet_schema,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
@@ -293,7 +298,6 @@ def build_parquet_map_style_dataloader(
drop_last=drop_last,
drop_first_row=drop_first_row,
text_padding_length=text_padding_length,
parquet_schema=parquet_schema,
seed=seed)
loader = StatefulDataLoader(
+34 -57
View File
@@ -22,7 +22,7 @@ logger = init_logger(__name__)
@dataclass
class PreprocessBatch:
class DatasetBatch:
"""
Batch information for dataset processing stages.
@@ -66,7 +66,7 @@ class DatasetStage(ABC):
"""
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""
Process the dataset batch.
@@ -88,7 +88,7 @@ class DatasetFilterStage(ABC):
"""
@abstractmethod
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
"""
Check if batch should be kept.
@@ -102,7 +102,7 @@ class DatasetFilterStage(ABC):
raise NotImplementedError
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""
Process the dataset batch (for non-filtering operations).
@@ -119,7 +119,7 @@ class DatasetFilterStage(ABC):
class DataValidationStage(DatasetFilterStage):
"""Stage for validating data items."""
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
"""
Validate data item.
@@ -142,7 +142,7 @@ class DataValidationStage(DatasetFilterStage):
return True
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""Process does nothing for validation - filtering is handled by should_keep."""
return batch
@@ -160,7 +160,7 @@ class ResolutionFilterStage(DatasetFilterStage):
self.max_height = max_height
self.max_width = max_width
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
"""
Check if data item passes resolution filtering.
@@ -193,7 +193,7 @@ class ResolutionFilterStage(DatasetFilterStage):
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
return batch
@@ -218,7 +218,7 @@ class FrameSamplingStage(DatasetFilterStage):
self.video_length_tolerance_range = video_length_tolerance_range
self.drop_short_ratio = drop_short_ratio
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
def should_keep(self, batch: DatasetBatch, **kwargs) -> bool:
"""
Check if video should be kept based on length constraints.
@@ -252,9 +252,9 @@ class FrameSamplingStage(DatasetFilterStage):
and random.random() < self.drop_short_ratio)
def process(self,
batch: PreprocessBatch,
batch: DatasetBatch,
temporal_sample_fn=None,
**kwargs) -> PreprocessBatch:
**kwargs) -> DatasetBatch:
"""
Process frame sampling for video data items.
@@ -298,7 +298,7 @@ class VideoTransformStage(DatasetStage):
def __init__(self, transform) -> None:
self.transform = transform
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""
Transform video data.
@@ -339,7 +339,7 @@ class ImageTransformStage(DatasetStage):
self.transform = transform
self.transform_topcrop = transform_topcrop
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""
Transform image data.
@@ -375,7 +375,7 @@ class TextEncodingStage(DatasetStage):
self.text_max_length = text_max_length
self.cfg_rate = cfg_rate
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
def process(self, batch: DatasetBatch, **kwargs) -> DatasetBatch:
"""
Process text data.
@@ -411,29 +411,6 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
torch.distributed.checkpoint.stateful.Stateful):
"""
Merged dataset for video and caption data with stage-based processing.
Assumes that data_merge_path is a txt file with the following format:
<folder_path>,<json_file_path>
The folder should contain videos.
The json file should be a list of dictionaries with the following format:
[
{
"path": "1gGQy4nxyUo-Scene-016.mp4",
"resolution": {
"width": 1920,
"height": 1080
},
"size": 2439112,
"fps": 25.0,
"duration": 6.88,
"num_frames": 172,
"cap": [
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
]
},
...
]
This dataset processes video and image data through a series of stages:
- Data validation
@@ -484,32 +461,32 @@ 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."""
all_data = []
# Read folder-annotation pairs
with open(self.data_merge_path) as f:
folder_anno_pairs = [
line.strip().split(",") for line in f if line.strip()
]
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)
# Process each folder-annotation pair
for folder, annotation_file in folder_anno_pairs:
with open(annotation_file) as f:
data_items = json.load(f)
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
return data_items
all_data.extend(data_items)
def _process_metadata(self) -> List[PreprocessBatch]:
return all_data[self.start_idx:]
def _process_metadata(self) -> List[DatasetBatch]:
"""Process the raw metadata through all filtering stages."""
raw_data = self._load_raw_data()
processed_batches = []
@@ -523,11 +500,11 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
sample_num_frames: List[int] = []
for item in raw_data:
batch = PreprocessBatch(path=item["path"],
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
batch = DatasetBatch(path=item["path"],
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
@@ -545,7 +522,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
len(raw_data), len(processed_batches))
return processed_batches
def _apply_filter_stages(self, batch: PreprocessBatch,
def _apply_filter_stages(self, batch: DatasetBatch,
filter_counts: Dict[str, int]) -> bool:
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
if not self.validation_stage.should_keep(batch):
+6 -141
View File
@@ -1,5 +1,4 @@
import random
from typing import Any, Dict, List, cast
from typing import Any, Dict, List
import numpy as np
import torch
@@ -21,7 +20,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.
"""
@@ -43,10 +42,7 @@ def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
if key == 'text_embedding' and random.random() < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
else:
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
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
@@ -57,11 +53,8 @@ def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
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]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
@@ -70,7 +63,7 @@ def collate_latents_embs_masks(
# Process each row individually
for i, row in enumerate(batch_to_process):
# 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"]
padded_emb, mask = pad(emb, text_padding_length)
@@ -90,131 +83,3 @@ def collate_latents_embs_masks(
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, caption_text
def collate_rows_from_parquet_schema(rows,
parquet_schema,
text_padding_length,
cfg_rate=0.0) -> Dict[str, Any]:
"""
Collate rows from parquet files based on the provided schema.
Dynamically processes tensor fields based on schema and returns batched data.
Args:
rows: List of row dictionaries from parquet files
parquet_schema: PyArrow schema defining the structure of the data
Returns:
Dict containing batched tensors and metadata
"""
if not rows:
return cast(Dict[str, Any], {})
# Initialize containers for different data types
batch_data: Dict[str, Any] = {}
# Get tensor and metadata field names from schema (fields ending with '_bytes')
tensor_fields = []
metadata_fields = []
for field in parquet_schema.names:
if field.endswith('_bytes'):
shape_field = field.replace('_bytes', '_shape')
dtype_field = field.replace('_bytes', '_dtype')
tensor_name = field.replace('_bytes', '')
tensor_fields.append(tensor_name)
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
elif not field.endswith('_shape') and not field.endswith('_dtype'):
# Only add actual metadata fields, not the shape/dtype helper fields
metadata_fields.append(field)
# Process each tensor field
for tensor_name in tensor_fields:
tensor_list = []
for row in rows:
# Get tensor data from row using the existing helper function pattern
shape_key = f"{tensor_name}_shape"
bytes_key = f"{tensor_name}_bytes"
if shape_key in row and bytes_key in row:
shape = row[shape_key]
bytes_data = row[bytes_key]
if len(bytes_data) == 0:
tensor = torch.zeros(0, dtype=torch.bfloat16)
else:
# Convert bytes to tensor using float32 as default
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()
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
# assert B == 1, "Batch size must be 1"
# tensor = tensor.squeeze(0)
tensor_list.append(tensor)
else:
# Handle missing tensor data
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_name == 'text_embedding':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
for tensor in tensor_list:
if tensor.numel() > 0:
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
else:
# Handle empty embeddings - assume default embedding dimension
padded_tensors.append(
torch.zeros(text_padding_length,
768,
dtype=torch.bfloat16))
attention_masks.append(torch.zeros(text_padding_length))
batch_data[tensor_name] = torch.stack(padded_tensors)
batch_data['text_attention_mask'] = torch.stack(attention_masks)
else:
# Stack 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
]
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
# Process metadata fields into info_list
info_list = []
for row in rows:
info = {}
for field in metadata_fields:
info[field] = row.get(field, "")
# Add prompt field for backward compatibility
info["prompt"] = info.get("caption", "")
info_list.append(info)
batch_data['info_list'] = info_list
# Add caption_text for backward compatibility
if info_list and 'caption' in info_list[0]:
batch_data['caption_text'] = [info['caption'] for info in info_list]
return batch_data
+11 -14
View File
@@ -1,6 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
import os
import pathlib
import datasets
@@ -18,8 +17,6 @@ class ValidationDataset(torch.utils.data.IterableDataset):
super().__init__()
self.filename = pathlib.Path(filename)
# get directory of filename
self.dir = os.path.abspath(self.filename.parent)
if not self.filename.exists():
raise FileNotFoundError(
@@ -63,41 +60,41 @@ class ValidationDataset(torch.utils.data.IterableDataset):
if sample.get("image_path", None) is not None:
image_path = sample["image_path"]
image_path = os.path.join(self.dir, image_path)
if not pathlib.Path(image_path).is_file(
) and not image_path.startswith("http"):
logger.warning("Image file %s does not exist.", image_path)
logger.warning("Image file %s does not exist.",
image_path.as_posix())
else:
sample["image"] = load_image(image_path)
sample["image"] = load_image(sample["image_path"])
if sample.get("video_path", None) is not None:
video_path = sample["video_path"]
video_path = os.path.join(self.dir, video_path)
if not pathlib.Path(video_path).is_file(
) and not video_path.startswith("http"):
logger.warning("Video file %s does not exist.", video_path)
logger.warning("Video file %s does not exist.",
video_path.as_posix())
else:
sample["video"] = load_video(video_path)
sample["video"] = load_video(sample["video_path"])
if sample.get("control_image_path", None) is not None:
control_image_path = sample["control_image_path"]
control_image_path = os.path.join(self.dir, control_image_path)
if not pathlib.Path(control_image_path).is_file(
) and not control_image_path.startswith("http"):
logger.warning("Control Image file %s does not exist.",
control_image_path)
control_image_path.as_posix())
else:
sample["control_image"] = load_image(control_image_path)
sample["control_image"] = load_image(
sample["control_image_path"])
if sample.get("control_video_path", None) is not None:
control_video_path = sample["control_video_path"]
control_video_path = os.path.join(self.dir, control_video_path)
if not pathlib.Path(control_video_path).is_file(
) and not control_video_path.startswith("http"):
logger.warning("Control Video file %s does not exist.",
control_video_path)
else:
sample["control_video"] = load_video(control_video_path)
sample["control_video"] = load_video(
sample["control_video_path"])
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+2 -2
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
@@ -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(
-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,
@@ -39,7 +39,6 @@ class ForwardBatch:
image_path: Optional[str] = None
image_embeds: List[torch.Tensor] = field(default_factory=list)
pil_image: Optional[PIL.Image.Image] = None
preprocessed_image: Optional[torch.Tensor] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
@@ -151,11 +150,7 @@ class TrainingBatch:
latents: Optional[torch.Tensor] = None
encoder_hidden_states: Optional[torch.Tensor] = None
encoder_attention_mask: Optional[torch.Tensor] = None
# i2v
preprocessed_image: Optional[torch.Tensor] = None
image_embeds: Optional[torch.Tensor] = None
image_latents: Optional[torch.Tensor] = None
infos: Optional[List[Dict[str, Any]]] = None
info: Optional[Dict[str, Any]] = None
# Transformer inputs
noisy_model_input: Optional[torch.Tensor] = None
@@ -165,9 +160,6 @@ class TrainingBatch:
attn_metadata: Optional[AttentionMetadata] = None
# input kwargs
input_kwargs: Optional[Dict[str, Any]] = None
# Training loss
loss: torch.Tensor | None = None
@@ -2,13 +2,11 @@
import gc
import multiprocessing
import os
from collections import defaultdict
from concurrent.futures import ProcessPoolExecutor
from itertools import chain
from typing import Any, Dict, List, Optional
import numpy as np
import PIL.Image
import pyarrow as pa
import pyarrow.parquet as pq
import torch
@@ -17,13 +15,12 @@ from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import ValidationDataset, getdataset
from fastvideo.v1.dataset.preprocessing_datasets import PreprocessBatch
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
from fastvideo.v1.pipelines.stages import EncodingStage, TextEncodingStage
logger = init_logger(__name__)
@@ -39,6 +36,9 @@ class BasePreprocessPipeline(ComposedPipelineBase):
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae"), ))
@torch.no_grad()
def forward(
self,
@@ -61,206 +61,39 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"""Get the schema fields for the pipeline type. Override in subclasses."""
raise NotImplementedError
def create_record_for_schema(self,
preprocess_batch: PreprocessBatch,
schema: pa.Schema,
strict: bool = False) -> Dict[str, Any]:
"""Create a record for the Parquet dataset using a generic schema-based approach.
Args:
preprocess_batch: The batch containing the data to extract
schema: PyArrow schema defining the expected fields
strict: If True, raises an exception when required fields are missing or unfilled
Returns:
Dictionary record matching the schema
Raises:
ValueError: If strict=True and required fields are missing or unfilled
"""
record = {}
unfilled_fields = []
for field in schema.names:
field_filled = False
if field.endswith('_bytes'):
# Handle binary tensor data - convert numpy array or tensor to bytes
tensor_name = field.replace('_bytes', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None:
try:
if hasattr(tensor_data, 'numpy'): # torch tensor
record[field] = tensor_data.cpu().numpy().tobytes()
field_filled = True
elif hasattr(tensor_data, 'tobytes'): # numpy array
record[field] = tensor_data.tobytes()
field_filled = True
else:
raise ValueError(
f"Unsupported tensor type for field {field}: {type(tensor_data)}"
)
except Exception as e:
if strict:
raise ValueError(
f"Failed to convert tensor {tensor_name} to bytes: {e}"
) from e
record[field] = b'' # Empty bytes for missing data
else:
record[field] = b'' # Empty bytes for missing data
elif field.endswith('_shape'):
# Handle tensor shape info
tensor_name = field.replace('_shape', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None and hasattr(tensor_data, 'shape'):
record[field] = list(tensor_data.shape)
field_filled = True
else:
record[field] = []
elif field.endswith('_dtype'):
# Handle tensor dtype info
tensor_name = field.replace('_dtype', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None and hasattr(tensor_data, 'dtype'):
record[field] = str(tensor_data.dtype)
field_filled = True
else:
record[field] = 'unknown'
elif field in ['width', 'height', 'num_frames']:
# Handle integer metadata fields
value = getattr(preprocess_batch, field, None)
if value is not None:
try:
record[field] = int(value)
field_filled = True
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to int: {e}"
) from e
record[field] = 0
else:
record[field] = 0
elif field in ['duration_sec', 'fps']:
# Handle float metadata fields
# Map schema field names to batch attribute names
attr_name = 'duration' if field == 'duration_sec' else field
value = getattr(preprocess_batch, attr_name, None)
if value is not None:
try:
record[field] = float(value)
field_filled = True
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to float: {e}"
) from e
record[field] = 0.0
else:
record[field] = 0.0
else:
# Handle string fields (id, file_name, caption, media_type, etc.)
# Map common schema field names to batch attribute names
attr_name = field
if field == 'caption':
attr_name = 'text'
elif field == 'file_name':
attr_name = 'path'
elif field == 'id':
# Generate ID from path if available
path_value = getattr(preprocess_batch, 'path', None)
if path_value:
import os
record[field] = os.path.basename(path_value).split(
'.')[0]
field_filled = True
else:
record[field] = ""
continue
elif field == 'media_type':
# Determine media type from path
path_value = getattr(preprocess_batch, 'path', None)
if path_value:
record[field] = 'video' if path_value.endswith(
'.mp4') else 'image'
field_filled = True
else:
record[field] = ""
continue
value = getattr(preprocess_batch, attr_name, None)
if value is not None:
record[field] = str(value)
field_filled = True
else:
record[field] = ""
# Track unfilled fields
if not field_filled:
unfilled_fields.append(field)
# Handle strict mode
if strict and unfilled_fields:
raise ValueError(
f"Required fields were not filled: {unfilled_fields}")
# Log unfilled fields as warning if not in strict mode
if unfilled_fields:
logger.warning(
"Some fields were not filled and got default values: %s",
unfilled_fields)
return record
def create_record(
self,
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."""
record = {
"id":
video_name,
"vae_latent_bytes":
vae_latent.tobytes(),
"vae_latent_shape":
list(vae_latent.shape),
"vae_latent_dtype":
str(vae_latent.dtype),
"text_embedding_bytes":
text_embedding.tobytes(),
"text_embedding_shape":
list(text_embedding.shape),
"text_embedding_dtype":
str(text_embedding.dtype),
"file_name":
video_name,
"caption":
valid_data["text"][idx] if len(valid_data["text"]) > 0 else "",
"media_type":
"video",
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx] if valid_data else "",
"media_type": "video",
"width":
valid_data["pixel_values"][idx].shape[-2]
if len(valid_data["pixel_values"]) > 0 else 0,
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
"height":
valid_data["pixel_values"][idx].shape[-1]
if len(valid_data["pixel_values"]) > 0 else 0,
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
"num_frames":
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
"duration_sec":
float(valid_data["duration"][idx])
if len(valid_data["duration"]) > 0 else 0.0,
"fps":
float(valid_data["fps"][idx])
if len(valid_data["fps"]) > 0 else 0.0,
float(valid_data["duration"][idx]) if valid_data else 0.0,
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
}
if extra_features:
record.update(extra_features)
@@ -380,6 +213,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 +231,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)
@@ -482,7 +318,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
for idx, sample in pbar:
with torch.inference_mode():
prompt = sample["caption"]
is_negative_prompt = idx == 0
# is_negative_prompt = idx == 0
# Text Encoder
batch = ForwardBatch(
@@ -510,43 +346,15 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"Shape after removing padding - Embeddings: %s, Mask: %s",
text_embedding.shape, text_attention_mask.shape)
extra_features = {}
if not is_negative_prompt:
height = sample["height"]
width = sample["width"]
if "image_path" in sample and "video_path" in sample:
raise ValueError(
"Only one of image_path or video_path should be provided"
)
if "image" in sample:
extra_features = self.preprocess_image(
sample["image"], height, width, fastvideo_args)
if "video" in sample:
extra_features = self.preprocess_video(
sample["video"], height, width, fastvideo_args)
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value.cpu().numpy()
else:
sample_extra_features[key] = value
valid_data = defaultdict(list)
valid_data["text"] = [prompt]
# Create record for Parquet dataset
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
valid_data=valid_data,
text_attention_mask=text_attention_mask,
valid_data=None,
idx=0,
extra_features=sample_extra_features)
extra_features=None)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
@@ -620,15 +428,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
del table
gc.collect() # Force garbage collection
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return {}
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return {}
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
@@ -8,7 +8,6 @@ using the modular pipeline architecture.
from typing import Any, Dict, List, Optional
import numpy as np
import PIL
import torch
from PIL import Image
@@ -16,13 +15,8 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.models.vision_utils import (get_default_height_width,
normalize, numpy_to_pt,
pil_to_numpy, resize)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.v1.pipelines.stages import ImageEncodingStage, TextEncodingStage
class PreprocessPipeline_I2V(BasePreprocessPipeline):
@@ -32,73 +26,18 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
assert hasattr(
self,
"image_encoding_stage"), "Image encoding stage must be created"
batch = ForwardBatch(
data_type="video",
pil_image=image,
)
result_batch = self.image_encoding_stage(batch, fastvideo_args)
clip_features = result_batch.image_embeds[0]
image = self.preprocess(
image,
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
height=height,
width=width)
return {
"clip_feature": clip_features[0],
"pil_image": image,
}
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return self.preprocess_image(video[0], height, width, fastvideo_args)
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for I2V pipeline."""
return [f.name for f in pyarrow_schema_i2v]
def get_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("image_encoder").to(get_torch_device())
self.get_module("vae").to(get_torch_device())
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
batch_size, _, num_frames, height, width = valid_data[
"pixel_values"].shape
latent_height = height // self.get_module(
"vae").spatial_compression_ratio
latent_width = width // self.get_module("vae").spatial_compression_ratio
processed_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
@@ -114,84 +53,22 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
features["clip_feature"] = clip_features
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0],
processed_img.shape[1], num_frames - 1,
height, width)
],
dim=2)
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor.to(latent_condition.device,
latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, list(range(1, num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
dim=2,
repeats=self.get_module("vae").temporal_compression_ratio)
mask_lat_size = torch.concat(
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
mask_lat_size = mask_lat_size.view(
batch_size, -1,
self.get_module("vae").temporal_compression_ratio, latent_height,
latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
features["first_frame_latent"] = image_latent
return features
return {"clip_feature": clip_features}
def create_record(
self,
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,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
@@ -210,69 +87,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes":
first_frame_latent.tobytes(),
"first_frame_latent_shape":
list(first_frame_latent.shape),
"first_frame_latent_dtype":
str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
return record
def pil_to_tensor(self, image: PIL.Image.Image) -> torch.Tensor:
image = image
image = np.array(image).astype(np.float32)
image = torch.from_numpy(image)
return image
def preprocess(self,
image: PIL.Image.Image,
vae_scale_factor: int,
height: int,
width: int,
resize_mode: str = "default") -> torch.Tensor:
image = [image]
height, width = get_default_height_width(image[0], vae_scale_factor,
height, width)
image = [
resize(i, height, width, resize_mode=resize_mode) for i in image
]
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
do_normalize = True
if image.min() < 0:
do_normalize = False
if do_normalize:
image = normalize(image)
return image
return record # type: ignore
EntryClass = PreprocessPipeline_I2V
@@ -59,6 +59,12 @@ if __name__ == "__main__":
default=2,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--preprocess_text_batch_size",
type=int,
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--samples_per_file", type=int, default=64)
parser.add_argument("--flush_frequency",
type=int,
@@ -73,6 +79,7 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
@@ -84,7 +91,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,
+12 -17
View File
@@ -51,27 +51,22 @@ class EncodingStage(PipelineStage):
"""
self.vae = self.vae.to(get_torch_device())
image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
raise ValueError("Image Path must be provided")
assert batch.height is not None
assert batch.width is not None
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
image = batch.preprocessed_image
# TODO(will)
if image is None:
assert batch.pil_image is not None
image = batch.pil_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
else:
# assumes image is loaded from parquet file and used for validation
image = image.transpose(1, 2)
logger.info("image: %s", image.shape)
image = batch.pil_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
@@ -186,7 +181,7 @@ class EncodingStage(PipelineStage):
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage inputs."""
result = VerificationResult()
# result.add_check("pil_image", batch.pil_image)
result.add_check("pil_image", batch.pil_image, V.not_none)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("generator", batch.generator,
@@ -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",
+61 -49
View File
@@ -4,43 +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", ""),
})
.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")
os.chdir("/FastVideo")
print(f"Cloning repository: {git_repo}")
print(f"Checking out commit: {git_commit}")
command = f"""
source $HOME/.local/bin/env &&
source /opt/venv/bin/activate &&
git clone {git_repo} /FastVideo &&
cd /FastVideo &&
git checkout {git_commit} &&
uv pip install -e .[test] &&
{pytest_command}
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/encoders -s
"""
result = subprocess.run([
@@ -49,38 +37,62 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_encoder_tests():
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
@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=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=3600)
def run_ssim_tests():
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=1800, 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=1800, 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=1800)
def run_inference_tests_STA():
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
@app.function(gpu="H100:1", image=image, timeout=1800)
def run_precision_tests_STA():
run_test("python csrc/attn/tests/test_sta.py")
@app.function(gpu="H100:1", image=image, timeout=1800)
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)
@@ -1 +0,0 @@
{"step_time":2.245914653001819,"_wandb":{"runtime":1434},"learning_rate":1e-05,"grad_norm":0.57421875,"avg_step_time":1.1814782944297622,"train_loss":0.07932619750499725,"vsa_sparsity":0,"_timestamp":1.750578625921253e+09,"validation_videos_40_steps":{"count":1,"videos":[{"size":420969,"path":"media/videos/validation_videos_40_steps_900_581ff5eae2909d3a7b36.mp4","_type":"video-file","sha256":"581ff5eae2909d3a7b362dcb24d060c006c09e4d4deb44b82f4aa697f6789ba7"}],"captions":false,"_type":"videos"},"_runtime":1434.62395329,"_step":901}
@@ -1,177 +0,0 @@
import os
from pathlib import Path
from huggingface_hub import snapshot_download
import shutil
import subprocess
import sys
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
NUM_NODES = "1"
MODEL_PATH = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
# preprocessing
DATA_DIR = "data"
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_preprocessed_data_i2v"))
# training
NUM_GPUS_PER_NODE_TRAINING = "8"
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_i2v_training_pipeline.py"
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
LOCAL_VALIDATION_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "validation_parquet_dataset")
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
# if data_dir.exists():
# print(f"Removing existing data directory at {data_dir}")
# shutil.rmtree(data_dir)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
# result = snapshot_download(
# repo_id="wlsaidhi/cats-overfit-merged",
# local_dir=str(LOCAL_RAW_DATA_DIR),
# repo_type="dataset",
# resume_download=True,
# token=os.environ.get("HF_TOKEN"), # In case authentication is needed
# )
print(f"Download completed successfully. Files downloaded to: {result}")
# Verify the download
if not LOCAL_RAW_DATA_DIR.exists():
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
# List downloaded files
print("Downloaded files:")
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
if file.is_file():
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
except Exception as e:
print(f"Error during download: {str(e)}")
raise
def run_preprocessing():
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
PREPROCESSING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge_1_sample.txt"),
"--preprocess_video_batch_size", "1",
"--max_height", "480",
"--max_width", "832",
"--num_frames", "77",
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_i2v_prompt_1_sample.json"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
"--preprocess_task", "i2v",
]
process = subprocess.run(cmd, check=True)
def run_training():
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
TRAINING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
"--tp_size", NUM_GPUS_PER_NODE_TRAINING,
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "40",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.1",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
]
print(f"Running training with command: {cmd}")
process = subprocess.run(cmd, check=True)
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
# download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
print(f"reference_video_file: {reference_video_file}")
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
print(f"final_validation_video_file: {final_validation_video_file}")
# Ensure both files exist
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
# Compute SSIM
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
reference_video_file,
final_validation_video_file,
use_ms_ssim=True # Using MS-SSIM for better quality assessment
)
print("\n===== SSIM Results for Step 900 Validation =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
if __name__ == "__main__":
test_e2e_overfit_single_sample()
@@ -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",
@@ -0,0 +1,51 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR="data/crush-smol_processed_main_t2v/latents/combined_parquet_dataset"
VALIDATION_DIR="data/crush-smol_processed_main_t2v/latents/validation_parquet_dataset"
NUM_GPUS=4
# 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
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/training/wan_training_pipeline.py\
--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_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1 \
--num_latent_t 8 \
--sp_size 4 \
--tp_size 4 \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim 4 \
--num_gpus $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--gradient_accumulation_steps=8 \
--max_train_steps=5000 \
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=6000 \
--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_finetune"\
--tracker_project_name wan_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 0.01 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -1,11 +1,10 @@
#!/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"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
OUTPUT_DIR="data/crush-smol_processed_main_t2v/latents"
VALIDATION_PATH="examples/training/finetune/wan_t2v_1.3b/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
@@ -22,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "t2v"
--preprocess_task "t2v"
+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}"
+74 -124
View File
@@ -22,8 +22,6 @@ from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadata)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_t2v, pyarrow_schema_t2v_validation)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -52,17 +50,14 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
_required_config_modules = ["scheduler", "transformer"]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[Dict[str, Any]]
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
Dict[str, Any]]]
current_epoch: int = 0
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self) -> None:
self.train_dataset_schema = pyarrow_schema_t2v
self.validation_dataset_schema = pyarrow_schema_t2v_validation
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.training_args = training_args
@@ -75,10 +70,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
assert training_args.seed is not None
self.seed = training_args.seed
assert self.transformer is not None
self.set_schemas()
self.transformer.requires_grad_(True)
self.transformer.train()
@@ -112,14 +104,12 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=self.seed)
seed=training_args.seed)
self.noise_scheduler = noise_scheduler
@@ -156,7 +146,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
@@ -169,20 +158,14 @@ 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 = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
latents, encoder_hidden_states, encoder_attention_mask, infos = batch
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.infos = infos
training_batch.info = infos
return training_batch
@@ -245,8 +228,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]
]
@@ -259,15 +243,25 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
# assert training_batch.attn_metadata is not None
assert training_batch.latents is not None
assert training_batch.noise is not None
assert training_batch.sigmas is not None
training_batch.input_kwargs = {
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
@@ -280,25 +274,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
"return_dict":
False,
}
return training_batch
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.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
else:
assert training_batch.attn_metadata is None
assert training_batch.input_kwargs is not None
input_kwargs = training_batch.input_kwargs
# if 'hunyuan' in self.training_args.model_type:
# input_kwargs["guidance"] = torch.tensor(
# [1000.0],
@@ -313,8 +288,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
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)
@@ -358,26 +331,15 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
for _ in range(self.training_args.gradient_accumulation_steps):
training_batch = self._get_next_batch(training_batch)
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
# Create noisy model input
training_batch = self._prepare_dit_inputs(training_batch)
# Shard latents across sp groups
training_batch.latents = shard_latents_across_sp(
training_batch.latents,
num_latent_t=self.training_args.num_latent_t)
# shard noisy_model_input to match
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
training_batch.noise = shard_latents_across_sp(
training_batch.noise,
num_latent_t=self.training_args.num_latent_t)
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
training_batch = self._prepare_dit_inputs(training_batch)
training_batch = self._build_attention_metadata(training_batch)
training_batch = self._build_input_kwargs(training_batch)
training_batch = self._transformer_forward_and_compute_loss(
training_batch)
@@ -410,10 +372,14 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
assert self.training_args is not None
# Set random seeds for deterministic training
set_random_seed(self.seed)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
assert self.training_args.seed is not None, "seed must be set"
seed = self.training_args.seed
set_random_seed(seed)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
@@ -539,52 +505,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
def _prepare_validation_inputs(
self, sampling_param: SamplingParam, training_args: TrainingArgs,
validation_batch: Dict[str, Any], num_inference_steps: int,
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
prompt = validation_batch['info_list'][0]['prompt']
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
prompt_embeds = prompt_embeds.to(get_torch_device())
prompt_attention_mask = prompt_attention_mask.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
prompt=prompt,
data_type="video",
latents=None,
seed=self.seed, # Use deterministic seed
generator=torch.Generator(device="cpu").manual_seed(self.seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
assert training_args is not None
@@ -601,8 +521,11 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
# Set deterministic seed for validation
set_random_seed(self.seed)
logger.info("Using validation seed: %s", self.seed)
validation_seed = training_args.seed if training_args.seed is not None else 42
torch.manual_seed(validation_seed)
torch.cuda.manual_seed_all(validation_seed)
logger.info("Using validation seed: %s", validation_seed)
# Prepare validation prompts
logger.info('fastvideo_args.validation_preprocessed_path: %s',
@@ -610,15 +533,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
training_args.validation_preprocessed_path,
batch_size=1,
parquet_schema=self.validation_dataset_schema,
num_data_workers=0,
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(
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
logger.info("Using negative_prompt: %s", negative_prompt)
transformer.eval()
@@ -631,13 +552,42 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
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)
for _, embeddings, masks, infos in validation_dataloader:
step_captions.extend([None]) # TODO(peiyuan): add caption
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
@@ -1,260 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any, Dict
import torch
import torch.distributed
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.distributed import get_torch_device, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
TrainingBatch)
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
WanImageToVideoValidationPipeline)
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (shard_latents_across_sp,
clip_grad_norm_while_handling_failing_dtensor_cases)
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanI2VTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_i2v
self.validation_dataset_schema = pyarrow_schema_i2v_validation
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus)
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents = 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']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.preprocessed_image = pil_image.to(get_torch_device())
training_batch.image_embeds = clip_features.to(get_torch_device())
training_batch.image_latents = image_latents.to(get_torch_device())
training_batch.infos = infos
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert self.noise_random_generator is not None
assert training_batch.image_latents is not None
# 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)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
assert training_batch.image_embeds is not None
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
"encoder_hidden_states_image":
encoder_hidden_states_image,
"return_dict":
False,
}
return training_batch
def _prepare_validation_inputs(
self, sampling_param: SamplingParam, training_args: TrainingArgs,
validation_batch: Dict[str, Any], num_inference_steps: int,
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
embeddings = validation_batch['text_embedding']
masks = validation_batch['text_attention_mask']
clip_features = validation_batch['clip_feature']
pil_image = validation_batch['pil_image']
infos = validation_batch['info_list']
prompt = infos[0]['prompt']
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
clip_features = clip_features.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
prompt=prompt,
data_type="video",
latents=None,
seed=self.seed, # Use deterministic seed
generator=torch.Generator(device="cpu").manual_seed(self.seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
image_embeds=[clip_features],
preprocessed_image=pil_image,
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
"""Override to add gradient synchronization across SP ranks."""
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
# CRITICAL FIX: Synchronize gradients across SP ranks before clipping
# Different SP ranks compute different gradients due to different noise patterns
# These gradients must be averaged across SP ranks for stable training
if self.training_args.sp_size > 1:
sp_group = get_sp_group()
for param in self.transformer.parameters():
if param.grad is not None:
# Average gradients across SP ranks
sp_group.all_reduce(param.grad, op=torch.distributed.ReduceOp.AVG)
if max_grad_norm is not None:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanI2VTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)
+1 -2
View File
@@ -1,4 +1,3 @@
# trigger test
[build-system]
requires = ["setuptools>=61.0"]
build-backend = "setuptools.build_meta"
@@ -20,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.15.0", "diffusers>=0.33.1", "bitsandbytes",
"torch==2.7.1", "torchvision",
# Acceleration & Optimization
+1 -1
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 \
+1 -1
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 \
-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/