Compare commits
25
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
94d45df2bc | ||
|
|
08a9958de2 | ||
|
|
4b76388f21 | ||
|
|
c69ab42cc4 | ||
|
|
8c967268a3 | ||
|
|
9ce7c8039e | ||
|
|
e1e25e95f9 | ||
|
|
490bde90e1 | ||
|
|
dc7596b973 | ||
|
|
335afa4457 | ||
|
|
3f77a6805a | ||
|
|
13d0aae706 | ||
|
|
404cbf4f3c | ||
|
|
958ffec844 | ||
|
|
31f000d1cc | ||
|
|
cd32b3e02f | ||
|
|
bf27908095 | ||
|
|
c5f9ea53b2 | ||
|
|
d32a7184da | ||
|
|
2930abe456 | ||
|
|
b93ef4289d | ||
|
|
401bdbd316 | ||
|
|
1048d79cf8 | ||
|
|
1e8406162d | ||
|
|
03edd35c83 |
+24
-1
@@ -104,6 +104,18 @@ steps:
|
||||
- TEST_TYPE=distillation_dmd
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/training/*self_forcing_distillation_pipeline.py"
|
||||
- "fastvideo/tests/training/self-forcing/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Self-Forcing Tests"
|
||||
env:
|
||||
- TEST_TYPE=self_forcing
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
@@ -198,4 +210,15 @@ steps:
|
||||
env:
|
||||
- TEST_TYPE=inference_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Unit Tests"
|
||||
env:
|
||||
- TEST_TYPE=unit_test
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
@@ -110,6 +110,10 @@ case "$TEST_TYPE" in
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
|
||||
;;
|
||||
# run_inference_tests_vmoba
|
||||
"self_forcing")
|
||||
log "Running self-forcing tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_self_forcing_tests"
|
||||
;;
|
||||
"inference_vmoba")
|
||||
log "Running V-MoBA inference tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
|
||||
@@ -118,6 +122,10 @@ case "$TEST_TYPE" in
|
||||
log "Running V-MoBA precision tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
|
||||
;;
|
||||
"unit_test")
|
||||
log "Running unit tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
|
||||
@@ -62,8 +62,8 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
run_unit_test:
|
||||
description: "Run unit-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
@@ -93,6 +93,7 @@ jobs:
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
unit-test: ${{ steps.filter.outputs.unit-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -102,6 +103,8 @@ jobs:
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.10'
|
||||
- 'docker/Dockerfile.python3.11'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
@@ -155,6 +158,9 @@ jobs:
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
unit-test:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -333,23 +339,42 @@ jobs:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
unit-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
job_id: "unit-test"
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
# nightly-test:
|
||||
# if: >-
|
||||
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
# uses: ./.github/workflows/runpod-test.yml
|
||||
# with:
|
||||
# job_id: "nightly-test"
|
||||
# gpu_type: "NVIDIA A40"
|
||||
# gpu_count: 4
|
||||
# volume_size: 100
|
||||
# disk_size: 100
|
||||
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
# timeout_minutes: 30
|
||||
# secrets:
|
||||
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
# Add other jobs to this list as you create them
|
||||
|
||||
@@ -64,3 +64,6 @@ docs/source/distillation/examples/
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
@@ -12,9 +12,6 @@ exclude: |
|
||||
scripts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/distill/.*|
|
||||
fastvideo/distill\.py|
|
||||
fastvideo/distill_adv\.py|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
@@ -44,10 +41,10 @@ repos:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 6.0.1
|
||||
hooks:
|
||||
- id: isort
|
||||
# - repo: https://github.com/PyCQA/isort
|
||||
# rev: 6.0.1
|
||||
# hooks:
|
||||
# - id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/HpTWYkNc" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -155,8 +155,8 @@ If you find FastVideo useful, please considering citing our work:
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
title={Vsa: Faster video diffusion with trainable sparse attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2505.13389},
|
||||
year={2025}
|
||||
}
|
||||
|
||||
@@ -20,5 +20,7 @@ setup(
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[]
|
||||
install_requires=[
|
||||
"flash-attn >= 2.7.1",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -6,8 +6,16 @@ import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
except ImportError:
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
_flash_attn_varlen_backward = _unsupported
|
||||
flash_attn_varlen_func = _unsupported
|
||||
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=prof_wan21_load
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1200G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_prof_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_prof_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# 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 TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
export FASTVIDEO_TORCH_PROFILER_DIR=/mnt/sharefs/users/hao.zhang/wl/torch_profiler/profiler_wan21_1_3_all_train_1n_8g/
|
||||
export FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_STACK=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1
|
||||
export FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_training_train_one_step"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
# GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo3/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_1.json"
|
||||
|
||||
# export CUDA_VISIBLE_DEVICES=0
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "prof_wan21_train"
|
||||
--max_train_steps 7
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--simulate_generator_forward
|
||||
--log_visualization
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
|
||||
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
|
||||
--init_weights_from_safetensors /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus 8 # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
)
|
||||
|
||||
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/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -0,0 +1,53 @@
|
||||
# Profiling FastVideo
|
||||
|
||||
!!! warning
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
|
||||
|
||||
## Profiling with PyTorch
|
||||
|
||||
FastVideo exposes a process-wide torch profiler that you can enable via environment variables. Set `FASTVIDEO_TORCH_PROFILER_DIR` to an absolute directory path to start collecting traces, and specify the regions you want recorded with `FASTVIDEO_TORCH_PROFILE_REGIONS`:
|
||||
|
||||
```bash
|
||||
FASTVIDEO_TORCH_PROFILER_DIR=/mnt/traces/fastvideo \
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_model_loading,profiler_region_training_step"
|
||||
```
|
||||
|
||||
All profiled regions must be registered in `fastvideo.profiler`; the current list includes:
|
||||
|
||||
- `profiler_region_model_loading` — pipeline/module loading
|
||||
- `profiler_region_inference_pre_denoising`
|
||||
- `profiler_region_inference_denoising`
|
||||
- `profiler_region_inference_post_denoising`
|
||||
- `profiler_region_training_checkpoint_saving`
|
||||
- `profiler_region_training_dit`
|
||||
- `profiler_region_training_validation`
|
||||
- `profiler_region_training_epoch`
|
||||
- `profiler_region_training_step`
|
||||
- `profiler_region_training_backward`
|
||||
- `profiler_region_training_optimizer`
|
||||
- `profiler_region_distillation_teacher_forward`
|
||||
- `profiler_region_distillation_student_forward`
|
||||
- `profiler_region_distillation_loss`
|
||||
- `profiler_region_distillation_update`
|
||||
|
||||
While profiling is enabled, FastVideo records additional annotations:
|
||||
|
||||
- `fastvideo.region::<name>` spans are emitted when entering a region.
|
||||
- `fastvideo.profiler.enable_collection` / `fastvideo.profiler.disable_collection` events mark when torch profiler collection is toggled on or off.
|
||||
|
||||
Only one profiler instance is created per process; subsequent pipelines reuse the same controller. If you set `FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo logs a warning and ignores that entry.
|
||||
|
||||
Additional knobs:
|
||||
|
||||
- `FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES`
|
||||
- `FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY`
|
||||
- `FASTVIDEO_TORCH_PROFILER_WITH_STACK`
|
||||
- `FASTVIDEO_TORCH_PROFILER_WITH_FLOPS`
|
||||
|
||||
Traces can be visualized using <https://ui.perfetto.dev/>.
|
||||
|
||||
### Best Practices
|
||||
|
||||
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
|
||||
- After profiling, clean up trace directories to avoid filling disks.
|
||||
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
|
||||
@@ -115,6 +115,7 @@ design/overview
|
||||
|
||||
contributing/overview
|
||||
contributing/developer_env/index
|
||||
contributing/profiling
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
|
||||
@@ -1,32 +1,40 @@
|
||||
(inference-optimizations)=
|
||||
|
||||
# Optimizations
|
||||
|
||||
This page describes the various options for speeding up generation times in FastVideo.
|
||||
|
||||
## Table of Contents
|
||||
|
||||
- Optimized Attention Backends
|
||||
|
||||
- [Flash Attention](#optimizations-flash)
|
||||
- [Sliding Tile Attention](#optimizations-sta)
|
||||
- [Sage Attention](#optimizations-sage)
|
||||
- [Sage Attention 3](#optimizations-sage3)
|
||||
|
||||
- Caching Techniques
|
||||
- [TeaCache](#optimizations-teacache)
|
||||
|
||||
(optimizations-backends)=
|
||||
|
||||
## Attention Backends
|
||||
|
||||
### Available Backends
|
||||
|
||||
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
|
||||
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
|
||||
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
|
||||
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
|
||||
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
|
||||
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
|
||||
|
||||
### Configuring Backends
|
||||
|
||||
There are two ways to configure the attention backend in FastVideo.
|
||||
|
||||
#### 1. In Python
|
||||
|
||||
In python, set the `FASTVIDEO_ATTENTION_BACKEND` environment variable before instantiating `VideoGenerator` like this:
|
||||
|
||||
```python
|
||||
@@ -34,6 +42,7 @@ os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
|
||||
```
|
||||
|
||||
#### 2. In CLI
|
||||
|
||||
You can also set the environment variable on the command line:
|
||||
|
||||
```bash
|
||||
@@ -41,6 +50,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
```
|
||||
|
||||
(optimizations-flash)=
|
||||
|
||||
### Flash Attention
|
||||
|
||||
**`FLASH_ATTN`**
|
||||
@@ -57,7 +67,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
|
||||
|
||||
cd hopper
|
||||
pip install ninja
|
||||
pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
@@ -66,7 +76,9 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
|
||||
:::
|
||||
|
||||
(optimizations-sta)=
|
||||
|
||||
### Sliding Tile Attention
|
||||
|
||||
**`SLIDING_TILE_ATTN`**
|
||||
|
||||
```bash
|
||||
@@ -76,7 +88,9 @@ pip install st_attn==0.0.4
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
|
||||
(optimizations-vsa)=
|
||||
|
||||
### Video Sparse Attention
|
||||
|
||||
**`VIDEO_SPARSE_ATTN`**
|
||||
|
||||
```bash
|
||||
@@ -87,19 +101,45 @@ python setup_vsa.py install
|
||||
Please see [this page](#vsa-installation) for more installation instructions.
|
||||
|
||||
(optimizations-sage)=
|
||||
|
||||
### Sage Attention
|
||||
|
||||
**`SAGE_ATTN`**
|
||||
|
||||
To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please compile from source:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/thu-ml/SageAttention.git
|
||||
cd sageattention
|
||||
cd sageattention
|
||||
python setup.py install # or pip install -e .
|
||||
```
|
||||
|
||||
(optimizations-sage3)=
|
||||
|
||||
### Sage Attention 3
|
||||
|
||||
**`SAGE_ATTN_THREE`**
|
||||
|
||||
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
|
||||
|
||||
#### Hardware Requirements
|
||||
|
||||
- RTX5090
|
||||
|
||||
#### Installation
|
||||
|
||||
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
|
||||
|
||||
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
|
||||
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
(optimizations-teacache)=
|
||||
|
||||
## Teacache
|
||||
|
||||
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
|
||||
|
||||
### What is TeaCache?
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
# VidProm Dataset
|
||||
|
||||
From [Self-Forcing](https://github.com/gdhe17/Self-Forcing) repository.
|
||||
|
||||
## Download the dataset
|
||||
|
||||
```bash
|
||||
./download_dataset.sh
|
||||
```
|
||||
@@ -0,0 +1,3 @@
|
||||
#! /bin/bash
|
||||
|
||||
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
|
||||
@@ -0,0 +1,140 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
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=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY=your_wandb_api_key
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_data_dir
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
|
||||
--output_dir your_output_dir
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors your_ode_init_weights_path
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -0,0 +1,180 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=sf_distill_profile
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=2
|
||||
#SBATCH --ntasks=2
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=sf_distill_output/sf_distill_profile.out
|
||||
#SBATCH --error=sf_distill_output/sf_distill_profile.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate wei-fv
|
||||
export HOME="/mnt/weka/home/hao.zhang/wei"
|
||||
|
||||
# Basic Info
|
||||
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=29503
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
export FASTVIDEO_TORCH_PROFILER_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/traces
|
||||
export FASTVIDEO_TORCH_PROFILE_REGIONS=profiler_region_model_loading
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=16
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
# DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_2.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan_t2v_finetune"
|
||||
--wandb_run_name "sf_distill_2e-6_4e-7_4n"
|
||||
# --use_sf_wan
|
||||
# --sf_ode_init_path "checkpoints/ode_init.pt"
|
||||
--max_train_steps 2
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_TOTAL_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_TOTAL_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 8
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
# --log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-6
|
||||
--fake_score_learning_rate 4e-7
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 10.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 200
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/FastVideo2/vidprom_8b16k_1e-5_gn1/checkpoint-3500/transformer/diffusion_pytorch_model.safetensors"
|
||||
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/sf2/checkpoint_v3_7k/model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
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/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -0,0 +1,157 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# 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 TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
|
||||
GENERATOR_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 448 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--simulate_generator_forward
|
||||
# --log_visualization
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
|
||||
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus 32 # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim 32
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
)
|
||||
|
||||
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/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -0,0 +1,165 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=moe_4n_sf_distill
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=moe_4n_sf_distill_output/moe_sf_distill_%j.out
|
||||
#SBATCH --error=moe_4n_sf_distill_output/moe_sf_distill_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# 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 TOKENIZERS_PARALLELISM=false
|
||||
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate wei-fv
|
||||
export HOME="/mnt/weka/home/hao.zhang/wei"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=32
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers"
|
||||
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
|
||||
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_dmd"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 448 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--simulate_generator_forward
|
||||
--log_visualization
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
--init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
|
||||
--init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
|
||||
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_TOTAL_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_TOTAL_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 40
|
||||
--validation_sampling_steps "8"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 10.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 12
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 200
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,850,700,550,350,275,200,125'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 4.0
|
||||
--real_score_guidance_scale_2 3.0
|
||||
--fake_score_learning_rate 4e-7
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
)
|
||||
|
||||
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/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir"checkpoints/wan_t2v_finetune"
|
||||
--output_dir $OUTPUT_DIR
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -72,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -91,7 +96,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
@@ -134,4 +139,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -72,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -91,7 +96,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
@@ -134,4 +139,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -72,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -91,7 +96,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
@@ -133,4 +138,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=moe_dmd_distill
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=2
|
||||
#SBATCH --ntasks=2
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=moe_dmd_distill_output/moe_dmd_distill_%j.out
|
||||
#SBATCH --error=moe_dmd_distill_output/moe_dmd_distill_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# 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 TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate wei-fv
|
||||
export HOME="/mnt/weka/home/hao.zhang/wei"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=16
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
|
||||
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name SFwan2.2_t2v_distill_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_dmd"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 448 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
# --log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frames 81
|
||||
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
|
||||
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_TOTAL_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_TOTAL_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--ema_start_step 100
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3.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/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -40,15 +40,18 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name Wan_distillation
|
||||
--output_dir "your_output_dir"
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -73,6 +76,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -92,11 +97,11 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "cosine_with_min_lr"
|
||||
--min_lr_ratio 0.5
|
||||
--lr_warmup_steps 100
|
||||
--fake_score_learning_rate 1e-5
|
||||
--fake_score_learning_rate 2e-6
|
||||
--fake_score_lr_scheduler "cosine_with_min_lr"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
@@ -141,4 +146,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -40,6 +40,8 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
@@ -73,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -92,11 +96,11 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "cosine_with_min_lr"
|
||||
--min_lr_ratio 0.5
|
||||
--lr_warmup_steps 100
|
||||
--fake_score_learning_rate 1e-5
|
||||
--fake_score_learning_rate 2e-6
|
||||
--fake_score_lr_scheduler "cosine_with_min_lr"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
@@ -142,4 +146,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -14,26 +14,29 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir="checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps=4000
|
||||
--train_batch_size=1
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 31
|
||||
--num_height 704
|
||||
--num_width 1280
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--training_state_checkpointing_steps=500
|
||||
--weight_only_checkpointing_steps=500
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
@@ -49,6 +52,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -68,8 +73,8 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
@@ -107,4 +112,4 @@ torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -14,6 +14,8 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
@@ -51,6 +53,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -109,4 +113,4 @@ torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
|
||||
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
--preprocess_task "t2v"
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
|
||||
The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.
|
||||
The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.
|
||||
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
|
||||
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
|
||||
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
|
||||
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
|
||||
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
|
||||
The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
|
||||
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
|
||||
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
|
||||
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
|
||||
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
|
||||
The video shows a large orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
|
||||
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
|
||||
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
|
||||
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
|
||||
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
|
||||
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
|
||||
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
|
||||
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
|
||||
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
|
||||
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
|
||||
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
|
||||
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
|
||||
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
|
||||
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.
|
||||
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
|
||||
@@ -0,0 +1,94 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--warp_denoising_step
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# 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_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--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/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
+136
@@ -0,0 +1,136 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=2e6B8_16kFV_ode_vidprom
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_vidprom16k/ode_vidprom8b16k_2e-6.out
|
||||
#SBATCH --error=ode_vidprom16k/ode_vidprom8b16k_2e-6.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your-conda-env
|
||||
|
||||
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_API_KEY=your-wandb-api-key
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="your-data-dir"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/causal_ode_init/validation.json"
|
||||
OUTPUT_DIR="your-output-dir"
|
||||
INIT_WEIGHTS_FROM_SAFETENSORS="your-init-weights-from-safetensors" # bidirectional weights from Wan2.1-T2V-1.3B-Diffusers
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir $OUTPUT_DIR
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "vidprom_8b16k_ode_init_2e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--warp_denoising_step
|
||||
--log_visualization
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--dmd_denoising_steps "1000,750,500,250"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# 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_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
--init_weights_from_safetensors $INIT_WEIGHTS_FROM_SAFETENSORS
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 500
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--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
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 1 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--flow_shift 5.0 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "ode_trajectory"
|
||||
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Elon Musk, dressed in a sleek white spacesuit with a reflective visor, walks confidently across the lunar surface. His posture is upright, and he moves steadily with purpose. The moon's rocky terrain and scattered boulders surround him, casting shadows under the dim sunlight. The background shows vast stretches of the moon's barren landscape with craters and dust clouds kicked up by his boots. The scene captures a wide shot, emphasizing the vastness and desolation of the lunar environment. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "In a dynamic action-packed sequence set in the Marvel multiverse, Spider-Man and Venom engage in an intense battle. Spider-Man, in his classic red and blue suit, swings and dodges venomous attacks from the black symbiote-covered Venom. Both characters display a range of acrobatic moves and powerful strikes. The environment is a chaotic urban landscape with crumbling buildings and neon lights, reflecting the multiversal theme. The camera captures the epic fight from various angles, including wide shots to show the scale of destruction and close-ups to highlight their fierce expressions and physical combat. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A warm, family-oriented scene depicting a father getting ready to leave the house to buy milk. The father, a middle-aged man with a kind face and a casual outfit, picks up a jacket from the coat rack. His posture is upright as he bends down slightly to put on his shoes. In the background, there are glimpses of a cozy living room with a family photograph on the wall. The camera focuses closely on the father, capturing his gentle smile and reassuring nod towards the camera before he opens the front door and steps outside. Static medium close-up shot. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Close-up shot of a man with a prosthetic hand that functions as a rocket launcher. He looks at his new hand with a mix of amazement and concern, his facial expression showing a blend of curiosity and apprehension. The prosthetic hand is sleek and metallic, with intricate details that resemble a high-tech weapon. The background is a dimly lit laboratory with various scientific equipment and monitors displaying data. The man stands in a relaxed posture, his other hand resting on his hip, as he inspects his new limb. The scene is rendered in a realistic sci-fi style, emphasizing the futuristic technology and the man's emotional response to his new appendage. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Realistic CCTV footage style, Kim Taehyung from the band BTS is involved in a drug deal, caught on camera. Kim Taehyung appears nervous and cautious, wearing casual clothing typical of a public space. He exchanges items discreetly with another person, who is partially obscured. Both individuals maintain a watchful demeanor, occasionally glancing around to ensure no one is watching them. The lighting is dim, with flickering fluorescent lights casting shadows on their faces. The background shows a typical urban setting with blurred figures moving in the distance. Static camera angle, medium close-up shot focusing on the interaction between Taehyung and the other individual. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Photorealistic studio setup with professional lighting, showcasing detailed cubic dissections of experimental plastic and felt-like materials on a pristine white background. Each cube reveals intricate layers and textures of the materials, emphasizing their unique properties. The scene has a shallow depth of field initially, then slowly pulls out to reveal the full arrangement of cubes, maintaining a wide depth of field throughout the transition. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -61,7 +61,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_only_checkpointing_steps 2000
|
||||
--training_state_checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -95,7 +95,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -93,7 +93,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=VSA_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=VSA_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# 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=VIDEO_SPARSE_ATTN
|
||||
# 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=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_VSA
|
||||
--output_dir "checkpoints/wan_t2v_finetune_VSA"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
# --enable_gradient_checkpointing_type "full" # if OOM enable this
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# 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 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "5.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-6
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 1
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# VSA arguments
|
||||
vsa_args=(
|
||||
--VSA_decay_rate 0.03 \
|
||||
--VSA_decay_interval_steps 50 \
|
||||
--VSA_sparsity 0.9 \
|
||||
)
|
||||
|
||||
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/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${vsa_args[@]}"
|
||||
@@ -91,9 +91,10 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 1e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=moe_finetune
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=2
|
||||
#SBATCH --ntasks=2
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=moe_output/moe_%j.out
|
||||
#SBATCH --error=moe_output/moe_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# 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 TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate wei-fv
|
||||
export HOME="/mnt/weka/home/hao.zhang/wei"
|
||||
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=16
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json"
|
||||
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name VSA_finetune # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/wei/Wan2.2-MoE-finetune"
|
||||
--max_train_steps 200
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 448 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
# --log_visualization
|
||||
--num_frames 81
|
||||
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
|
||||
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_TOTAL_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_TOTAL_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
# --log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "5.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--dit_cpu_offload True
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
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/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -61,7 +61,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -95,7 +95,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -61,7 +61,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -61,7 +61,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -92,7 +92,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 400
|
||||
--weight_only_checkpointing_steps 400
|
||||
--training_state_checkpointing_steps 400
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -61,7 +61,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 400
|
||||
--weight_only_checkpointing_steps 400
|
||||
--training_state_checkpointing_steps 400
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -28,4 +28,4 @@
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SageAttention3Backend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_ATTN_THREE"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SageAttention3Impl"]:
|
||||
return SageAttention3Impl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
# @staticmethod
|
||||
# def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
class SageAttention3Impl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
@@ -5,7 +5,6 @@ from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from flash_attn.bert_padding import pad_input
|
||||
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
@@ -134,6 +133,8 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
**extra_impl_args) -> None:
|
||||
self.prefix = prefix
|
||||
self.layer_idx = self._get_layer_idx(prefix)
|
||||
from flash_attn.bert_padding import pad_input
|
||||
self.pad_input = pad_input
|
||||
|
||||
def _get_layer_idx(self, prefix: str) -> int | None:
|
||||
match = re.search(r"blocks\.(\d+)", prefix)
|
||||
@@ -169,7 +170,6 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
moba_chunk_size = attn_metadata.st_chunk_size
|
||||
moba_topk = attn_metadata.st_topk
|
||||
|
||||
# torch.distributed.breakpoint()
|
||||
query, chunk_size = process_moba_input(query,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
simsum_threshold=attn_metadata.moba_threshold,
|
||||
threshold_type=attn_metadata.moba_threshold_type,
|
||||
)
|
||||
hidden_states = pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = process_moba_output(hidden_states,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
|
||||
@@ -15,18 +15,16 @@ class DiTArchConfig(ArchConfig):
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
)
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -92,6 +92,9 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
pos_embed_seq_len: int | None = None
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
# Wan MoE
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Causal Wan
|
||||
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
|
||||
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
|
||||
|
||||
@@ -4,12 +4,13 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -87,6 +87,7 @@ class PipelineConfig:
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
# isort: off
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
|
||||
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -48,6 +48,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
@@ -60,6 +61,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
@@ -54,6 +54,9 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp32", ))
|
||||
|
||||
# self-forcing params
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
@@ -82,7 +85,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
|
||||
default_factory=CLIPVisionConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
def __post_init__(self):
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
@@ -108,19 +111,17 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 757, 522])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
expand_timesteps: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -132,12 +133,26 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
# self-forcing params
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
class Wan2_2_I2V_A14B_Config(WanI2V480PConfig):
|
||||
flow_shift: float | None = 5.0
|
||||
boundary_ratio: float | None = 0.900
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -40,6 +40,7 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -47,6 +48,8 @@ class SamplingParam:
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.data_type = "video" if self.num_frames > 1 else "image"
|
||||
@@ -167,6 +170,12 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
default=SamplingParam.boundary_ratio,
|
||||
help="Boundary timestep ratio",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save-video",
|
||||
action="store_true",
|
||||
@@ -198,6 +207,18 @@ class SamplingParam:
|
||||
help=
|
||||
"Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-latents",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_latents,
|
||||
help="Whether to return the trajectory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-decoded",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_decoded,
|
||||
help="Whether to return the decoded trajectory",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -144,18 +144,22 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 4.0
|
||||
guidance_scale_2: float = 3.0
|
||||
guidance_scale: float = 4.0 # high_noise
|
||||
guidance_scale_2: float = 3.0 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 3.5
|
||||
guidance_scale_2: float = 3.5
|
||||
guidance_scale: float = 3.5 # high_noise
|
||||
guidance_scale_2: float = 3.5 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -4,7 +4,7 @@ from torchvision.transforms import Lambda
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
@@ -37,9 +37,15 @@ def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
def gettextdataset(args) -> TextDataset:
|
||||
return TextDataset(data_merge_path=args.data_merge_path,
|
||||
args=args,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset"
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
"""
|
||||
Utilities for converting preprocessing records (dicts) into Arrow tables and
|
||||
writing Parquet datasets in fixed-size chunks.
|
||||
|
||||
This module centralizes table construction and Parquet file writing so
|
||||
pipelines only need to define their PyArrow schema and produce per-sample
|
||||
record dictionaries.
|
||||
|
||||
Key APIs:
|
||||
- records_to_table(records, schema): Safely convert a list of dictionaries into
|
||||
a pa.Table, casting to the provided schema.
|
||||
- ParquetDatasetWriter: Buffer tables and flush to a directory as multiple
|
||||
Parquet files with a fixed number of rows per file. Uses temporary files and
|
||||
atomic rename to avoid partially written outputs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
|
||||
def records_to_table(records: list[dict[str, Any]], schema: pa.Schema) -> pa.Table:
|
||||
"""Build a PyArrow table from Python record dicts using an explicit schema.
|
||||
|
||||
Arrow will cast values to the target schema when possible (e.g., promoting
|
||||
Python ints/floats to pa.int64/pa.float64), eliminating hand-written per-
|
||||
field array construction.
|
||||
|
||||
Args:
|
||||
records: List of dictionaries, each representing one row. Keys must
|
||||
match schema field names.
|
||||
schema: Target PyArrow schema. Controls field names and types.
|
||||
|
||||
Returns:
|
||||
pa.Table: In-memory table matching the provided schema. If ``records``
|
||||
is empty, returns an empty table with the given schema.
|
||||
"""
|
||||
if not records:
|
||||
return pa.table({}, schema=schema)
|
||||
return pa.Table.from_pylist(records, schema=schema)
|
||||
|
||||
|
||||
class ParquetDatasetWriter:
|
||||
"""Accumulate tables and flush them to a Parquet directory in fixed-size chunks.
|
||||
|
||||
Behavior:
|
||||
- Writes files under worker-specific subdirectories for parallelism.
|
||||
- Uses temporary files and atomic rename to avoid partial files being left
|
||||
behind on failure.
|
||||
- Only full chunks of ``samples_per_file`` rows are written on each flush;
|
||||
any remainder rows are re-buffered for the next flush.
|
||||
|
||||
Note:
|
||||
- Instances are not meant to be shared across processes. Create one writer
|
||||
per process if using multiprocessing.
|
||||
"""
|
||||
|
||||
def __init__(self, out_dir: str, samples_per_file: int, compression: str = "zstd") -> None:
|
||||
"""Initialize the dataset writer.
|
||||
|
||||
Args:
|
||||
out_dir: Output directory where Parquet files will be written.
|
||||
samples_per_file: Fixed number of rows per Parquet file.
|
||||
compression: Compression codec passed to ``pyarrow.parquet.write_table``
|
||||
(e.g., ``"zstd"``, ``"snappy"``, ``"gzip"``).
|
||||
"""
|
||||
self.out_dir = out_dir
|
||||
self.samples_per_file = max(int(samples_per_file), 1)
|
||||
self.compression = compression
|
||||
os.makedirs(self.out_dir, exist_ok=True)
|
||||
self._tables: list[pa.Table] = []
|
||||
|
||||
def append_table(self, table: pa.Table) -> None:
|
||||
"""Append a non-empty table to the internal buffer.
|
||||
|
||||
Args:
|
||||
table: A ``pa.Table`` to buffer. Empty or ``None`` tables are ignored.
|
||||
"""
|
||||
if table is None or len(table) == 0:
|
||||
return
|
||||
self._tables.append(table)
|
||||
|
||||
def _combine(self) -> pa.Table | None:
|
||||
"""Combine all buffered tables into a single table, if any.
|
||||
|
||||
Returns:
|
||||
A concatenated table, a single table if only one was buffered, or
|
||||
``None`` if no tables are buffered.
|
||||
"""
|
||||
if not self._tables:
|
||||
return None
|
||||
if len(self._tables) == 1:
|
||||
return self._tables[0]
|
||||
return pa.concat_tables(self._tables, promote_options='none')
|
||||
|
||||
def flush(self, num_workers: int | None = None, write_remainder: bool = False) -> int:
|
||||
"""Write accumulated tables to disk and clear the written portion.
|
||||
|
||||
Only complete chunks of size ``samples_per_file`` are written. Any
|
||||
remainder rows are kept buffered for the next flush.
|
||||
|
||||
Args:
|
||||
num_workers: Optional override for the number of parallel workers
|
||||
used to write chunks. Defaults to ``min(cpu_count, chunks)``.
|
||||
write_remainder: If True, also write any leftover rows (< samples_per_file)
|
||||
as a final small Parquet file (useful for the last flush at the
|
||||
end of preprocessing).
|
||||
|
||||
Returns:
|
||||
int: Number of rows successfully written in this flush call.
|
||||
"""
|
||||
combined = self._combine()
|
||||
self._tables = []
|
||||
if combined is None or len(combined) == 0:
|
||||
return 0
|
||||
|
||||
num_samples = len(combined)
|
||||
total_chunks = num_samples // self.samples_per_file
|
||||
if total_chunks == 0:
|
||||
if not write_remainder:
|
||||
# Not enough to form a full chunk; keep buffered for next round
|
||||
# Re-buffer and return 0 written
|
||||
self._tables = [combined]
|
||||
return 0
|
||||
# Last flush: write the small remainder as a final file in worker_0
|
||||
worker_dir = os.path.join(self.out_dir, "worker_0")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
# Determine next index
|
||||
num_parquets = 0
|
||||
for _, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
chunk_path = os.path.join(worker_dir, f"data_chunk_{num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
pq.write_table(combined, temp_path, compression=self.compression)
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(chunk_path)
|
||||
os.rename(temp_path, chunk_path)
|
||||
return num_samples
|
||||
|
||||
# Only write full chunks; keep remainder for next flush
|
||||
written_rows = total_chunks * self.samples_per_file
|
||||
remainder = num_samples - written_rows
|
||||
|
||||
table_to_write = combined.slice(0, written_rows)
|
||||
remainder_table = combined.slice(written_rows, remainder) if remainder > 0 else None
|
||||
if remainder_table is not None and len(remainder_table) > 0:
|
||||
if write_remainder:
|
||||
# Write the remainder as a final small file (worker_0)
|
||||
worker_dir = os.path.join(self.out_dir, "worker_0")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
num_parquets = 0
|
||||
for _, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
remainder_path = os.path.join(worker_dir,
|
||||
f"data_chunk_{num_parquets}.parquet")
|
||||
temp_path = remainder_path + '.tmp'
|
||||
pq.write_table(remainder_table,
|
||||
temp_path,
|
||||
compression=self.compression)
|
||||
if os.path.exists(remainder_path):
|
||||
os.remove(remainder_path)
|
||||
os.rename(temp_path, remainder_path)
|
||||
else:
|
||||
self._tables = [remainder_table]
|
||||
|
||||
# Parallel write by chunk ranges
|
||||
if num_workers is None:
|
||||
num_workers = min(multiprocessing.cpu_count(), max(total_chunks, 1))
|
||||
num_workers = max(int(num_workers), 1)
|
||||
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
|
||||
|
||||
work_ranges: list[tuple[int, int, pa.Table, int, str, int, str]] = []
|
||||
for worker_id in range(num_workers):
|
||||
start_chunk = worker_id * chunks_per_worker
|
||||
end_chunk = min((worker_id + 1) * chunks_per_worker, total_chunks)
|
||||
if start_chunk < end_chunk:
|
||||
work_ranges.append(
|
||||
(
|
||||
start_chunk,
|
||||
end_chunk,
|
||||
table_to_write,
|
||||
worker_id,
|
||||
self.out_dir,
|
||||
self.samples_per_file,
|
||||
self.compression,
|
||||
)
|
||||
)
|
||||
|
||||
written_total = 0
|
||||
if len(work_ranges) == 1:
|
||||
written_total += _process_chunk_range(work_ranges[0])
|
||||
return written_total
|
||||
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = [executor.submit(_process_chunk_range, args) for args in work_ranges]
|
||||
for f in futures:
|
||||
written_total += f.result()
|
||||
return written_total + (len(remainder_table) if write_remainder and remainder_table is not None else 0)
|
||||
|
||||
|
||||
def _process_chunk_range(args: Any) -> int:
|
||||
"""Worker function to write a contiguous range of chunk files.
|
||||
|
||||
Args:
|
||||
args: Tuple containing
|
||||
- start_chunk (int): inclusive start chunk index
|
||||
- end_chunk (int): exclusive end chunk index
|
||||
- table (pa.Table): concatenated table containing all rows to write
|
||||
- worker_id (int): numeric worker identifier
|
||||
- output_dir (str): base output directory
|
||||
- samples_per_file (int): rows per chunk file
|
||||
- compression (str): compression codec for Parquet
|
||||
|
||||
Returns:
|
||||
int: Total number of rows written by this worker.
|
||||
"""
|
||||
start_chunk, end_chunk, table, worker_id, output_dir, samples_per_file, compression = args
|
||||
total_written = 0
|
||||
num_samples = len(table)
|
||||
|
||||
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
|
||||
# Offset to continue numbering if files exist
|
||||
num_parquets = 0
|
||||
for root, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
|
||||
for i in range(start_chunk, end_chunk):
|
||||
start_sample = i * samples_per_file
|
||||
end_sample = min((i + 1) * samples_per_file, num_samples)
|
||||
if end_sample <= start_sample:
|
||||
continue
|
||||
chunk = table.slice(start_sample, end_sample - start_sample)
|
||||
|
||||
chunk_path = os.path.join(worker_dir, f"data_chunk_{i + num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
try:
|
||||
pq.write_table(chunk, temp_path, compression=compression)
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(chunk_path)
|
||||
os.rename(temp_path, chunk_path)
|
||||
total_written += len(chunk)
|
||||
except Exception:
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
raise
|
||||
|
||||
return total_written
|
||||
|
||||
|
||||
|
||||
+68
@@ -1,5 +1,7 @@
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
@@ -120,3 +122,69 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
|
||||
})
|
||||
|
||||
return records
|
||||
|
||||
|
||||
def ode_text_only_record_creator(
|
||||
video_name: str, text_embedding: np.ndarray, caption: str,
|
||||
trajectory_latents: np.ndarray,
|
||||
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
|
||||
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
|
||||
|
||||
Args:
|
||||
video_name: Base name/id for the sample (without extension).
|
||||
text_embedding: Text encoder output array [SeqLen, Dim].
|
||||
caption: Original text prompt.
|
||||
trajectory_latents: Collected trajectory latents array.
|
||||
trajectory_timesteps: Collected timesteps array.
|
||||
|
||||
Returns:
|
||||
dict suitable for records_to_table(…, pyarrow_schema_ode_trajectory_text_only)
|
||||
"""
|
||||
assert trajectory_latents is not None, "trajectory_latents is required"
|
||||
assert trajectory_timesteps is not None, "trajectory_timesteps is required"
|
||||
|
||||
record = {
|
||||
"id": f"text_{video_name}",
|
||||
"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": caption,
|
||||
"media_type": "text",
|
||||
}
|
||||
|
||||
record.update({
|
||||
"trajectory_latents_bytes": trajectory_latents.tobytes(),
|
||||
"trajectory_latents_shape": list(trajectory_latents.shape),
|
||||
"trajectory_latents_dtype": str(trajectory_latents.dtype),
|
||||
})
|
||||
|
||||
record.update({
|
||||
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
|
||||
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
|
||||
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
|
||||
def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
|
||||
caption: str) -> dict[str, Any]:
|
||||
"""Create a text-only record matching pyarrow_schema_text_only.
|
||||
|
||||
Args:
|
||||
text_name: Base id/name for the text sample.
|
||||
text_embedding: Text encoder output array [SeqLen, Dim].
|
||||
caption: Original text prompt.
|
||||
|
||||
Returns:
|
||||
dict suitable for records_to_table(…, pyarrow_schema_text_only)
|
||||
"""
|
||||
record = {
|
||||
"id": f"text_{text_name}",
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"caption": caption,
|
||||
}
|
||||
return record
|
||||
@@ -50,6 +50,7 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
@@ -78,3 +79,40 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_ode_trajectory_text_only = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- ODE Trajectory ---
|
||||
pa.field("trajectory_latents_bytes", pa.binary()),
|
||||
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
|
||||
pa.field("trajectory_latents_dtype", pa.string()),
|
||||
pa.field("trajectory_timesteps_bytes", pa.binary()),
|
||||
pa.field("trajectory_timesteps_shape", pa.list_(pa.int64())),
|
||||
pa.field("trajectory_timesteps_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # Always 'text' for text-only
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_text_only = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("caption", pa.string()),
|
||||
])
|
||||
|
||||
@@ -628,3 +628,134 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
|
||||
|
||||
class TextDataset(torch.utils.data.IterableDataset,
|
||||
torch.distributed.checkpoint.stateful.Stateful):
|
||||
"""
|
||||
Text-only dataset for processing prompts from a simple text file.
|
||||
|
||||
Assumes that data_merge_path is a text file with one prompt per line:
|
||||
A cat playing with a ball
|
||||
A dog running in the park
|
||||
A person cooking dinner
|
||||
...
|
||||
|
||||
This dataset processes text data through text encoding stages only.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
data_merge_path: str,
|
||||
args,
|
||||
start_idx: int = 0,
|
||||
seed: int = 42):
|
||||
self.data_merge_path = data_merge_path
|
||||
self.start_idx = start_idx
|
||||
self.args = args
|
||||
self.seed = seed
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
# Initialize text encoding stage
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=getattr(args, 'training_cfg_rate', 0.0),
|
||||
seed=self.seed)
|
||||
|
||||
# Process text data
|
||||
self.processed_batches = self._process_text_data()
|
||||
|
||||
def _load_text_data(self) -> list[str]:
|
||||
"""Load text prompts from file."""
|
||||
prompts = []
|
||||
with open(self.data_merge_path, 'r', encoding='utf-8') as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line: # Skip empty lines
|
||||
prompts.append(line)
|
||||
|
||||
logger.info(f"Loaded {len(prompts)} text prompts from {self.data_merge_path}")
|
||||
return prompts
|
||||
|
||||
def _process_text_data(self) -> list[PreprocessBatch]:
|
||||
"""Process the text prompts through text encoding stage."""
|
||||
raw_prompts = self._load_text_data()
|
||||
processed_batches = []
|
||||
|
||||
for idx, prompt in enumerate(raw_prompts):
|
||||
# Create a text-only batch with dummy path
|
||||
batch = PreprocessBatch(
|
||||
path=f"text_prompt_{idx}",
|
||||
cap=[prompt], # TextEncodingStage expects a list
|
||||
resolution=None,
|
||||
fps=None,
|
||||
duration=None,
|
||||
num_frames=0,
|
||||
sample_frame_index=None,
|
||||
sample_num_frames=0
|
||||
)
|
||||
|
||||
processed_batches.append(batch)
|
||||
|
||||
logger.info(f"Processed {len(processed_batches)} text batches")
|
||||
return processed_batches
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterator for the dataset."""
|
||||
# Set up distributed sampling if needed
|
||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||
rank = torch.distributed.get_rank()
|
||||
world_size = torch.distributed.get_world_size()
|
||||
else:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
|
||||
# Calculate chunk for this rank
|
||||
total_items = len(self.processed_batches)
|
||||
items_per_rank = math.ceil(total_items / world_size)
|
||||
start_idx = rank * items_per_rank + self.start_idx
|
||||
end_idx = min(start_idx + items_per_rank, total_items)
|
||||
|
||||
# Yield items for this rank
|
||||
for idx in range(start_idx, end_idx):
|
||||
if idx < len(self.processed_batches):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def _get_item(self, idx: int) -> dict:
|
||||
"""Get a single processed text item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply text encoding stage
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary for text-only processing with required schema fields
|
||||
result = {
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
# Required schema fields for ODE trajectory processing
|
||||
"id": f"text_{idx}",
|
||||
"file_name": batch.path,
|
||||
"caption": batch.text,
|
||||
"media_type": "text",
|
||||
"width": 1,
|
||||
"height": 1,
|
||||
"num_frames": 0,
|
||||
"duration_sec": 0.0,
|
||||
"fps": 0.0,
|
||||
}
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
|
||||
@@ -5,7 +5,7 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
|
||||
@@ -344,6 +344,9 @@ class VideoGenerator:
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
"logging_info": logging_info,
|
||||
"trajectory": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
def set_lora_adapter(self,
|
||||
|
||||
@@ -26,6 +26,12 @@ if TYPE_CHECKING:
|
||||
NVCC_THREADS: str | None = None
|
||||
CMAKE_BUILD_TYPE: str | None = None
|
||||
VERBOSE: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
|
||||
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
FASTVIDEO_STAGE_LOGGING: bool = False
|
||||
|
||||
@@ -175,6 +181,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "SLIDING_TILE_ATTN" : use Sliding Tile Attention
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
@@ -196,6 +203,34 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os.
|
||||
path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
|
||||
|
||||
# Enable torch profiler to record shapes if set
|
||||
# FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1. If not set, torch profiler will
|
||||
# not record shapes.
|
||||
"FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES":
|
||||
lambda: bool(
|
||||
os.getenv("FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES", "0") != "0"),
|
||||
|
||||
# Enable torch profiler to profile memory if set
|
||||
# FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1. If not set, torch profiler
|
||||
# will not profile memory.
|
||||
"FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY":
|
||||
lambda: bool(
|
||||
os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY", "0") != "0"),
|
||||
|
||||
# Enable torch profiler to profile stack if set
|
||||
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. If not set, torch profiler WILL
|
||||
# profile stack by default.
|
||||
"FASTVIDEO_TORCH_PROFILER_WITH_STACK":
|
||||
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "1") != "0"),
|
||||
|
||||
# Enable torch profiler to profile flops if set
|
||||
# FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1. If not set, torch profiler will
|
||||
# not profile flops.
|
||||
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
|
||||
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS":
|
||||
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
|
||||
|
||||
# If set, fastvideo will run in development mode, which will enable
|
||||
# some additional endpoints for developing and debugging,
|
||||
# e.g. `/reset_prefix_cache`
|
||||
|
||||
+136
-6
@@ -133,6 +133,7 @@ class FastVideoArgs:
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
torch_compile_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
disable_autocast: bool = False
|
||||
|
||||
@@ -158,12 +159,15 @@ class FastVideoArgs:
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
override_transformer_cls_name: str | None = None
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
|
||||
|
||||
# # DMD parameters
|
||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||
|
||||
# MoE parameters used by Wan2.2
|
||||
boundary_ratio: float | None = None
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
@@ -329,6 +333,13 @@ class FastVideoArgs:
|
||||
help="Use torch.compile to speed up DiT inference." +
|
||||
"However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--torch-compile-kwargs",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--dit-cpu-offload",
|
||||
@@ -396,6 +407,21 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-transformer-cls-name",
|
||||
type=str,
|
||||
default=FastVideoArgs.override_transformer_cls_name,
|
||||
help="Override transformer cls name",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--init-weights-from-safetensors",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
parser.add_argument(
|
||||
"--init-weights-from-safetensors-2",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -424,6 +450,21 @@ class FastVideoArgs:
|
||||
mode_value = getattr(args, attr, FastVideoArgs.mode.value)
|
||||
kwargs['mode'] = ExecutionMode.from_string(
|
||||
mode_value) if isinstance(mode_value, str) else mode_value
|
||||
elif attr == 'torch_compile_kwargs':
|
||||
# Parse JSON string for torch.compile kwargs
|
||||
torch_compile_kwargs_str = getattr(args, 'torch_compile_kwargs',
|
||||
None)
|
||||
if torch_compile_kwargs_str:
|
||||
try:
|
||||
import json
|
||||
kwargs['torch_compile_kwargs'] = json.loads(
|
||||
torch_compile_kwargs_str)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(
|
||||
f"Invalid JSON for torch_compile_kwargs: {e}"
|
||||
) from e
|
||||
else:
|
||||
kwargs['torch_compile_kwargs'] = {}
|
||||
elif attr == 'workload_type':
|
||||
# Convert string to WorkloadType enum
|
||||
workload_type_value = getattr(args, 'workload_type',
|
||||
@@ -603,7 +644,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
real_score_model_path: str = "" # path for real score (teacher) model
|
||||
fake_score_model_path: str = "" # path for fake score (critic) model
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
@@ -625,7 +669,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
# output
|
||||
output_dir: str = ""
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
|
||||
|
||||
# optimizer & scheduler
|
||||
@@ -658,6 +701,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
linear_quadratic_threshold: float = 0.0
|
||||
linear_range: float = 0.0
|
||||
weight_decay: float = 0.0
|
||||
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
|
||||
use_ema: bool = False
|
||||
multi_phased_distill_schedule: str = ""
|
||||
pred_decay_weight: float = 0.0
|
||||
@@ -678,16 +722,29 @@ class TrainingArgs(FastVideoArgs):
|
||||
|
||||
# distillation args
|
||||
generator_update_interval: int = 5
|
||||
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
|
||||
min_timestep_ratio: float = 0.2
|
||||
max_timestep_ratio: float = 0.98
|
||||
real_score_guidance_scale: float = 3.5
|
||||
real_score_guidance_scale_2: float = 3.5
|
||||
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
|
||||
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
|
||||
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
|
||||
training_state_checkpointing_steps: int = 0 # for resuming training
|
||||
weight_only_checkpointing_steps: int = 0 # for inference
|
||||
log_visualization: bool = False
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
warp_denoising_step: bool = False
|
||||
|
||||
# Self-forcing specific arguments
|
||||
num_frame_per_block: int = 3
|
||||
independent_first_frame: bool = False
|
||||
enable_gradient_masking: bool = True
|
||||
gradient_mask_last_n_frames: int = 21
|
||||
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -789,6 +846,20 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
help="Directory to cache models")
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
parser.add_argument(
|
||||
"--generator-model-path",
|
||||
type=str,
|
||||
help="Path to generator (student) model for DMD distillation")
|
||||
parser.add_argument(
|
||||
"--real-score-model-path",
|
||||
type=str,
|
||||
help="Path to real score (teacher) model for DMD distillation")
|
||||
parser.add_argument(
|
||||
"--fake-score-model-path",
|
||||
type=str,
|
||||
help="Path to fake score (critic) model for DMD distillation")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
@@ -844,9 +915,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--checkpoints-total-limit",
|
||||
type=int,
|
||||
help="Maximum number of checkpoints to keep")
|
||||
parser.add_argument("--checkpointing-steps",
|
||||
type=int,
|
||||
help="Steps between checkpoints")
|
||||
parser.add_argument(
|
||||
"--training-state-checkpointing-steps",
|
||||
type=int,
|
||||
@@ -963,6 +1031,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Linear quadratic threshold")
|
||||
parser.add_argument("--linear-range", type=float, help="Linear range")
|
||||
parser.add_argument("--weight-decay", type=float, help="Weight decay")
|
||||
parser.add_argument("--betas",
|
||||
type=str,
|
||||
default=TrainingArgs.betas,
|
||||
help="Betas for optimizer (format: 'beta1,beta2')")
|
||||
parser.add_argument("--use-ema",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA")
|
||||
@@ -1013,6 +1085,13 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=int,
|
||||
default=TrainingArgs.generator_update_interval,
|
||||
help="Ratio of student updates to critic updates.")
|
||||
parser.add_argument(
|
||||
"--dfake-gen-update-ratio",
|
||||
type=int,
|
||||
default=TrainingArgs.dfake_gen_update_ratio,
|
||||
help=
|
||||
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
|
||||
)
|
||||
parser.add_argument("--min-timestep-ratio",
|
||||
type=float,
|
||||
default=TrainingArgs.min_timestep_ratio,
|
||||
@@ -1025,10 +1104,19 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=float,
|
||||
default=TrainingArgs.real_score_guidance_scale,
|
||||
help="Teacher guidance scale")
|
||||
parser.add_argument("--real-score-guidance-scale-2",
|
||||
type=float,
|
||||
default=TrainingArgs.real_score_guidance_scale_2,
|
||||
help="Teacher guidance scale")
|
||||
parser.add_argument("--fake-score-learning-rate",
|
||||
type=float,
|
||||
default=TrainingArgs.fake_score_learning_rate,
|
||||
help="Learning rate for fake score transformer")
|
||||
parser.add_argument(
|
||||
"--fake-score-betas",
|
||||
type=str,
|
||||
default=TrainingArgs.fake_score_betas,
|
||||
help="Betas for fake score optimizer (format: 'beta1,beta2')")
|
||||
parser.add_argument(
|
||||
"--fake-score-lr-scheduler",
|
||||
type=str,
|
||||
@@ -1041,6 +1129,48 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--simulate-generator-forward",
|
||||
action=StoreBoolean,
|
||||
help="Whether to simulate generator forward to match inference")
|
||||
parser.add_argument(
|
||||
"--warp-denoising-step",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Whether to warp denoising step according to the scheduler time shift"
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
parser.add_argument(
|
||||
"--num-frame-per-block",
|
||||
type=int,
|
||||
default=TrainingArgs.num_frame_per_block,
|
||||
help="Number of frames per block for causal generation")
|
||||
parser.add_argument(
|
||||
"--independent-first-frame",
|
||||
action=StoreBoolean,
|
||||
help="Whether the first frame is independent in causal generation")
|
||||
parser.add_argument(
|
||||
"--enable-gradient-masking",
|
||||
action=StoreBoolean,
|
||||
help="Whether to enable frame-level gradient masking")
|
||||
parser.add_argument(
|
||||
"--gradient-mask-last-n-frames",
|
||||
type=int,
|
||||
default=TrainingArgs.gradient_mask_last_n_frames,
|
||||
help="Number of last frames to enable gradients for")
|
||||
parser.add_argument(
|
||||
"--validate-cache-structure",
|
||||
action=StoreBoolean,
|
||||
help="Whether to validate KV cache structure (debug flag)")
|
||||
parser.add_argument(
|
||||
"--same-step-across-blocks",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use the same exit timestep for all blocks")
|
||||
parser.add_argument(
|
||||
"--last-step-only",
|
||||
action=StoreBoolean,
|
||||
help="Whether to only use the last timestep for training")
|
||||
parser.add_argument("--context-noise",
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@@ -77,9 +77,11 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
lora_A = self.lora_A.to_local()
|
||||
|
||||
if not self.merged and not self.disable_lora:
|
||||
delta = x @ (
|
||||
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
|
||||
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
|
||||
lora_A_sliced = self.slice_lora_a_weights(
|
||||
lora_A.to(x, non_blocking=True))
|
||||
lora_B_sliced = self.slice_lora_b_weights(
|
||||
lora_B.to(x, non_blocking=True))
|
||||
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
|
||||
if self.lora_alpha != self.lora_rank:
|
||||
delta = delta * (
|
||||
self.lora_alpha / self.lora_rank # type: ignore
|
||||
|
||||
@@ -147,6 +147,9 @@ class CausalWanSelfAttention(nn.Module):
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
kv_cache["k"] = kv_cache["k"].detach()
|
||||
kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = self.attn(
|
||||
@@ -176,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -209,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
# Only T2V for now
|
||||
@@ -223,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -249,29 +250,29 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
num_frames = temb.shape[1]
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
e = self.scale_shift_table + temb
|
||||
# e.shape: [batch_size, num_frames, 6, inner_dim]
|
||||
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
assert shift_msa.dtype == torch.float32
|
||||
# assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
|
||||
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -285,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -295,13 +294,10 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
crossattn_cache=crossattn_cache)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -364,8 +360,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -375,7 +370,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.num_frame_per_block = 1
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
assert self.num_frame_per_block <= 3
|
||||
self.independent_first_frame = False
|
||||
|
||||
self.__post_init__()
|
||||
@@ -487,12 +483,16 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -539,14 +539,9 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output
|
||||
return torch.stack(output)
|
||||
|
||||
def _forward_train(self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -587,8 +582,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.block_mask is None:
|
||||
@@ -601,8 +596,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -637,14 +636,9 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output
|
||||
return torch.stack(output)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -655,3 +649,30 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
return self._forward_inference(*args, **kwargs)
|
||||
else:
|
||||
return self._forward_train(*args, **kwargs)
|
||||
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
@@ -416,6 +416,11 @@ class TransformerLoader(ComponentLoader):
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
|
||||
logger.info("transformer cls_name: %s", cls_name)
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
cls_name = fastvideo_args.override_transformer_cls_name
|
||||
logger.info("Overriding transformer cls_name to %s", cls_name)
|
||||
|
||||
fastvideo_args.model_paths["transformer"] = model_path
|
||||
|
||||
# Config from Diffusers supersedes fastvideo's model config
|
||||
@@ -430,8 +435,24 @@ class TransformerLoader(ComponentLoader):
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
# Check if we should use custom initialization weights
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
|
||||
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
|
||||
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
|
||||
|
||||
if use_custom_weights:
|
||||
if 'transformer_2' in model_path:
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors_2', None)
|
||||
assert custom_weights_path is not None, "Custom initialization weights must be provided"
|
||||
if os.path.isdir(custom_weights_path):
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(custom_weights_path), "*.safetensors"))
|
||||
else:
|
||||
assert custom_weights_path.endswith(".safetensors"), "Custom initialization weights must be a safetensors file"
|
||||
safetensors_list = [custom_weights_path]
|
||||
|
||||
logger.info("Loading model from %s safetensors files: %s",
|
||||
len(safetensors_list), safetensors_list)
|
||||
|
||||
default_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
@@ -454,18 +475,20 @@ class TransformerLoader(ComponentLoader):
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
|
||||
fsdp_inference=fastvideo_args.use_fsdp_inference,
|
||||
# TODO(will): make these configurable
|
||||
default_dtype=default_dtype,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
training_mode=fastvideo_args.training_mode)
|
||||
training_mode=fastvideo_args.training_mode,
|
||||
enable_torch_compile=fastvideo_args.enable_torch_compile,
|
||||
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs)
|
||||
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
dtypes = set(param.dtype for param in model.parameters())
|
||||
if len(dtypes) > 1:
|
||||
model = model.to(default_dtype)
|
||||
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
|
||||
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -54,7 +54,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
|
||||
torch.set_default_dtype(old_dtype)
|
||||
|
||||
|
||||
# TODO(PY): add compile option
|
||||
# Supports optional torch.compile for FSDP-wrapped models during training
|
||||
def maybe_load_fsdp_model(
|
||||
model_cls: type[nn.Module],
|
||||
init_params: dict[str, Any],
|
||||
@@ -62,6 +62,7 @@ def maybe_load_fsdp_model(
|
||||
device: torch.device,
|
||||
hsdp_replicate_dim: int,
|
||||
hsdp_shard_dim: int,
|
||||
default_dtype: torch.dtype,
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
cpu_offload: bool = False,
|
||||
@@ -69,6 +70,8 @@ def maybe_load_fsdp_model(
|
||||
output_dtype: torch.dtype | None = None,
|
||||
training_mode: bool = True,
|
||||
pin_cpu_memory: bool = True,
|
||||
enable_torch_compile: bool = False,
|
||||
torch_compile_kwargs: dict[str, Any] | None = None,
|
||||
) -> torch.nn.Module:
|
||||
"""
|
||||
Load the model with FSDP if is training, else load the model without FSDP.
|
||||
@@ -87,7 +90,8 @@ def maybe_load_fsdp_model(
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
|
||||
with set_default_dtype(param_dtype), torch.device("meta"):
|
||||
logger.info("Loading model with default_dtype: %s", default_dtype)
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
# Check if we should use FSDP
|
||||
@@ -125,7 +129,7 @@ def maybe_load_fsdp_model(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
param_dtype,
|
||||
default_dtype,
|
||||
strict=True,
|
||||
cpu_offload=cpu_offload,
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
@@ -137,6 +141,14 @@ def maybe_load_fsdp_model(
|
||||
# Avoid unintended computation graph accumulation during inference
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
p.requires_grad = False
|
||||
|
||||
compile_in_loader = enable_torch_compile and training_mode
|
||||
if compile_in_loader:
|
||||
compile_kwargs = torch_compile_kwargs or {}
|
||||
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s",
|
||||
compile_kwargs)
|
||||
model = torch.compile(model, **compile_kwargs)
|
||||
logger.info("torch.compile enabled for %s", type(model).__name__)
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
|
||||
"""
|
||||
Args:
|
||||
clean_latent: the clean latent with shape [B, C, H, W],
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise: the noise with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
|
||||
Returns:
|
||||
the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
@@ -22,8 +22,10 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
|
||||
|
||||
config_name = "scheduler_config.json"
|
||||
order = 1
|
||||
@register_to_config
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
@@ -62,8 +64,15 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
elif timestep.ndim == 0:
|
||||
# handles the case where timestep is a scalar, this occurs when we
|
||||
# use this scheduler for ODE trajectory
|
||||
timestep = timestep.unsqueeze(0)
|
||||
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep = timestep.to(model_output.device)
|
||||
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -171,10 +171,10 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
# timestep shape should be [B]
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
pred_noise = pred_noise.float().to(device)
|
||||
noise_input_latent = noise_input_latent.float().to(device)
|
||||
sigmas = scheduler.sigmas.float().to(device)
|
||||
timesteps = scheduler.timesteps.float().to(device)
|
||||
pred_noise = pred_noise.double().to(device)
|
||||
noise_input_latent = noise_input_latent.double().to(device)
|
||||
sigmas = scheduler.sigmas.double().to(device)
|
||||
timesteps = scheduler.timesteps.double().to(device)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -7,8 +7,6 @@ This module wires the causal DMD denoising stage into the modular pipeline.
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
@@ -28,10 +26,6 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
@@ -55,6 +49,7 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=CausalDMDDenosingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
|
||||
@@ -14,12 +14,14 @@ import torch
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.distributed import (
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
maybe_init_distributed_environment_and_model_parallel, get_world_group)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.profiler import get_or_create_profiler
|
||||
from fastvideo.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages import PipelineStage
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.utils import (maybe_download_model,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
@@ -40,7 +42,7 @@ class ComposedPipelineBase(ABC):
|
||||
_extra_config_module_map: dict[str, str] = {}
|
||||
training_args: TrainingArgs | None = None
|
||||
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
|
||||
modules: dict[str, torch.nn.Module] = {}
|
||||
modules: dict[str, Any] = {}
|
||||
post_init_called: bool = False
|
||||
|
||||
# TODO(will): args should support both inference args and training args
|
||||
@@ -69,9 +71,19 @@ class ComposedPipelineBase(ABC):
|
||||
maybe_init_distributed_environment_and_model_parallel(
|
||||
fastvideo_args.tp_size, fastvideo_args.sp_size)
|
||||
|
||||
# Torch profiler. Enabled and configured through env vars:
|
||||
# FASTVIDEO_TORCH_PROFILER_DIR=/path/to/save/trace
|
||||
trace_dir = envs.FASTVIDEO_TORCH_PROFILER_DIR
|
||||
self.profiler_controller = get_or_create_profiler(trace_dir)
|
||||
self.profiler = self.profiler_controller.profiler
|
||||
|
||||
self.local_rank = get_world_group().local_rank
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
self.modules = self.load_modules(fastvideo_args, loaded_modules)
|
||||
fastvideo_args.dit_cpu_offload = False
|
||||
with self.profiler_controller.region("profiler_region_model_loading"):
|
||||
self.modules = self.load_modules(fastvideo_args, loaded_modules)
|
||||
|
||||
def set_trainable(self) -> None:
|
||||
# Only train DiT
|
||||
@@ -99,9 +111,30 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
self.initialize_pipeline(self.fastvideo_args)
|
||||
if self.fastvideo_args.enable_torch_compile:
|
||||
self.modules["transformer"] = torch.compile(
|
||||
self.modules["transformer"])
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
transformer_module = self.modules["transformer"]
|
||||
if self.fastvideo_args.training_mode:
|
||||
logger.info(
|
||||
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
|
||||
)
|
||||
else:
|
||||
fsdp_module_cls = None
|
||||
try:
|
||||
from torch.distributed.fsdp import FSDPModule # type: ignore
|
||||
fsdp_module_cls = FSDPModule
|
||||
except Exception: # pragma: no cover - FSDP not always available
|
||||
fsdp_module_cls = None
|
||||
if fsdp_module_cls is not None and isinstance(
|
||||
transformer_module, fsdp_module_cls):
|
||||
logger.info(
|
||||
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
|
||||
)
|
||||
else:
|
||||
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
|
||||
logger.info("Enabling torch.compile for DiT with kwargs=%s",
|
||||
compile_kwargs)
|
||||
self.modules["transformer"] = torch.compile(
|
||||
transformer_module, **compile_kwargs)
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
|
||||
if not self.fastvideo_args.training_mode:
|
||||
logger.info("Creating pipeline stages...")
|
||||
@@ -237,20 +270,19 @@ class ComposedPipelineBase(ABC):
|
||||
# remove keys that are not pipeline modules
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
# @TODO(Wei): Temporary hack
|
||||
if "boundary_ratio" in model_index and model_index[
|
||||
"boundary_ratio"] is not None:
|
||||
logger.info(
|
||||
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
|
||||
)
|
||||
self.required_config_modules.append("transformer_2")
|
||||
if fastvideo_args.boundary_ratio is None:
|
||||
logger.info(
|
||||
"MoE pipeline detected. Setting boundary ratio to %s",
|
||||
model_index["boundary_ratio"])
|
||||
fastvideo_args.boundary_ratio = model_index["boundary_ratio"]
|
||||
logger.info("MoE pipeline detected. Setting boundary ratio to %s",
|
||||
model_index["boundary_ratio"])
|
||||
fastvideo_args.pipeline_config.dit_config.boundary_ratio = model_index[
|
||||
"boundary_ratio"]
|
||||
|
||||
model_index.pop("boundary_ratio", None)
|
||||
# used by Wan2.2 ti2v
|
||||
model_index.pop("expand_timesteps", None)
|
||||
|
||||
# some sanity checks
|
||||
@@ -283,8 +315,8 @@ class ComposedPipelineBase(ABC):
|
||||
architecture) in model_index.items():
|
||||
if transformers_or_diffusers is None:
|
||||
logger.warning(
|
||||
"Module in model_index.json has null value, removing from required_config_modules"
|
||||
)
|
||||
"Module %s in model_index.json has null value, removing from required_config_modules",
|
||||
module_name)
|
||||
if module_name in self.required_config_modules:
|
||||
self.required_config_modules.remove(module_name)
|
||||
continue
|
||||
@@ -332,6 +364,18 @@ class ComposedPipelineBase(ABC):
|
||||
self._stage_name_mapping[stage_name] = stage
|
||||
setattr(self, stage_name, stage)
|
||||
|
||||
def profile(self, is_start: bool = True):
|
||||
if self.profiler is None:
|
||||
raise RuntimeError("Profiler is not enabled.")
|
||||
if is_start:
|
||||
self.profiler.start()
|
||||
else:
|
||||
self.profiler.stop()
|
||||
# only print profiler results on rank 0
|
||||
if self.local_rank == 0:
|
||||
print(self.profiler.key_averages().table(
|
||||
sort_by="self_cuda_time_total"))
|
||||
|
||||
# TODO(will): don't hardcode no_grad
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
|
||||
@@ -129,6 +129,7 @@ class ForwardBatch:
|
||||
timesteps: torch.Tensor | None = None
|
||||
timestep: torch.Tensor | float | int | None = None
|
||||
step_index: int | None = None
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Scheduler parameters
|
||||
num_inference_steps: int = 50
|
||||
@@ -147,7 +148,12 @@ class ForwardBatch:
|
||||
modules: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Final output (after pipeline completion)
|
||||
output: Any = None
|
||||
output: torch.Tensor | None = None
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
trajectory_timesteps: list[torch.Tensor] | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_decoded: list[torch.Tensor] | None = None
|
||||
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
@@ -206,6 +212,10 @@ class TrainingBatch:
|
||||
infos: list[dict[str, Any]] | None = None
|
||||
mask_lat_size: torch.Tensor | None = None
|
||||
|
||||
# ODE trajectory supervision
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: torch.Tensor | None = None
|
||||
timesteps: torch.Tensor | None = None
|
||||
@@ -236,6 +246,7 @@ class TrainingBatch:
|
||||
fake_score_loss: float = 0.0
|
||||
|
||||
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -12,6 +10,8 @@ from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import getdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.preprocessing_datasets import PreprocessBatch
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -54,10 +54,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"""Get additional features specific to the pipeline type. Override in subclasses."""
|
||||
return {}
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for the pipeline type. Override in subclasses."""
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for this pipeline. Must be overridden."""
|
||||
raise NotImplementedError
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for the pipeline type."""
|
||||
return [f.name for f in self.get_pyarrow_schema()]
|
||||
|
||||
def create_record_for_schema(self,
|
||||
preprocess_batch: PreprocessBatch,
|
||||
schema: pa.Schema,
|
||||
@@ -400,166 +404,22 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
# Add progress bar for writing to Parquet dataset
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
# Convert batch data to PyArrow arrays
|
||||
arrays = []
|
||||
for field in self.get_schema_fields():
|
||||
if field.endswith('_bytes'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.binary()))
|
||||
elif field.endswith('_shape'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.list_(pa.int32())))
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.int32()))
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.float32()))
|
||||
else:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data]))
|
||||
|
||||
table = pa.Table.from_arrays(arrays,
|
||||
names=self.get_schema_fields())
|
||||
table = records_to_table(batch_data, self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
# Store the table in a list for later processing
|
||||
if not hasattr(self, 'all_tables'):
|
||||
self.all_tables = []
|
||||
self.all_tables.append(table)
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if num_processed_samples >= args.flush_frequency:
|
||||
self._flush_tables(num_processed_samples, args,
|
||||
combined_parquet_dir)
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
num_processed_samples = 0
|
||||
self.all_tables = []
|
||||
|
||||
def _flush_tables(self, num_processed_samples: int, args,
|
||||
combined_parquet_dir: str):
|
||||
"""Flush collected tables to disk."""
|
||||
assert hasattr(self, 'all_tables') and self.all_tables
|
||||
print(f"Combining {len(self.all_tables)} batches...")
|
||||
combined_table = pa.concat_tables(self.all_tables)
|
||||
assert len(combined_table) == num_processed_samples
|
||||
print(f"Total samples collected: {len(combined_table)}")
|
||||
|
||||
# Calculate total number of chunks needed, discarding remainder
|
||||
total_chunks = max(num_processed_samples // args.samples_per_file, 1)
|
||||
|
||||
print(f"Fixed samples per parquet file: {args.samples_per_file}")
|
||||
print(f"Total number of parquet files: {total_chunks}")
|
||||
print(
|
||||
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
|
||||
)
|
||||
|
||||
# Split work among processes
|
||||
num_workers = int(min(multiprocessing.cpu_count(), total_chunks))
|
||||
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
|
||||
|
||||
print(f"Using {num_workers} workers to process {total_chunks} chunks")
|
||||
logger.info("Chunks per worker: %s", chunks_per_worker)
|
||||
|
||||
# Prepare work ranges
|
||||
work_ranges = []
|
||||
for i in range(num_workers):
|
||||
start_idx = i * chunks_per_worker
|
||||
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
|
||||
if start_idx < total_chunks:
|
||||
work_ranges.append(
|
||||
(start_idx, end_idx, combined_table, i,
|
||||
combined_parquet_dir, args.samples_per_file))
|
||||
|
||||
total_written = 0
|
||||
failed_ranges = []
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = {
|
||||
executor.submit(self.process_chunk_range, work_range):
|
||||
work_range
|
||||
for work_range in work_ranges
|
||||
}
|
||||
for future in tqdm(futures, desc="Processing chunks"):
|
||||
try:
|
||||
written = future.result()
|
||||
total_written += written
|
||||
logger.info("Processed chunk with %s samples", written)
|
||||
except Exception as e:
|
||||
work_range = futures[future]
|
||||
failed_ranges.append(work_range)
|
||||
logger.error("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
# Retry failed ranges sequentially
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.process_chunk_range(work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
logger.info("Total samples written: %s", total_written)
|
||||
|
||||
@staticmethod
|
||||
def process_chunk_range(args: Any) -> int:
|
||||
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
|
||||
try:
|
||||
total_written = 0
|
||||
num_samples = len(table)
|
||||
|
||||
# Create worker-specific subdirectory
|
||||
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
|
||||
# Check how many files there are already in the dir, and update i accordingly
|
||||
num_parquets = 0
|
||||
for root, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
|
||||
for i in range(start_idx, end_idx):
|
||||
start_sample = i * samples_per_file
|
||||
end_sample = min((i + 1) * samples_per_file, num_samples)
|
||||
chunk = table.slice(start_sample, end_sample - start_sample)
|
||||
|
||||
# Create chunk file in worker's directory
|
||||
chunk_path = os.path.join(
|
||||
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
|
||||
try:
|
||||
# Write to temporary file
|
||||
pq.write_table(chunk, temp_path, compression='zstd')
|
||||
|
||||
# Rename temporary file to final file
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(
|
||||
chunk_path) # Remove existing file if it exists
|
||||
os.rename(temp_path, chunk_path)
|
||||
|
||||
total_written += len(chunk)
|
||||
except Exception as e:
|
||||
# Clean up temporary file if it exists
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
raise e
|
||||
|
||||
return total_written
|
||||
except Exception as e:
|
||||
logger.error("Error processing chunks %s-%s for worker %s: %s",
|
||||
start_idx, end_idx, worker_id, str(e))
|
||||
raise
|
||||
|
||||
@@ -40,9 +40,9 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
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_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for I2V pipeline."""
|
||||
return pyarrow_schema_i2v
|
||||
|
||||
def get_extra_features(self, valid_data: dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
|
||||
@@ -0,0 +1,323 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
ODE Trajectory Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
|
||||
using the modular pipeline architecture.
|
||||
|
||||
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
ode_text_only_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
"""ODE Trajectory preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int
|
||||
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for ODE Trajectory pipeline."""
|
||||
return pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
assert fastvideo_args.pipeline_config.flow_shift == 5
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
|
||||
denoising_strength=1.0)
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
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="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None)))
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
pipeline=self,
|
||||
))
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
|
||||
args):
|
||||
"""Preprocess text-only data and generate trajectory information."""
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, text in enumerate(data["text"]):
|
||||
if text and text.strip(): # Check if text is not empty
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples (text-only)
|
||||
valid_data = {
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
# Add fps and duration if available in data
|
||||
if "fps" in data:
|
||||
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
|
||||
if "duration" in data:
|
||||
valid_data["duration"] = [
|
||||
data["duration"][i] for i in valid_indices
|
||||
]
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
sampling_params = SamplingParam.from_pretrained(args.model_path)
|
||||
|
||||
# encode negative prompt for trajectory collection
|
||||
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
|
||||
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
sampling_params.negative_prompt,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
negative_prompt_embed = negative_prompt_embeds_list[0][0]
|
||||
negative_prompt_attention_mask = negative_prompt_masks_list[
|
||||
0][0]
|
||||
else:
|
||||
negative_prompt_embed = None
|
||||
negative_prompt_attention_mask = None
|
||||
|
||||
trajectory_latents = []
|
||||
trajectory_timesteps = []
|
||||
trajectory_decoded = []
|
||||
|
||||
for i, (prompt_embed, prompt_attention_mask) in enumerate(
|
||||
zip(prompt_embeds, prompt_attention_masks,
|
||||
strict=False)):
|
||||
prompt_embed = prompt_embed.unsqueeze(0)
|
||||
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
|
||||
|
||||
# Collect the trajectory data (text-to-video generation)
|
||||
batch = ForwardBatch(**shallow_asdict(sampling_params), )
|
||||
batch.prompt_embeds = [prompt_embed]
|
||||
batch.prompt_attention_mask = [prompt_attention_mask]
|
||||
batch.negative_prompt_embeds = [negative_prompt_embed]
|
||||
batch.negative_attention_mask = [
|
||||
negative_prompt_attention_mask
|
||||
]
|
||||
batch.num_inference_steps = 48
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
batch.return_trajectory_decoded = False
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
batch.fps = args.train_fps
|
||||
batch.guidance_scale = 6.0
|
||||
batch.do_classifier_free_guidance = True
|
||||
|
||||
result_batch = self.input_validation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.timestep_preparation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.latent_preparation_stage(
|
||||
result_batch, fastvideo_args)
|
||||
result_batch = self.denoising_stage(result_batch,
|
||||
fastvideo_args)
|
||||
result_batch = self.decoding_stage(result_batch,
|
||||
fastvideo_args)
|
||||
|
||||
trajectory_latents.append(
|
||||
result_batch.trajectory_latents.cpu())
|
||||
trajectory_timesteps.append(
|
||||
result_batch.trajectory_timesteps.cpu())
|
||||
trajectory_decoded.append(result_batch.trajectory_decoded)
|
||||
|
||||
# Prepare extra features for text-only processing
|
||||
extra_features = {
|
||||
"trajectory_latents": trajectory_latents,
|
||||
"trajectory_timesteps": trajectory_timesteps
|
||||
}
|
||||
|
||||
if batch.return_trajectory_decoded:
|
||||
for i, decoded_frames in enumerate(trajectory_decoded):
|
||||
for j, decoded_frame in enumerate(decoded_frames):
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data: list[dict[str, Any]] = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
for idx, video_path in save_pbar:
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Get extra features for this sample
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).numpy()
|
||||
else:
|
||||
assert isinstance(value, list)
|
||||
if isinstance(value[idx], torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).float().numpy()
|
||||
else:
|
||||
sample_extra_features[key] = value[idx]
|
||||
|
||||
# Create record for Parquet dataset (text-only ODE schema)
|
||||
record: dict[str, Any] = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
trajectory_latents=sample_extra_features[
|
||||
"trajectory_latents"],
|
||||
trajectory_timesteps=sample_extra_features[
|
||||
"trajectory_timesteps"],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=self.combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
self.num_processed_samples = 0
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
if written:
|
||||
logger.info("Final flush wrote %s samples", written)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for video preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing videos",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.video_data: dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: dict[str, Any] = {} # Store latent tensors
|
||||
self.preprocess_text_and_trajectory(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_ODE_Trajectory
|
||||
@@ -15,9 +15,9 @@ class PreprocessPipeline_T2V(BasePreprocessPipeline):
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
|
||||
|
||||
def get_schema_fields(self):
|
||||
"""Get the schema fields for T2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_t2v]
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for T2V pipeline."""
|
||||
return pyarrow_schema_t2v
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_T2V
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Text-only Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Text-only Data Preprocessing pipeline
|
||||
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import text_only_record_creator
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import TextEncodingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PreprocessPipeline_Text(BasePreprocessPipeline):
|
||||
"""Text-only preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer"]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int = 0
|
||||
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for text-only pipeline."""
|
||||
return pyarrow_schema_text_only
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
def preprocess_text_only(self, fastvideo_args: FastVideoArgs, args):
|
||||
"""Preprocess text-only data."""
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, text in enumerate(data["text"]):
|
||||
if text and text.strip(): # Check if text is not empty
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples (text-only)
|
||||
valid_data = {
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
logger.info("===== prompt_embeds: %s", prompt_embeds.shape)
|
||||
logger.info("===== prompt_attention_masks: %s",
|
||||
prompt_attention_masks.shape)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
for idx, text_path in save_pbar:
|
||||
text_name = os.path.basename(text_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Create record for Parquet dataset (text-only schema)
|
||||
record = text_only_record_creator(
|
||||
text_name=text_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
pyarrow_schema_text_only)
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=self.combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
self.num_processed_samples = 0
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
if written:
|
||||
logger.info("Final flush wrote %s samples", written)
|
||||
|
||||
# Text-only record creation moved to fastvideo.dataset.dataloader.record_schema
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading text dataset
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for text preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing text",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.text_data: dict[str, Any] = {} # Store text metadata and paths
|
||||
|
||||
self.preprocess_text_only(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_Text
|
||||
@@ -1,5 +1,6 @@
|
||||
import argparse
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
@@ -9,8 +10,12 @@ from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
|
||||
PreprocessPipeline_I2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
|
||||
PreprocessPipeline_ODE_Trajectory)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
|
||||
PreprocessPipeline_Text)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -21,12 +26,22 @@ def main(args) -> None:
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if args.preprocess_task == "text_only":
|
||||
kwargs = {
|
||||
"text_encoder_cpu_offload": False,
|
||||
}
|
||||
else:
|
||||
# Full config for video/image processing
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
|
||||
}
|
||||
pipeline_config.update_config_from_dict(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=args.model_path,
|
||||
num_gpus=get_world_size(),
|
||||
@@ -35,7 +50,23 @@ def main(args) -> None:
|
||||
text_encoder_cpu_offload=False,
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
|
||||
if args.preprocess_task == "t2v":
|
||||
PreprocessPipeline = PreprocessPipeline_T2V
|
||||
elif args.preprocess_task == "i2v":
|
||||
PreprocessPipeline = PreprocessPipeline_I2V
|
||||
elif args.preprocess_task == "text_only":
|
||||
PreprocessPipeline = PreprocessPipeline_Text
|
||||
elif args.preprocess_task == "ode_trajectory":
|
||||
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
|
||||
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
|
||||
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
|
||||
else:
|
||||
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only")
|
||||
|
||||
logger.info("Preprocess task: %s using %s", args.preprocess_task,
|
||||
PreprocessPipeline.__name__)
|
||||
|
||||
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
@@ -74,7 +105,12 @@ 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("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--flow_shift", type=float, default=None)
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
|
||||
help="Type of preprocessing task to run")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import torch # type: ignore
|
||||
import gc
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -34,13 +35,15 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
Denoising stage for causal diffusion.
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__(transformer, scheduler)
|
||||
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
|
||||
super().__init__(transformer, scheduler, transformer_2)
|
||||
# KV and cross-attention cache state (initialized on first forward)
|
||||
self.kv_cache1: list | None = None
|
||||
self.crossattn_cache: list | None = None
|
||||
self.transformer = transformer
|
||||
self.transformer_2 = transformer_2
|
||||
# self.kv_cache1: list | None = None
|
||||
# self.crossattn_cache: list | None = None
|
||||
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
||||
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
|
||||
self.num_transformer_blocks = len(self.transformer.blocks)
|
||||
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
|
||||
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
|
||||
|
||||
@@ -65,21 +68,18 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
-1] * self.transformer.config.arch_config.patch_size[-2]
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
# TODO(will): make this a parameter once we add i2v support
|
||||
independent_first_frame = self.transformer.independent_first_frame
|
||||
|
||||
independent_first_frame = self.transformer.independent_first_frame if hasattr(
|
||||
self.transformer, 'independent_first_frame') else False
|
||||
# Timesteps for DMD
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
logger.info("Warping timesteps...")
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
logger.info("Using timesteps: %s", timesteps)
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_kwargs: dict = {}
|
||||
@@ -104,34 +104,43 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
# Initialize or reset caches
|
||||
if self.kv_cache1 is None:
|
||||
self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
self._initialize_crossattn_cache(
|
||||
crossattn_cache = self._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=fastvideo_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
else:
|
||||
assert self.crossattn_cache is not None
|
||||
# reset cross-attention cache
|
||||
for block_index in range(self.num_transformer_blocks):
|
||||
self.crossattn_cache[block_index][
|
||||
"is_init"] = False # type: ignore
|
||||
# reset kv cache pointers
|
||||
for block_index in range(len(self.kv_cache1)):
|
||||
self.kv_cache1[block_index][
|
||||
"global_end_index"] = torch.tensor( # type: ignore
|
||||
[0],
|
||||
dtype=torch.long,
|
||||
device=latents.device)
|
||||
self.kv_cache1[block_index][
|
||||
"local_end_index"] = torch.tensor( # type: ignore
|
||||
[0],
|
||||
dtype=torch.long,
|
||||
device=latents.device)
|
||||
# if self.kv_cache1 is None:
|
||||
# self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
# dtype=target_dtype,
|
||||
# device=latents.device)
|
||||
# self._initialize_crossattn_cache(
|
||||
# batch_size=latents.shape[0],
|
||||
# max_text_len=fastvideo_args.pipeline_config.
|
||||
# text_encoder_configs[0].arch_config.text_len,
|
||||
# dtype=target_dtype,
|
||||
# device=latents.device)
|
||||
# else:
|
||||
# assert self.crossattn_cache is not None
|
||||
# # reset cross-attention cache
|
||||
# for block_index in range(self.num_transformer_blocks):
|
||||
# self.crossattn_cache[block_index][
|
||||
# "is_init"] = False # type: ignore
|
||||
# # reset kv cache pointers
|
||||
# for block_index in range(len(self.kv_cache1)):
|
||||
# self.kv_cache1[block_index][
|
||||
# "global_end_index"] = torch.tensor( # type: ignore
|
||||
# [0],
|
||||
# dtype=torch.long,
|
||||
# device=latents.device)
|
||||
# self.kv_cache1[block_index][
|
||||
# "local_end_index"] = torch.tensor( # type: ignore
|
||||
# [0],
|
||||
# dtype=torch.long,
|
||||
# device=latents.device)
|
||||
|
||||
# Optional: cache context features from provided image latents prior to generation
|
||||
current_start_frame = 0
|
||||
@@ -154,8 +163,8 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
image_first_btchw,
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=current_start_frame *
|
||||
self.frame_seq_length,
|
||||
**image_kwargs,
|
||||
@@ -180,8 +189,8 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
ref_btchw,
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=current_start_frame *
|
||||
self.frame_seq_length,
|
||||
**image_kwargs,
|
||||
@@ -223,6 +232,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
video_raw_latent_shape = noise_latents_btchw.shape
|
||||
|
||||
for i, t_cur in enumerate(timesteps):
|
||||
if self.transformer_2 is not None and fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None and t_cur < fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps:
|
||||
current_model = self.transformer_2
|
||||
else:
|
||||
current_model = self.transformer
|
||||
# Copy for pred conversion
|
||||
noise_latents = noise_latents_btchw.clone()
|
||||
latent_model_input = current_latents.to(target_dtype)
|
||||
@@ -273,12 +286,12 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
pred_noise_btchw = self.transformer(
|
||||
pred_noise_btchw = current_model(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
@@ -338,12 +351,12 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
_ = self.transformer(
|
||||
_ = current_model(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
@@ -352,6 +365,11 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
del kv_cache1
|
||||
del crossattn_cache
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
@@ -389,7 +407,8 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
self.kv_cache1 = kv_cache1
|
||||
# self.kv_cache1 = kv_cache1
|
||||
return kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
|
||||
device) -> None:
|
||||
@@ -418,7 +437,8 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
self.crossattn_cache = crossattn_cache
|
||||
# self.crossattn_cache = crossattn_cache
|
||||
return crossattn_cache
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
|
||||
@@ -50,6 +50,63 @@ class DecodingStage(PipelineStage):
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
"""
|
||||
Decode latent representations into pixel space using VAE.
|
||||
|
||||
Args:
|
||||
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
|
||||
fastvideo_args: Configuration containing:
|
||||
- disable_autocast: Whether to disable automatic mixed precision (default: False)
|
||||
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
|
||||
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
|
||||
|
||||
Returns:
|
||||
Decoded video tensor with shape (batch, channels, frames, height, width),
|
||||
normalized to [0, 1] range and moved to CPU as float32
|
||||
"""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
@@ -59,13 +116,28 @@ class DecodingStage(PipelineStage):
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
This method processes the batch through the VAE decoder, converting latent
|
||||
representations to pixel-space video/images. It also optionally decodes
|
||||
trajectory latents for visualization purposes.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
batch: The current batch containing:
|
||||
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
|
||||
- return_trajectory_decoded (optional): Flag to decode trajectory latents
|
||||
- trajectory_latents (optional): Latents at different timesteps
|
||||
- trajectory_timesteps (optional): Corresponding timesteps
|
||||
fastvideo_args: Configuration containing:
|
||||
- output_type: "latent" to skip decoding, otherwise decode to pixels
|
||||
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
|
||||
- model_loaded: Track VAE loading state
|
||||
- model_paths: Path to VAE model if loading needed
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
Modified batch with:
|
||||
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
|
||||
- trajectory_decoded (if requested): List of decoded frames per timestep
|
||||
"""
|
||||
# load vae if not already loaded (used for memory constrained devices)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
loader = VAELoader()
|
||||
@@ -75,58 +147,29 @@ class DecodingStage(PipelineStage):
|
||||
pipeline.add_module("vae", self.vae)
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
latents = batch.latents
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if latents is None:
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
frames = batch.latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
frames = self.decode(batch.latents, fastvideo_args)
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
batch.trajectory_decoded = []
|
||||
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
|
||||
for idx in range(batch.trajectory_latents.shape[1]):
|
||||
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
|
||||
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
|
||||
cur_timestep = batch.trajectory_timesteps[idx]
|
||||
logger.info("decoding trajectory latent for timestep: %s",
|
||||
cur_timestep)
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
image = image.cpu().float()
|
||||
frames = frames.cpu().float()
|
||||
|
||||
# Update batch with decoded image
|
||||
batch.output = image
|
||||
batch.output = frames
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
|
||||
@@ -85,8 +85,9 @@ class DenoisingStage(PipelineStage):
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
|
||||
) # hack
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
|
||||
)
|
||||
|
||||
def forward(
|
||||
@@ -204,8 +205,14 @@ class DenoisingStage(PipelineStage):
|
||||
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||
if fastvideo_args.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if batch.boundary_ratio is not None:
|
||||
logger.info("Overriding boundary ratio from %s to %s",
|
||||
boundary_ratio, batch.boundary_ratio)
|
||||
boundary_ratio = batch.boundary_ratio
|
||||
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
@@ -247,6 +254,10 @@ class DenoisingStage(PipelineStage):
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -427,6 +438,11 @@ class DenoisingStage(PipelineStage):
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# save trajectory latents if needed
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(t)
|
||||
trajectory_latents.append(latents)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
@@ -434,9 +450,28 @@ class DenoisingStage(PipelineStage):
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
trajectory_tensor: torch.Tensor | None = None
|
||||
if trajectory_latents:
|
||||
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
|
||||
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
|
||||
dim=0)
|
||||
else:
|
||||
trajectory_tensor = None
|
||||
trajectory_timesteps_tensor = None
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
if sp_group:
|
||||
latents = sequence_model_parallel_all_gather(latents, dim=2)
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_tensor = trajectory_tensor.to(
|
||||
get_local_torch_device())
|
||||
trajectory_tensor = sequence_model_parallel_all_gather(
|
||||
trajectory_tensor, dim=3)
|
||||
|
||||
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
|
||||
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
@@ -115,6 +115,7 @@ class CudaPlatformBase(Platform):
|
||||
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
logger.info("Selected backend: %s", selected_backend)
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
@@ -144,6 +145,20 @@ class CudaPlatformBase(Platform):
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
|
||||
try:
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend)
|
||||
logger.info("Using Sage Attention 3 backend.")
|
||||
|
||||
return "fastvideo.attention.backends.sage_attn3.SageAttention3Backend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention 3 backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
|
||||
@@ -18,6 +18,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
@@ -0,0 +1,459 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utilities for managing the PyTorch profiler within FastVideo.
|
||||
|
||||
The profiler is shared across the process; this module adds a light-weight
|
||||
controller that gates collection based on named *regions*. Regions may be
|
||||
enabled through dedicated environment variables (e.g.
|
||||
``FASTVIDEO_TORCH_PROFILE_MODEL_LOADING=1``) or via the consolidated
|
||||
``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list (e.g.
|
||||
``FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit``).
|
||||
|
||||
Typical usage from client code::
|
||||
|
||||
controller = TorchProfilerController(profiler, activities)
|
||||
with controller.region("training_dit"):
|
||||
run_training_step()
|
||||
|
||||
To introduce a new region, register it via :func:`register_profiler_region`
|
||||
and wrap the corresponding code in :meth:`TorchProfilerController.region`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
import functools
|
||||
from collections.abc import Iterable
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_GLOBAL_PROFILER: torch.profiler.profile | None = None
|
||||
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProfilerRegion:
|
||||
"""Metadata describing a profiler region."""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
default_enabled: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.name or self.name.strip() != self.name:
|
||||
raise ValueError(
|
||||
f"Profiler region name must be non-empty without surrounding whitespace: {self.name!r}"
|
||||
)
|
||||
if not self.name.islower():
|
||||
raise ValueError(
|
||||
f"Profiler region name must be lower-case: {self.name!r}")
|
||||
|
||||
|
||||
_REGISTERED_REGIONS: dict[str, ProfilerRegion] = {}
|
||||
|
||||
|
||||
def _normalize_token(token: str) -> str:
|
||||
return token.strip().lower()
|
||||
|
||||
|
||||
def register_profiler_region(
|
||||
name: str,
|
||||
description: str,
|
||||
*,
|
||||
default_enabled: bool = False,
|
||||
) -> None:
|
||||
"""Register a profiler region so configuration can validate inputs."""
|
||||
|
||||
canonical = _normalize_token(name)
|
||||
if canonical in _REGISTERED_REGIONS:
|
||||
raise ValueError(f"Profiler region {name!r} is already registered")
|
||||
|
||||
region = ProfilerRegion(
|
||||
name=canonical,
|
||||
description=description,
|
||||
default_enabled=bool(default_enabled),
|
||||
)
|
||||
_REGISTERED_REGIONS[canonical] = region
|
||||
|
||||
|
||||
def resolve_profiler_region(name: str) -> ProfilerRegion | None:
|
||||
"""Return the registered region matching ``name`` or ``None`` if absent."""
|
||||
|
||||
canonical = _normalize_token(name)
|
||||
return _REGISTERED_REGIONS.get(canonical)
|
||||
|
||||
|
||||
def list_profiler_regions() -> list[ProfilerRegion]:
|
||||
"""Return all registered profiler regions sorted by canonical name."""
|
||||
|
||||
return [_REGISTERED_REGIONS[name] for name in sorted(_REGISTERED_REGIONS)]
|
||||
|
||||
|
||||
_DEFAULT_ACTIVITIES: tuple[torch.profiler.ProfilerActivity, ...] = (
|
||||
torch.profiler.ProfilerActivity.CPU,
|
||||
torch.profiler.ProfilerActivity.CUDA,
|
||||
)
|
||||
|
||||
|
||||
def get_global_profiler() -> torch.profiler.profile | None:
|
||||
"""Return the global profiler instance if one was created."""
|
||||
|
||||
return _GLOBAL_PROFILER
|
||||
|
||||
|
||||
def set_global_profiler(profiler: torch.profiler.profile | None) -> None:
|
||||
global _GLOBAL_PROFILER
|
||||
_GLOBAL_PROFILER = profiler
|
||||
|
||||
|
||||
def get_global_controller() -> TorchProfilerController | None:
|
||||
return _GLOBAL_CONTROLLER
|
||||
|
||||
|
||||
def set_global_controller(controller: TorchProfilerController | None) -> None:
|
||||
global _GLOBAL_CONTROLLER
|
||||
_GLOBAL_CONTROLLER = controller
|
||||
|
||||
|
||||
register_profiler_region(
|
||||
name="profiler_region_model_loading",
|
||||
description="Module/model loading during pipeline initialization.",
|
||||
default_enabled=False,
|
||||
)
|
||||
# register_profiler_region(
|
||||
# name="profiler_region_inference_pre_denoising",
|
||||
# description="Pre-denoising inference steps (conditioning, preprocessing).",
|
||||
# )
|
||||
# register_profiler_region(
|
||||
# name="profiler_region_inference_denoising",
|
||||
# description="The main inference denoising loop.",
|
||||
# )
|
||||
# register_profiler_region(
|
||||
# name="profiler_region_inference_post_denoising",
|
||||
# description=
|
||||
# "Post-processing after denoising (decoder, conditioning restores).",
|
||||
# )
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_save_checkpoint",
|
||||
description="Training save checkpoint operations.",
|
||||
)
|
||||
|
||||
# general training related regions
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_validation",
|
||||
description="Validation loop during training.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_train_one_step",
|
||||
description="High-level step orchestration in the training loop.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_train",
|
||||
description="Single optimizer step including forward/backward passes.",
|
||||
)
|
||||
|
||||
# distillation specific regions
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_teacher_forward",
|
||||
description="Teacher model forward pass in distillation pipelines.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_student_forward",
|
||||
description="Student model forward pass in distillation pipelines.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_loss",
|
||||
description="Distillation loss computation and aggregation.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_update",
|
||||
description="Parameter updates specific to distillation workflows.",
|
||||
)
|
||||
|
||||
|
||||
def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
|
||||
"""Create or reuse the process-wide torch profiler controller."""
|
||||
|
||||
existing = get_global_controller()
|
||||
if existing is not None:
|
||||
if trace_dir:
|
||||
logger.info("Reusing existing global torch profiler controller")
|
||||
return existing
|
||||
|
||||
if not trace_dir:
|
||||
logger.info("Torch profiler disabled; returning no-op controller")
|
||||
return TorchProfilerController(None, _DEFAULT_ACTIVITIES, disabled=True)
|
||||
|
||||
logger.info("Profiling enabled. Traces will be saved to: %s", trace_dir)
|
||||
logger.info(
|
||||
"Profiler config: record_shapes=%s, profile_memory=%s, with_stack=%s, with_flops=%s",
|
||||
envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
|
||||
envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
|
||||
envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
|
||||
envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
|
||||
)
|
||||
logger.info("FASTVIDEO_TORCH_PROFILE_REGIONS=%s",
|
||||
envs.FASTVIDEO_TORCH_PROFILE_REGIONS)
|
||||
|
||||
def trace_handler(prof):
|
||||
# print(prof.key_averages().table(
|
||||
# sort_by="self_cuda_time_total", row_limit=-1))
|
||||
logger.info("Profiling trace saved to: %s", "test_trace_" + str(prof.step_num) + ".json")
|
||||
prof.export_chrome_trace("test2_trace_" + str(prof.step_num) + ".json")
|
||||
|
||||
profiler = torch.profiler.profile(
|
||||
activities=_DEFAULT_ACTIVITIES,
|
||||
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
|
||||
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
|
||||
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
|
||||
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
|
||||
schedule=torch.profiler.schedule(wait=2, warmup=1, active=2),
|
||||
# on_trace_ready=trace_handler,
|
||||
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir,
|
||||
use_gzip=True),
|
||||
)
|
||||
controller = TorchProfilerController(profiler, _DEFAULT_ACTIVITIES)
|
||||
controller.start()
|
||||
logger.info("Torch profiler started")
|
||||
return controller
|
||||
|
||||
|
||||
@dataclass
|
||||
class TorchProfilerConfig:
|
||||
"""Configuration for torch profiler region control.
|
||||
|
||||
Use :meth:`from_env` to construct an instance with defaults inherited from
|
||||
registered regions and optional overrides from the
|
||||
``FASTVIDEO_TORCH_PROFILE_REGIONS`` environment variable. The resulting
|
||||
``regions`` map is consumed by :class:`TorchProfilerController` to decide
|
||||
when collection should be enabled.
|
||||
"""
|
||||
|
||||
regions: dict[str, bool]
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> TorchProfilerConfig:
|
||||
"""Build a configuration from process environment variables."""
|
||||
|
||||
requested_regions = {
|
||||
token.strip()
|
||||
for token in (getattr(envs, "FASTVIDEO_TORCH_PROFILE_REGIONS", "")
|
||||
or "").split(",") if token.strip()
|
||||
}
|
||||
|
||||
if not requested_regions:
|
||||
available = ", ".join(region.name
|
||||
for region in list_profiler_regions())
|
||||
raise ValueError(
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS must list at least one region; "
|
||||
f"available regions: {available}")
|
||||
|
||||
regions: dict[str, bool] = {}
|
||||
available_regions = list_profiler_regions()
|
||||
available_names = ", ".join(region.name for region in available_regions)
|
||||
|
||||
for token in requested_regions:
|
||||
resolved = resolve_profiler_region(token)
|
||||
if resolved is None:
|
||||
logger.warning(
|
||||
"Unknown profiler region '%s'; available regions: %s",
|
||||
token, available_names)
|
||||
continue
|
||||
regions[resolved.name] = True
|
||||
|
||||
if not regions:
|
||||
raise ValueError(
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS did not match any known regions; "
|
||||
f"requested={sorted(requested_regions)}, available={available_names}"
|
||||
)
|
||||
|
||||
return cls(regions=regions)
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"TorchProfilerConfig(regions={self.regions})"
|
||||
|
||||
|
||||
class TorchProfilerController:
|
||||
"""Helper that toggles torch profiler collection for named regions.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
profiler:
|
||||
The shared :class:`torch.profiler.profile` instance, or ``None`` if
|
||||
profiling is disabled.
|
||||
activities:
|
||||
Iterable of :class:`torch.profiler.ProfilerActivity` recorded by the
|
||||
profiler.
|
||||
config:
|
||||
Optional :class:`TorchProfilerConfig`. If omitted, :meth:`from_env`
|
||||
constructs one during initialization.
|
||||
|
||||
Examples
|
||||
--------
|
||||
Enabling an existing region from the command line::
|
||||
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit \
|
||||
python fastvideo/training/wan_training_pipeline.py ...
|
||||
|
||||
Wrapping a code block in a custom region::
|
||||
|
||||
controller = TorchProfilerController(profiler, activities)
|
||||
with controller.region("training_validation"):
|
||||
run_validation_epoch()
|
||||
|
||||
Adding a new region requires three steps:
|
||||
1. Define an env var in ``envs.py``.
|
||||
2. Add a default entry to ``register_profiler_region`` in this module.
|
||||
3. Wrap the target code in :meth:`region` using the new name.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
profiler: Any,
|
||||
activities: Iterable[torch.profiler.ProfilerActivity],
|
||||
config: TorchProfilerConfig | None = None,
|
||||
disabled: bool = False,
|
||||
) -> None:
|
||||
activities_tuple = tuple(activities)
|
||||
existing = get_global_controller()
|
||||
if existing is not None and not disabled:
|
||||
raise RuntimeError(
|
||||
"TorchProfilerController already initialized globally. Use get_global_controller()."
|
||||
)
|
||||
if disabled:
|
||||
self._profiler = None
|
||||
return
|
||||
|
||||
self._profiler = profiler
|
||||
self._activities = activities_tuple
|
||||
self._config = config or TorchProfilerConfig.from_env()
|
||||
self._collection_enabled = False
|
||||
self._active_region_depth = 0
|
||||
logger.info(
|
||||
"PROFILER: TorchProfilerController initialized with config: %s",
|
||||
self._config)
|
||||
set_global_profiler(self._profiler)
|
||||
set_global_controller(self)
|
||||
|
||||
@property
|
||||
def is_enabled(self) -> bool:
|
||||
"""Return ``True`` when the underlying profiler is collecting."""
|
||||
|
||||
if self._profiler is None:
|
||||
return False
|
||||
return self._collection_enabled
|
||||
|
||||
def is_region_enabled(self, region: str) -> bool:
|
||||
"""Return ``True`` if ``region`` should be collected."""
|
||||
|
||||
if self._profiler is None:
|
||||
return False
|
||||
return self._config.regions.get(region, False)
|
||||
|
||||
def _set_collection(self, enabled: bool) -> None:
|
||||
if self._profiler is None:
|
||||
return
|
||||
if self._collection_enabled == enabled:
|
||||
return
|
||||
event = ("fastvideo.profiler.enable_collection"
|
||||
if enabled else "fastvideo.profiler.disable_collection")
|
||||
with torch.profiler.record_function(event):
|
||||
self._profiler.toggle_collection_dynamic(enabled, self._activities)
|
||||
self._collection_enabled = enabled
|
||||
|
||||
@contextlib.contextmanager
|
||||
def region(self, region: str):
|
||||
"""Context manager that enables profiling for ``region`` if configured."""
|
||||
|
||||
if self._profiler is None:
|
||||
yield
|
||||
return
|
||||
|
||||
if not self.is_region_enabled(region):
|
||||
yield
|
||||
return
|
||||
|
||||
with torch.profiler.record_function(f"fastvideo.region::{region}"):
|
||||
self._active_region_depth += 1
|
||||
if self._active_region_depth == 1:
|
||||
logger.info(
|
||||
"PROFILER: Setting collection to True (depth=%s) for region %s",
|
||||
self._active_region_depth, region)
|
||||
self._set_collection(True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._active_region_depth -= 1
|
||||
logger.info("PROFILER: Decreasing active region depth to %s",
|
||||
self._active_region_depth)
|
||||
if self._active_region_depth == 0:
|
||||
logger.info(
|
||||
"PROFILER: Setting collection to False upon exiting region %s",
|
||||
region)
|
||||
self._set_collection(False)
|
||||
|
||||
def start(self) -> None:
|
||||
"""Start the profiler and pause collection until a region is entered."""
|
||||
|
||||
logger.info("PROFILER: Starting profiler...")
|
||||
if self._profiler is None:
|
||||
return
|
||||
self._profiler.start()
|
||||
logger.info("PROFILER: Profiler started")
|
||||
# Profiler starts with collection disabled by default.
|
||||
logger.info("PROFILER: Setting collection to False")
|
||||
self._set_collection(False)
|
||||
logger.info("PROFILER: Profiler started with collection disabled")
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the profiler after disabling collection and clearing state."""
|
||||
|
||||
if self._profiler is None:
|
||||
return
|
||||
|
||||
logger.info("PROFILER: Stopping profiler...")
|
||||
self._profiler.stop()
|
||||
logger.info("PROFILER: Profiler stopped")
|
||||
self._active_region_depth = 0
|
||||
set_global_profiler(None)
|
||||
set_global_controller(None)
|
||||
|
||||
@property
|
||||
def has_profiler(self) -> bool:
|
||||
"""Return ``True`` when a profiler instance is available."""
|
||||
|
||||
return self._profiler is not None
|
||||
|
||||
@property
|
||||
def activities(self) -> tuple[torch.profiler.ProfilerActivity, ...]:
|
||||
return tuple(self._activities)
|
||||
|
||||
@property
|
||||
def profiler(self) -> torch.profiler.profile | None:
|
||||
return self._profiler
|
||||
|
||||
|
||||
def profile_region(
|
||||
region: str) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
|
||||
"""Wrap a bound method so it runs inside a profiler region if available."""
|
||||
|
||||
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
|
||||
|
||||
@functools.wraps(fn)
|
||||
def wrapped(self, *args, **kwargs):
|
||||
controller = getattr(self, "profiler_controller", None)
|
||||
if controller is None or not controller.has_profiler:
|
||||
return fn(self, *args, **kwargs)
|
||||
with controller.region(region):
|
||||
return fn(self, *args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
||||
return decorator
|
||||
@@ -0,0 +1,108 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from fastvideo.dataset.dataloader.parquet_io import (
|
||||
ParquetDatasetWriter,
|
||||
records_to_table,
|
||||
)
|
||||
|
||||
|
||||
def test_records_to_table_types():
|
||||
schema = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("width", pa.int64()),
|
||||
])
|
||||
records = [{
|
||||
"id": "a",
|
||||
"vae_latent_bytes": b"\x00\x01",
|
||||
"vae_latent_shape": [1, 2, 3],
|
||||
"duration_sec": 1.5,
|
||||
"width": 640,
|
||||
}]
|
||||
|
||||
table = records_to_table(records, schema)
|
||||
assert table.schema == schema
|
||||
assert table.num_rows == 1
|
||||
cols = {name: table.column(name).to_pylist()[0] for name in schema.names}
|
||||
assert cols["id"] == "a"
|
||||
assert isinstance(cols["vae_latent_bytes"], (bytes, bytearray))
|
||||
assert cols["vae_latent_shape"] == [1, 2, 3]
|
||||
assert abs(cols["duration_sec"] - 1.5) < 1e-6
|
||||
assert cols["width"] == 640
|
||||
|
||||
|
||||
def test_writer_flush_and_remainder(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
records = [{"id": str(i)} for i in range(25)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
written = writer.flush(num_workers=1)
|
||||
assert written == 20
|
||||
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 2
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 20
|
||||
|
||||
# Append remainder to complete another chunk
|
||||
extra = records_to_table([{"id": str(i)} for i in range(5)], schema)
|
||||
writer.append_table(extra)
|
||||
written2 = writer.flush(num_workers=1)
|
||||
assert written2 == 10
|
||||
files2 = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files2) == 3
|
||||
total_rows2 = sum(pq.read_table(str(f)).num_rows for f in files2)
|
||||
assert total_rows2 == 30
|
||||
|
||||
|
||||
def test_writer_flush_write_remainder(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
# 25 rows, 10 per file => 2 full files + 1 remainder(5)
|
||||
records = [{"id": str(i)} for i in range(25)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out_last"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
# First flush writes 20
|
||||
written1 = writer.flush(num_workers=1)
|
||||
assert written1 == 20
|
||||
# Final flush with remainder
|
||||
written2 = writer.flush(num_workers=1, write_remainder=True)
|
||||
assert written2 == 5
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 3
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 25
|
||||
|
||||
|
||||
def test_writer_parallel_workers(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
# 40 rows, 10 per file => 4 files
|
||||
records = [{"id": str(i)} for i in range(40)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out_parallel"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
written = writer.flush(num_workers=2)
|
||||
assert written == 40
|
||||
|
||||
# Ensure files exist under worker subdirs
|
||||
worker_dirs = [p for p in out_dir.iterdir() if p.is_dir() and p.name.startswith("worker_")]
|
||||
assert len(worker_dirs) >= 1
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 4
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 40
|
||||
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
basic_t2v_record_creator,
|
||||
i2v_record_creator,
|
||||
ode_text_only_record_creator,
|
||||
text_only_record_creator,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
def _mk_basic_batch(N: int) -> PreprocessBatch:
|
||||
batch = PreprocessBatch(data_type="video")
|
||||
batch.video_file_name = [f"vid_{i}" for i in range(N)]
|
||||
batch.prompt = [f"caption_{i}" for i in range(N)]
|
||||
batch.width = [640 for _ in range(N)]
|
||||
batch.height = [360 for _ in range(N)]
|
||||
batch.fps = [4 for _ in range(N)]
|
||||
batch.num_frames = [2 for _ in range(N)]
|
||||
# Latents: shape (N, C, T, H, W); per-record use latents[idx]
|
||||
batch.latents = np.zeros((N, 4, 2, 8, 8), dtype=np.float32)
|
||||
# Prompt embeds: list of per-record arrays [Seq, Dim]
|
||||
batch.prompt_embeds = [np.ones((6, 16), dtype=np.float32) for _ in range(N)]
|
||||
return batch
|
||||
|
||||
|
||||
def test_basic_t2v_record_creator_fields():
|
||||
N = 2
|
||||
batch = _mk_basic_batch(N)
|
||||
|
||||
records = basic_t2v_record_creator(batch)
|
||||
assert isinstance(records, list) and len(records) == N
|
||||
|
||||
for i, rec in enumerate(records):
|
||||
assert rec["id"] == batch.video_file_name[i]
|
||||
# Latents bytes/shape/dtype
|
||||
assert isinstance(rec["vae_latent_bytes"], (bytes, bytearray))
|
||||
assert rec["vae_latent_shape"] == list(batch.latents[i].shape)
|
||||
assert rec["vae_latent_dtype"] == str(batch.latents[i].dtype)
|
||||
# Text embedding
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(batch.prompt_embeds[i].shape)
|
||||
assert rec["text_embedding_dtype"] == str(batch.prompt_embeds[i].dtype)
|
||||
# Meta
|
||||
assert rec["caption"] == batch.prompt[i]
|
||||
assert rec["media_type"] == "video"
|
||||
assert rec["width"] == int(batch.width[i])
|
||||
assert rec["height"] == int(batch.height[i])
|
||||
assert rec["num_frames"] == batch.latents[i].shape[1]
|
||||
|
||||
|
||||
def test_i2v_record_creator_additional_fields():
|
||||
N = 3
|
||||
batch = _mk_basic_batch(N)
|
||||
# image_embeds is a list of length 1, with an array of shape [N, D]
|
||||
batch.image_embeds = [np.ones((N, 32), dtype=np.float32)]
|
||||
# first frame latent per record
|
||||
batch.image_latent = np.zeros((N, 4, 1, 8, 8), dtype=np.float32)
|
||||
# pil image per record
|
||||
batch.pil_image = np.zeros((N, 8, 8, 3), dtype=np.uint8)
|
||||
|
||||
records = i2v_record_creator(batch)
|
||||
assert isinstance(records, list) and len(records) == N
|
||||
|
||||
for i, rec in enumerate(records):
|
||||
# clip feature
|
||||
assert isinstance(rec["clip_feature_bytes"], (bytes, bytearray))
|
||||
assert rec["clip_feature_shape"] == list(batch.image_embeds[0][i].shape)
|
||||
assert rec["clip_feature_dtype"] == str(batch.image_embeds[0][i].dtype)
|
||||
# first frame latent
|
||||
assert isinstance(rec["first_frame_latent_bytes"], (bytes, bytearray))
|
||||
assert rec["first_frame_latent_shape"] == list(batch.image_latent[i].shape)
|
||||
assert rec["first_frame_latent_dtype"] == str(batch.image_latent[i].dtype)
|
||||
# pil image
|
||||
assert isinstance(rec["pil_image_bytes"], (bytes, bytearray))
|
||||
assert rec["pil_image_shape"] == list(batch.pil_image[i].shape)
|
||||
assert rec["pil_image_dtype"] == str(batch.pil_image[i].dtype)
|
||||
|
||||
|
||||
def test_ode_text_only_record_creator():
|
||||
video_name = "ex"
|
||||
caption = "a prompt"
|
||||
text_embedding = np.ones((6, 16), dtype=np.float32)
|
||||
traj = np.ones((5, 4, 2, 2), dtype=np.float32)
|
||||
tsteps = np.arange(5, dtype=np.float32)
|
||||
|
||||
rec = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=caption,
|
||||
trajectory_latents=traj,
|
||||
trajectory_timesteps=tsteps,
|
||||
)
|
||||
assert rec["id"] == f"text_{video_name}"
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(text_embedding.shape)
|
||||
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
|
||||
assert rec["file_name"] == video_name
|
||||
assert rec["caption"] == caption
|
||||
assert rec["media_type"] == "text"
|
||||
# Trajectory fields
|
||||
assert isinstance(rec["trajectory_latents_bytes"], (bytes, bytearray))
|
||||
assert rec["trajectory_latents_shape"] == list(traj.shape)
|
||||
assert rec["trajectory_latents_dtype"] == str(traj.dtype)
|
||||
assert isinstance(rec["trajectory_timesteps_bytes"], (bytes, bytearray))
|
||||
assert rec["trajectory_timesteps_shape"] == list(tsteps.shape)
|
||||
assert rec["trajectory_timesteps_dtype"] == str(tsteps.dtype)
|
||||
|
||||
|
||||
def test_text_only_record_creator():
|
||||
text_name = "note1"
|
||||
caption = "a prompt"
|
||||
text_embedding = np.ones((7, 16), dtype=np.float32)
|
||||
rec = text_only_record_creator(
|
||||
text_name=text_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=caption,
|
||||
)
|
||||
assert rec["id"] == f"text_{text_name}"
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(text_embedding.shape)
|
||||
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
|
||||
assert rec["caption"] == caption
|
||||
@@ -62,7 +62,7 @@ def run_test(pytest_command: str):
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
def run_encoder_tests():
|
||||
run_test("pytest ./fastvideo/tests/encoders -vs")
|
||||
|
||||
@@ -117,3 +117,11 @@ def run_inference_lora_tests():
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900)
|
||||
def run_distill_dmd_tests():
|
||||
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_self_forcing_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
|
||||
|
||||
Binary file not shown.
@@ -0,0 +1,184 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from huggingface_hub import snapshot_download
|
||||
import subprocess
|
||||
import sys
|
||||
from fastvideo.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
|
||||
import shutil
|
||||
|
||||
# Import the training pipeline
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
|
||||
NUM_NODES = "1"
|
||||
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
|
||||
# preprocessing
|
||||
DATA_DIR = "data"
|
||||
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol"))
|
||||
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
|
||||
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/pipelines/preprocess/v1_preprocess.py"
|
||||
|
||||
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol_processed_t2v"))
|
||||
|
||||
|
||||
# training
|
||||
NUM_GPUS_PER_NODE_TRAINING = "4"
|
||||
TRAINING_ENTRY_FILE_PATH = "fastvideo/training/wan_distillation_pipeline.py"
|
||||
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
|
||||
LOCAL_VALIDATION_DATASET_FILE = "examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
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)
|
||||
|
||||
print(f"Creating data directory at {data_dir}")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
|
||||
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
|
||||
try:
|
||||
result = snapshot_download(
|
||||
repo_id="wlsaidhi/crush-smol-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():
|
||||
# remove the local_preprocessed_data_dir if it exists
|
||||
if LOCAL_PREPROCESSED_DATA_DIR.exists():
|
||||
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
|
||||
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
|
||||
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
|
||||
PREPROCESSING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--seed", "42",
|
||||
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge.txt"),
|
||||
"--preprocess_video_batch_size", "1",
|
||||
"--max_height", "480",
|
||||
"--max_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--dataloader_num_workers", "0",
|
||||
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
|
||||
"--train_fps", "16",
|
||||
"--samples_per_file", "1",
|
||||
"--flush_frequency", "1",
|
||||
"--video_length_tolerance_range", "5",
|
||||
"--preprocess_task", "t2v",
|
||||
]
|
||||
|
||||
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_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", "1",
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "10",
|
||||
"--gradient_accumulation_steps", "1",
|
||||
"--max_train_steps", "501",
|
||||
"--learning_rate", "2e-6",
|
||||
"--fake_score_learning_rate", "2e-6",
|
||||
"--mixed_precision", "bf16",
|
||||
"--training_state_checkpointing_steps", "1000",
|
||||
"--weight_only_checkpointing_steps", "1000",
|
||||
"--validation_steps", "50",
|
||||
"--validation_sampling_steps", "3",
|
||||
"--log_validation",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--ema_start_step", "0",
|
||||
"--training_cfg_rate", "0.0",
|
||||
"--output_dir", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "ci_wan_t2v_dmd_overfit",
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--flow_shift", "8",
|
||||
"--validation_guidance_scale", "6.0",
|
||||
"--weight_decay", "0.01",
|
||||
"--generator_update_interval", "5",
|
||||
"--dmd_denoising_steps", "1000,757,522",
|
||||
"--min_timestep_ratio", "0.02",
|
||||
"--max_timestep_ratio", "0.98",
|
||||
"--seed", "1000",
|
||||
"--real_score_guidance_scale", "3.5",
|
||||
"--dit_precision", "fp32",
|
||||
"--max_grad_norm", "1.0",
|
||||
"--enable_gradient_checkpointing_type", "full",
|
||||
]
|
||||
|
||||
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()
|
||||
@@ -111,7 +111,8 @@ def run_training():
|
||||
"--max_train_steps", "901",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--checkpointing_steps", "6000",
|
||||
"--weight_only_checkpointing_steps", "6000",
|
||||
"--training_state_checkpointing_steps", "6000",
|
||||
"--validation_steps", "100",
|
||||
"--validation_sampling_steps", "50",
|
||||
"--log_validation",
|
||||
|
||||
@@ -62,6 +62,11 @@ def download_data():
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
# remove the local_preprocessed_data_dir if it exists
|
||||
if LOCAL_PREPROCESSED_DATA_DIR.exists():
|
||||
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
|
||||
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
|
||||
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
@@ -111,7 +116,8 @@ def run_training():
|
||||
"--max_train_steps", "901",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--checkpointing_steps", "6000",
|
||||
"--weight_only_checkpointing_steps", "6000",
|
||||
"--training_state_checkpointing_steps", "6000",
|
||||
"--validation_steps", "100",
|
||||
"--validation_sampling_steps", "50",
|
||||
"--log_validation",
|
||||
|
||||
@@ -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}
|
||||
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":1.260593056678772,"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.2620866410434246,"_runtime":107.325113071}
|
||||
@@ -46,7 +46,8 @@ def run_worker():
|
||||
"--max_train_steps", "5",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--checkpointing_steps", "30",
|
||||
"--weight_only_checkpointing_steps", "30",
|
||||
"--training_state_checkpointing_steps", "30",
|
||||
"--validation_steps", "10",
|
||||
"--validation_sampling_steps", "50",
|
||||
"--log_validation",
|
||||
@@ -110,7 +111,7 @@ def test_distributed_training():
|
||||
'avg_step_time': 1.0,
|
||||
'grad_norm': 0.1,
|
||||
'step_time': 1.0,
|
||||
'train_loss': 0.001
|
||||
'train_loss': 0.005
|
||||
}
|
||||
|
||||
failures = []
|
||||
|
||||
@@ -1 +1 @@
|
||||
{"train_loss":0.10545400530099869,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.53125,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"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.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"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.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","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."],"_type":"videos","count":3}}
|
||||
{"train_loss":0.1021774671971798,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.11582941561937332,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"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.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"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.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","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."],"_type":"videos","count":3}}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user