Compare commits

...
34 Commits
Author SHA1 Message Date
Will Lin 141a1140f6 refactor sampling pipeline 2026-01-20 15:40:53 -08:00
Shijie Wang 6294015389 Debug transformer output misalignment 2026-01-20 14:38:47 -08:00
Shijie Wang e31b6c9e90 Fix OCR Rewards 2026-01-20 14:38:22 -08:00
Tamoghno Kandar bf0ff21eeb Fix OCR Rewards 2026-01-20 14:38:22 -08:00
Shijie Wang 6f937102ad Enable validation videos 2026-01-20 14:38:22 -08:00
Tamoghno Kandar 0164e93019 Add Validation Loop 2026-01-20 14:38:21 -08:00
Shijie Wang 67e457aa92 resolved cuda OOM error 2026-01-20 14:38:21 -08:00
Shijie Wang d795f0c443 remove additional sampling pipeline 2026-01-20 14:38:21 -08:00
loaydatrain 02452dd6e7 fixed dtype mismatch 2026-01-20 14:38:21 -08:00
Shijie Wang 91ef24bc14 update run script 2026-01-20 14:38:20 -08:00
Shijie Wang e76e9fda15 minor fix 2026-01-20 14:38:20 -08:00
Shijie Wang 3b17f5a621 fix trajectory collection & reward computation 2026-01-20 14:38:20 -08:00
Shijie Wang d758878705 minor fix 2026-01-20 14:38:19 -08:00
Shijie Wang 689e629420 Add entry point script 2026-01-20 14:38:19 -08:00
Shijie Wang 873dc9695f Complete train_one_step and grpo policy loss 2026-01-20 14:38:19 -08:00
Shijie Wang bfc0f46d61 Implement trajectories collection, reward and advantage computing 2026-01-20 14:38:18 -08:00
Shijie Wang 39907dbe4d Port per-prompt stat tracker 2026-01-20 14:38:18 -08:00
Shijie Wang abdd0c9b6a Implement SDE step & SDE pipeline with log prob 2026-01-20 14:38:18 -08:00
Shijie (Jacob) Wang f32a12200d Refactor and trim down unnecessary RL args 2026-01-20 14:38:18 -08:00
Shijie Wang f1d2c9e6b7 Add RL dataset & dataloader 2026-01-20 14:38:18 -08:00
Jiali Chen 450579cb42 init algorithm backbone and refactor rl_pipeline 2026-01-20 14:38:17 -08:00
Jiali Chen 26d7d6cc08 minor bug fix 2026-01-20 14:38:17 -08:00
Jiali Chen 44f0124eaa refactor and add ocr reward model 2026-01-20 14:38:17 -08:00
Jiali Chen d3ace51394 Phase 1 minor fixes 2026-01-20 14:38:17 -08:00
Jiali Chen 58954c660b implement Phase 1 backbone code 2026-01-20 14:38:16 -08:00
alexzmsandWilliam Lin 31f44110b5 [kernel] [bugfix] [ci] bump v0.2.4. Fix STA output handling, TurboDiffusion CUDA norm dtypes for fastvideo-kernel unit tests. (#1020)
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2026-01-19 17:42:42 -08:00
William Lin 21f3ce6577 [kernel] Fix fastvideo-kernel release workflow (#1019) 2026-01-17 15:57:46 -08:00
XOR-op 785d123e36 [feat] Hooks API and layerwise offloading for all DiTs (#1006) 2026-01-17 11:22:02 -08:00
William Lin d58c551c11 [chore] release fastvideo-kernel 0.2.3 (#1018) 2026-01-17 02:24:23 -08:00
alexzms 560628709c [Bug Fix] Add autograd wrapper for block-sparse attention in fastvideo-kernel + fix CMake extension linking (#1015) 2026-01-16 21:16:43 -08:00
William Lin 0f53b51e6c [CI] Fix OOM issues in ssim tests (#1011) 2026-01-16 21:15:20 -08:00
alexzmsandWill Lin 06093a9c4e [CI] SSIM tests optimization: load all model weights from Modal persistent Volume (#958)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-16 11:50:43 -08:00
KyleShao dbddfab6d2 [feat] Introduce Cosmos 2.5 Text2World pipeline (#974) 2026-01-15 15:09:05 -08:00
William Lin 7188170277 [misc] [bugfix] unpin 'av' in pyproject (#1009) 2026-01-13 15:40:46 -08:00
96 changed files with 11160 additions and 635 deletions
+1 -1
View File
@@ -61,7 +61,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- TEST_TYPE=ssim
+16 -1
View File
@@ -156,8 +156,23 @@ jobs:
# Fix the wheel to be manylinux compliant
pip install auditwheel
# Point auditwheel at torch libs, but do not vendor them into the wheel.
TORCH_LIB_DIR=$(python - <<'PY'
import os
import torch
print(os.path.join(os.path.dirname(torch.__file__), "lib"))
PY
)
export LD_LIBRARY_PATH="${TORCH_LIB_DIR}:${LD_LIBRARY_PATH}"
# Target manylinux_2_35 (Ubuntu 22.04 native)
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist \
--exclude libtorch_cuda.so \
--exclude libtorch_cpu.so \
--exclude libtorch.so \
--exclude libc10.so \
--exclude libc10_cuda.so \
--exclude libtorch_python.so
# Move fixed wheels back to dist for upload consistency
rm dist/*.whl
mv fixed_dist/*.whl dist/
+1 -1
View File
@@ -68,7 +68,7 @@ repos:
entry: bash
args:
- -c
- 'git ls-files | grep -v "^fastvideo/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
- 'git ls-files | grep -v "^\"*fastvideo/tests/ssim/" | grep -v "^\"*fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
language: system
always_run: true
pass_filenames: false
+3 -1
View File
@@ -13,11 +13,13 @@ from fastvideo_kernel import video_sparse_attn
# q, k, v: [batch_size, num_heads, seq_len, head_dim]
# variable_block_sizes: Number of valid tokens per block
# q_variable_block_sizes: Number of valid tokens per q block (can differ from KV for q/k of different lengths)
# topk: Number of blocks to attend
output = video_sparse_attn(
q, k, v,
variable_block_sizes=block_sizes,
block_sizes,
block_sizes,
topk=32
)
```
@@ -0,0 +1,42 @@
from fastvideo import VideoGenerator
def main():
# Point this to your local diffusers model dir (or replace with a HF model ID).
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
generator = VideoGenerator.from_pretrained(
model_path,
num_gpus=1,
use_fsdp_inference=False, # set True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
prompt = (
"A high-definition video captures the precision of robotic welding in an industrial setting. The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. The welding process is in full swing, with bright sparks and intense light illuminating the scene, creating a vivid display of blue and white hues. A significant amount of smoke billows around the welding area, partially obscuring the view but emphasizing the heat and activity. The background reveals parts of the workshop environment, including a ventilation system and various pieces of machinery, indicating a busy and functional industrial workspace. As the video progresses, the robotic arm maintains its steady position, continuing the welding process and moving to its left. The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. The metal surface beneath the torch shows ongoing signs of heating and melting. The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, underscoring the ongoing nature of the welding operation."
)
video = generator.generate_video(
prompt,
negative_prompt="",
height=704,
width=1280,
num_frames=77,
num_inference_steps=35,
guidance_scale=7.0,
fps=24,
output_path="outputs_video/cosmos2_5_t2w.mp4",
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
+129
View File
@@ -0,0 +1,129 @@
#!/bin/bash
# Change to FastVideo root directory (3 levels up from this script)
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
FASTVIDEO_ROOT="$(cd "$SCRIPT_DIR/../../.." && pwd)"
cd "$FASTVIDEO_ROOT"
# Add FastVideo root to PYTHONPATH so Python can find the fastvideo package
export PYTHONPATH="$FASTVIDEO_ROOT${PYTHONPATH:+:$PYTHONPATH}"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
RL_DATASET_DIR="data/ocr/" # Path to RL prompt dataset directory (should contain train.txt and test.txt)
VALIDATION_DATASET_FILE="$SCRIPT_DIR/validation.json"
NUM_GPUS=1
# use GPU 3
export CUDA_VISIBLE_DEVICES=3
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_grpo"
--output_dir "checkpoints/wan_t2v_grpo"
--max_train_steps 5000
--train_batch_size 4
# --train_sp_batch_size 4
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 5
--num_height 240
--num_width 416
--num_frames 33
--lora_rank 32
--lora_training True
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
# --use-fsdp-inference False
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments (for RL prompt dataset)
dataset_args=(
--data_path $RL_DATASET_DIR # Used as fallback if rl_dataset_path not set
--rl_dataset_path $RL_DATASET_DIR # RL prompt dataset directory
--rl_dataset_type "text" # "text" or "geneval"
--rl_num_image_per_prompt 4 # k parameter (number of samples per prompt)
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation True
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 5
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 10
--training_state_checkpointing_steps 10
--weight_decay 1e-4
--max_grad_norm 1.0
)
# RL-specific arguments
rl_args=(
--inference_mode False
--rl_mode True
--rl_algorithm "grpo"
--rl_kl_beta 0.004 # KL regularization coefficient
--rl_policy_clip_range 0.2 # Policy clipping range for GRPO
--rl_kl_reward 0.0 # KL reward coefficient (typically 0)
--rl_global_std False # Use per-prompt std (recommended for GRPO)
--rl_per_prompt_stat_tracking True # Enable per-prompt stat tracking
--rl_warmup_steps 0 # Number of warmup steps (SFT before RL)
--reward-models "{\"paddle_ocr\": 1.0}" # use video_ocr reward function
)
# CFG arguments
cfg_args=(
--guidance_scale 1.0 # use guidance_scale > 1.0 to enable CFG
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0 # No CFG during training (CFG used in sampling)
--dit_precision "fp32"
# --dit_precision "bf16"
--num_euler_timesteps 50
--ema_start_step 0
# --resume_from_checkpoint "checkpoints/wan_t2v_grpo/checkpoint-XXX"
--enable-gradient-checkpointing-type "full"
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port 29501 \
"$FASTVIDEO_ROOT/fastvideo/training/wan_rl_training_pipeline.py" \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${rl_args[@]}" \
"${miscellaneous_args[@]}"
+26
View File
@@ -10,6 +10,8 @@ if(GPU_BACKEND STREQUAL "ROCM")
enable_language(HIP)
else()
enable_language(CUDA)
# Ensure CUDA toolkit targets (CUDA::cudart, CUDA::cuda_driver, etc.) are available.
find_package(CUDAToolkit REQUIRED)
endif()
# Import common utils if needed, but we keep it simple for now
@@ -153,6 +155,30 @@ if(BUILD_CXX_KERNELS)
$<$<COMPILE_LANGUAGE:CUDA>:${CUDA_FLAGS}>
)
# Link against Torch libraries to avoid undefined symbols at import time
# (e.g., torch::autograd vtables) when loading the extension module.
target_link_libraries(fastvideo_kernel_ops PRIVATE ${TORCH_LIBRARIES})
# Also link against libtorch_python to satisfy Python-binding symbols
# (e.g., torch::PyWarningHandler) required by torch/extension.h.
execute_process(
COMMAND "${Python_EXECUTABLE}" -c "import torch; from pathlib import Path; p=Path(torch.__file__).parent/'lib'; m=sorted(p.glob('libtorch_python*')); print(str(m[0]) if m else '')"
OUTPUT_VARIABLE TORCH_PYTHON_LIBRARY_PATH
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(TORCH_PYTHON_LIBRARY_PATH)
message(STATUS "TORCH_PYTHON_LIBRARY_PATH: ${TORCH_PYTHON_LIBRARY_PATH}")
target_link_libraries(fastvideo_kernel_ops PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
else()
message(WARNING "Could not locate libtorch_python; fastvideo_kernel_ops may fail to import.")
endif()
# Link CUDA runtime + driver explicitly (fixes missing symbols like cuGetErrorString at import time)
if(NOT GPU_BACKEND STREQUAL "ROCM")
target_link_libraries(fastvideo_kernel_ops PRIVATE CUDA::cudart CUDA::cuda_driver)
endif()
# We install it to fastvideo_kernel/_C so we can load it to register the ops
install(TARGETS fastvideo_kernel_ops LIBRARY DESTINATION fastvideo_kernel/_C)
endif()
+1 -1
View File
@@ -34,7 +34,7 @@ from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_att
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
# Example: Video Sparse Attention (with Triton fallback)
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
out = video_sparse_attn(q, k, v, block_sizes, block_sizes, topk=5)
# Example: VMoBA
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
@@ -639,7 +639,8 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
// store kq and vq
// ensuring all writes are finished
// ! the following two line seems unnecessary.
// tma::store_async_wait(); // ensure qg is finished
__syncthreads();
warpgroup::store(kg_smem[0], kg_reg);
@@ -660,145 +661,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
tma::store_async_wait();
}
template<int D>
void block_sparse_attention_forward_impl(
bf16* d_q, bf16* d_k, bf16* d_v, float* d_l, bf16* d_o,
int batch, int qo_heads, int kv_heads, int seq_len, int hr,
int max_kv_blocks_per_q,
int32_t* q2k_block_sparse_index_ptr,
int32_t* q2k_block_sparse_num_ptr,
int32_t* block_size_ptr,
cudaStream_t stream
) {
using K = fwd_attend_ker_tile_dims<D>;
using q_tile = st_bf<K::qo_height, K::tile_width>;
using k_tile = st_bf<K::kv_height, K::tile_width>;
using v_tile = st_bf<K::kv_height, K::tile_width>;
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
using o_tile = st_bf<K::qo_height, K::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<D>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
globals g{
qg_arg, kg_arg, vg_arg, lg_arg, og_arg,
static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q),
q2k_block_sparse_index_ptr, q2k_block_sparse_num_ptr, block_size_ptr
};
// Shared memory size for the kernel
// 54000 bytes is calibrated for H100 shared memory constraints for these tile sizes
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<D>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<D><<<grid, (128), mem_size, stream>>>(g);
}
template<int D>
void block_sparse_attention_backward_impl(
bf16* d_q, bf16* d_k, bf16* d_v, bf16* d_o, bf16* d_og, float* d_l, float* d_d, float* d_qg, float* d_kg, float* d_vg,
int batch, int qo_heads, int kv_heads, int seq_len, int hr, int max_q_blocks_per_kv,
int32_t* k2q_block_sparse_index_ptr,
int32_t* k2q_block_sparse_num_ptr,
int32_t* block_size_ptr,
cudaStream_t stream
) {
using G = bwd_attend_ker_tile_dims<D>;
using og_tile = st_bf<4*16, D>;
using o_tile = st_bf<4*16, D>;
using d_tile = col_vec<st_fl<4*16, D>>;
using og_global = gl<bf16, -1, -1, -1, -1, og_tile>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using d_global = gl<float, -1, -1, -1, -1, d_tile>;
using prep_globals = bwd_prep_globals<D>;
constexpr int mem_size_prep = kittens::MAX_SHARED_MEMORY;
int threads_prep = PREP_NUM_WARPS * kittens::WARP_THREADS;
dim3 grid_bwd_prep(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
cudaFuncSetAttribute(
bwd_attend_prep_ker<D>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size_prep
);
bwd_attend_prep_ker<D><<<grid_bwd_prep, threads_prep, mem_size_prep, stream>>>(bwd_g);
using bwd_q_tile = st_bf<G::tile_h_qo, G::tile_width>;
using bwd_k_tile = st_bf<G::tile_h, G::tile_width>;
using bwd_v_tile = st_bf<G::tile_h, G::tile_width>;
using bwd_og_tile = st_bf<G::tile_h_qo, G::tile_width>;
using bwd_qg_tile = st_fl<G::tile_h_qo, G::tile_width>;
using bwd_kg_tile = st_fl<G::tile_h, G::tile_width>;
using bwd_vg_tile = st_fl<G::tile_h, G::tile_width>;
using bwd_l_tile = row_vec<st_fl<G::tile_h_qo, G::tile_h>>;
using bwd_d_tile = row_vec<st_fl<G::tile_h_qo, G::tile_h>>;
using bwd_q_global = gl<bf16, -1, -1, -1, -1, bwd_q_tile>;
using bwd_k_global = gl<bf16, -1, -1, -1, -1, bwd_k_tile>;
using bwd_v_global = gl<bf16, -1, -1, -1, -1, bwd_v_tile>;
using bwd_og_global = gl<bf16, -1, -1, -1, -1, bwd_og_tile>;
using bwd_qg_global = gl<float, -1, -1, -1, -1, bwd_qg_tile>;
using bwd_kg_global = gl<float, -1, -1, -1, -1, bwd_kg_tile>;
using bwd_vg_global = gl<float, -1, -1, -1, -1, bwd_vg_tile>;
using bwd_l_global = gl<float, -1, -1, -1, -1, bwd_l_tile>;
using bwd_d_global = gl<float, -1, -1, -1, -1, bwd_d_tile>;
using bwd_global_args = bwd_globals<D>;
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_global_args bwd_global{bwd_q_arg, bwd_k_arg, bwd_v_arg, bwd_og_arg, bwd_qg_arg, bwd_kg_arg, bwd_vg_arg, bwd_l_arg, bwd_d_arg,
static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_q_blocks_per_kv),
k2q_block_sparse_index_ptr, k2q_block_sparse_num_ptr, block_size_ptr};
dim3 grid_bwd_main(seq_len/64, qo_heads, batch);
int threads_main = 128;
// Calibrated shared memory sizes for different head dimensions
int bwd_mem_size = (D == 64) ? 72000 : 113000;
cudaFuncSetAttribute(
bwd_attend_ker<D>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
bwd_mem_size
);
bwd_attend_ker<D><<<grid_bwd_main, threads_main, bwd_mem_size, stream>>>(bwd_global);
}
#include "pyutils/torch_helpers.cuh"
#include <ATen/cuda/CUDAContext.h>
#include <iostream>
@@ -810,23 +672,32 @@ block_sparse_attention_forward(
torch::Tensor v,
torch::Tensor q2k_block_sparse_index,
torch::Tensor q2k_block_sparse_num,
torch::Tensor block_size
torch::Tensor kv_block_size
)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
CHECK_INPUT(v);
// q shape: (batch, qo_heads, q_seq_len, head_dim)
// k shape: (batch, kv_heads, kv_seq_len, head_dim)
// v shape: (batch, kv_heads, kv_seq_len, head_dim)
// q2k_block_sparse_index shape: (batch, qo_heads, num_q_blocks, max_kv_blocks_per_q)
// q2k_block_sparse_num shape: (batch, qo_heads, num_q_blocks)
// kv_block_size shape: (num_kv_blocks) This does not need other dimensions because across all batch/heads the padding is the same.
auto batch = q.size(0);
auto seq_len = q.size(2);
auto q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto head_dim = q.size(3);
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
auto num_q_blocks = block_size.size(0);
auto num_q_blocks = q2k_block_sparse_index.size(2);
auto num_kv_blocks = kv_block_size.size(0);
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
TORCH_CHECK(num_q_blocks * BLOCK_M == q_seq_len, "This kernel supports variable q block size, but it assumes the input sequence is properly padded.");
TORCH_CHECK(num_kv_blocks * BLOCK_M == kv_seq_len, "This kernel supports variable kv block size, but it assumes the input sequence is properly padded.");
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
@@ -834,11 +705,9 @@ block_sparse_attention_forward(
TORCH_CHECK(q2k_block_sparse_index.size(0) == batch, "q2k_block_sparse_index batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q2k_block_sparse_num.size(0) == batch, "q2k_block_sparse_num batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(q2k_block_sparse_index.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_index idx 2 - must match seq_len / BLOCK_M");
TORCH_CHECK(q2k_block_sparse_num.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_num idx 2 - must match seq_len / BLOCK_M");
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K inputs");
TORCH_CHECK(q2k_block_sparse_num.size(2) == num_q_blocks, "q2k_block_sparse_num idx 2 - must match num_q_blocks");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
@@ -864,12 +733,12 @@ block_sparse_attention_forward(
// for the returned outputs
torch::Tensor o = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(head_dim)}, v.options());
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
@@ -880,32 +749,110 @@ block_sparse_attention_forward(
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
// Temporated implementation to avoid code duplication between head_dim=64 and 128
if (head_dim == 64) {
block_sparse_attention_forward_impl<64>(
d_q, d_k, d_v, d_l, d_o,
batch, qo_heads, kv_heads, seq_len, hr,
max_kv_blocks_per_q,
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
using q_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<64>::kv_height, fwd_attend_ker_tile_dims<64>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<64>::kv_height, fwd_attend_ker_tile_dims<64>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<64>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
globals g{
qg_arg,
kg_arg,
vg_arg,
lg_arg,
og_arg,
static_cast<int>(q_seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr()),
stream
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
};
constexpr int mem_size = 54000;
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<64>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
} else if (head_dim == 128) {
block_sparse_attention_forward_impl<128>(
d_q, d_k, d_v, d_l, d_o,
batch, qo_heads, kv_heads, seq_len, hr,
max_kv_blocks_per_q,
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
fwd_attend_ker<64><<<grid, (128), mem_size, stream>>>(g);
CHECK_CUDA_ERROR(cudaGetLastError());
// cudaStreamSynchronize(stream);
}
if (head_dim == 128) {
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<128>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
globals g{
qg_arg,
kg_arg,
vg_arg,
lg_arg,
og_arg,
static_cast<int>(q_seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr()),
stream
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
};
constexpr int mem_size = 54000;
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<128>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
} else {
TORCH_CHECK(false, "Unsupported head_dim: ", head_dim, ". Only 64 and 128 are supported.");
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
CHECK_CUDA_ERROR(cudaGetLastError());
// cudaStreamSynchronize(stream);
}
return {o, l_vec};
@@ -921,7 +868,7 @@ block_sparse_attention_backward(torch::Tensor q,
torch::Tensor og,
torch::Tensor k2q_block_sparse_index,
torch::Tensor k2q_block_sparse_num,
torch::Tensor block_size)
torch::Tensor kv_block_size)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -930,11 +877,23 @@ block_sparse_attention_backward(torch::Tensor q,
CHECK_INPUT(o);
CHECK_INPUT(og);
// q: [batch, qo_heads, q_seq_len, head_dim]
// k: [batch, kv_heads, kv_seq_len, head_dim]
// v: [batch, kv_heads, kv_seq_len, head_dim]
// o: [batch, qo_heads, q_seq_len, head_dim]
// l_vec: [batch, qo_heads, q_seq_len, 1]
// og: [batch, qo_heads, q_seq_len, head_dim]
// k2q_block_sparse_index: [batch, kv_heads, num_kv_blocks, max_num_q_blocks]
// k2q_block_sparse_num: [batch, kv_heads, num_kv_blocks]
// kv_block_size: [num_kv_blocks]
auto batch = q.size(0);
auto seq_len = q.size(2);
auto q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto head_dim = q.size(3);
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
auto num_kv_blocks = kv_block_size.size(0);
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index.size(2) must match num_kv_blocks (kv_block_size.size(0))");
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
@@ -945,23 +904,18 @@ block_sparse_attention_backward(torch::Tensor q,
TORCH_CHECK(k2q_block_sparse_index.size(0) == batch, "k2q_block_sparse_index batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k2q_block_sparse_num.size(0) == batch, "k2q_block_sparse_num batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(l_vec.size(2) == seq_len, "L sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(o.size(2) == seq_len, "O sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(og.size(2) == seq_len, "OG sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k2q_block_sparse_index.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_index idx 2 - must match seq_len / BLOCK_N");
TORCH_CHECK(k2q_block_sparse_num.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_num idx 2 - must match seq_len / BLOCK_N");
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K sequence length");
TORCH_CHECK(l_vec.size(2) == q_seq_len, "L sequence length dimension - idx 2 - must match Q sequence length");
TORCH_CHECK(o.size(2) == q_seq_len, "O sequence length dimension - idx 2 - must match Q sequence length");
TORCH_CHECK(og.size(2) == q_seq_len, "OG sequence length dimension - idx 2 - must match Q sequence length");
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
TORCH_CHECK(k2q_block_sparse_num.size(2) == num_kv_blocks, "k2q_block_sparse_num idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(o.size(3) == head_dim, "O head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(og.size(3) == head_dim, "OG head dimension - idx 3 - must match for all non-vector inputs");
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
@@ -988,20 +942,20 @@ block_sparse_attention_backward(torch::Tensor q,
torch::Tensor qg = torch::zeros({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(head_dim)}, l_vec.options());
torch::Tensor kg = torch::zeros({static_cast<const uint>(batch),
static_cast<const uint>(kv_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(kv_seq_len),
static_cast<const uint>(head_dim)}, l_vec.options());
torch::Tensor vg = torch::zeros({static_cast<const uint>(batch),
static_cast<const uint>(kv_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(kv_seq_len),
static_cast<const uint>(head_dim)}, l_vec.options());
torch::Tensor d_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(1)}, l_vec.options());
float* qg_ptr = qg.data_ptr<float>();
@@ -1030,7 +984,7 @@ block_sparse_attention_backward(torch::Tensor q,
// cudaStreamSynchronize(stream);
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
dim3 grid_bwd(q_seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (head_dim == 64) {
using og_tile = st_bf<4*16, 64>;
@@ -1043,9 +997,9 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_prep_globals = bwd_prep_globals<64>;
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1082,15 +1036,15 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_global_args = bwd_globals<64>;
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1101,14 +1055,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(seq_len),
static_cast<int>(kv_seq_len), // N is not used in the kernel
static_cast<int>(hr),
static_cast<int>(max_q_blocks_per_kv),
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())};
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
@@ -1147,9 +1101,9 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_prep_globals = bwd_prep_globals<128>;
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1186,15 +1140,15 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_global_args = bwd_globals<128>;
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1205,14 +1159,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(seq_len),
static_cast<int>(kv_seq_len), // N is not used in the kernel
static_cast<int>(hr),
static_cast<int>(max_q_blocks_per_kv),
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())};
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
@@ -1233,4 +1187,4 @@ block_sparse_attention_backward(torch::Tensor q,
return {qg, kg, vg};
//cudadevicesynchronize();
}
}
@@ -4,6 +4,7 @@
#include <torch/all.h>
#include <torch/python.h>
#include <cutlass/cutlass.h>
#include <cutlass/numeric_types.h>
#include "common/common.hpp"
#include "norm/layernorm.hpp"
@@ -14,10 +15,6 @@ auto layer_norm(
std::optional<at::Tensor const> const B,
std::optional<at::Tensor> Output
) {
using ElementIn = float;
using ElementOut = float;
using ElementWeight = float;
int64_t const m = Input.size(0);
int64_t const n = Input.size(1);
torch::Device const input_device = Input.device();
@@ -26,31 +23,70 @@ auto layer_norm(
Output.emplace(
torch::empty(
{m, n},
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
torch::TensorOptions().device(input_device).dtype(Input.scalar_type())
)
);
}
TORCH_CHECK(Output.value().scalar_type() == Input.scalar_type(),
"Output dtype must match Input dtype. Got Output=",
Output.value().scalar_type(), ", Input=", Input.scalar_type());
if (W.has_value()) {
TORCH_CHECK(W.value().scalar_type() == Input.scalar_type(),
"W dtype must match Input dtype. Got W=",
W.value().scalar_type(), ", Input=", Input.scalar_type());
}
if (B.has_value()) {
TORCH_CHECK(B.value().scalar_type() == Input.scalar_type(),
"B dtype must match Input dtype. Got B=",
B.value().scalar_type(), ", Input=", Input.scalar_type());
}
void *Iptr = Input.data_ptr();
void *Wptr = W.has_value() ? W.value().data_ptr() : nullptr;
void *Bptr = B.has_value() ? B.value().data_ptr() : nullptr;
void *Optr = Output.value().data_ptr();
BOOL_SWITCH(B.has_value(), BIAS, [&]{
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
CONFIG_SWITCH(n, [&]{
layernorm<
ElementIn, ElementOut, ElementWeight,
AFFINE, BIAS,
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA> (
Iptr, Wptr, Bptr,
Optr, eps, m, n,
at::cuda::getCurrentCUDAStream().stream()
);
auto stream = at::cuda::getCurrentCUDAStream().stream();
if (Input.scalar_type() == at::kHalf) {
using ElementIn = cutlass::half_t;
using ElementOut = cutlass::half_t;
using ElementWeight = cutlass::half_t;
BOOL_SWITCH(B.has_value(), BIAS, [&]{
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
CONFIG_SWITCH(n, [&]{
layernorm<ElementIn, ElementOut, ElementWeight, AFFINE, BIAS, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
Iptr, Wptr, Bptr, Optr, eps, m, n, stream);
});
});
});
});
} else if (Input.scalar_type() == at::kBFloat16) {
using ElementIn = cutlass::bfloat16_t;
using ElementOut = cutlass::bfloat16_t;
using ElementWeight = cutlass::bfloat16_t;
BOOL_SWITCH(B.has_value(), BIAS, [&]{
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
CONFIG_SWITCH(n, [&]{
layernorm<ElementIn, ElementOut, ElementWeight, AFFINE, BIAS, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
Iptr, Wptr, Bptr, Optr, eps, m, n, stream);
});
});
});
} else if (Input.scalar_type() == at::kFloat) {
using ElementIn = float;
using ElementOut = float;
using ElementWeight = float;
BOOL_SWITCH(B.has_value(), BIAS, [&]{
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
CONFIG_SWITCH(n, [&]{
layernorm<ElementIn, ElementOut, ElementWeight, AFFINE, BIAS, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
Iptr, Wptr, Bptr, Optr, eps, m, n, stream);
});
});
});
} else {
TORCH_CHECK(false, "Unsupported dtype for layer_norm_cuda: ", Input.scalar_type());
}
@@ -68,9 +68,22 @@ public:
// mean reduction
float u = _reduce_sum(x, shared_data) / params.n;
// IMPORTANT:
// Loader pads out-of-range lanes with 0. That is OK for the sum, but after
// subtracting mean, those padded lanes become -u and would incorrectly
// contribute to the variance. Mask them back to 0 before variance reduction.
// We launch exactly NumThrPerCta threads for a 1xMaxHiddenSize tile,
// so each thread is responsible for a contiguous chunk in N.
int thr_n_offset = tidx * NumElementPerThread;
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < NumElementPerThread; ++i)
x[i] -= u;
for (int i = 0; i < NumElementPerThread; ++i) {
int idx = thr_n_offset + i;
if (idx < params.n) {
x[i] -= u;
} else {
x[i] = 0.f;
}
}
__syncthreads();
// var reduction
@@ -4,6 +4,7 @@
#include <torch/all.h>
#include <torch/python.h>
#include <cutlass/cutlass.h>
#include <cutlass/numeric_types.h>
#include <pybind11/pybind11.h>
#include "common/common.hpp"
@@ -16,10 +17,6 @@ auto rms_norm(
std::optional<at::Tensor>& Output
) {
using ElementIn = float;
using ElementOut = float;
using ElementWeight = float;
int64_t const m = Input.size(0);
int64_t const n = Input.size(1);
torch::Device const input_device = Input.device();
@@ -28,27 +25,51 @@ auto rms_norm(
Output.emplace(
torch::empty(
{m, n},
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
torch::TensorOptions().device(input_device).dtype(Input.scalar_type())
)
);
}
TORCH_CHECK(Output.value().scalar_type() == Input.scalar_type(),
"Output dtype must match Input dtype. Got Output=",
Output.value().scalar_type(), ", Input=", Input.scalar_type());
if (Weight.has_value()) {
TORCH_CHECK(Weight.value().scalar_type() == Input.scalar_type(),
"Weight dtype must match Input dtype. Got Weight=",
Weight.value().scalar_type(), ", Input=", Input.scalar_type());
}
void *Iptr = Input.data_ptr();
void *Wptr = Weight.has_value() ? Weight.value().data_ptr() : nullptr;
void *Optr = Output.value().data_ptr();
CONFIG_SWITCH(n, [&]{
rmsnorm<
ElementIn, ElementOut, ElementWeight,
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA
> (
Iptr, Wptr,
Optr,
eps, m, n,
at::cuda::getCurrentCUDAStream().stream()
);
});
if (Input.scalar_type() == at::kHalf) {
using ElementIn = cutlass::half_t;
using ElementOut = cutlass::half_t;
using ElementWeight = cutlass::half_t;
CONFIG_SWITCH(n, [&]{
rmsnorm<ElementIn, ElementOut, ElementWeight, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
Iptr, Wptr, Optr, eps, m, n, at::cuda::getCurrentCUDAStream().stream());
});
} else if (Input.scalar_type() == at::kBFloat16) {
using ElementIn = cutlass::bfloat16_t;
using ElementOut = cutlass::bfloat16_t;
using ElementWeight = cutlass::bfloat16_t;
CONFIG_SWITCH(n, [&]{
rmsnorm<ElementIn, ElementOut, ElementWeight, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
Iptr, Wptr, Optr, eps, m, n, at::cuda::getCurrentCUDAStream().stream());
});
} else if (Input.scalar_type() == at::kFloat) {
using ElementIn = float;
using ElementOut = float;
using ElementWeight = float;
CONFIG_SWITCH(n, [&]{
rmsnorm<ElementIn, ElementOut, ElementWeight, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
Iptr, Wptr, Optr, eps, m, n, at::cuda::getCurrentCUDAStream().stream());
});
} else {
TORCH_CHECK(false, "Unsupported dtype for rms_norm_cuda: ", Input.scalar_type());
}
return Output;
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.2"
version = "0.2.4"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -0,0 +1,298 @@
from __future__ import annotations
import os
from typing import Tuple
import torch
def _get_sm90_ops():
try:
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
except Exception:
return None, None
return (
getattr(fastvideo_kernel_ops, "block_sparse_fwd", None),
getattr(fastvideo_kernel_ops, "block_sparse_bwd", None),
)
def _is_sm90() -> bool:
if not torch.cuda.is_available():
return False
major, minor = torch.cuda.get_device_capability(0)
return major == 9 and minor == 0
def _force_triton() -> bool:
# Force Triton even on SM90 and even if the compiled extension is available.
# Useful for CI / debugging / parity testing.
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Pure-torch (no triton) conversion:
block_map: [B, H, Q, KV] bool (or [H, Q, KV] which will be treated as B=1)
returns:
index: [B, H, Q, KV] int32 (packed KV indices, -1 padding)
num: [B, H, Q] int32 (#kv blocks per q block)
"""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
if block_map.dim() != 4:
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
B, H, Q, KV = block_map.shape
index = torch.full((B, H, Q, KV), -1, dtype=torch.int32, device=block_map.device)
num = torch.zeros((B, H, Q), dtype=torch.int32, device=block_map.device)
# Small sizes in practice (B=1, H<=16, Q/KV<=64), so a Python loop is fine.
for b in range(B):
for h in range(H):
for q in range(Q):
kv_idx = torch.nonzero(block_map[b, h, q], as_tuple=False).flatten().to(torch.int32)
n = int(kv_idx.numel())
if n:
index[b, h, q, :n] = kv_idx
num[b, h, q] = n
return index, num
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_triton",
mutates_args=(),
device_types="cuda",
)
def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_forward,
)
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
return o, M
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_backward_triton",
mutates_args=(),
device_types="cuda",
)
def block_sparse_attn_backward_triton(
grad_output: torch.Tensor,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_backward,
)
dq, dk, dv = triton_block_sparse_attn_backward(
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
)
return dq, dk, dv
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_triton")
def _block_sparse_attn_backward_triton_fake(
grad_output: torch.Tensor,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
return dq, dk, dv
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def _setup_context_triton(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_sm90",
mutates_args=(),
device_types="cuda",
)
def block_sparse_attn_sm90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
block_sparse_fwd, _ = _get_sm90_ops()
if block_sparse_fwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
)
return o_padded, lse_padded
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_sm90")
def _block_sparse_attn_sm90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q_padded)
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
return o, lse
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_backward_sm90",
mutates_args=(),
device_types="cuda",
)
def block_sparse_attn_backward_sm90(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
_, block_sparse_bwd = _get_sm90_ops()
if block_sparse_bwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
dq, dk, dv = block_sparse_bwd(
q_padded,
k_padded,
v_padded,
o_padded,
lse_padded,
grad_output_padded,
k2q_idx,
k2q_num,
variable_block_sizes.int(),
)
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
def _block_sparse_attn_backward_sm90_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q_padded)
dk = torch.empty_like(k_padded)
dv = torch.empty_like(v_padded)
return dq, dk, dv
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
)
return dq, dk, dv, None, None
def _setup_context_sm90(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
o, lse = output
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
def block_sparse_attn(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Unified block-sparse attention op with autograd support.
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
"""
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
# Triton path: generally assumes q/k/v share the same padded length
if q.shape[2] != k.shape[2] or q.shape[2] != v.shape[2]:
raise RuntimeError("Triton fallback requires q/k/v to have the same padded length.")
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
+58 -15
View File
@@ -1,5 +1,6 @@
import math
import torch
from .block_sparse_attn import block_sparse_attn
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
from .triton_kernels.index import map_to_index
@@ -45,14 +46,22 @@ def sliding_tile_attention(
flag = shape_map[seq_shape]
for head_idx, (t, h, w) in enumerate(window_size):
# Per-head slices are not contiguous in the batch dimension when batch>1
# (they keep the original head-stride). The TK kernel assumes contiguous
# [B, H, S, D] layout, so we materialize a contiguous [B,1,S,D] view.
q_h = q[:, head_idx:head_idx + 1].contiguous()
k_h = k[:, head_idx:head_idx + 1].contiguous()
v_h = v[:, head_idx:head_idx + 1].contiguous()
o_h = torch.empty_like(q_h)
sta_fwd(
q[:, head_idx:head_idx + 1], k[:, head_idx:head_idx + 1],
v[:, head_idx:head_idx + 1], output[:, head_idx:head_idx + 1],
q_h, k_h,
v_h, o_h,
t, h, w, text_length, False, has_text, flag
)
output[:, head_idx:head_idx + 1] = o_h
if has_text:
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
sta_fwd(q.contiguous(), k.contiguous(), v.contiguous(), output, 3, 3, 3, text_length, True, True, flag)
return output[:, :, :seq_length]
@@ -62,6 +71,7 @@ def video_sparse_attn(
k: torch.Tensor,
v: torch.Tensor,
variable_block_sizes: torch.Tensor,
q_variable_block_sizes: torch.Tensor,
topk: int,
block_size: int | tuple = 64,
compress_attn_weight: torch.Tensor = None,
@@ -70,14 +80,42 @@ def video_sparse_attn(
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
batch, heads, seq_len, dim = q.shape
batch, heads, q_seq_len, dim = q.shape
kv_seq_len = k.shape[2]
if v.shape[2] != kv_seq_len:
raise ValueError(
f"Expected k and v to have the same sequence length, got "
f"k.shape[2]={kv_seq_len}, v.shape[2]={v.shape[2]}"
)
if k.shape[0] != batch or v.shape[0] != batch or k.shape[1] != heads or v.shape[1] != heads:
raise ValueError("Expected q/k/v to have the same batch and head dimensions.")
if q_seq_len % block_elements != 0 or kv_seq_len % block_elements != 0:
raise ValueError(
f"q_seq_len and kv_seq_len must be divisible by block_elements={block_elements}, "
f"got q_seq_len={q_seq_len}, kv_seq_len={kv_seq_len}"
)
q_num_blocks = q_seq_len // block_elements
kv_num_blocks = kv_seq_len // block_elements
if variable_block_sizes.numel() != kv_num_blocks:
raise ValueError(
f"variable_block_sizes must have length kv_num_blocks={kv_num_blocks}, "
f"got {variable_block_sizes.numel()}"
)
if q_variable_block_sizes.numel() != q_num_blocks:
raise ValueError(
f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
f"got {q_variable_block_sizes.numel()}"
)
# Compression branch
q_c = q.view(batch, heads, seq_len // block_elements, block_elements, dim)
k_c = k.view(batch, heads, seq_len // block_elements, block_elements, dim)
v_c = v.view(batch, heads, seq_len // block_elements, block_elements, dim)
q_c = q.view(batch, heads, q_num_blocks, block_elements, dim)
k_c = k.view(batch, heads, kv_num_blocks, block_elements, dim)
v_c = v.view(batch, heads, kv_num_blocks, block_elements, dim)
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
q_c = (q_c.float().sum(dim=3) / q_variable_block_sizes.view(1, 1, -1, 1)).to(
q.dtype)
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
k.dtype)
@@ -88,9 +126,9 @@ def video_sparse_attn(
attn = torch.softmax(scores, dim=-1)
out_c = torch.matmul(attn, v_c)
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
out_c = out_c.view(batch, heads, q_num_blocks, 1, dim)
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, seq_len, dim)
1).view(batch, heads, q_seq_len, dim)
# Sparse branch
topk_idx = torch.topk(scores, topk, dim=-1).indices
@@ -100,12 +138,17 @@ def video_sparse_attn(
idx, num = map_to_index(mask)
if block_sparse_fwd is not None:
out_s = block_sparse_fwd(
q, k, v, idx, num, variable_block_sizes.int()
)[0] # block_sparse_fwd returns vector<Tensor>
# Use autograd-enabled wrapper so backward works (and still uses SM90 kernel when available)
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
else:
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num,
variable_block_sizes)
if q_seq_len != kv_seq_len:
raise RuntimeError(
"q/k have different lengths, but the compiled CUDA kernel (block_sparse_fwd) "
"is not available. The Triton fallback currently requires q and k/v to have "
"the same padded length."
)
# Triton-only forward (kept for environments without the wrapper deps)
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
@@ -1 +1 @@
__version__ = "0.2.2"
__version__ = "0.2.4"
+6 -31
View File
@@ -42,37 +42,13 @@ def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, q
q_padded = vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
k_padded = vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
# Use raw kernel or triton
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
except ImportError:
raw_kernel = None
# Use autograd-enabled wrapper (internally dispatches to SM90 kernel or Triton)
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
output_padded, _aux = block_sparse_attn(
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
)
from fastvideo_kernel.triton_kernels.index import map_to_index
# Convert mask to indices
# block_sparse_mask is [H, M, N] bool
# We need to map it to index.
# block_sparse_mask needs to be expanded/reshaped?
# generate_block_sparse_mask_for_function returns [H, NumBlocksQ, NumBlocksKV]
# Ops.py logic:
# mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
# idx, num = map_to_index(mask)
idx, num = map_to_index(block_sparse_mask.unsqueeze(0)) # Add batch dim [1, H, M, N]
if raw_kernel:
out_s = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())
output = out_s[0]
else:
# Fallback to triton testing if C++ not available
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
output, _ = triton_block_sparse_attn_forward(q_padded, k_padded, v_padded, idx, num, variable_block_sizes)
output = output[:, :, q_non_pad_index, :]
output = output_padded[:, :, q_non_pad_index, :]
output.backward(dO)
return output, Q.grad, K.grad, V.grad
@@ -264,7 +240,6 @@ def generate_error_graphs_qkdiff(h, d, error_mode='all'):
print("-" * 150)
@pytest.mark.skip()
def test_video_sparse_attention_backward():
if not torch.cuda.is_available():
return
+7 -20
View File
@@ -3,6 +3,7 @@ import sys
from typing import Tuple
import torch
import pytest
from .utils import (
generate_block_sparse_mask_for_function,
@@ -57,23 +58,14 @@ def block_sparse_forward_test(
k_padded = ref.vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = ref.vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
# Use raw kernel or triton
# Use autograd-enabled wrapper (internally dispatches SM90 C++ vs Triton)
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
except ImportError:
raw_kernel = None
from fastvideo_kernel.triton_kernels.index import map_to_index
idx, num = map_to_index(block_sparse_mask)
if raw_kernel:
out_padded = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())[0]
else:
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
out_padded, _ = triton_block_sparse_attn_forward(
q_padded, k_padded, v_padded, idx, num, variable_block_sizes
out_padded, _aux = block_sparse_attn(
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
)
except RuntimeError as e:
pytest.skip(str(e))
# Remove padding on the query side
out = out_padded[:, :, q_non_pad_index, :]
@@ -156,11 +148,6 @@ def run_forward_qk_diff(
) -> Tuple[float, float]:
"""
Forward-only correctness test for the case S_q != S_kv.
NOTE:
- The Triton backend supports different Q/KV logical lengths via padding.
- The SM90 (H100) CUDA backend currently assumes the same number of blocks
for Q and KV, so we skip this test there.
"""
assert torch.cuda.is_available(), "VSA kernels require CUDA"
@@ -276,8 +276,9 @@ class VideoSparseAttentionImpl(AttentionImpl):
query,
key,
value,
variable_block_sizes=attn_metadata.variable_block_sizes,
topk=cur_topk,
attn_metadata.variable_block_sizes,
attn_metadata.variable_block_sizes,
cur_topk,
block_size=VSA_TILE_SIZE,
compress_attn_weight=gate_compress).transpose(1, 2)
@@ -7,10 +7,11 @@ from fastvideo.configs.models.encoders.clip import (
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig"
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config"
]
@@ -0,0 +1,72 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for Reason1 (Qwen2.5-VL) text encoder."""
from dataclasses import dataclass, field
from typing import Any
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig
@dataclass
class Reason1ArchConfig(TextEncoderArchConfig):
"""Architecture settings (defaults match Qwen2.5-VL-7B-Instruct)."""
architectures: list[str] = field(
default_factory=lambda: ["Qwen2_5_VLForConditionalGeneration"])
model_type: str = "qwen2_5_vl"
vocab_size: int = 152064
hidden_size: int = 3584
num_hidden_layers: int = 28
num_attention_heads: int = 28
num_key_value_heads: int = 4
intermediate_size: int = 18944
text_len: int = 512
hidden_state_skip_layer: int = 0
bos_token_id: int = 151643
pad_token_id: int = 151643
eos_token_id: int = 151645
image_token_id: int = 151655
video_token_id: int = 151656
vision_token_id: int = 151654
vision_start_token_id: int = 151652
vision_end_token_id: int = 151653
vision_config: dict[str, Any] | None = None
rope_theta: float = 1000000.0
rope_scaling: dict[str, Any] | None = field(default_factory=lambda: {
"type": "mrope",
"mrope_section": [16, 24, 24]
})
max_position_embeddings: int = 128000
max_window_layers: int = 28
embedding_concat_strategy: str = "mean_pooling"
n_layers_per_group: int = 5
num_embedding_padding_tokens: int = 512
attention_dropout: float = 0.0
hidden_act: str = "silu"
initializer_range: float = 0.02
rms_norm_eps: float = 1e-6
use_sliding_window: bool = False
sliding_window: int = 32768
tie_word_embeddings: bool = False
use_cache: bool = False
output_hidden_states: bool = True
torch_dtype: str = "bfloat16"
_attn_implementation: str = "flash_attention_2"
@dataclass
class Reason1Config(TextEncoderConfig):
"""Reason1 text encoder config."""
arch_config: Reason1ArchConfig = field(default_factory=Reason1ArchConfig)
tokenizer_type: str = "Qwen/Qwen2.5-VL-7B-Instruct"
@@ -1,4 +1,5 @@
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
@@ -9,5 +10,6 @@ __all__ = [
"WanVAEConfig",
"StepVideoVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Hunyuan15VAEConfig",
]
@@ -0,0 +1,223 @@
"""Cosmos 2.5 (Wan2.1-style) VAE config and checkpoint-key mapping."""
from __future__ import annotations
import re
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class Cosmos25VAEArchConfig(VAEArchConfig):
_name_or_path: str = ""
base_dim: int = 96
decoder_base_dim: int | None = None
z_dim: int = 16
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
is_residual: bool = False
in_channels: int = 3
out_channels: int = 3
patch_size: int | None = None
scale_factor_temporal: int = 4
scale_factor_spatial: int = 8
clip_output: bool = True
latents_mean: tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
# Simple 1:1 renames. More complex decoder remapping is handled by
# `map_official_key()`.
param_names_mapping: dict[str, str] = field(
default_factory=lambda: {
r"^conv1\.(.*)$": r"quant_conv.\1",
r"^conv2\.(.*)$": r"post_quant_conv.\1",
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
})
@staticmethod
def map_official_key(key: str) -> str | None:
"""Map a single official checkpoint key into FastVideo key space."""
def map_residual_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^residual\.0\.gamma$", sub):
return f"{prefix}.norm1.gamma"
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv1.{m.group(1)}"
if re.match(r"^residual\.3\.gamma$", sub):
return f"{prefix}.norm2.gamma"
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv2.{m.group(1)}"
m = re.match(r"^shortcut\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv_shortcut.{m.group(1)}"
return None
def map_attn_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^norm\.gamma$", sub):
return f"{prefix}.norm.gamma"
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
if m:
return f"{prefix}.to_qkv.{m.group(1)}"
m = re.match(r"^proj\.(weight|bias)$", sub)
if m:
return f"{prefix}.proj.{m.group(1)}"
return None
def map_resample_subkey(prefix: str, sub: str) -> str | None:
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
if m:
return f"{prefix}.resample.1.{m.group(1)}"
m = re.match(r"^time_conv\.(weight|bias)$", sub)
if m:
return f"{prefix}.time_conv.{m.group(1)}"
return None
m = re.match(r"^conv1\.(weight|bias)$", key)
if m:
return f"quant_conv.{m.group(1)}"
m = re.match(r"^conv2\.(weight|bias)$", key)
if m:
return f"post_quant_conv.{m.group(1)}"
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_in.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
if m:
return f"{m.group(1)}.norm_out.gamma"
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_out.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0",
m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if m:
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0",
m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1",
m.group(2))
m = re.match(r"^encoder\.downsamples\.(\d+)\.(.*)$", key)
if m:
idx = int(m.group(1))
sub = m.group(2)
if sub.startswith("residual.") or sub.startswith("shortcut."):
return map_residual_subkey(f"encoder.down_blocks.{idx}", sub)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return map_resample_subkey(f"encoder.down_blocks.{idx}", sub)
return None
m = re.match(r"^decoder\.upsamples\.(\d+)\.(.*)$", key)
if m:
uidx = int(m.group(1))
sub = m.group(2)
if uidx in (0, 1, 2):
block_i, res_i = 0, uidx
elif uidx == 3:
block_i, res_i = 0, None
elif uidx in (4, 5, 6):
block_i, res_i = 1, uidx - 4
elif uidx == 7:
block_i, res_i = 1, None
elif uidx in (8, 9, 10):
block_i, res_i = 2, uidx - 8
elif uidx == 11:
block_i, res_i = 2, None
elif uidx in (12, 13, 14):
block_i, res_i = 3, uidx - 12
else:
return None
if res_i is None:
return map_resample_subkey(
f"decoder.up_blocks.{block_i}.upsamplers.0",
sub,
)
return map_residual_subkey(
f"decoder.up_blocks.{block_i}.resnets.{res_i}",
sub,
)
return None
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
def __post_init__(self):
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(
self.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
self.temporal_compression_ratio = self.scale_factor_temporal
self.spatial_compression_ratio = self.scale_factor_spatial
@dataclass
class Cosmos25VAEConfig(VAEConfig):
"""Cosmos2.5 VAE config."""
arch_config: Cosmos25VAEArchConfig = field(
default_factory=Cosmos25VAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames -
self.tile_sample_stride_num_frames) * 2
+2 -1
View File
@@ -1,6 +1,7 @@
from fastvideo.configs.pipelines.base import (PipelineConfig,
SlidingTileAttnConfig)
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.registry import (
@@ -15,5 +16,5 @@ __all__ = [
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "get_pipeline_config_cls_from_name"
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
]
+89
View File
@@ -0,0 +1,89 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import Cosmos25VideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25ArchConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.reason1 import Reason1Config, Reason1ArchConfig
from fastvideo.configs.models.vaes import Cosmos25VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
def _identity_preprocess_text(prompt: str) -> str:
return prompt
def reason1_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
hidden_states = getattr(outputs, "hidden_states", None)
if hidden_states is None:
raise ValueError("Reason1 postprocess requires outputs.hidden_states")
hs = list(hidden_states)[1:]
normed = []
for h in hs:
h = h.float()
h = (h - h.mean(dim=-1, keepdim=True)) / (h.std(dim=-1, keepdim=True) +
1e-8)
normed.append(h)
return torch.cat(normed, dim=-1).to(hidden_states[0].dtype)
@dataclass
class Cosmos25Config(PipelineConfig):
"""Configuration for Cosmos 2.5 (Predict2.5) video generation pipeline."""
dit_config: DiTConfig = field(default_factory=lambda: Cosmos25VideoConfig(
arch_config=Cosmos25ArchConfig(
num_attention_heads=16,
attention_head_dim=128,
in_channels=16,
out_channels=16,
num_layers=28,
patch_size=[1, 2, 2],
max_size=[128, 240, 240],
rope_scale=[1.0, 3.0, 3.0],
text_embed_dim=1024,
mlp_ratio=4.0,
adaln_lora_dim=256,
use_adaln_lora=True,
concat_padding_mask=True,
extra_pos_embed_type=None,
use_crossattn_projection=True,
rope_enable_fps_modulation=False,
qk_norm="rms_norm",
)))
vae_config: VAEConfig = field(default_factory=Cosmos25VAEConfig)
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (Reason1Config(arch_config=Reason1ArchConfig(
embedding_concat_strategy="full_concat")), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (_identity_preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(reason1_postprocess_text, ))
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", ))
embedded_cfg_scale: float = 0.0
flow_shift: float = 5.0
vae_tiling: bool = False
vae_sp: bool = False
STA_mode: STA_Mode = STA_Mode.NONE
skip_time_steps: int = 0
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self._vae_latent_dim = 16
+7 -1
View File
@@ -6,6 +6,7 @@ from collections.abc import Callable
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -55,6 +56,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
"nvidia/Cosmos-Predict2-2B-Video2World": CosmosConfig,
"KyleShao/Cosmos-Predict2.5-2B-Diffusers": Cosmos25Config,
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGameI2V480PConfig,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGameI2V480PConfig,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGameI2V480PConfig,
@@ -94,7 +96,10 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"stepvideo":
lambda id: "stepvideo" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower(),
lambda id: "cosmos" in id.lower() and ("2.5" not in id.lower(
) and "2_5" not in id.lower() and "25" not in id.lower()),
"cosmos25":
lambda id: "cosmos25" in id.lower(),
"turbodiffusion":
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
# Add other pipeline architecture detectors
@@ -105,6 +110,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"longcatimagetovideo": LongCatT2V480PConfig,
"longcatvideocontinuation": LongCatT2V480PConfig,
"longcat": LongCatT2V480PConfig,
"cosmos25": Cosmos25Config,
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"matrixgame": MatrixGameI2V480PConfig,
+19
View File
@@ -0,0 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Cosmos_Predict2_5_2B_Diffusers_SamplingParam(SamplingParam):
"""Defaults for Cosmos 2.5 (Predict2.5) text-to-video diffusers-format model."""
height: int = 480
width: int = 832
num_frames: int = 121
fps: int = 24
guidance_scale: float = 7.0
# Official Cosmos2.5 sampling uses empty string as unconditional.
negative_prompt: str = ""
num_inference_steps: int = 35
+11 -3
View File
@@ -9,6 +9,7 @@ from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hun
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
from fastvideo.configs.sample.cosmos2_5 import Cosmos_Predict2_5_2B_Diffusers_SamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
@@ -96,6 +97,10 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"nvidia/Cosmos-Predict2-2B-Video2World":
Cosmos_Predict2_2B_Video2World_SamplingParam,
# Cosmos2.5
"KyleShao/Cosmos-Predict2.5-2B-Diffusers":
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
# MatrixGame2.0 models
"FastVideo/Matrix-Game-2.0-Base-Diffusers":
MatrixGame2_SamplingParam,
@@ -135,6 +140,10 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "matrixgame" in id.lower() or "matrix-game" in id.lower(),
"turbodiffusion":
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
"cosmos25":
lambda id: "cosmos2_5" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower() and "2_5" not in id.lower(),
# Add other pipeline architecture detectors
}
@@ -153,6 +162,8 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"matrixgame": MatrixGame2_SamplingParam,
"turbodiffusion":
TurboDiffusionT2V_1_3B_SamplingParam, # Default to T2V for fallback
"cosmos25": Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
"cosmos": Cosmos_Predict2_2B_Video2World_SamplingParam,
# Other fallbacks by architecture
}
@@ -176,9 +187,6 @@ def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal sampling param for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
+3 -1
View File
@@ -8,6 +8,7 @@ from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset,
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.dataset.rl_prompt_dataset import build_rl_prompt_dataloader
def getdataset(args) -> VideoCaptionMergedDataset:
@@ -47,5 +48,6 @@ def gettextdataset(args) -> TextDataset:
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
"VideoCaptionMergedDataset", "TextDataset",
"build_rl_prompt_dataloader"
]
+174
View File
@@ -0,0 +1,174 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from torch.utils.data import Dataset, DataLoader, Sampler
import json
import os
class TextPromptDataset(Dataset):
"""Dataset for loading text prompts from a simple text file (one prompt per line)."""
def __init__(self, dataset, split='train'):
self.file_path = os.path.join(dataset, f'{split}.txt')
with open(self.file_path, 'r') as f:
self.prompts = [line.strip() for line in f.readlines()]
def __len__(self):
return len(self.prompts)
def __getitem__(self, idx):
return {"prompt": self.prompts[idx], "metadata": {}}
@staticmethod
def collate_fn(examples):
prompts = [example["prompt"] for example in examples]
metadatas = [example["metadata"] for example in examples]
return prompts, metadatas
class GenevalPromptDataset(Dataset):
"""Dataset for loading prompts with metadata from JSONL files (e.g., GenEval format)."""
def __init__(self, dataset, split='train'):
self.file_path = os.path.join(dataset, f'{split}_metadata.jsonl')
with open(self.file_path, 'r', encoding='utf-8') as f:
self.metadatas = [json.loads(line) for line in f]
self.prompts = [item['prompt'] for item in self.metadatas]
def __len__(self):
return len(self.prompts)
def __getitem__(self, idx):
return {"prompt": self.prompts[idx], "metadata": self.metadatas[idx]}
@staticmethod
def collate_fn(examples):
prompts = [example["prompt"] for example in examples]
metadatas = [example["metadata"] for example in examples]
return prompts, metadatas
class KRepeatSampler(Sampler):
"""Sampler that repeats each sample k times, ensuring synchronized random selection. For single-node training, set num_replicas=1 and rank=0."""
def __init__(self, dataset, batch_size, k, num_replicas, rank, seed=0):
self.dataset = dataset
self.batch_size = batch_size # Batch size per GPU/card
self.k = k # Number of repetitions per sample
self.num_replicas = num_replicas # Total number of GPUs/cards
self.rank = rank # Current GPU/card rank
self.seed = seed # Random seed for synchronization
# Calculate the number of unique samples needed for each iteration
self.total_samples = self.num_replicas * self.batch_size
assert self.total_samples % self.k == 0, f"k can not div n*b, k{k}-num_replicas{num_replicas}-batch_size{batch_size}"
self.m = self.total_samples // self.k # different number of samples
self.step = 0
def __iter__(self):
while True:
# Generate a deterministic random sequence to ensure all cards are synchronized
g = torch.Generator()
g.manual_seed(self.seed + self.step)
# Randomly select m unique samples
indices = torch.randperm(len(self.dataset), generator=g)[:self.m].tolist()
# Repeat each sample k times to generate a total of n*b samples
repeated_indices = [idx for idx in indices for _ in range(self.k)]
# Shuffle the order to ensure even distribution
shuffled_indices = torch.randperm(len(repeated_indices), generator=g).tolist()
shuffled_samples = [repeated_indices[i] for i in shuffled_indices]
# Split samples among all cards
per_card_samples = []
for i in range(self.num_replicas):
start = i * self.batch_size
end = start + self.batch_size
per_card_samples.append(shuffled_samples[start:end])
# Return the sample indices for the current card
yield per_card_samples[self.rank]
def __len__(self):
return len(self.dataset) // self.batch_size
def set_step(self, step):
"""Used to synchronize the random state for different epochs."""
self.step = step
def build_rl_prompt_dataloader(
dataset_path: str,
dataset_type: str = "text",
split: str = "train",
train_batch_size: int = 8,
test_batch_size: int = 8,
k: int = 1,
seed: int = 42,
train_num_workers: int = 1,
test_num_workers: int = 8,
num_replicas: int = 1,
rank: int = 0,
) -> tuple[DataLoader, DataLoader]:
"""
Factory function to create train and test dataloaders for RL prompt datasets.
Args:
dataset_path: Path to dataset directory
dataset_type: "text" for TextPromptDataset or "geneval" for GenevalPromptDataset
split: Dataset split ("train" or "test")
train_batch_size: Batch size per GPU for training
test_batch_size: Batch size for testing
k: Number of times to repeat each sample (num_image_per_prompt)
seed: Random seed for sampler synchronization
train_num_workers: Number of workers for training dataloader
test_num_workers: Number of workers for test dataloader
num_replicas: Number of replicas (default 1 for single-node)
rank: Rank of current process (default 0 for single-node)
Returns:
Tuple of (train_dataloader, test_dataloader)
"""
# Create datasets based on type
if dataset_type == "text":
train_dataset = TextPromptDataset(dataset_path, 'train')
test_dataset = TextPromptDataset(dataset_path, 'test')
collate_fn = TextPromptDataset.collate_fn
elif dataset_type == "geneval":
train_dataset = GenevalPromptDataset(dataset_path, 'train')
test_dataset = GenevalPromptDataset(dataset_path, 'test')
collate_fn = GenevalPromptDataset.collate_fn
else:
raise ValueError(f"Unknown dataset_type: {dataset_type}. Must be 'text' or 'geneval'")
# Create infinite-loop training sampler
train_sampler = KRepeatSampler(
dataset=train_dataset,
batch_size=train_batch_size,
k=k,
num_replicas=num_replicas,
rank=rank,
seed=seed
)
# Create training dataloader with batch_sampler (infinite loop)
train_dataloader = DataLoader(
train_dataset,
batch_sampler=train_sampler,
num_workers=train_num_workers,
collate_fn=collate_fn,
)
# Create standard test dataloader
test_dataloader = DataLoader(
test_dataset,
batch_size=test_batch_size,
collate_fn=collate_fn,
shuffle=False,
num_workers=test_num_workers,
)
return train_dataloader, test_dataloader, train_dataset, test_dataset
+312 -2
View File
@@ -133,7 +133,7 @@ class FastVideoArgs:
# CPU offload parameters
dit_cpu_offload: bool = True
use_fsdp_inference: bool = False
dit_layerwise_offload: bool = False
dit_layerwise_offload: bool = True
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
vae_cpu_offload: bool = True
@@ -740,6 +740,271 @@ def get_current_fastvideo_args() -> FastVideoArgs:
return _current_fastvideo_args
@dataclasses.dataclass
class RLArgs:
"""
Reinforcement Learning (RL) specific arguments
"""
# ============================================================================
# SHARED RL CONFIGURATION
rl_mode: bool = False # Enable RL training mode
rl_algorithm: str = "grpo" # RL algorithm to use: "grpo", "ppo", "dpo"
# Trajectory collection
num_rollouts: int = 4 # Number of rollouts to collect per training step
rollout_steps: str = "20,30" # Random intermediate steps for sampling (comma-separated)
noise_injection_min: int = 10 # Minimum timestep for noise injection
noise_injection_max: int = 40 # Maximum timestep for noise injection
use_sde_sampling: bool = True # Use SDE sampling (Flow-GRPO-Fast)
num_denoising_steps: int = 2 # Number of denoising steps per trajectory (1-2 for fast)
# Advantage estimation
gamma: float = 0.99 # Discount factor for returns
lambda_param: float = 0.95 # GAE lambda parameter
use_gae: bool = True # Use Generalized Advantage Estimation
normalize_advantages: bool = True # Normalize advantages before policy update
# Reward models
reward_models: dict[str, float] = field(default_factory=lambda: {"dummy": 1.0}) # reward models (names, weight)
value_model_path: str = "" # Path to value model (can be empty to train from scratch)
value_model_share_backbone: bool = False # Share transformer backbone between policy and value
# Training schedule
warmup_steps: int = 1000 # Collect SFT-style data before starting RL
collect_on_policy: bool = True # Collect fresh rollouts each step (on-policy)
timestep_fraction: float = 0.99 # Fraction of timesteps to train on
num_inner_epochs: int = 1 # Number of inner epochs per outer epoch
# KL regularization
kl_beta: float = 0.004 # KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)
kl_reward: float = 0.0 # KL reward coefficient (alternative to KL loss, typically 0)
# SFT integration
sft_weight: float = 0.0 # SFT loss weight for supervised learning in RL training
sft_batch_size: int = 3 # Batch size for SFT data
# CFG
guidance_scale = 1.0 # use guidance_scale > 1.0 to enable CFG
# Statistics tracking
global_std: bool = False # Use global std across all samples vs per-group std
per_prompt_stat_tracking: bool = True # Track statistics per prompt
# Training options
use_diffusion_loss: bool = True # Use diffusion loss in training
# ============================================================================
# GRPO-SPECIFIC CONFIGURATION
# Policy optimization
grpo_policy_clip_range: float = 0.001 # PPO-style clipping range for policy ratio
grpo_value_clip_range: float = 0.2 # Value function clipping range
grpo_num_policy_epochs: int = 1 # Number of policy update epochs (GRPO typically uses 1)
grpo_num_value_epochs: int = 1 # Number of value function update epochs
grpo_target_kl: float = 0.01 # Target KL divergence for early stopping
grpo_entropy_coef: float = 0.0 # Entropy coefficient for exploration
grpo_value_loss_coef: float = 0.5 # Value loss coefficient
# GRPO-Guard safety mechanisms
grpo_use_grpo_guard: bool = True # Enable GRPO-Guard safety mechanisms
grpo_ratio_norm_correction: bool = True # RatioNorm: correct importance ratio bias
grpo_gradient_reweighting: bool = True # Reweight gradients across denoising steps
grpo_max_importance_ratio: float = 10.0 # Clip importance ratios above this value
# ============================================================================
# DPO-SPECIFIC CONFIGURATION
dpo_beta: float = 100.0 # DPO regularization parameter (typically much larger than GRPO beta)
dpo_ref_update_step: int = 10000000 # Reference model update frequency for OnlineDPO
dpo_label_smoothing: float = 0.0 # Label smoothing for DPO loss
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
"""Add RL-specific CLI arguments to the parser."""
# RL (Reinforcement Learning) arguments
parser.add_argument("--rl-mode",
action=StoreBoolean,
help="Enable RL training mode")
parser.add_argument("--rl-algorithm",
type=str,
default=RLArgs.rl_algorithm,
choices=["grpo", "ppo", "dpo"],
help="RL algorithm to use (grpo, ppo, dpo)")
# Trajectory collection (Flow-GRPO-Fast)
parser.add_argument("--rl-num-rollouts",
type=int,
default=RLArgs.num_rollouts,
help="Number of rollouts to collect per training step")
parser.add_argument("--rl-rollout-steps",
type=str,
default=RLArgs.rollout_steps,
help="Random intermediate steps for sampling (comma-separated)")
parser.add_argument("--rl-noise-injection-min",
type=int,
default=RLArgs.noise_injection_min,
help="Minimum timestep for noise injection")
parser.add_argument("--rl-noise-injection-max",
type=int,
default=RLArgs.noise_injection_max,
help="Maximum timestep for noise injection")
parser.add_argument("--rl-use-sde-sampling",
action=StoreBoolean,
help="Use SDE sampling (Flow-GRPO-Fast)")
parser.add_argument("--rl-num-denoising-steps",
type=int,
default=RLArgs.num_denoising_steps,
help="Number of denoising steps per trajectory (1-2 for fast)")
# Advantage estimation
parser.add_argument("--rl-gamma",
type=float,
default=RLArgs.gamma,
help="Discount factor for returns")
parser.add_argument("--rl-lambda",
type=float,
default=RLArgs.lambda_param,
help="GAE lambda parameter")
parser.add_argument("--rl-use-gae",
action=StoreBoolean,
help="Use Generalized Advantage Estimation")
parser.add_argument("--rl-normalize-advantages",
action=StoreBoolean,
help="Normalize advantages before policy update")
# Policy optimization (GRPO/PPO)
parser.add_argument("--rl-policy-clip-range",
type=float,
default=RLArgs.grpo_policy_clip_range,
dest="grpo_policy_clip_range", # Map to RLArgs field name
help="PPO-style clipping range for policy ratio")
parser.add_argument("--rl-value-clip-range",
type=float,
default=RLArgs.grpo_value_clip_range,
help="Value function clipping range")
parser.add_argument("--rl-num-policy-epochs",
type=int,
default=RLArgs.grpo_num_policy_epochs,
help="Number of policy update epochs (GRPO typically uses 1)")
parser.add_argument("--rl-num-value-epochs",
type=int,
default=RLArgs.grpo_num_value_epochs,
help="Number of value function update epochs")
parser.add_argument("--rl-target-kl",
type=float,
default=RLArgs.grpo_target_kl,
help="Target KL divergence for early stopping")
parser.add_argument("--rl-entropy-coef",
type=float,
default=RLArgs.grpo_entropy_coef,
help="Entropy coefficient for exploration")
parser.add_argument("--rl-value-loss-coef",
type=float,
default=RLArgs.grpo_value_loss_coef,
help="Value loss coefficient")
# GRPO-Guard (safety mechanisms)
parser.add_argument("--rl-use-grpo-guard",
action=StoreBoolean,
help="Enable GRPO-Guard safety mechanisms")
parser.add_argument("--rl-ratio-norm-correction",
action=StoreBoolean,
help="RatioNorm: correct importance ratio bias")
parser.add_argument("--rl-gradient-reweighting",
action=StoreBoolean,
help="Reweight gradients across denoising steps")
parser.add_argument("--rl-max-importance-ratio",
type=float,
default=RLArgs.grpo_max_importance_ratio,
help="Clip importance ratios above this value")
# Reward models
parser.add_argument("--reward-models",
type=str,
default='{"dummy": 1.0}',
help="Reward models as JSON dict (e.g., '{\"video_ocr\": 1.0, \"pickscore\": 0.5}')")
parser.add_argument("--value-model-path",
type=str,
default=RLArgs.value_model_path,
help="Path to value model (can be empty to train from scratch)")
parser.add_argument("--value-model-share-backbone",
action=StoreBoolean,
help="Share transformer backbone between policy and value")
# Training schedule
parser.add_argument("--rl-warmup-steps",
type=int,
default=RLArgs.warmup_steps,
help="Collect SFT-style data before starting RL")
parser.add_argument("--rl-collect-on-policy",
action=StoreBoolean,
help="Collect fresh rollouts each step (on-policy)")
parser.add_argument("--rl-timestep-fraction",
type=float,
default=RLArgs.timestep_fraction,
help="Fraction of timesteps to train on")
parser.add_argument("--rl-num-inner-epochs",
type=int,
default=RLArgs.num_inner_epochs,
help="Number of inner epochs per outer epoch")
# KL regularization
parser.add_argument("--rl-kl-beta",
type=float,
default=RLArgs.kl_beta,
dest="kl_beta", # Map CLI arg to RLArgs field name
help="KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)")
parser.add_argument("--rl-kl-reward",
type=float,
default=RLArgs.kl_reward,
help="KL reward coefficient (alternative to KL loss, typically 0)")
# SFT integration
parser.add_argument("--rl-sft-weight",
type=float,
default=RLArgs.sft_weight,
help="SFT loss weight for supervised learning in RL training")
parser.add_argument("--rl-sft-batch-size",
type=int,
default=RLArgs.sft_batch_size,
help="Batch size for SFT data")
# CFG settings
parser.add_argument("--guidance-scale",
type=float,
default=1.0,
help="Guidance scale for CFG")
# Statistics tracking
parser.add_argument("--rl-global-std",
action=StoreBoolean,
help="Use global std across all samples vs per-group std")
parser.add_argument("--rl-per-prompt-stat-tracking",
action=StoreBoolean,
help="Track statistics per prompt")
# Training options
parser.add_argument("--rl-use-diffusion-loss",
action=StoreBoolean,
help="Use diffusion loss in training")
# DPO-specific
parser.add_argument("--dpo-beta",
type=float,
default=RLArgs.dpo_beta,
help="DPO regularization parameter (typically much larger than GRPO beta)")
parser.add_argument("--dpo-ref-update-step",
type=int,
default=RLArgs.dpo_ref_update_step,
help="Reference model update frequency for OnlineDPO")
parser.add_argument("--dpo-label-smoothing",
type=float,
default=RLArgs.dpo_label_smoothing,
help="Label smoothing for DPO loss")
return parser
@dataclasses.dataclass
class TrainingArgs(FastVideoArgs):
"""
@@ -752,6 +1017,11 @@ class TrainingArgs(FastVideoArgs):
num_height: int = 0
num_width: int = 0
num_frames: int = 0
# RL dataset configuration (for RL prompt datasets)
rl_dataset_path: str = "" # Path to RL prompt dataset directory (defaults to data_path if not set)
rl_dataset_type: str = "text" # "text" or "geneval"
rl_num_image_per_prompt: int = 4 # k parameter for KRepeatSampler (num_image_per_prompt)
train_batch_size: int = 0
num_latent_t: int = 0
@@ -862,6 +1132,9 @@ class TrainingArgs(FastVideoArgs):
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
# Nested RL configuration
rl_args: RLArgs = dataclasses.field(default_factory=RLArgs)
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
@@ -886,6 +1159,25 @@ class TrainingArgs(FastVideoArgs):
kwargs[attr] = WorkloadType.from_string(
workload_type_value) if isinstance(
workload_type_value, str) else workload_type_value
elif attr == 'rl_args':
# Construct nested RLArgs from CLI arguments
rl_kwargs = {}
for rl_field in dataclasses.fields(RLArgs):
rl_attr = rl_field.name
if hasattr(args, rl_attr):
value = getattr(args, rl_attr)
# Special handling for reward_models: parse JSON string to dict
if rl_attr == 'reward_models' and isinstance(value, str):
rl_kwargs[rl_attr] = json.loads(value) if value else {}
else:
rl_kwargs[rl_attr] = value
else:
# Use default value from RLArgs
if rl_field.default_factory is not dataclasses.MISSING:
rl_kwargs[rl_attr] = rl_field.default_factory()
elif rl_field.default is not dataclasses.MISSING:
rl_kwargs[rl_attr] = rl_field.default
kwargs[attr] = RLArgs(**rl_kwargs)
# Use getattr with default value from the dataclass for potentially missing attributes
else:
# Get the field to check its default value
@@ -915,11 +1207,26 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--data-path",
type=str,
required=True,
help="Path to parquet files")
help="Path to parquet files (or RL prompt dataset directory for RL training)")
parser.add_argument("--dataloader-num-workers",
type=int,
required=True,
help="Number of workers for dataloader")
# RL dataset arguments (optional, defaults to data_path)
parser.add_argument("--rl-dataset-path",
type=str,
default="",
help="Path to RL prompt dataset directory (defaults to --data-path if not set)")
parser.add_argument("--rl-dataset-type",
type=str,
default="text",
choices=["text", "geneval"],
help="RL dataset type: 'text' for TextPromptDataset or 'geneval' for GenevalPromptDataset")
parser.add_argument("--rl-num-image-per-prompt",
type=int,
default=4,
help="Number of times to repeat each prompt (k parameter for KRepeatSampler)")
parser.add_argument("--num-height",
type=int,
required=True,
@@ -1284,6 +1591,9 @@ class TrainingArgs(FastVideoArgs):
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
# RL (Reinforcement Learning) arguments
RLArgs.add_cli_args(parser)
return parser
View File
+122
View File
@@ -0,0 +1,122 @@
# SPDX-License-Identifier: Apache-2.0
import functools
from typing import Any
from torch import nn
class ForwardHook:
"""
Base class for forward hooks.
Hooks are used in the way:
modified_args, modified_kwargs = hook.pre_forward(module, *args, **kwargs)
output = module.forward(*modified_args, **modified_kwargs)
modified_output = hook.post_forward(module, output)
"""
@classmethod
def name(cls) -> str:
raise NotImplementedError
def on_attach(self, module: nn.Module): # noqa: B027
"""Called once when the hook is attached to the module."""
pass
def on_detach(self, module: nn.Module): # noqa: B027
"""
Called once when the hook is detached from the module.
Note: this function is not guaranteed to be called if the module is
deleted before the hook is detached.
"""
pass
def pre_forward(self, module: nn.Module, *args,
**kwargs) -> tuple[tuple[Any, ...], dict[str, Any]]:
"""Called before the module's forward method is executed."""
return args, kwargs
def post_forward(self, module: nn.Module, output: Any) -> Any:
"""Called after the module's forward method is executed."""
return output
class ModuleHookManager:
module_hook_attribute = "_hook_manager"
def __init__(self, module: nn.Module):
self.module = module
self.forward_hooks: dict[str, ForwardHook] = {}
self.original_forward = module.forward
@classmethod
def get_from(cls, module: nn.Module) -> "ModuleHookManager | None":
if hasattr(module, cls.module_hook_attribute):
return getattr(module, cls.module_hook_attribute)
return None
@classmethod
def get_from_or_default(cls, module: nn.Module) -> "ModuleHookManager":
if not hasattr(module, cls.module_hook_attribute):
setattr(module, cls.module_hook_attribute, cls(module))
def forward_hook_wrapper(mod: nn.Module, *args, **kwargs):
manager: ModuleHookManager = getattr(mod,
cls.module_hook_attribute)
for hook in manager.forward_hooks.values():
args, kwargs = hook.pre_forward(mod, *args, **kwargs)
output = manager.original_forward(*args, **kwargs)
for hook in reversed(manager.forward_hooks.values()):
output = hook.post_forward(mod, output)
return output
module.forward = functools.partial(forward_hook_wrapper, module)
return getattr(module, cls.module_hook_attribute)
@staticmethod
def remove_from_manager(module: nn.Module) -> None:
if hasattr(module, ModuleHookManager.module_hook_attribute):
manager: ModuleHookManager = getattr(
module, ModuleHookManager.module_hook_attribute)
module.forward = manager.original_forward
delattr(module, ModuleHookManager.module_hook_attribute)
def _check_manager_attached(self) -> None:
if not hasattr(self.module, self.module_hook_attribute):
raise ValueError("ModuleHookManager is not attached to the module.")
if getattr(self.module, self.module_hook_attribute) is not self:
raise ValueError(
"ModuleHookManager attached to the module is different.")
def append_forward_hook(self, hook: ForwardHook):
self._check_manager_attached()
if hook.name() in self.forward_hooks:
raise ValueError(
f"Hook with name {hook.name()} is already registered.")
# after python 3.7, dicts maintain insertion order
self.forward_hooks[hook.name()] = hook
hook.on_attach(self.module)
def replace_forward_hook(self,
hook_name: str,
new_hook: ForwardHook,
run_on_attach: bool = True):
self._check_manager_attached()
if hook_name not in self.forward_hooks:
raise ValueError(f"No hook with name {hook_name} found.")
old_hook = self.forward_hooks[hook_name]
if run_on_attach:
old_hook.on_detach(self.module)
self.forward_hooks[hook_name] = new_hook
new_hook.on_attach(self.module)
def remove_forward_hook(self, hook_name: str, run_detach: bool = True):
self._check_manager_attached()
if hook_name not in self.forward_hooks:
raise ValueError(f"No hook with name {hook_name} found.")
if run_detach:
self.forward_hooks[hook_name].on_detach(self.module)
del self.forward_hooks[hook_name]
def get_forward_hook(self, hook_name: str) -> ForwardHook | None:
return self.forward_hooks.get(hook_name, None)
+164
View File
@@ -0,0 +1,164 @@
from contextlib import contextmanager
from typing import Any
import torch
from torch import nn
from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def _tensor_placeholder(tensor: torch.Tensor,
device: torch.device) -> torch.Tensor:
"""Create a rank-preserving empty placeholder on the specified device."""
shape = (0, ) if tensor.ndim <= 0 else (0, ) * tensor.ndim
return torch.empty(shape, device=device, dtype=tensor.dtype)
class LayerwiseOffloadState:
def __init__(
self,
async_copy_stream: torch.cuda.Stream,
device: torch.device,
next_state: "LayerwiseOffloadState | None" = None,
) -> None:
self.async_copy_stream = async_copy_stream
self.next_state = next_state
self.gpu_named_parameters: dict[str, torch.Tensor] = {}
self.cpu_named_parameters: dict[str, torch.Tensor] = {}
self.module_ref: nn.Module = None # type: ignore
self.device: torch.device = device
def _will_offload(self, name: str) -> bool:
return True
@torch.compiler.disable
def on_init(self, module: nn.Module):
self.module_ref = module
for name, param in self.module_ref.named_parameters():
if self._will_offload(name):
self.cpu_named_parameters[name] = (
param.data.detach().to("cpu").pin_memory())
param.data = _tensor_placeholder(param.data, self.device)
@torch.compiler.disable
def wait_and_replace_params(self):
torch.cuda.current_stream().wait_stream(self.async_copy_stream)
# now gpu_named_parameters are ready
for name, param in self.module_ref.named_parameters():
if not self._will_offload(name):
continue
if name not in self.gpu_named_parameters:
# first load with blocking load
self.gpu_named_parameters[name] = self.cpu_named_parameters[
name].to(self.device)
param.data = self.gpu_named_parameters[name]
@torch.compiler.disable
def prefetch_params(self):
compute_stream = torch.cuda.current_stream()
with torch.cuda.stream(self.async_copy_stream):
for name, param in self.module_ref.named_parameters():
if not self._will_offload(name):
continue
assert name not in self.gpu_named_parameters
gpu_param = self.cpu_named_parameters[name].to(
self.device, non_blocking=True)
gpu_param.record_stream(
compute_stream
) # ensure tensor will not be freed until forward is completed
self.gpu_named_parameters[name] = gpu_param
@torch.compiler.disable
def release_gpu_params(self):
for name, param in self.module_ref.named_parameters():
if self._will_offload(name):
param.data = _tensor_placeholder(param.data, self.device)
del self.gpu_named_parameters[name]
assert len(self.gpu_named_parameters) == 0
class LayerwiseOffloadHook(ForwardHook):
"""A hook that enables layerwise CPU offloading during forward pass."""
def __init__(self, state: LayerwiseOffloadState) -> None:
self.state = state
def on_attach(self, module: nn.Module):
self.state.on_init(module) # pyright: ignore
def on_detach(self, module: nn.Module):
named_parameters = dict(module.named_parameters())
for name, cpu_tensor in self.state.cpu_named_parameters.items():
if name not in self.state.gpu_named_parameters:
if name in named_parameters:
named_parameters[name].data = cpu_tensor.to(
device=self.state.device)
else:
logger.warning(
"Parameter {} not found in module during detachment.",
name,
)
@classmethod
def name(cls) -> str:
return "LayerwiseOffloadHook"
def pre_forward(self, module: nn.Module, *args, **kwargs):
self.state.wait_and_replace_params() # pyright: ignore
if self.state.next_state is not None:
self.state.next_state.prefetch_params() # pyright: ignore
return args, kwargs
def post_forward(self, module: torch.nn.Module, output: Any):
self.state.release_gpu_params() # pyright: ignore
return output
@contextmanager
def mutate_params_scope(self):
try:
# load params to GPU and keep them there
self.state.wait_and_replace_params() # pyright: ignore
yield
finally:
# instead of releasing, we should overwrite the original params since they have been modified
self.state.cpu_named_parameters.clear()
self.state.gpu_named_parameters.clear()
self.state.on_init(self.state.module_ref) # pyright: ignore
def enable_layerwise_offload(model: nn.Module, is_replace: bool = False):
if torch.cuda.is_available():
device = torch.device("cuda", torch.cuda.current_device())
else:
logger.warning(
"CUDA is not available. Layerwise offloading is disabled.")
return
state_list = []
async_stream = torch.cuda.Stream()
for name, submodule in model.named_children():
if isinstance(submodule, nn.ModuleList):
for idx, module_entry in enumerate(submodule):
state = LayerwiseOffloadState(async_copy_stream=async_stream,
device=device)
state_list.append(state)
hook_mgr = ModuleHookManager.get_from_or_default(module_entry)
hook = LayerwiseOffloadHook(state)
if is_replace:
existing_hook = hook_mgr.forward_hooks.get(hook.name())
if existing_hook is not None:
hook_mgr.replace_forward_hook(hook.name(), hook)
else:
raise AssertionError(
f"Expect hook exists in {name} for replacement.")
else:
hook_mgr.append_forward_hook(hook)
break
if len(state_list) == 0:
raise ValueError(
"No nn.ModuleList found in the model for layerwise offloading.")
# circular linking of states
for i in range(len(state_list)):
state_list[i].next_state = state_list[(i + 1) % len(state_list)]
+2 -12
View File
@@ -734,18 +734,8 @@ class WanTransformer3DModel(CachableDiT):
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
else:
offload_mgr = getattr(self, "_layerwise_offload_manager", None)
use_offload = offload_mgr is not None and getattr(offload_mgr, "enabled", False)
for i, block in enumerate(self.blocks):
scope = offload_mgr.layer_scope(
prefetch_layer_idx=i + 1 if i + 1 < len(self.blocks) else None,
release_layer_idx=i,
non_blocking=True,
) if use_offload else nullcontext()
with scope:
hidden_states = block(hidden_states, encoder_hidden_states,
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
# if teacache is enabled, we need to cache the original hidden states
File diff suppressed because it is too large Load Diff
+353
View File
@@ -0,0 +1,353 @@
# SPDX-License-Identifier: Apache-2.0
"""Reason1 (Qwen2.5-VL) text encoder."""
import os
from dataclasses import dataclass
from collections.abc import Iterable
import torch
from transformers import AutoProcessor
from fastvideo.configs.models.encoders import BaseEncoderOutput, Reason1Config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.models.encoders.qwen2_5_vl_custom import (
Qwen2_5_VLForConditionalGenerationSimple,
Qwen2_5_VLConfig,
get_rope_index,
)
logger = init_logger(__name__)
@dataclass(frozen=True)
class _WeightsSource:
"""Mimic `TextEncoderLoader.Source` (avoid import cycles)."""
model_or_path: str
prefix: str = ""
fall_back_to_pt: bool = True
allow_patterns_overrides: list[str] | None = None
class Reason1TextEncoder(TextEncoder):
"""Reason1 (Qwen2.5-VL) text encoder."""
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
def __init__(self, config: Reason1Config, prefix: str = "", checkpoint_path: str | None = None):
super().__init__(config)
self.prefix = prefix
self.quant_config = None # For future quantization support
self.embedding_concat_strategy = config.arch_config.embedding_concat_strategy
self.n_layers_per_group = config.arch_config.n_layers_per_group
self.num_embedding_padding_tokens = config.arch_config.num_embedding_padding_tokens
config_path = checkpoint_path if checkpoint_path else config.tokenizer_type
logger.info("Initializing Reason1TextEncoder (Qwen2.5-VL) from %s", config_path)
try:
from transformers import AutoConfig as HFAutoConfig
hf_config = HFAutoConfig.from_pretrained(
config_path,
trust_remote_code=True,
)
except Exception as e:
logger.warning("Failed to load HF config from %s (%s). Using default Qwen2.5-VL-7B config.",
config_path, e)
hf_config = Qwen2_5_VLConfig(
hidden_size=3584,
intermediate_size=18944,
max_window_layers=28,
num_attention_heads=28,
num_hidden_layers=28,
num_key_value_heads=4,
tie_word_embeddings=False,
vocab_size=152064,
)
hf_config.output_hidden_states = True
if hasattr(config.arch_config, '_attn_implementation') and config.arch_config._attn_implementation:
hf_config._attn_implementation = config.arch_config._attn_implementation
else:
hf_config._attn_implementation = "flash_attention_2"
logger.info("Reason1 attention implementation: %s", getattr(hf_config, "_attn_implementation", None))
with torch.device("meta"):
self.model = Qwen2_5_VLForConditionalGenerationSimple(hf_config)
self.processor = AutoProcessor.from_pretrained(
config_path,
trust_remote_code=True,
)
weights_override = os.getenv("FASTVIDEO_REASON1_WEIGHTS_PATH")
if weights_override:
self.secondary_weights = (
_WeightsSource(
model_or_path=weights_override,
prefix="",
fall_back_to_pt=True,
allow_patterns_overrides=None,
),
)
logger.info("Reason1TextEncoder: overlaying weights from %s", weights_override)
self._weights_loaded = False
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
# Cosmos2.5 alignment: keep attention_mask=None.
outputs = self.model(
input_ids=input_ids,
attention_mask=None,
position_ids=position_ids,
inputs_embeds=inputs_embeds,
output_hidden_states=True,
return_dict=True,
pixel_values=kwargs.get('pixel_values', None),
pixel_values_videos=kwargs.get('pixel_values_videos', None),
image_grid_thw=kwargs.get('image_grid_thw', None),
video_grid_thw=kwargs.get('video_grid_thw', None),
)
hidden_states = outputs.hidden_states
last_hidden_state = hidden_states[-1]
return BaseEncoderOutput(
last_hidden_state=last_hidden_state,
hidden_states=hidden_states if output_hidden_states else None,
attention_mask=None,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
first_weight = None
weights_list = []
for name, weight in weights:
if first_weight is None:
first_weight = weight
self.model = self.model.to_empty(device=weight.device)
self.model.init_weights(buffer_device=weight.device)
weights_list.append((name, weight))
params_dict = dict(self.model.named_parameters())
loaded_params: set[str] = set()
skipped_weights = {"lm_head": 0, "visual": 0, "decoder": 0}
for name, loaded_weight in weights_list:
if "lm_head" in name:
skipped_weights["lm_head"] += 1
continue
if "visual" in name:
skipped_weights["visual"] += 1
continue
if "decoder" in name:
skipped_weights["decoder"] += 1
continue
# Handle stacked params mapping (for quantized models)
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
param_name_with_prefix = f"model.{name}" if self.prefix == "" else f"{self.prefix}.{name}"
loaded_params.add(param_name_with_prefix)
break
else:
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
param_name_with_prefix = f"model.{name}" if self.prefix == "" else f"{self.prefix}.{name}"
loaded_params.add(param_name_with_prefix)
if first_weight is not None:
self.model = self.model.to(first_weight.device)
all_params = set(f"model.{name}" if self.prefix == "" else f"{self.prefix}.{name}"
for name in params_dict.keys())
loaded_params.update(all_params)
# Mark weights as loaded
self._weights_loaded = True
return loaded_params
def compute_text_embeddings_online(
self,
data_batch: dict[str, list[str]],
input_caption_key: str,
) -> torch.Tensor:
prompts = data_batch[input_caption_key]
return self.compute_text_embeddings(prompts)
def compute_text_embeddings(
self,
prompts: list[str],
device: str | torch.device = "cuda",
) -> torch.Tensor:
"""Compute embeddings for a list of prompts."""
input_ids_batch = []
tok = getattr(self.processor, "tokenizer", None)
if tok is None:
raise RuntimeError("Reason1TextEncoder requires processor.tokenizer")
pad_id = getattr(tok, "pad_id", None)
if pad_id is None:
pad_id = getattr(tok, "pad_token_id", None)
if pad_id is None:
pad_id = getattr(self.model.config, "pad_token_id", None)
if pad_id is None:
pad_id = 0
for prompt in prompts:
conversations = [
{
"role": "system",
"content": [
{
"type": "text",
"text": "You are a helpful assistant who will provide prompts to an image generator.",
}
],
},
{
"role": "user",
"content": [
{
"type": "text",
"text": prompt,
}
],
},
]
try:
tokenizer_output = tok.apply_chat_template(
conversations,
tokenize=True,
add_generation_prompt=False,
add_vision_id=False,
)
except TypeError:
tokenizer_output = tok.apply_chat_template(
conversations,
tokenize=True,
add_generation_prompt=False,
)
if isinstance(tokenizer_output, dict) and "input_ids" in tokenizer_output:
input_ids = tokenizer_output["input_ids"]
if hasattr(input_ids, "tolist"):
input_ids = input_ids.tolist()
else:
input_ids = tokenizer_output
if hasattr(input_ids, "tolist"):
input_ids = input_ids.tolist()
if isinstance(input_ids, list) and len(input_ids) == 1 and isinstance(
input_ids[0], list):
input_ids = input_ids[0]
if not isinstance(input_ids, list):
raise RuntimeError(
f"Unexpected chat_template output type: {type(tokenizer_output)}"
)
if self.num_embedding_padding_tokens > len(input_ids):
pad_len = self.num_embedding_padding_tokens - len(input_ids)
input_ids = input_ids + [pad_id] * pad_len
else:
input_ids = input_ids[:self.num_embedding_padding_tokens]
input_ids = torch.LongTensor(input_ids).to(device=device)
input_ids_batch.append(input_ids)
input_ids_batch = torch.stack(input_ids_batch, dim=0)
# Cosmos2.5 alignment: keep attention_mask=None.
target_device = input_ids_batch.device
try:
embed_device = self.model.model.embed_tokens.weight.device # type: ignore[attr-defined]
except Exception:
embed_device = None
if embed_device is not None and embed_device != target_device:
self.model = self.model.to(target_device)
with torch.no_grad():
position_ids, _ = get_rope_index(
self.model.config,
input_ids_batch,
image_grid_thw=None,
video_grid_thw=None,
second_per_grid_ts=None,
attention_mask=None,
)
position_ids = position_ids.to(target_device)
outputs = self.model.model(
input_ids=input_ids_batch,
position_ids=position_ids,
attention_mask=None,
output_hidden_states=True,
return_dict=True,
use_cache=False,
)
hidden_states = outputs.hidden_states
normalized_hidden_states = []
for layer_idx in range(1, len(hidden_states)):
normalized_state = self._mean_normalize(hidden_states[layer_idx])
normalized_hidden_states.append(normalized_state)
if self.embedding_concat_strategy == "full_concat":
text_embeddings = torch.cat(normalized_hidden_states, dim=-1)
elif self.embedding_concat_strategy == "mean_pooling":
text_embeddings = torch.stack(normalized_hidden_states).mean(dim=0)
elif self.embedding_concat_strategy == "pool_every_n_layers_and_concat":
pooled_embeddings = []
for i in range(0, len(normalized_hidden_states), self.n_layers_per_group):
group = normalized_hidden_states[i : i + self.n_layers_per_group]
pooled = torch.stack(group).mean(dim=0)
pooled_embeddings.append(pooled)
text_embeddings = torch.cat(pooled_embeddings, dim=-1)
else:
raise ValueError(
f"Unknown embedding_concat_strategy: {self.embedding_concat_strategy}"
)
return text_embeddings
@staticmethod
def _mean_normalize(tensor: torch.Tensor) -> torch.Tensor:
return (tensor - tensor.mean(dim=-1, keepdim=True)) / (
tensor.std(dim=-1, keepdim=True) + 1e-8
)
+90 -50
View File
@@ -15,8 +15,7 @@ import torch.distributed as dist
import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file
from torch.distributed import init_device_mesh
from transformers import AutoImageProcessor, AutoTokenizer
from transformers import UMT5EncoderModel
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.configs.models import EncoderConfig
@@ -35,8 +34,8 @@ from fastvideo.models.loader.weight_utils import (
safetensors_weights_iterator,
)
from fastvideo.models.registry import ModelRegistry
from fastvideo.utils import PRECISION_TO_TYPE
from fastvideo.models.layerwise_offload import LayerwiseOffloadManager
from fastvideo.utils import PRECISION_TO_TYPE, is_pin_memory_available
from fastvideo.hooks.layerwise_offload import enable_layerwise_offload
logger = init_logger(__name__)
@@ -91,10 +90,12 @@ class ComponentLoader(ABC):
if module_type in module_loaders:
loader_cls, expected_library = module_loaders[module_type]
# Assert that the library matches what's expected for this module type
assert transformers_or_diffusers == expected_library, (
f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
)
# Allow fastvideo.* libraries for custom implementations (e.g. Cosmos2_5Pipeline)
# that aren't available in diffusers/transformers yet
is_fastvideo_module = transformers_or_diffusers.startswith("fastvideo.")
if not is_fastvideo_module:
# Assert that the library matches what's expected for this module type
assert transformers_or_diffusers == expected_library, f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
return loader_cls()
# For unknown module types, use a generic loader
@@ -279,7 +280,7 @@ class TextEncoderLoader(ComponentLoader):
target_device: torch.device,
fastvideo_args: FastVideoArgs,
dtype: str = "fp16",
use_text_encoder_override: bool = False, # prevent subclasses from misusing
use_text_encoder_override: bool = False, # prevent subclasses from misusing
):
use_cpu_offload = (
fastvideo_args.text_encoder_cpu_offload
@@ -296,7 +297,10 @@ class TextEncoderLoader(ComponentLoader):
)
# Set quantization config if specified
if use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if (
use_text_encoder_override
and fastvideo_args.override_text_encoder_quant is not None
):
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError(
"override_text_encoder_quant is set but override_text_encoder_safetensors is None"
@@ -313,7 +317,10 @@ class TextEncoderLoader(ComponentLoader):
model: TextEncoder = model_cls(model_config) # type: ignore
weights_to_load = {name for name, _ in model.named_parameters()}
if use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None:
if (
use_text_encoder_override
and fastvideo_args.override_text_encoder_safetensors is not None
):
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
@@ -340,6 +347,7 @@ class TextEncoderLoader(ComponentLoader):
from fastvideo.platforms import current_platform
if use_cpu_offload:
pin_cpu_memory = fastvideo_args.pin_cpu_memory and is_pin_memory_available()
# Disable FSDP for MPS as it's not compatible
if current_platform.is_mps():
logger.info(
@@ -357,7 +365,7 @@ class TextEncoderLoader(ComponentLoader):
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
pin_cpu_memory=pin_cpu_memory,
)
else:
mesh = init_device_mesh(
@@ -371,7 +379,7 @@ class TextEncoderLoader(ComponentLoader):
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
pin_cpu_memory=pin_cpu_memory,
)
# We only enable strict check for non-quantized models
# that have loaded weights tracking currently.
@@ -449,6 +457,33 @@ class TokenizerLoader(ComponentLoader):
"""Load the tokenizer based on the model path, and inference args."""
logger.info("Loading tokenizer from %s", model_path)
# Cosmos2.5 stores an AutoProcessor config in `tokenizer/config.json` (not a tokenizer
# config). Use its `_name_or_path` (e.g. Qwen/Qwen2.5-VL-7B-Instruct) as the source.
tokenizer_cfg_path = os.path.join(model_path, "config.json")
if os.path.exists(tokenizer_cfg_path):
try:
with open(tokenizer_cfg_path, "r") as f:
tokenizer_cfg = json.load(f)
if isinstance(tokenizer_cfg, dict) and (
tokenizer_cfg.get("_class_name") == "AutoProcessor"
or "processor_type" in tokenizer_cfg
):
src = tokenizer_cfg.get("_name_or_path", "")
if isinstance(src, str) and src.strip():
processor = AutoProcessor.from_pretrained(
src.strip(),
trust_remote_code=True,
)
logger.info(
"Loaded tokenizer/processor from %s: %s",
src,
processor.__class__.__name__,
)
return processor
except Exception:
# If parsing fails, fall through to AutoTokenizer below.
pass
tokenizer = AutoTokenizer.from_pretrained(
model_path, # "<path to model>/tokenizer"
# in v0, this was same string as encoder_name "ClipTextModel"
@@ -491,19 +526,40 @@ class VAELoader(ComponentLoader):
if fastvideo_args.pipeline_config.vae_precision
else torch.bfloat16
):
# Cosmos2.5 uses a Wan2.1 VAE stored as `tokenizer.safetensors` under the VAE folder.
is_cosmos25 = fastvideo_args.pipeline_config.__class__.__name__ == "Cosmos25Config"
if class_name == "AutoencoderKLWan" and is_cosmos25:
from fastvideo.models.vaes.cosmos25wanvae import Cosmos25WanVAE
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
vae = Cosmos25WanVAE(device=target_device, dtype=dtype)
weight_path = os.path.join(model_path, "tokenizer.safetensors")
if not os.path.exists(weight_path):
raise FileNotFoundError(
f"Missing Cosmos2.5 VAE weights: {weight_path}"
)
sd = safetensors_load_file(weight_path)
vae.load_state_dict(sd, strict=False)
return vae.eval()
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
vae.load_state_dict(
loaded, strict=False
) # We might only load encoder or decoder
os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Common case: a single `.safetensors` checkpoint file.
# Some models may be sharded into multiple files; in that case we merge.
if len(safetensors_list) == 1:
loaded = safetensors_load_file(safetensors_list[0])
else:
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
vae.load_state_dict(loaded, strict=False)
return vae.eval()
@@ -581,7 +637,17 @@ class TransformerLoader(ComponentLoader):
]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
default_dtype)
assert fastvideo_args.hsdp_shard_dim is not None
# Cosmos2.5 checkpoints can include extra entries not present in the
# instantiated model (e.g. pos_embedder ranges / *_extra_state). Load
# non-strictly for Cosmos2.5 only; keep upstream strict behavior for others.
strict_load = not (
cls_name.startswith("Cosmos25")
or cls_name == "Cosmos25Transformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25"
)
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={"config": dit_config, "hf_config": hf_config},
@@ -589,6 +655,7 @@ class TransformerLoader(ComponentLoader):
device=get_local_torch_device(),
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
strict=strict_load,
cpu_offload=fastvideo_args.dit_cpu_offload,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
fsdp_inference=fastvideo_args.use_fsdp_inference,
@@ -611,35 +678,8 @@ class TransformerLoader(ComponentLoader):
model = model.eval()
if fastvideo_args.dit_layerwise_offload and hasattr(model, "blocks"):
# Check if this is a Wan model (only Wan models support layerwise offload)
is_wan_model = "Wan" in cls_name
if not is_wan_model:
logger.warning(
"Layerwise offload is currently only supported for Wan models. "
"Model class '%s' does not support layerwise offload. "
"Disabling layerwise offload for this model.",
cls_name
)
else:
try:
num_layers = len(getattr(model, "blocks"))
except TypeError:
num_layers = None
if isinstance(num_layers, int) and num_layers > 0:
# Ensure model is on the correct device (CUDA) before initializing manager
# This ensures non-managed parameters (embeddings, final norms) are on GPU
model = model.to(get_local_torch_device())
mgr = LayerwiseOffloadManager(
model,
module_list_attr="blocks",
num_layers=num_layers,
enabled=True,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
auto_initialize=True,
)
setattr(model, "_layerwise_offload_manager", mgr)
if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload:
enable_layerwise_offload(model)
return model
+25 -2
View File
@@ -23,7 +23,7 @@ from fastvideo.logger import init_logger
from fastvideo.models.loader.utils import (get_param_names_mapping,
hf_to_custom_state_dict)
from fastvideo.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.utils import set_mixed_precision_policy
from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available
logger = init_logger(__name__)
@@ -67,6 +67,7 @@ def maybe_load_fsdp_model(
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
strict: bool = True,
cpu_offload: bool = False,
fsdp_inference: bool = False,
output_dtype: torch.dtype | None = None,
@@ -106,6 +107,7 @@ def maybe_load_fsdp_model(
logger.info("Disabling FSDP for MPS platform as it's not compatible")
if use_fsdp:
pin_cpu_memory = pin_cpu_memory and is_pin_memory_available()
world_size = hsdp_replicate_dim * hsdp_shard_dim
if not training_mode and not fsdp_inference:
hsdp_replicate_dim = world_size
@@ -141,7 +143,7 @@ def maybe_load_fsdp_model(
weight_iterator,
device,
default_dtype,
strict=True,
strict=strict,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
@@ -151,6 +153,7 @@ def maybe_load_fsdp_model(
f"Unexpected param or buffer {n} on meta device.")
# 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
@@ -293,6 +296,26 @@ def load_model_from_full_model_state_dict(
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
# Some checkpoints include extra entries that are not part of the
# instantiated model's state_dict (e.g. `_extra_state` keys from
# some FSDP checkpoint formats). These can be safely skipped.
if (target_param_name.endswith("._extra_state")
or target_param_name.endswith("_extra_state")):
logger.warning(
"Skipping non-parameter checkpoint key: %s",
target_param_name,
)
continue
# For non-strict loads, treat this as an "unexpected key" and skip it
# (mirrors torch.nn.Module.load_state_dict(strict=False)).
if not strict:
logger.warning(
"Skipping unexpected checkpoint key (not present in model): %s",
target_param_name,
)
continue
raise ValueError(
f"Parameter {target_param_name} not found in custom model state dict. The hf to custom mapping may be incorrect."
)
+8 -2
View File
@@ -30,8 +30,9 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"),
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -50,6 +51,9 @@ _TEXT_ENCODER_MODELS = {
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
"Qwen2_5_VLForConditionalGeneration":
("encoders", "reason1", "Reason1TextEncoder"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
@@ -72,6 +76,8 @@ _SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
"FlowUniPCMultistepScheduler":
("schedulers", "scheduling_flow_unipc_multistep", "FlowUniPCMultistepScheduler"),
"SelfForcingFlowMatchScheduler":
("schedulers", "scheduling_self_forcing_flow_match",
"SelfForcingFlowMatchScheduler"),
@@ -109,6 +109,9 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
sigmas = 1.0 - alphas
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32)
# Needed when final_sigmas_type == "sigma_min" (kept for compatibility).
self.alphas_cumprod = torch.from_numpy(alphas).to(dtype=torch.float32)
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
assert shift is not None
@@ -171,6 +174,8 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
sigmas: list[float] | None = None,
mu: float | None | None = None,
shift: float | None | None = None,
use_karras_sigmas: bool | None = None,
use_kerras_sigma: bool | None = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -186,21 +191,44 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
" you have to pass a value for `mu` when `use_dynamic_shifting` is set to be `True`"
)
if sigmas is None:
assert num_inference_steps is not None
sigmas = np.linspace(self.sigma_max, self.sigma_min,
num_inference_steps +
1).copy()[:-1] # pyright: ignore
# Cosmos official uses `use_kerras_sigma=True` and a specific EDM sigma schedule.
# Some external code uses the misspelling `use_kerras_sigma`; support both.
if use_karras_sigmas is None and use_kerras_sigma is not None:
use_karras_sigmas = use_kerras_sigma
if use_karras_sigmas:
# Force to use the exact sigma used in official EDM sampler:
# sigma_max=200, sigma_min=0.01, rho=7
sigma_max = 200.0
sigma_min = 0.01
rho = 7.0
# Match the official Cosmos implementation: Karras/EDM schedule with
# `num_inference_steps + 1` points, then `final_sigmas_type="zero"`
# appends the terminal sigma (0.0).
ramp = np.arange(num_inference_steps + 1,
dtype=np.float32) / float(num_inference_steps)
min_inv_rho = sigma_min**(1 / rho)
max_inv_rho = sigma_max**(1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
# Convert EDM sigma to flow-matching sigma in [0, 1).
sigmas = sigmas / (1.0 + sigmas)
else:
if sigmas is None:
assert num_inference_steps is not None
sigmas = np.linspace(self.sigma_max, self.sigma_min,
num_inference_steps +
1).copy()[:-1] # pyright: ignore
if self.config.use_dynamic_shifting:
assert mu is not None
sigmas = self.time_shift(mu, 1.0, sigmas) # pyright: ignore
else:
if shift is None:
shift = self.config.shift
assert isinstance(sigmas, np.ndarray)
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
if not use_karras_sigmas:
if shift is None:
shift = self.config.shift
assert isinstance(sigmas, np.ndarray)
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas) # pyright: ignore
if self.config.final_sigmas_type == "sigma_min":
sigma_last = ((1 - self.alphas_cumprod[0]) /
@@ -418,8 +446,12 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
# Numerical safety
eps = 1e-12
lambda_t = torch.log(torch.clamp(alpha_t, min=eps)) - torch.log(
torch.clamp(sigma_t, min=eps))
lambda_s0 = torch.log(torch.clamp(alpha_s0, min=eps)) - torch.log(
torch.clamp(sigma_s0, min=eps))
h = lambda_t - lambda_s0
device = sample.device
@@ -430,7 +462,8 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
si = self.step_index - i # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
lambda_si = torch.log(torch.clamp(alpha_si, min=eps)) - torch.log(
torch.clamp(sigma_si, min=eps))
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
@@ -563,8 +596,11 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
eps = 1e-12
lambda_t = torch.log(torch.clamp(alpha_t, min=eps)) - torch.log(
torch.clamp(sigma_t, min=eps))
lambda_s0 = torch.log(torch.clamp(alpha_s0, min=eps)) - torch.log(
torch.clamp(sigma_s0, min=eps))
h = lambda_t - lambda_s0
device = this_sample.device
@@ -575,7 +611,8 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
si = self.step_index - (i + 1) # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
lambda_si = torch.log(torch.clamp(alpha_si, min=eps)) - torch.log(
torch.clamp(sigma_si, min=eps))
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
+735
View File
@@ -0,0 +1,735 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""
Cosmos 2.5 / Wan2.1 VAE adapter.
Why this exists:
- Cosmos2.5 uses a Wan2.1-style VAE, but the *diffusion model* operates in a
**normalized latent space**:
z_norm = (z - mean) / std
Meanwhile, FastVideo's `AutoencoderKLWan` operates in the VAE's native latent
space (denormalized):
z = z_norm * std + mean
This adapter provides a single, stable interface for FastVideo pipelines:
- `encode(x)` returns an object with `.mean` / `.sample()` / `.mode()`
- `decode(z)` returns a tensor in pixel space
It also exposes flags used by pipeline stages to avoid double (de)normalization:
- `handles_latent_norm = True` -> stages should NOT normalize encoder latents
- `handles_latent_denorm = True` -> stages should NOT denormalize before decode
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
@dataclass
class _TensorLatentDist:
"""Minimal distribution-like wrapper used by pipeline stages."""
mean: torch.Tensor
def mode(self) -> torch.Tensor:
return self.mean
def sample(self, generator: Any | None = None) -> torch.Tensor: # generator for API compatibility
# The official interface encodes deterministically; for compatibility we
# return the mean. (Stochastic posterior sampling isn't required for
# Cosmos2.5 inference.)
_ = generator
return self.mean
class Cosmos25WanVAEAdapter(nn.Module):
"""
Adapter that makes a Wan2.1-style VAE follow Cosmos2.5's latent contract:
- `encode()` returns **normalized** latents
- `decode()` expects **normalized** latents
"""
# Pipeline stage hints (see latent_preparation.py / decoding.py / image_encoding.py)
handles_latent_norm: bool = True
handles_latent_denorm: bool = True
latent_norm_mode: str = "internal" # informational
def __init__(
self,
inner: Any,
*,
latents_mean: Optional[torch.Tensor] = None,
latents_std: Optional[torch.Tensor] = None,
) -> None:
super().__init__()
self.inner = inner
# Preserve `config` when available; some pipeline utilities expect it.
self.config = getattr(inner, "config", None)
# If not provided, try to derive from `config.latents_mean/std`.
cfg = self.config
if latents_mean is None and cfg is not None and hasattr(cfg, "latents_mean"):
latents_mean = torch.tensor(cfg.latents_mean, dtype=torch.float32).view(1, -1, 1, 1, 1)
if latents_std is None and cfg is not None and hasattr(cfg, "latents_std"):
latents_std = torch.tensor(cfg.latents_std, dtype=torch.float32).view(1, -1, 1, 1, 1)
if latents_mean is None or latents_std is None:
raise RuntimeError(
"Cosmos25WanVAEAdapter requires latents_mean/latents_std (either passed explicitly or available on inner.config)."
)
# Register as buffers so `.to(...)` moves them with the module.
self.register_buffer("_latents_mean", latents_mean, persistent=False)
self.register_buffer("_latents_std", latents_std, persistent=False)
def _to_latent_stats(self, like: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
mean = self._latents_mean.to(device=like.device, dtype=like.dtype)
std = self._latents_std.to(device=like.device, dtype=like.dtype)
return mean, std
def get_latent_num_frames(self, num_pixel_frames: int) -> int:
# Keep parity with official interface.
if hasattr(self.inner, "get_latent_num_frames"):
return int(self.inner.get_latent_num_frames(num_pixel_frames))
return 1 + (num_pixel_frames - 1) // 4
def encode(self, x: torch.Tensor) -> _TensorLatentDist:
"""
Returns *normalized* latents (Cosmos contract).
"""
enc_out = self.inner.encode(x)
# Support common encoder output shapes:
# - DiagonalGaussianDistribution (FastVideo VAE): has `.mean` / `.sample()` / `.mode()`
# - diffusers EncoderOutput: has `.latent_dist`
# - raw tensor
if hasattr(enc_out, "latent_dist"):
dist = enc_out.latent_dist
z_mean = dist.mode() if hasattr(dist, "mode") else dist.mean
elif hasattr(enc_out, "mode") and hasattr(enc_out, "mean"):
z_mean = enc_out.mode()
elif isinstance(enc_out, torch.Tensor):
z_mean = enc_out
else:
attrs = [a for a in dir(enc_out) if not a.startswith("_")]
raise RuntimeError(
f"Unsupported VAE encoder output type: {type(enc_out)}. attrs={attrs}"
)
mean, std = self._to_latent_stats(z_mean)
z_norm = (z_mean - mean) / std
return _TensorLatentDist(z_norm)
def decode(self, z: torch.Tensor) -> torch.Tensor:
"""
Expects *normalized* latents (Cosmos contract).
"""
mean, std = self._to_latent_stats(z)
z_denorm = z * std + mean
out = self.inner.decode(z_denorm)
return out.sample if hasattr(out, "sample") else out
#
# Official-like Wan2.1 VAE implementation (ported from cosmos_predict2 wan2pt1.py)
# -------------------------------------------------------------------------------
# Motivation:
# - We already solved checkpoint *key mapping* and can load official weights.
# - Remaining output drift vs the official tokenizer is largely decoder-side.
# - FastVideo's `AutoencoderKLWan` uses a different temporal upsample path
# (`DupUp3D` + `first_chunk` slicing), while the official tokenizer uses
# `Resample(mode="upsample3d")` with a time-conv + interleave reshape.
#
# This section ports the core modules (CausalConv3d/Resample/etc.) so we can run
# a VAE that is behaviorally closer to the official implementation WITHOUT
# importing any official repo classes at runtime.
#
CACHE_T = 2
class Cosmos25CausalConv3d(nn.Conv3d):
"""
Official-like causal 3D convolution.
Matches `CausalConv3d` in the official tokenizer: uses explicit F.pad and
supports a `cache_x` prefix for causal chunking.
"""
def __init__(self, *args: Any, **kwargs: Any) -> None:
super().__init__(*args, **kwargs)
# padding order for F.pad: (W_left, W_right, H_left, H_right, T_left, T_right)
self._padding: tuple[int, ...] = (
self.padding[2],
self.padding[2],
self.padding[1],
self.padding[1],
2 * self.padding[0],
0,
)
self.padding = (0, 0, 0)
def forward(self, x: torch.Tensor, cache_x: torch.Tensor | None = None) -> torch.Tensor:
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
return super().forward(x)
class Cosmos25RMSNorm(nn.Module):
"""Official-like RMS_norm (uses learnable gamma and optional bias)."""
def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None:
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
def forward(self, x: torch.Tensor) -> torch.Tensor:
dim = 1 if self.channel_first else -1
return F.normalize(x, dim=dim) * self.scale * self.gamma + self.bias
class Cosmos25Upsample(nn.Upsample):
"""Official-like Upsample that is safe for bf16 (casts to fp32 internally)."""
def forward(self, x: torch.Tensor) -> torch.Tensor: # type: ignore[override]
return super().forward(x.float()).type_as(x)
class Cosmos25Resample(nn.Module):
"""
Official-like Resample used for both spatial and temporal up/downsampling.
"""
def __init__(self, dim: int, mode: str) -> None:
assert mode in ("none", "upsample2d", "upsample3d", "downsample2d", "downsample3d")
super().__init__()
self.dim = dim
self.mode = mode
if mode == "upsample2d":
self.resample = nn.Sequential(
Cosmos25Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
nn.Conv2d(dim, dim // 2, 3, padding=1),
)
elif mode == "upsample3d":
self.resample = nn.Sequential(
Cosmos25Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
nn.Conv2d(dim, dim // 2, 3, padding=1),
)
self.time_conv = Cosmos25CausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
elif mode == "downsample2d":
self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
elif mode == "downsample3d":
self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
self.time_conv = Cosmos25CausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
else:
self.resample = nn.Identity()
def forward(self, x: torch.Tensor, feat_cache: list[Any] | None = None, feat_idx: list[int] = [0]) -> torch.Tensor:
b, c, t, h, w = x.size()
# Temporal upsample uses a time-conv and then interleaves frames.
if self.mode == "upsample3d" and feat_cache is not None:
idx = feat_idx[0]
if feat_cache[idx] is None:
feat_cache[idx] = "Rep"
feat_idx[0] += 1
else:
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] != "Rep":
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] == "Rep":
cache_x = torch.cat([torch.zeros_like(cache_x).to(cache_x.device), cache_x], dim=2)
if feat_cache[idx] == "Rep":
x = self.time_conv(x)
else:
x = self.time_conv(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
x = x.reshape(b, c, t * 2, h, w)
t = x.shape[2]
x = rearrange(x, "b c t h w -> (b t) c h w")
x = self.resample(x)
x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
# Temporal downsample: time_conv consumes last-frame cache.
if self.mode == "downsample3d" and feat_cache is not None:
idx = feat_idx[0]
if feat_cache[idx] is None:
feat_cache[idx] = x.clone()
feat_idx[0] += 1
else:
cache_x = x[:, :, -1:, :, :].clone()
x = self.time_conv(torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
feat_cache[idx] = cache_x
feat_idx[0] += 1
return x
class Cosmos25ResidualBlock(nn.Module):
def __init__(self, in_dim: int, out_dim: int, dropout: float = 0.0) -> None:
super().__init__()
self.in_dim = in_dim
self.out_dim = out_dim
self.residual = nn.Sequential(
Cosmos25RMSNorm(in_dim, images=False),
nn.SiLU(),
Cosmos25CausalConv3d(in_dim, out_dim, 3, padding=1),
Cosmos25RMSNorm(out_dim, images=False),
nn.SiLU(),
nn.Dropout(dropout),
Cosmos25CausalConv3d(out_dim, out_dim, 3, padding=1),
)
self.shortcut = Cosmos25CausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity()
def forward(self, x: torch.Tensor, feat_cache: list[Any] | None = None, feat_idx: list[int] = [0]) -> torch.Tensor:
h = self.shortcut(x)
for layer in self.residual:
if isinstance(layer, Cosmos25CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x + h
class Cosmos25AttentionBlock(nn.Module):
"""Official-like causal self-attention with a single head."""
def __init__(self, dim: int) -> None:
super().__init__()
self.dim = dim
self.norm = Cosmos25RMSNorm(dim)
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
self.proj = nn.Conv2d(dim, dim, 1)
nn.init.zeros_(self.proj.weight)
def forward(self, x: torch.Tensor) -> torch.Tensor:
identity = x
b, c, t, h, w = x.size()
x2 = rearrange(x, "b c t h w -> (b t) c h w")
x2 = self.norm(x2)
q, k, v = (
self.to_qkv(x2)
.reshape(b * t, 1, c * 3, -1)
.permute(0, 1, 3, 2)
.contiguous()
.chunk(3, dim=-1)
)
x2 = F.scaled_dot_product_attention(q, k, v)
x2 = x2.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
x2 = self.proj(x2)
x2 = rearrange(x2, "(b t) c h w-> b c t h w", t=t)
return x2 + identity
class Cosmos25Encoder3d(nn.Module):
def __init__(
self,
dim: int = 96,
z_dim: int = 32,
dim_mult: list[int] = [1, 2, 4, 4],
num_res_blocks: int = 2,
attn_scales: list[float] = [],
temperal_downsample: list[bool] = [False, True, True],
dropout: float = 0.0,
) -> None:
super().__init__()
dims = [dim * u for u in [1] + dim_mult]
scale = 1.0
self.conv1 = Cosmos25CausalConv3d(3, dims[0], 3, padding=1)
downsamples: list[nn.Module] = []
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
for _ in range(num_res_blocks):
downsamples.append(Cosmos25ResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
downsamples.append(Cosmos25AttentionBlock(out_dim))
in_dim = out_dim
if i != len(dim_mult) - 1:
mode = "downsample3d" if temperal_downsample[i] else "downsample2d"
downsamples.append(Cosmos25Resample(out_dim, mode=mode))
scale /= 2.0
self.downsamples = nn.Sequential(*downsamples)
self.middle = nn.Sequential(
Cosmos25ResidualBlock(out_dim, out_dim, dropout),
Cosmos25AttentionBlock(out_dim),
Cosmos25ResidualBlock(out_dim, out_dim, dropout),
)
self.head = nn.Sequential(
Cosmos25RMSNorm(out_dim, images=False),
nn.SiLU(),
Cosmos25CausalConv3d(out_dim, z_dim, 3, padding=1),
)
def forward(self, x: torch.Tensor, feat_cache: list[Any] | None = None, feat_idx: list[int] = [0]) -> torch.Tensor:
if feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = self.conv1(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = self.conv1(x)
for layer in self.downsamples:
if feat_cache is not None:
x = layer(x, feat_cache, feat_idx) # type: ignore[misc]
else:
x = layer(x) # type: ignore[misc]
for layer in self.middle:
if isinstance(layer, Cosmos25ResidualBlock) and feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x) # type: ignore[misc]
for layer in self.head:
if isinstance(layer, Cosmos25CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x) # type: ignore[misc]
return x
class Cosmos25Decoder3d(nn.Module):
def __init__(
self,
dim: int = 96,
z_dim: int = 16,
dim_mult: list[int] = [1, 2, 4, 4],
num_res_blocks: int = 2,
attn_scales: list[float] = [],
temperal_upsample: list[bool] = [False, True, True],
dropout: float = 0.0,
) -> None:
super().__init__()
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
scale = 1.0 / 2 ** (len(dim_mult) - 2)
self.conv1 = Cosmos25CausalConv3d(z_dim, dims[0], 3, padding=1)
self.middle = nn.Sequential(
Cosmos25ResidualBlock(dims[0], dims[0], dropout),
Cosmos25AttentionBlock(dims[0]),
Cosmos25ResidualBlock(dims[0], dims[0], dropout),
)
upsamples: list[nn.Module] = []
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
if i in (1, 2, 3):
in_dim = in_dim // 2
for _ in range(num_res_blocks + 1):
upsamples.append(Cosmos25ResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
upsamples.append(Cosmos25AttentionBlock(out_dim))
in_dim = out_dim
if i != len(dim_mult) - 1:
mode = "upsample3d" if temperal_upsample[i] else "upsample2d"
upsamples.append(Cosmos25Resample(out_dim, mode=mode))
scale *= 2.0
self.upsamples = nn.Sequential(*upsamples)
self.head = nn.Sequential(
Cosmos25RMSNorm(out_dim, images=False),
nn.SiLU(),
Cosmos25CausalConv3d(out_dim, 3, 3, padding=1),
)
def forward(self, x: torch.Tensor, feat_cache: list[Any] | None = None, feat_idx: list[int] = [0]) -> torch.Tensor:
if feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = self.conv1(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = self.conv1(x)
for layer in self.middle:
if isinstance(layer, Cosmos25ResidualBlock) and feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x) # type: ignore[misc]
for layer in self.upsamples:
if feat_cache is not None:
x = layer(x, feat_cache, feat_idx) # type: ignore[misc]
else:
x = layer(x) # type: ignore[misc]
for layer in self.head:
if isinstance(layer, Cosmos25CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x) # type: ignore[misc]
return x
def _count_cosmos25_conv3d(model: nn.Module) -> int:
return sum(1 for m in model.modules() if isinstance(m, Cosmos25CausalConv3d))
class Cosmos25WanVAE(nn.Module):
"""
A FastVideo-native copy of the *official-like* Wan2.1 VAE core.
Key properties:
- Module naming matches official tokenizer (`encoder`, `decoder`, `conv1`, `conv2`)
so it can consume `tokenizer.pth` keys directly.
- `encode()` returns **normalized** latents and `decode()` expects **normalized**
latents (Cosmos2.5 contract), matching `Wan2pt1VAEInterface`.
"""
handles_latent_norm: bool = True
handles_latent_denorm: bool = True
def __init__(
self,
*,
device: torch.device | str = "cpu",
dtype: torch.dtype = torch.float32,
temporal_window: int = 4,
latents_mean: Optional[torch.Tensor] = None,
latents_std: Optional[torch.Tensor] = None,
) -> None:
super().__init__()
# Official hyperparams for Cosmos2.5 tokenizer (Wan2.1 VAE).
cfg = dict(
dim=96,
z_dim=16,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0,
temporal_window=temporal_window,
)
self.z_dim = 16
self.temporal_window = temporal_window
self.encoder = Cosmos25Encoder3d(
dim=cfg["dim"],
z_dim=cfg["z_dim"] * 2,
dim_mult=cfg["dim_mult"],
num_res_blocks=cfg["num_res_blocks"],
attn_scales=cfg["attn_scales"],
temperal_downsample=cfg["temperal_downsample"],
dropout=cfg["dropout"],
)
self.conv1 = Cosmos25CausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
self.conv2 = Cosmos25CausalConv3d(self.z_dim, self.z_dim, 1)
self.decoder = Cosmos25Decoder3d(
dim=cfg["dim"],
z_dim=cfg["z_dim"],
dim_mult=cfg["dim_mult"],
num_res_blocks=cfg["num_res_blocks"],
attn_scales=cfg["attn_scales"],
temperal_upsample=list(cfg["temperal_downsample"])[::-1],
dropout=cfg["dropout"],
)
# Default Cosmos2.5 latent stats (shared with configs).
if latents_mean is None:
latents_mean = torch.tensor(
[
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
],
dtype=torch.float32,
).view(1, 16, 1, 1, 1)
if latents_std is None:
latents_std = torch.tensor(
[
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
],
dtype=torch.float32,
).view(1, 16, 1, 1, 1)
self.register_buffer("_latents_mean", latents_mean, persistent=False)
self.register_buffer("_latents_std", latents_std, persistent=False)
self.to(device=device, dtype=dtype)
self.clear_cache()
def clear_cache(self) -> None:
# Decoder cache
self._conv_num = _count_cosmos25_conv3d(self.decoder)
self._conv_idx = [0]
self._feat_map: list[Any] = [None] * self._conv_num
# Encoder cache
self._enc_conv_num = _count_cosmos25_conv3d(self.encoder)
self._enc_conv_idx = [0]
self._enc_feat_map: list[Any] = [None] * self._enc_conv_num
def _scale(self, like: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
mean = self._latents_mean.to(device=like.device, dtype=like.dtype)
std = self._latents_std.to(device=like.device, dtype=like.dtype)
return mean, 1.0 / std
def _i0_encode(self, x: torch.Tensor) -> torch.Tensor:
return self.encoder(x[:, :, :1, :, :], feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx)
def _i0_decode(self, x: torch.Tensor) -> torch.Tensor:
return self.decoder(x[:, :, 0:1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx)
def encode(self, x: torch.Tensor) -> _TensorLatentDist:
"""
Encode to *normalized* latents (Cosmos contract).
"""
self.clear_cache()
t = x.shape[2]
iters = 1 + (t - 1) // self.temporal_window
for i in range(iters):
self._enc_conv_idx = [0]
if i == 0:
out = self._i0_encode(x)
else:
out_ = self.encoder(
x[:, :, 1 + self.temporal_window * (i - 1) : 1 + self.temporal_window * i, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx,
)
out = torch.cat([out, out_], 2)
if (t - 1) % self.temporal_window:
self._enc_conv_idx = [0]
out_ = self.encoder(
x[:, :, 1 + self.temporal_window * (iters - 1) :, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx,
)
out = torch.cat([out, out_], 2)
mu, _log_var = self.conv1(out).chunk(2, dim=1)
mean, inv_std = self._scale(mu)
z_norm = (mu - mean) * inv_std
self.clear_cache()
return _TensorLatentDist(z_norm)
def decode(self, latent: torch.Tensor) -> torch.Tensor:
"""
Decode from *normalized* latents (Cosmos contract).
"""
self.clear_cache()
mean, inv_std = self._scale(latent)
z = latent / inv_std + mean # z = z_norm * std + mean
iter_ = z.shape[2]
x = self.conv2(z)
for i in range(iter_):
self._conv_idx = [0]
if i == 0:
out = self._i0_decode(x)
else:
out_ = self.decoder(x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2)
self.clear_cache()
return out
# --- Interface helpers (match official Wan2pt1VAEInterface) ---
def get_latent_num_frames(self, num_pixel_frames: int) -> int:
return 1 + (int(num_pixel_frames) - 1) // 4
def get_pixel_num_frames(self, num_latent_frames: int) -> int:
return (int(num_latent_frames) - 1) * 4 + 1
@property
def spatial_compression_factor(self) -> int:
return 8
@property
def temporal_compression_factor(self) -> int:
return 4
@property
def latent_ch(self) -> int:
return 16
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos 2.5 pipeline entry (staged pipeline)."""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage,
Cosmos25DenoisingStage,
Cosmos25LatentPreparationStage,
DecodingStage, InputValidationStage,
Cosmos25TextEncodingStage,
Cosmos25TimestepPreparationStage)
logger = init_logger(__name__)
class Cosmos2_5Pipeline(ComposedPipelineBase):
"""Cosmos 2.5 video generation pipeline."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler",
"safety_checker"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
logger.info("Creating Cosmos 2.5 pipeline stages...")
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(
stage_name="prompt_encoding_stage",
stage=Cosmos25TextEncodingStage(
text_encoder=self.get_module("text_encoder"), ),
)
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=Cosmos25TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=Cosmos25LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=Cosmos25DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info("Cosmos 2.5 pipeline stages created")
# Entry point for pipeline registry
EntryClass = Cosmos2_5Pipeline
@@ -201,8 +201,6 @@ class ComposedPipelineBase(ABC):
# fwd, bwd, and other operations' precision.
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
pipe = cls(model_path,
fastvideo_args,
required_config_modules=required_config_modules,
+274 -93
View File
@@ -1,7 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
from collections import defaultdict
from collections.abc import Hashable
from contextlib import nullcontext
from typing import Any
from collections.abc import Generator
import torch
import torch.distributed as dist
@@ -12,8 +14,13 @@ from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.layers.lora.linear import (BaseLayerWithLoRA, get_lora_layer,
replace_submodule)
from fastvideo.hooks.hooks import ModuleHookManager
from fastvideo.hooks.layerwise_offload import LayerwiseOffloadHook
from fastvideo.layers.lora.linear import (
BaseLayerWithLoRA,
get_lora_layer,
replace_submodule,
)
from fastvideo.logger import init_logger
from fastvideo.models.loader.utils import get_param_names_mapping
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -22,17 +29,89 @@ from fastvideo.utils import maybe_download_lora
logger = init_logger(__name__)
def _get_hook_ctx(module: nn.Module | None):
if module is None:
return nullcontext()
hook_mgr = ModuleHookManager.get_from(module)
if hook_mgr is not None:
offload_hook = hook_mgr.forward_hooks.get(LayerwiseOffloadHook.name())
if offload_hook is not None:
return offload_hook.mutate_params_scope() # type: ignore
return nullcontext()
def _named_module_by_prefix(
module: nn.Module, prefixes: list[str]
) -> list[tuple[str | None, list[tuple[str, nn.Module]]]]:
none_list: list[tuple[str, nn.Module]] = []
prefix_list: list[tuple[str, list[tuple[str, nn.Module]]]] = [
(prefix, []) for prefix in prefixes
]
for name, submodule in module.named_modules():
for cur_prefix, cur_list in prefix_list:
# we should exclude e.g. block.1 and block.12.attn
if name.startswith(cur_prefix + "."):
cur_list.append((name, submodule))
break
else:
none_list.append((name, submodule))
return prefix_list + [(None, none_list)] # type: ignore
class LoRAModelLayers:
def __init__(self, block_list: list[tuple[str, nn.Module]]) -> None:
# block_name -> {layer_name -> layer}
self.block_to_lora_layers: dict[str, dict[str, BaseLayerWithLoRA]] = {}
# layer_name -> block_name
self.lora_layers_to_block: dict[str, str | None] = {}
self.other_lora_layers: dict[str, BaseLayerWithLoRA] = {}
self.block_mapping = dict(block_list)
def add_lora_layer(self, block_name: str | None, layer_name: str,
layer: BaseLayerWithLoRA):
if block_name is None:
self.other_lora_layers[layer_name] = layer
self.lora_layers_to_block[layer_name] = None
else:
if block_name not in self.block_to_lora_layers:
self.block_to_lora_layers[block_name] = {}
self.block_to_lora_layers[block_name][layer_name] = layer
self.lora_layers_to_block[layer_name] = block_name
def all_lora_layers(
self, ) -> Generator[tuple[str, BaseLayerWithLoRA], Any, None]:
for block_layers in self.block_to_lora_layers.values():
for name, layer in block_layers.items():
yield name, layer
for name, layer in self.other_lora_layers.items():
yield name, layer
def lora_layers_by_block(
self,
) -> Generator[
tuple[nn.Module | None, dict[str, BaseLayerWithLoRA]],
Any,
None,
]:
for block_name, layers in self.block_to_lora_layers.items():
yield self.block_mapping[block_name], layers
yield None, self.other_lora_layers
class LoRAPipeline(ComposedPipelineBase):
"""
Pipeline that supports injecting LoRA adapters into the diffusion transformer.
TODO: support training.
"""
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
dict
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
cur_adapter_name: str = ""
cur_adapter_path: str = ""
lora_layers: dict[str, dict[str, BaseLayerWithLoRA]] = {}
# model_name -> layers
lora_layers: dict[str, LoRAModelLayers] = {}
fastvideo_args: FastVideoArgs | TrainingArgs
exclude_lora_layers: dict[str, list[str]] = {}
device: torch.device = get_local_torch_device()
@@ -48,10 +127,10 @@ class LoRAPipeline(ComposedPipelineBase):
self.device = get_local_torch_device()
# build list of trainable transformers
for transformer_name in self.trainable_transformer_names:
if transformer_name in self.modules and self.modules[
transformer_name] is not None:
self.trainable_transformer_modules[
transformer_name] = self.modules[transformer_name]
if (transformer_name in self.modules
and self.modules[transformer_name] is not None):
self.trainable_transformer_modules[transformer_name] = (
self.modules[transformer_name])
# check for transformer_2 in case of Wan2.2 MoE or fake_score_transformer_2
if transformer_name.endswith("_2"):
raise ValueError(
@@ -59,19 +138,23 @@ class LoRAPipeline(ComposedPipelineBase):
)
secondary_transformer_name = transformer_name + "_2"
if secondary_transformer_name in self.modules and self.modules[
secondary_transformer_name] is not None:
if (secondary_transformer_name in self.modules
and self.modules[secondary_transformer_name] is not None):
self.trainable_transformer_modules[
secondary_transformer_name] = self.modules[
secondary_transformer_name]
logger.info("trainable_transformer_modules: %s",
self.trainable_transformer_modules.keys())
logger.info(
"trainable_transformer_modules: %s",
self.trainable_transformer_modules.keys(),
)
for transformer_name, transformer_module in self.trainable_transformer_modules.items(
):
self.exclude_lora_layers[
transformer_name] = transformer_module.config.arch_config.exclude_lora_layers
for (
transformer_name,
transformer_module,
) in self.trainable_transformer_modules.items():
self.exclude_lora_layers[transformer_name] = (
transformer_module.config.arch_config.exclude_lora_layers)
self.lora_target_modules = self.fastvideo_args.lora_target_modules
self.lora_path = self.fastvideo_args.lora_path
self.lora_nickname = self.fastvideo_args.lora_nickname
@@ -83,20 +166,33 @@ class LoRAPipeline(ComposedPipelineBase):
self.fastvideo_args.lora_alpha = self.fastvideo_args.lora_rank
self.lora_rank = self.fastvideo_args.lora_rank # type: ignore
self.lora_alpha = self.fastvideo_args.lora_alpha # type: ignore
logger.info("Using LoRA training with rank %d and alpha %d",
self.lora_rank, self.lora_alpha)
logger.info(
"Using LoRA training with rank %d and alpha %d",
self.lora_rank,
self.lora_alpha,
)
if self.lora_target_modules is None:
self.lora_target_modules = [
"q_proj", "k_proj", "v_proj", "o_proj", "to_q", "to_k",
"to_v", "to_out", "to_qkv", "to_gate_compress"
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"to_q",
"to_k",
"to_v",
"to_out",
"to_qkv",
"to_gate_compress",
]
logger.info(
"Using default lora_target_modules for all transformers: %s",
self.lora_target_modules)
self.lora_target_modules,
)
else:
logger.warning(
"Using custom lora_target_modules for all transformers, which may not be intended: %s",
self.lora_target_modules)
self.lora_target_modules,
)
self.convert_to_lora_layers()
# Inference
@@ -104,7 +200,8 @@ class LoRAPipeline(ComposedPipelineBase):
self.convert_to_lora_layers()
self.set_lora_adapter(
self.lora_nickname, # type: ignore
self.lora_path) # type: ignore
self.lora_path,
) # type: ignore
def is_target_layer(self, module_name: str) -> bool:
if self.lora_target_modules is None:
@@ -114,9 +211,9 @@ class LoRAPipeline(ComposedPipelineBase):
def set_trainable(self) -> None:
def set_lora_grads(lora_layers: dict[str, BaseLayerWithLoRA],
def set_lora_grads(lora_layers: LoRAModelLayers,
device_mesh: DeviceMesh):
for name, layer in lora_layers.items():
for name, layer in lora_layers.all_lora_layers():
layer.lora_A.requires_grad_(True)
layer.lora_B.requires_grad_(True)
layer.base_layer.requires_grad_(False)
@@ -131,10 +228,15 @@ class LoRAPipeline(ComposedPipelineBase):
super().set_trainable()
return
device_mesh = init_device_mesh("cuda", (dist.get_world_size(), 1),
mesh_dim_names=["fake", "replicate"])
for transformer_name, transformer_module in self.trainable_transformer_modules.items(
):
device_mesh = init_device_mesh(
"cuda",
(dist.get_world_size(), 1),
mesh_dim_names=["fake", "replicate"],
)
for (
transformer_name,
transformer_module,
) in self.trainable_transformer_modules.items():
transformer_module.train()
transformer_module.requires_grad_(False)
if transformer_name in self.lora_layers:
@@ -151,32 +253,71 @@ class LoRAPipeline(ComposedPipelineBase):
if self.lora_initialized:
return
self.lora_initialized = True
for transformer_name, transformer_module in self.trainable_transformer_modules.items(
):
for (
transformer_name,
transformer_module,
) in self.trainable_transformer_modules.items():
converted_count = 0
# init bookkeeping structures
if transformer_name not in self.lora_layers:
self.lora_layers[transformer_name] = {}
logger.info("Converting %s to LoRA Transformer", transformer_name)
for name, layer in transformer_module.named_modules():
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers[transformer_name]:
if exclude_layer in name:
excluded = True
# get block list
block_list = []
for name, submodule in transformer_module.named_children():
if isinstance(submodule, nn.ModuleList):
block_list = [(f"{name}.{i}", m)
for i, m in enumerate(submodule)]
break
if excluded:
continue
self.lora_layers[transformer_name] = LoRAModelLayers(block_list)
logger.info("Converting %s to LoRA Transformer", transformer_name)
# scan every module and convert to LoRA layer if applicable
layer = get_lora_layer(layer,
lora_rank=self.lora_rank,
lora_alpha=self.lora_alpha,
training_mode=self.training_mode)
if layer is not None:
self.lora_layers[transformer_name][name] = layer
replace_submodule(transformer_module, name, layer)
converted_count += 1
for block_name, block_modules in _named_module_by_prefix(
transformer_module,
list(self.lora_layers[transformer_name].block_mapping),
):
if block_name is not None and (
not self.fastvideo_args.training_mode
and self.fastvideo_args.dit_layerwise_offload):
scope_ctx = _get_hook_ctx(
self.lora_layers[transformer_name].
block_mapping[block_name])
else:
scope_ctx = nullcontext()
with scope_ctx:
for name, layer in block_modules:
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers[
transformer_name]:
if exclude_layer in name:
excluded = True
break
if excluded:
continue
layer = get_lora_layer(
layer,
lora_rank=self.lora_rank,
lora_alpha=self.lora_alpha,
training_mode=self.training_mode,
)
if layer is not None:
block_name_split = name.split(".", 2)
if len(block_name_split) > 2:
block_name = (block_name_split[0] + "." +
block_name_split[1])
else:
block_name = None
if (block_name
not in self.lora_layers[transformer_name].
block_mapping):
block_name = None
self.lora_layers[transformer_name].add_lora_layer(
block_name, name, layer)
replace_submodule(transformer_module, name, layer)
converted_count += 1
logger.info("Converted %d layers to LoRA layers", converted_count)
def set_lora_adapter(self,
@@ -209,7 +350,7 @@ class LoRAPipeline(ComposedPipelineBase):
# Extract alpha values and weights in a single pass
to_merge_params: defaultdict[Hashable,
dict[Any, Any]] = defaultdict(dict)
dict[Any, Any]] = (defaultdict(dict))
for name, weight in lora_state_dict.items():
# Extract weights (lora_A, lora_B, and lora_alpha)
name = name.replace("diffusion_model.", "")
@@ -223,13 +364,14 @@ class LoRAPipeline(ComposedPipelineBase):
target_name, _, _ = param_names_mapping_fn(layer_name)
# Store alpha alongside weights with same target_name base
alpha_key = target_name + ".lora_alpha"
self.lora_adapters[lora_nickname][alpha_key] = weight.item(
) if weight.numel() == 1 else float(weight.mean())
self.lora_adapters[lora_nickname][alpha_key] = (
weight.item()
if weight.numel() == 1 else float(weight.mean()))
continue
name, _, _ = lora_param_names_mapping_fn(name)
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
name)
target_name, merge_index, num_params_to_merge = (
param_names_mapping_fn(name))
# for (in_dim, r) @ (r, out_dim), we only merge (r, out_dim * n) where n is the number of linear layers to fuse
# see param mapping in HunyuanVideoArchConfig
if merge_index is not None and "lora_B" in name:
@@ -261,45 +403,84 @@ class LoRAPipeline(ComposedPipelineBase):
# Merge the new adapter
adapted_count = 0
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = self.lora_adapters[lora_nickname].get(
lora_alpha_name) if adapter_updated else None
layer.set_lora_weights(
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path, name)
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
for (
transformer_name,
transformer_lora_layers,
) in self.lora_layers.items():
for (
module,
layers,
) in transformer_lora_layers.lora_layers_by_block():
with _get_hook_ctx(module):
for name, layer in layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if (lora_A_name in self.lora_adapters[lora_nickname]
and lora_B_name
in self.lora_adapters[lora_nickname]):
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][
lora_A_name]
lora_B = self.lora_adapters[lora_nickname][
lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = (self.lora_adapters[lora_nickname].get(
lora_alpha_name) if adapter_updated else None)
try:
layer.set_lora_weights(
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.
training_mode,
lora_path=lora_path,
)
except Exception as e:
logger.error(
"Error setting LoRA weights for layer %s: %s",
name,
str(e),
)
raise e
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path,
name,
)
layer.disable_lora = True
logger.info(
"Rank %d: LoRA adapter %s applied to %d layers",
rank,
lora_path,
adapted_count,
)
def merge_lora_weights(self) -> None:
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
layer.merge_lora_weights()
for (
transformer_name,
transformer_lora_layers,
) in self.lora_layers.items():
for (
module,
layers,
) in transformer_lora_layers.lora_layers_by_block():
with _get_hook_ctx(module):
for name, layer in layers.items():
layer.merge_lora_weights()
def unmerge_lora_weights(self) -> None:
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
layer.unmerge_lora_weights()
for (
transformer_name,
transformer_lora_layers,
) in self.lora_layers.items():
for (
module,
layers,
) in transformer_lora_layers.lora_layers_by_block():
with _get_hook_ctx(module):
for name, layer in layers.items():
layer.unmerge_lora_weights()
@@ -67,6 +67,21 @@ class ForwardBatch:
execution, allowing methods to update specific components without needing
to manage numerous individual parameters.
"""
@dataclass
class RLData:
"""RL-specific data collection options and outputs."""
enabled: bool = False
collect_log_probs: bool = True
collect_kl: bool = False
kl_reward: float = 0.0
store_trajectory: bool = True
keep_trajectory_on_cpu: bool = False
log_probs: torch.Tensor | None = None
kl: torch.Tensor | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_timesteps: torch.Tensor | None = None
# TODO(will): double check that args are separate from fastvideo_args
# properly. Also maybe think about providing an abstraction for pipeline
# specific arguments.
@@ -197,6 +212,9 @@ class ForwardBatch:
logging_info: PipelineLoggingInfo = field(
default_factory=PipelineLoggingInfo)
# RL data collection
rl_data: "ForwardBatch.RLData" = field(default_factory=RLData)
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
@@ -267,6 +285,36 @@ class TrainingBatch:
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
# RL/GRPO-specific attributes
reward_scores: torch.Tensor | None = None # Computed rewards from reward models
log_probs: torch.Tensor | None = None # Current policy log probabilities [B, num_steps] or [B]
old_log_probs: torch.Tensor | None = None # Old policy log probs (for importance ratio) [B, num_steps] or [B]
advantages: torch.Tensor | None = None # GAE advantages [B, num_steps] or [B]
returns: torch.Tensor | None = None # TD returns (advantages + values) [B, num_steps] or [B]
values: torch.Tensor | None = None # Value function predictions [B]
old_values: torch.Tensor | None = None # Old value predictions (for clipping) [B]
# GRPO sampling-specific attributes
kl: torch.Tensor | None = None # KL divergences from sampling [B, num_steps] (if kl_reward > 0)
prompt_ids: torch.Tensor | None = None # Prompt token IDs for stat tracking [B, seq_len]
prompt_embeds: torch.Tensor | None = None # Prompt embeddings used in sampling [B, seq_len, hidden_dim]
negative_prompt_embeds: torch.Tensor | None = None # Negative prompt embeddings for CFG [B, seq_len, hidden_dim]
# RL loss components
policy_loss: float = 0.0 # GRPO/PPO policy loss
value_loss: float = 0.0 # Value function loss
kl_divergence: float = 0.0 # KL(new_policy || old_policy)
importance_ratio: float = 1.0 # exp(log_prob - old_log_prob)
clip_fraction: float = 0.0 # Fraction of ratios that were clipped
# RL metrics
advantage_mean: float = 0.0 # Mean advantage (should be ~0 after normalization)
advantage_std: float = 1.0 # Std of advantages
reward_mean: float = 0.0 # Mean reward across batch
reward_std: float = 0.0 # Std of rewards
value_mean: float = 0.0 # Mean value prediction
entropy: float = 0.0 # Policy entropy (for exploration)
@dataclass
class PreprocessBatch(ForwardBatch):
+1
View File
@@ -29,6 +29,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"Cosmos2_5Pipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
"MatrixGameCausalDMDPipeline": "matrixgame",
"LongCatPipeline": "longcat",
+11 -4
View File
@@ -10,7 +10,8 @@ from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
CosmosDenoisingStage,
DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
@@ -19,14 +20,16 @@ from fastvideo.pipelines.stages.image_encoding import (
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import (
CosmosLatentPreparationStage, LatentPreparationStage)
Cosmos25LatentPreparationStage, CosmosLatentPreparationStage,
LatentPreparationStage)
from fastvideo.pipelines.stages.matrixgame_denoising import (
MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.text_encoding import (Cosmos25TextEncodingStage,
TextEncodingStage)
from fastvideo.pipelines.stages.timestep_preparation import (
TimestepPreparationStage)
Cosmos25TimestepPreparationStage, TimestepPreparationStage)
# LongCat stages
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
@@ -37,14 +40,17 @@ __all__ = [
"PipelineStage",
"InputValidationStage",
"TimestepPreparationStage",
"Cosmos25TimestepPreparationStage",
"LatentPreparationStage",
"CosmosLatentPreparationStage",
"Cosmos25LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"MatrixGameCausalDenoisingStage",
"CosmosDenoisingStage",
"Cosmos25DenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
@@ -54,6 +60,7 @@ __all__ = [
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
"TextEncodingStage",
"Cosmos25TextEncodingStage",
"StepvideoPromptEncodingStage",
# LongCat stages
"LongCatVideoVAEEncodingStage",
+22 -19
View File
@@ -51,39 +51,42 @@ class DecodingStage(PipelineStage):
return result
def _denormalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
# denormalization for MatrixGame VAE
# z = z * std + mean during decode
if (hasattr(self.vae.config, 'latents_mean')
and hasattr(self.vae.config, 'latents_std')):
# Convert config values to tensors
latents_mean = torch.tensor(self.vae.config.latents_mean,
"""Convert normalized latents into the VAE's expected latent space."""
# Some VAEs handle latent (de)normalization internally.
if bool(getattr(self.vae, "handles_latent_denorm", False)):
return latents
cfg = getattr(self.vae, "config", None)
# MatrixGame-style: z = z * std + mean
if (cfg is not None and hasattr(cfg, "latents_mean")
and hasattr(cfg, "latents_std")):
latents_mean = torch.tensor(cfg.latents_mean,
device=latents.device,
dtype=latents.dtype).view(
1, -1, 1, 1, 1)
latents_std = torch.tensor(self.vae.config.latents_std,
latents_std = torch.tensor(cfg.latents_std,
device=latents.device,
dtype=latents.dtype).view(
1, -1, 1, 1, 1)
return latents * latents_std + latents_mean
# Apply denormalization: z = z * std + mean
latents = latents * latents_std + latents_mean
elif hasattr(self.vae, 'scaling_factor'):
# Standard VAE scaling
# Diffusers-style: scaling_factor (+ optional shift_factor)
if hasattr(self.vae, "scaling_factor"):
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 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)
latents = latents + self.vae.shift_factor.to(
latents.device, latents.dtype)
else:
latents += self.vae.shift_factor
latents = latents + self.vae.shift_factor
return latents
@torch.no_grad()
@@ -273,4 +276,4 @@ class DecodingStage(PipelineStage):
del pipeline.modules["vae"]
fastvideo_args.model_loaded["vae"] = False
return batch
return batch
+327 -33
View File
@@ -4,11 +4,14 @@ Denoising stage for diffusion pipelines.
"""
import inspect
import math
import weakref
from collections.abc import Iterable
from contextlib import nullcontext
from typing import Any
import torch
from diffusers.utils.torch_utils import randn_tensor
from tqdm.auto import tqdm
from fastvideo.attention import get_attn_backend
@@ -52,6 +55,84 @@ except ImportError:
logger = init_logger(__name__)
def sde_step_with_logprob(
scheduler,
model_output: torch.FloatTensor,
timestep: float | torch.FloatTensor,
sample: torch.FloatTensor,
prev_sample: torch.FloatTensor | None = None,
generator: torch.Generator | None = None,
deterministic: bool = False,
return_pixel_log_prob: bool = False,
return_dt_and_std_dev_t: bool = False
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ...]:
"""
Predict the sample from the previous timestep by reversing the SDE and
compute log probabilities for the transition.
"""
if isinstance(timestep, torch.Tensor):
if timestep.ndim == 0:
timestep = timestep.unsqueeze(0)
step_indices = [
scheduler.index_for_timestep(t.item()) for t in timestep
]
else:
step_indices = [scheduler.index_for_timestep(timestep)]
prev_step_indices = [step + 1 for step in step_indices]
sigmas = scheduler.sigmas.to(sample.device, sample.dtype)
sigma = sigmas[step_indices].view(-1, 1, 1, 1, 1)
sigma_prev = sigmas[prev_step_indices].view(-1, 1, 1, 1, 1)
sigma_max = sigmas[0].item()
sigma_min = sigmas[-1].item()
dt = sigma_prev - sigma
std_dev_t = sigma_min + (sigma_max - sigma_min) * sigma
prev_sample_mean = (sample * (1 + std_dev_t**2 / (2 * sigma) * dt) +
model_output * (1 + std_dev_t**2 * (1 - sigma) /
(2 * sigma)) * dt)
if prev_sample is not None and generator is not None:
raise ValueError(
"Cannot pass both generator and prev_sample. Please make sure that either `generator` or"
" `prev_sample` stays `None`.")
if prev_sample is None:
variance_noise = randn_tensor(
model_output.shape,
generator=generator,
device=model_output.device,
dtype=model_output.dtype,
)
sqrt_dt = torch.sqrt(-1 * dt)
prev_sample = prev_sample_mean + std_dev_t * sqrt_dt * variance_noise
else:
sqrt_dt = torch.sqrt(-1 * dt)
if deterministic:
prev_sample = sample + dt * model_output
sqrt_dt = torch.sqrt(-1 * dt)
if return_pixel_log_prob:
raise NotImplementedError(
"Pixel-level log prob is not supported in this helper.")
std_dev_sqrt_dt = std_dev_t * sqrt_dt
log_prob = (
-((prev_sample.detach() - prev_sample_mean)**2) /
(2 *
(std_dev_sqrt_dt**2)) - torch.log(std_dev_sqrt_dt + 1e-8) - torch.log(
torch.sqrt(2 * torch.as_tensor(math.pi, device=sample.device))))
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
if return_dt_and_std_dev_t:
return prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt
return prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt
class DenoisingStage(PipelineStage):
"""
Stage for running the denoising loop in diffusion pipelines.
@@ -203,9 +284,11 @@ class DenoisingStage(PipelineStage):
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
assert latent_model_input.shape[0] == 1, "only support batch size 1"
rl_data = batch.rl_data if batch.rl_data and batch.rl_data.enabled else None
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
assert latent_model_input.shape[
0] == 1, "TI2V task only supports batch size 1"
# TI2V directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
@@ -243,6 +326,12 @@ class DenoisingStage(PipelineStage):
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
rl_timesteps: list[torch.Tensor] = []
rl_latents: list[torch.Tensor] = []
rl_log_probs: list[torch.Tensor] = []
rl_kl: list[torch.Tensor] = []
if rl_data is not None and rl_data.store_trajectory:
rl_latents.append(latents)
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
@@ -329,6 +418,24 @@ class DenoisingStage(PipelineStage):
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
def run_transformer(model, encoder_hidden_states, cond_kwargs,
is_cfg_negative: bool):
batch.is_cfg_negative = is_cfg_negative
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
):
return model(
latent_model_input,
encoder_hidden_states,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**cond_kwargs,
**action_kwargs,
)
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
@@ -390,40 +497,13 @@ class DenoisingStage(PipelineStage):
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# fastvideo_args or training_args, and attn_metadata.
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred = current_model(
latent_model_input,
prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
)
noise_pred = run_transformer(current_model, prompt_embeds,
pos_cond_kwargs, False)
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
):
noise_pred_uncond = current_model(
latent_model_input,
neg_prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
)
noise_pred_uncond = run_transformer(
current_model, neg_prompt_embeds, neg_cond_kwargs,
True)
noise_pred_text = noise_pred
noise_pred = noise_pred_uncond + current_guidance_scale * (
@@ -438,11 +518,58 @@ class DenoisingStage(PipelineStage):
guidance_rescale=batch.guidance_rescale,
)
# Compute the previous noisy sample
prev_latents = latents
latents = self.scheduler.step(noise_pred,
t,
latents,
**extra_step_kwargs,
return_dict=False)[0]
if rl_data is not None:
if rl_data.collect_log_probs:
_, log_prob, prev_latents_mean, std_dev_t, _ = sde_step_with_logprob(
self.scheduler,
noise_pred.float(),
t,
prev_latents.float(),
prev_sample=latents.float(),
deterministic=False,
return_dt_and_std_dev_t=True,
)
rl_log_probs.append(log_prob)
if rl_data.collect_kl and rl_data.kl_reward > 0:
adapter_ctx = nullcontext()
if hasattr(current_model, "disable_adapter"):
adapter_ctx = current_model.disable_adapter()
with adapter_ctx:
noise_pred_ref = run_transformer(
current_model, prompt_embeds,
pos_cond_kwargs, False)
if batch.do_classifier_free_guidance:
noise_pred_uncond_ref = run_transformer(
current_model, neg_prompt_embeds,
neg_cond_kwargs, True)
noise_pred_text_ref = noise_pred_ref
noise_pred_ref = noise_pred_uncond_ref + current_guidance_scale * (
noise_pred_text_ref -
noise_pred_uncond_ref)
_, _, prev_latents_mean_ref, std_dev_t_ref, _ = sde_step_with_logprob(
self.scheduler,
noise_pred_ref.float(),
t,
prev_latents.float(),
prev_sample=latents.float(),
deterministic=False,
return_dt_and_std_dev_t=True,
)
if not torch.allclose(std_dev_t, std_dev_t_ref):
logger.warning(
"std_dev_t mismatch in RL KL computation at step %s",
i)
kl = (prev_latents_mean -
prev_latents_mean_ref)**2 / (2 * std_dev_t**2)
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
rl_kl.append(kl)
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
latents = latents.squeeze(0)
latents = (1. - mask2[0]) * z + mask2[0] * latents
@@ -452,6 +579,15 @@ class DenoisingStage(PipelineStage):
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
trajectory_latents.append(latents)
if rl_data is not None:
rl_timesteps.append(t)
if rl_data.store_trajectory:
rl_latents.append(latents)
if rl_data.collect_kl and rl_data.kl_reward <= 0:
rl_kl.append(
torch.zeros(latents.shape[0],
device=latents.device,
dtype=latents.dtype))
# Update progress bar
if i == len(timesteps) - 1 or (
@@ -472,6 +608,25 @@ class DenoisingStage(PipelineStage):
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()
if rl_data is not None:
if rl_timesteps:
rl_data.trajectory_timesteps = torch.stack(rl_timesteps, dim=0)
if rl_data.keep_trajectory_on_cpu:
rl_data.trajectory_timesteps = rl_data.trajectory_timesteps.cpu(
)
if rl_data.store_trajectory and rl_latents:
rl_data.trajectory_latents = torch.stack(rl_latents, dim=1)
if rl_data.keep_trajectory_on_cpu:
rl_data.trajectory_latents = rl_data.trajectory_latents.cpu(
)
if rl_log_probs:
rl_data.log_probs = torch.stack(rl_log_probs, dim=1)
if rl_data.keep_trajectory_on_cpu:
rl_data.log_probs = rl_data.log_probs.cpu()
if rl_kl:
rl_data.kl = torch.stack(rl_kl, dim=1)
if rl_data.keep_trajectory_on_cpu:
rl_data.kl = rl_data.kl.cpu()
# Update batch with final latents
batch.latents = latents
@@ -1018,6 +1173,145 @@ class CosmosDenoisingStage(DenoisingStage):
return result
class Cosmos25DenoisingStage(CosmosDenoisingStage):
"""Denoising stage for Cosmos 2.5 DiT (expects 1D/2D timestep, not 5D)."""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
if hasattr(self.transformer, 'module'):
transformer_dtype = next(self.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latents = batch.latents
if latents is None:
raise ValueError(
"latents must be provided for Cosmos25DenoisingStage")
guidance_scale = batch.guidance_scale
# Use timesteps prepared by Cosmos25TimestepPreparationStage when available.
if batch.timesteps is None:
self.scheduler.set_timesteps(batch.num_inference_steps,
device=latents.device)
timesteps = self.scheduler.timesteps
else:
timesteps = batch.timesteps.to(latents.device)
# Match official behavior: pass fps as a tensor.
fps_val = batch.fps if isinstance(batch.fps, int | float) else 24
fps_tensor = torch.tensor([fps_val],
device=latents.device,
dtype=target_dtype)
# Cosmos2.5 denoises a 4D latent (C,T,H,W) and the scheduler.step expects (B,C,T,H,W).
latents_4d = latents[0]
# Masks from latent prep stage
condition_mask = batch.cond_mask.to(target_dtype) if hasattr(
batch, 'cond_mask') else None
padding_mask = batch.padding_mask.to(target_dtype) if hasattr(
batch, 'padding_mask') else None
if condition_mask is None:
_, t, h, w = latents_4d.shape
condition_mask = torch.zeros(1,
1,
t,
h,
w,
device=latents.device,
dtype=target_dtype)
if padding_mask is None:
_, _, h, w = latents_4d.shape
padding_mask = torch.ones(1,
1,
h,
w,
device=latents.device,
dtype=target_dtype)
# Cosmos2.5 timestep scaling (see compare_pipelines.py): t * 0.001
timestep_scale = 0.001
with self.progress_bar(total=len(timesteps)) as progress_bar:
for i, t in enumerate(timesteps):
t_val = float(t)
timestep_val = t_val * timestep_scale
timestep = torch.tensor([[timestep_val]],
device=latents.device,
dtype=target_dtype)
model_hidden_states = latents_4d.unsqueeze(0)
with (
set_forward_context(current_timestep=int(t_val),
attn_metadata=None,
forward_batch=batch),
torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled),
):
cond_v = self.transformer(
hidden_states=model_hidden_states.to(target_dtype),
encoder_hidden_states=batch.prompt_embeds[0].to(
target_dtype),
timestep=timestep,
fps=fps_tensor,
condition_mask=condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds:
uncond_v = self.transformer(
hidden_states=model_hidden_states.to(target_dtype),
encoder_hidden_states=batch.
negative_prompt_embeds[0].to(target_dtype),
timestep=timestep,
fps=fps_tensor,
condition_mask=condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
v = uncond_v + guidance_scale * (cond_v - uncond_v)
else:
v = cond_v
prev = self.scheduler.step(v.unsqueeze(0),
t,
latents_4d.unsqueeze(0),
**extra_step_kwargs,
return_dict=False)[0]
latents_4d = prev.squeeze(0)
progress_bar.update()
batch.latents = latents_4d.unsqueeze(0)
return batch
class DmdDenoisingStage(DenoisingStage):
"""
Denoising stage for DMD.
@@ -5,6 +5,7 @@ Latent preparation stage for diffusion pipelines.
from typing import Any
import numpy as np
import torch
from diffusers.utils.torch_utils import randn_tensor
@@ -427,6 +428,210 @@ class CosmosLatentPreparationStage(PipelineStage):
return batch
class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
"""Latent preparation for Cosmos 2.5 DiT input conventions."""
@staticmethod
def _arch_invariant_randn(
shape: tuple[int, ...],
*,
seed: int,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
"""Architecture-invariant RNG (matches cosmos_predict2.misc.arch_invariant_rand)."""
rng = np.random.RandomState(seed)
arr = rng.standard_normal(shape).astype(np.float32)
return torch.from_numpy(arr).to(device=device, dtype=dtype)
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
# Differences vs `CosmosLatentPreparationStage`: channel convention, seed usage,
# and `padding_mask` for concat_padding_mask=True.
# Determine batch size
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
else:
batch_size = batch.prompt_embeds[0].shape[0]
batch_size *= batch.num_videos_per_prompt
# Match `compare_pipelines.py`: initialize noise in fp32, then run the
# denoising computation in bf16.
dtype = torch.float32
device = get_local_torch_device()
generator = batch.generator
latents = batch.latents
num_frames = batch.num_frames
height = batch.height
width = batch.width
if height is None or width is None:
raise ValueError("Height and width must be provided")
vae_scale_factor_spatial = 8
vae_scale_factor_temporal = 4
latent_height = height // 8
latent_width = width // vae_scale_factor_spatial
num_latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
# Cosmos 2.5 convention: transformer in_channels == latent channels
num_channels_latents = self.transformer.config.in_channels
shape = (batch_size, num_channels_latents, num_latent_frames,
latent_height, latent_width)
init_latents = None
conditioning_latents = None
video = None
if hasattr(batch, 'video') and batch.video is not None:
video = batch.video
elif hasattr(batch, 'pil_image') and batch.pil_image is not None:
vae_scale_factor_spatial = 8
image_processor = ImageProcessor(
vae_scale_factor=vae_scale_factor_spatial)
processed_image = image_processor.preprocess(
batch.pil_image, height, width)
video = processed_image.unsqueeze(2)
video = video.to(device=device, dtype=torch.bfloat16)
elif hasattr(
batch,
'preprocessed_image') and batch.preprocessed_image is not None:
if isinstance(batch.preprocessed_image, torch.Tensor):
if batch.preprocessed_image.dim() == 4:
video = batch.preprocessed_image.unsqueeze(2)
elif batch.preprocessed_image.dim() == 5:
video = batch.preprocessed_image
else:
logger.info(
"CosmosLatentPreparationStage - No video input sources found")
if video is not None:
num_cond_frames = video.size(2)
if num_cond_frames >= num_frames:
num_cond_latent_frames = (num_frames -
1) // vae_scale_factor_temporal + 1
video = video[:, :, -num_frames:]
else:
num_cond_latent_frames = (num_cond_frames -
1) // vae_scale_factor_temporal + 1
num_padding_frames = num_frames - num_cond_frames
last_frame = video[:, :, -1:]
padding = last_frame.repeat(1, 1, num_padding_frames, 1, 1)
video = torch.cat([video, padding], dim=2)
if self.vae is not None:
self.vae = self.vae.to(device)
self.vae = self.vae.to(dtype=video.dtype)
def retrieve_latents(
encoder_output: Any,
generator: Any | None = None) -> torch.Tensor:
if hasattr(encoder_output, "latent_dist"):
return encoder_output.latent_dist.sample(generator)
elif hasattr(encoder_output, "latents"):
return encoder_output.latents
elif hasattr(encoder_output, "sample"):
return encoder_output.sample(generator)
elif isinstance(encoder_output, torch.Tensor):
return encoder_output
else:
attrs = [
attr for attr in dir(encoder_output)
if not attr.startswith('_')
]
raise AttributeError(
f"Could not access latents of provided encoder_output. Available attributes: {attrs}"
)
if isinstance(generator, list):
init_latents = [
retrieve_latents(self.vae.encode(video[i].unsqueeze(0)),
generator=torch.Generator(
device="cpu").manual_seed(100))
for i in range(batch_size)
]
else:
init_latents = [
retrieve_latents(
self.vae.encode(vid.unsqueeze(0)),
torch.Generator(device="cpu").manual_seed(100))
for vid in video
]
init_latents = torch.cat(init_latents, dim=0).to(dtype)
cfg = getattr(self.vae, "config", None)
if (not bool(getattr(self.vae, "handles_latent_norm", False))
and cfg is not None and hasattr(cfg, 'latents_mean')
and hasattr(cfg, 'latents_std')):
latents_mean = torch.tensor(cfg.latents_mean).view(
1, cfg.z_dim, 1, 1, 1).to(device, dtype)
latents_std = torch.tensor(cfg.latents_std).view(
1, cfg.z_dim, 1, 1, 1).to(device, dtype)
init_latents = (init_latents - latents_mean
) / latents_std * self.scheduler.sigma_data
conditioning_latents = init_latents
self.vae.to("cpu")
else:
num_cond_latent_frames = 0
if latents is None:
seed = int(batch.seed if batch.seed is not None else 0)
# Use arch-invariant RNG to match Cosmos2.5 reference sampling.
latents_fp32 = self._arch_invariant_randn(shape,
seed=seed,
device=device,
dtype=torch.float32)
latents = latents_fp32.to(torch.bfloat16)
else:
# If latents are supplied, keep compute dtype consistent with Cosmos sampling.
latents = latents.to(device=device, dtype=torch.bfloat16)
# Cosmos2.5 starts from unit Gaussian noise (no extra sigma_max scaling).
padding_shape = (batch_size, 1, num_latent_frames, latent_height,
latent_width)
ones_padding = latents.new_ones(padding_shape)
zeros_padding = latents.new_zeros(padding_shape)
cond_indicator = latents.new_zeros(1, 1, latents.size(2), 1, 1)
cond_indicator[:, :, :num_cond_latent_frames] = 1.0
cond_mask = cond_indicator * ones_padding + (
1 - cond_indicator) * zeros_padding
uncond_indicator = None
uncond_mask = None
if batch.do_classifier_free_guidance:
uncond_indicator = latents.new_zeros(1, 1, latents.size(2), 1, 1)
uncond_indicator[:, :, :num_cond_latent_frames] = 1.0
uncond_mask = uncond_indicator * ones_padding + (
1 - uncond_indicator) * zeros_padding
# Cosmos 2.5 requires a spatial padding mask when concat_padding_mask=True
padding_mask = latents.new_ones(batch_size, 1, latent_height,
latent_width)
batch.latents = latents
batch.raw_latent_shape = latents.shape
batch.conditioning_latents = conditioning_latents
batch.cond_indicator = cond_indicator
batch.uncond_indicator = uncond_indicator
batch.cond_mask = cond_mask
batch.uncond_mask = uncond_mask
batch.padding_mask = padding_mask
return batch
def adjust_video_length(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> int:
"""
@@ -328,3 +328,69 @@ class TextEncodingStage(PipelineStage):
lambda x: not batch.do_classifier_free_guidance or V.
list_of_tensors_with_min_dims(x, 2))
return result
class Cosmos25TextEncodingStage(PipelineStage):
"""Cosmos 2.5 text encoding stage.
Cosmos 2.5 uses Reason1 (Qwen2.5-VL) and relies on the encoder's
`compute_text_embeddings_online()`.
"""
def __init__(self, text_encoder) -> None:
super().__init__()
self.text_encoder = text_encoder
@torch.no_grad()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.prompt is not None
prompts = [batch.prompt] if isinstance(batch.prompt,
str) else batch.prompt
encoder = self.text_encoder
if not hasattr(encoder, "compute_text_embeddings_online"):
raise RuntimeError(
"Cosmos25TextEncodingStage requires text_encoder.compute_text_embeddings_online()"
)
with set_forward_context(current_timestep=0, attn_metadata=None):
prompt_embeds = encoder.compute_text_embeddings_online(
{"text": prompts}, "text")
batch.prompt_embeds = [prompt_embeds]
if batch.do_classifier_free_guidance:
neg = batch.negative_prompt
neg_prompts = ([neg] *
len(prompts)) if isinstance(neg, str) else neg
with set_forward_context(current_timestep=0, attn_metadata=None):
neg_embeds = encoder.compute_text_embeddings_online(
{"text": neg_prompts}, "text")
batch.negative_prompt_embeds = [neg_embeds]
else:
batch.negative_prompt_embeds = []
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
result.add_check(
"negative_prompt",
batch.negative_prompt,
lambda x:
(not batch.do_classifier_free_guidance) or isinstance(x, str),
)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_of_tensors_min_dims(2))
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
@@ -123,3 +123,32 @@ class TimestepPreparationStage(PipelineStage):
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.with_dims(1)])
return result
class Cosmos25TimestepPreparationStage(TimestepPreparationStage):
"""Cosmos 2.5 timestep preparation with scheduler-specific kwargs."""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
scheduler = self.scheduler
device = get_local_torch_device()
num_inference_steps = batch.num_inference_steps
extra_kwargs: dict = {}
sig = inspect.signature(scheduler.set_timesteps)
if "shift" in sig.parameters:
extra_kwargs["shift"] = fastvideo_args.pipeline_config.flow_shift
# Prefer the canonical diffusers kwarg name if available.
if "use_karras_sigmas" in sig.parameters:
extra_kwargs["use_karras_sigmas"] = True
elif "use_kerras_sigma" in sig.parameters:
extra_kwargs["use_kerras_sigma"] = True
scheduler.set_timesteps(num_inference_steps,
device=device,
**extra_kwargs)
batch.timesteps = scheduler.timesteps
return batch
+60
View File
@@ -0,0 +1,60 @@
from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager
from torch import nn
from typing import Any
import torch
class EventHook(ForwardHook):
def __init__(self, content: str, event_list: list[str]):
self.content = content
self.event_list = event_list
def name(self) -> str:
return f"EventHook_{self.content}"
def pre_forward(self, module: nn.Module, *args, **kwargs):
print(
f"[{self.content}] Pre-forward called with args[0].shape: {args[0].shape}"
)
self.event_list.append(f"[pre]{self.content}")
return args, kwargs
def post_forward(self, module: nn.Module, output: Any):
print(
f"[{self.content}] Post-forward called with outputs.shape: {output.shape}"
)
self.event_list.append(f"[post]{self.content}")
return output
def test_hook_execution_order():
"""Test that hooks are executed in the correct order: LIFO for pre-hooks, FIFO for post-hooks."""
# Create a simple model
model = nn.Linear(10, 20)
# Create event list to track hook execution order
events = []
# Create and push hooks in order: A then B
manager = ModuleHookManager.get_from_or_default(model)
hook_a = EventHook("A", events)
hook_b = EventHook("B", events)
manager.append_forward_hook(hook_a)
manager.append_forward_hook(hook_b)
# Perform a forward pass
input_tensor = torch.randn(2, 10)
model(input_tensor)
# Verify the execution order is [pre_a, pre_b, post_b, post_a]
# Pre-hooks should be FILO (First In Last Out): A then B
# Post-hooks should be LIFO (Last In First Out): B then A
expected_events = ["[pre]A", "[pre]B", "[post]B", "[post]A"]
assert events == expected_events, (
f"Expected {expected_events}, but got {events}"
)
print(f"✓ Hook execution order test passed: {events}")
@@ -0,0 +1,395 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
import torch.nn as nn
from fastvideo.hooks.layerwise_offload import (
LayerwiseOffloadHook,
enable_layerwise_offload,
)
from fastvideo.hooks.hooks import ModuleHookManager
class SimpleBlock(nn.Module):
"""A simple block with linear layers for testing."""
def __init__(self, hidden_size: int, dtype: torch.dtype = torch.float32):
super().__init__()
self.linear1 = nn.Linear(hidden_size, hidden_size, dtype=dtype)
self.linear2 = nn.Linear(hidden_size, hidden_size, dtype=dtype)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.linear1(x)
x = torch.relu(x)
x = self.linear2(x)
return x
class SimpleModelWithModuleList(nn.Module):
"""A simple model with ModuleList for testing layerwise offloading."""
def __init__(
self,
num_blocks: int = 4,
hidden_size: int = 128,
dtype: torch.dtype = torch.float32,
):
super().__init__()
self.blocks = nn.ModuleList(
[SimpleBlock(hidden_size, dtype=dtype) for _ in range(num_blocks)]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for block in self.blocks:
x = block(x)
return x
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_basic():
"""Test basic functionality of layerwise offloading."""
device = torch.device("cuda")
hidden_size = 128
batch_size = 2
seq_len = 16
num_blocks = 4
# Create model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
# Get reference output without offloading
input_tensor = torch.randn(batch_size, seq_len, hidden_size, device=device)
with torch.no_grad():
reference_output = model(input_tensor.clone())
# Enable layerwise offloading
enable_layerwise_offload(model)
# Verify parameters are offloaded to CPU
for block in model.blocks:
for param in block.parameters():
# Parameters should be placeholder tensors (empty)
assert param.numel() == 0, (
"Parameters should be offloaded (empty tensors)"
)
# Run forward pass with offloading
with torch.no_grad():
offloaded_output = model(input_tensor.clone())
# Check output correctness
assert torch.allclose(
reference_output, offloaded_output, rtol=1e-4, atol=1e-5
), "Output with offloading should match reference output"
print(
f" Layerwise offload basic test passed: max diff = {(reference_output - offloaded_output).abs().max().item()}"
)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_bf16():
"""Test layerwise offloading with bfloat16 precision."""
device = torch.device("cuda")
hidden_size = 256
batch_size = 1
seq_len = 32
num_blocks = 3
# Create model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.bfloat16
).to(device)
# Get reference output without offloading
input_tensor = torch.randn(
batch_size, seq_len, hidden_size, device=device, dtype=torch.bfloat16
)
with torch.no_grad():
reference_output = model(input_tensor.clone())
# Enable layerwise offloading
enable_layerwise_offload(model)
# Run forward pass with offloading
with torch.no_grad():
offloaded_output = model(input_tensor.clone())
# Check output correctness (looser tolerance for bf16)
assert torch.allclose(
reference_output, offloaded_output, rtol=1e-2, atol=1e-3
), "Output with offloading should match reference output for bf16"
print(
f" Layerwise offload bf16 test passed: max diff = {(reference_output - offloaded_output).abs().max().item()}"
)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_multiple_forward_passes():
"""Test that layerwise offloading works correctly across multiple forward passes."""
device = torch.device("cuda")
hidden_size = 64
batch_size = 2
seq_len = 8
num_blocks = 3
num_iterations = 5
# Create model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
# Get reference outputs without offloading
torch.manual_seed(42)
reference_outputs = []
for i in range(num_iterations):
input_tensor = torch.randn(
batch_size, seq_len, hidden_size, device=device
)
with torch.no_grad():
reference_outputs.append(model(input_tensor.clone()))
# Enable layerwise offloading
enable_layerwise_offload(model)
# Run multiple forward passes with offloading
torch.manual_seed(42)
for i in range(num_iterations):
input_tensor = torch.randn(
batch_size, seq_len, hidden_size, device=device
)
with torch.no_grad():
offloaded_output = model(input_tensor.clone())
# Check output correctness for each iteration
assert torch.allclose(
reference_outputs[i], offloaded_output, rtol=1e-4, atol=1e-5
), f"Output mismatch in iteration {i}"
print(
f" Multiple forward passes test passed for {num_iterations} iterations"
)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_parameter_integrity():
"""Test that parameters are correctly restored during forward pass."""
device = torch.device("cuda")
hidden_size = 128
num_blocks = 2
# Create model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
# Store original parameter values per block
original_params_per_block = []
for block in model.blocks:
block_params = {}
for name, param in block.named_parameters():
block_params[name] = param.data.clone()
original_params_per_block.append(block_params)
# Enable layerwise offloading
enable_layerwise_offload(model)
# Create hook managers and verify they exist
for block_idx, block in enumerate(model.blocks):
manager = ModuleHookManager.get_from(block)
assert manager is not None, "Hook manager should be attached to blocks"
hook: LayerwiseOffloadHook | None = manager.get_forward_hook(
"LayerwiseOffloadHook"
)
assert hook is not None, "LayerwiseOffloadHook should be registered"
# Verify parameters are stored in CPU
state = hook.state
assert len(state.cpu_named_parameters) > 0, (
"CPU parameters should be stored"
)
# Verify CPU parameters match original values for this block
original_params = original_params_per_block[block_idx]
for name, cpu_param in state.cpu_named_parameters.items():
assert name in original_params, (
f"CPU parameter {name} not found in original params for block {block_idx}"
)
assert torch.allclose(
cpu_param.cpu(), original_params[name].cpu(), rtol=1e-5, atol=1e-6
), f"CPU parameter {name} should match original for block {block_idx}"
print(" Parameter integrity test passed")
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_no_modulelist_error():
"""Test that enabling offload on a model without ModuleList raises an error."""
class ModelWithoutModuleList(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(128, 128)
def forward(self, x):
return self.linear(x)
model = ModelWithoutModuleList().to("cuda")
with pytest.raises(
ValueError,
match="No nn.ModuleList found in the model for layerwise offloading",
):
enable_layerwise_offload(model)
print(" No ModuleList error test passed")
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_memory_reduction():
"""Test that layerwise offloading reduces GPU memory usage."""
device = torch.device("cuda")
hidden_size = 512
num_blocks = 8
batch_size = 1
seq_len = 64
# Create model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
# Measure initial GPU memory
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
input_tensor = torch.randn(batch_size, seq_len, hidden_size, device=device)
with torch.no_grad():
_ = model(input_tensor)
memory_without_offload = torch.cuda.max_memory_allocated()
# Reset model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
# Enable offloading
enable_layerwise_offload(model)
# Measure GPU memory with offloading
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
input_tensor = torch.randn(batch_size, seq_len, hidden_size, device=device)
with torch.no_grad():
_ = model(input_tensor)
memory_with_offload = torch.cuda.max_memory_allocated()
# Memory with offload should be less (parameters are offloaded)
# Note: This is a weak check as memory usage depends on many factors
print(
f"Memory without offload: {memory_without_offload / 1024**2:.2f} MB, "
f"with offload: {memory_with_offload / 1024**2:.2f} MB"
)
# We expect some reduction but this test is more informational
assert memory_with_offload < memory_without_offload * 1.5, (
"Memory usage should not increase significantly with offloading"
)
print(" Memory reduction test passed")
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_gradient_disabled():
"""Test that layerwise offloading works correctly with gradients disabled."""
device = torch.device("cuda")
hidden_size = 128
batch_size = 2
seq_len = 16
num_blocks = 3
# Create model
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
model.eval()
# Enable layerwise offloading
enable_layerwise_offload(model)
input_tensor = torch.randn(batch_size, seq_len, hidden_size, device=device)
# Forward pass should work without gradients
with torch.no_grad():
output = model(input_tensor)
assert output.requires_grad is False, "Output should not require gradients"
assert output.shape == (
batch_size,
seq_len,
hidden_size,
), "Output shape should match input"
print(" Gradient disabled test passed")
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_different_batch_sizes():
"""Test layerwise offloading with different batch sizes."""
device = torch.device("cuda")
hidden_size = 128
seq_len = 16
num_blocks = 3
model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
# Get reference model without offloading
reference_model = SimpleModelWithModuleList(
num_blocks=num_blocks, hidden_size=hidden_size, dtype=torch.float32
).to(device)
reference_model.load_state_dict(model.state_dict())
# Enable layerwise offloading
enable_layerwise_offload(model)
# Test different batch sizes
for batch_size in [1, 2, 4, 8]:
input_tensor = torch.randn(
batch_size, seq_len, hidden_size, device=device
)
with torch.no_grad():
reference_output = reference_model(input_tensor.clone())
offloaded_output = model(input_tensor.clone())
assert torch.allclose(
reference_output, offloaded_output, rtol=1e-4, atol=1e-5
), f"Output mismatch for batch_size={batch_size}"
print(" Different batch sizes test passed")
if __name__ == "__main__":
# Run tests manually for debugging
if torch.cuda.is_available():
print("Running layerwise offloading tests...")
test_layerwise_offload_basic()
test_layerwise_offload_bf16()
test_layerwise_offload_multiple_forward_passes()
test_layerwise_offload_parameter_integrity()
test_layerwise_offload_no_modulelist_error()
test_layerwise_offload_memory_reduction()
test_layerwise_offload_gradient_disabled()
test_layerwise_offload_different_batch_sizes()
print("\n All tests passed!")
else:
print("CUDA not available, skipping tests")
@@ -93,8 +93,11 @@ def test_merge_lora_weights(model_id):
lora_nickname = lora_config["lora_nickname"]
lora_path = lora_config["lora_path"]
# When layerwise offload is enabled, placeholder tensors cannot be compared directly.
args = FastVideoArgs.from_kwargs(
model_path=model_id,
dit_layerwise_offload=False,
use_fsdp_inference=True,
dit_cpu_offload=True,
dit_precision="bf16",
)
+2 -2
View File
@@ -82,12 +82,12 @@ def run_transformer_tests():
@app.function(
gpu="L40S:4",
image=image,
timeout=3600,
timeout=6000,
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol}
)
def run_ssim_tests():
run_test("export MODEL_PATH='/root/data/weights' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
run_test("export HF_HOME='/root/data/.cache' && export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
+3
View File
@@ -0,0 +1,3 @@
# SPDX-License-Identifier: Apache-2.0
@@ -1,11 +0,0 @@
{
"mean_ssim": 0.7606927804004999,
"min_ssim": 0.7035917639732361,
"max_ssim": 0.7920367121696472,
"reference_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
"generated_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
"parameters": {
"num_inference_steps": 4,
"prompt": "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
}
}
@@ -26,17 +26,17 @@ else:
# Base parameters from the shell script
HUNYUAN_PARAMS = {
"num_gpus": 2,
"num_gpus": 4,
"model_path": "FastVideo/FastHunyuan-diffusers",
"height": 720,
"width": 1280,
"num_frames": 45,
"num_inference_steps": 6,
"num_inference_steps": 2,
"guidance_scale": 1,
"embedded_cfg_scale": 6,
"flow_shift": 17,
"seed": 1024,
"sp_size": 2,
"sp_size": 4,
"tp_size": 1,
"vae_sp": True,
"fps": 24,
@@ -48,7 +48,7 @@ WAN_T2V_PARAMS = {
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 20,
"num_inference_steps": 4,
"guidance_scale": 3,
"embedded_cfg_scale": 6,
"flow_shift": 7.0,
@@ -67,7 +67,7 @@ WAN_I2V_PARAMS = {
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 6,
"num_inference_steps": 2,
"guidance_scale": 5.0,
"embedded_cfg_scale": 6,
"flow_shift": 7.0,
@@ -232,7 +232,9 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
"flow_shift": BASE_PARAMS["flow_shift"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
"dit_cpu_offload": True,
"use_fsdp_inference": True,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
}
if BASE_PARAMS.get("vae_sp"):
init_kwargs["vae_sp"] = True
@@ -0,0 +1,429 @@
# SPDX-License-Identifier: Apache-2.0
"""
SSIM-based similarity tests for LongCat video generation.
Tests three LongCat modes:
- T2V (Text-to-Video): 480p video from text prompt
- I2V (Image-to-Video): 480p video from image + text prompt
- VC (Video Continuation): 480p video continuation from input video + text prompt
Sampling parameters are derived from:
- examples/inference/basic/basic_longcat_t2v.py
- examples/inference/basic/basic_longcat_i2v.py
- examples/inference/basic/basic_longcat_vc.py
Note: num_inference_steps is reduced for CI speed (4 steps vs 50 in examples).
"""
import os
import pytest
import torch
from fastvideo import VideoGenerator
from fastvideo.logger import init_logger
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
logger = init_logger(__name__)
# Device-specific reference folder
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = "_reference_videos"
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
else:
logger.warning(f"Unsupported device for ssim tests: {device_name}")
# Common negative prompt from example scripts
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
# =============================================================================
# LongCat T2V Parameters (from basic_longcat_t2v.py)
# =============================================================================
LONGCAT_T2V_PARAMS = {
"num_gpus": 1,
"model_path": "FastVideo/LongCat-Video-T2V-Diffusers",
"height": 480,
"width": 480,
"num_frames": 43,
"num_inference_steps": 4, # Reduced from 50 for CI speed
"guidance_scale": 4.0,
"fps": 15,
"seed": 42,
"negative_prompt": NEGATIVE_PROMPT,
}
# =============================================================================
# LongCat I2V Parameters (from basic_longcat_i2v.py)
# =============================================================================
LONGCAT_I2V_PARAMS = {
"num_gpus": 1,
"model_path": "FastVideo/LongCat-Video-I2V-Diffusers",
"height": 480,
"width": 480, # Square for I2V
"num_frames": 43,
"num_inference_steps": 4, # Reduced from 50 for CI speed
"guidance_scale": 4.0,
"fps": 15,
"seed": 42,
"negative_prompt": NEGATIVE_PROMPT,
}
# =============================================================================
# LongCat VC Parameters (from basic_longcat_vc.py)
# =============================================================================
LONGCAT_VC_PARAMS = {
"num_gpus": 1,
"model_path": "FastVideo/LongCat-Video-VC-Diffusers",
"height": 480,
"width": 480,
"num_frames": 43,
"num_inference_steps": 4, # Reduced from 50 for CI speed
"guidance_scale": 4.0,
"fps": 15,
"seed": 42,
"num_cond_frames": 13,
"negative_prompt": NEGATIVE_PROMPT,
}
# Test prompts
T2V_TEST_PROMPTS = [
"In a realistic photography style, a white boy around seven or eight years old "
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
"features a green lawn and several tall trees, creating a warm and loving scene.",
]
I2V_TEST_PROMPTS = [
"A woman sits at a wooden table by the window in a cozy café. She reaches out "
"with her right hand, picks up the white coffee cup from the saucer, and gently "
"brings it to her lips to take a sip. After drinking, she places the cup back on "
"the table and looks out the window, enjoying the peaceful atmosphere.",
]
I2V_IMAGE_PATHS = [
"assets/girl.png",
]
VC_TEST_PROMPTS = [
"A person rides a motorcycle along a long, straight road that stretches between "
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
"the motorcycle centered between the guardrails, while the scenery passes by on "
"both sides. The video captures the journey from the rider's perspective, emphasizing "
"the sense of motion and adventure.",
]
VC_VIDEO_PATHS = [
"assets/motorcycle.mp4",
]
def _resolve_asset_path(asset_path: str) -> str:
"""Resolve asset path relative to FastVideo root."""
# Check if absolute or already exists
if os.path.isabs(asset_path) or os.path.exists(asset_path):
return asset_path
# Try relative to workspace root
script_dir = os.path.dirname(os.path.abspath(__file__))
repo_root = os.path.abspath(os.path.join(script_dir, "..", "..", ".."))
return os.path.join(repo_root, asset_path)
@pytest.mark.parametrize("prompt", T2V_TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_longcat_t2v_similarity(prompt: str, ATTENTION_BACKEND: str):
"""
Test LongCat T2V inference and compare output to reference videos using SSIM.
Parameters derived from examples/inference/basic/basic_longcat_t2v.py
"""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
model_id = "LongCat-Video-T2V"
output_dir = os.path.join(script_dir, "generated_videos", model_id, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
init_kwargs = {
"num_gpus": LONGCAT_T2V_PARAMS["num_gpus"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"enable_bsa": False,
}
generation_kwargs = {
"output_path": output_dir,
"height": LONGCAT_T2V_PARAMS["height"],
"width": LONGCAT_T2V_PARAMS["width"],
"num_frames": LONGCAT_T2V_PARAMS["num_frames"],
"num_inference_steps": LONGCAT_T2V_PARAMS["num_inference_steps"],
"guidance_scale": LONGCAT_T2V_PARAMS["guidance_scale"],
"fps": LONGCAT_T2V_PARAMS["fps"],
"seed": LONGCAT_T2V_PARAMS["seed"],
"negative_prompt": LONGCAT_T2V_PARAMS["negative_prompt"],
}
generator = VideoGenerator.from_pretrained(
model_path=LONGCAT_T2V_PARAMS["model_path"], **init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
generator.shutdown()
generated_video_path = os.path.join(output_dir, output_video_name)
assert os.path.exists(generated_video_path), (
f"Output video was not generated at {generated_video_path}"
)
# Find reference video
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
raise FileNotFoundError(
f"Reference video not found for prompt: {prompt[:50]}... with backend: {ATTENTION_BACKEND}"
)
reference_video_path = os.path.join(reference_folder, reference_video_name)
logger.info(f"Computing SSIM between {reference_video_path} and {generated_video_path}")
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
write_ssim_results(
output_dir, ssim_values, reference_video_path, generated_video_path,
LONGCAT_T2V_PARAMS["num_inference_steps"], prompt
)
min_acceptable_ssim = 0.90
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}"
)
@pytest.mark.parametrize("prompt", I2V_TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_longcat_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
"""
Test LongCat I2V inference and compare output to reference videos using SSIM.
Parameters derived from examples/inference/basic/basic_longcat_i2v.py
"""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
model_id = "LongCat-Video-I2V"
output_dir = os.path.join(script_dir, "generated_videos", model_id, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
# Get image path for this prompt
prompt_idx = I2V_TEST_PROMPTS.index(prompt)
image_path = _resolve_asset_path(I2V_IMAGE_PATHS[prompt_idx])
init_kwargs = {
"num_gpus": LONGCAT_I2V_PARAMS["num_gpus"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"enable_bsa": False,
}
generation_kwargs = {
"output_path": output_dir,
"image_path": image_path,
"height": LONGCAT_I2V_PARAMS["height"],
"width": LONGCAT_I2V_PARAMS["width"],
"num_frames": LONGCAT_I2V_PARAMS["num_frames"],
"num_inference_steps": LONGCAT_I2V_PARAMS["num_inference_steps"],
"guidance_scale": LONGCAT_I2V_PARAMS["guidance_scale"],
"fps": LONGCAT_I2V_PARAMS["fps"],
"seed": LONGCAT_I2V_PARAMS["seed"],
"negative_prompt": LONGCAT_I2V_PARAMS["negative_prompt"],
}
generator = VideoGenerator.from_pretrained(
model_path=LONGCAT_I2V_PARAMS["model_path"], **init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
generator.shutdown()
generated_video_path = os.path.join(output_dir, output_video_name)
assert os.path.exists(generated_video_path), (
f"Output video was not generated at {generated_video_path}"
)
# Find reference video
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
raise FileNotFoundError(
f"Reference video not found for prompt: {prompt[:50]}... with backend: {ATTENTION_BACKEND}"
)
reference_video_path = os.path.join(reference_folder, reference_video_name)
logger.info(f"Computing SSIM between {reference_video_path} and {generated_video_path}")
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
write_ssim_results(
output_dir, ssim_values, reference_video_path, generated_video_path,
LONGCAT_I2V_PARAMS["num_inference_steps"], prompt
)
min_acceptable_ssim = 0.90
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}"
)
@pytest.mark.parametrize("prompt", VC_TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_longcat_vc_similarity(prompt: str, ATTENTION_BACKEND: str):
"""
Test LongCat VC (Video Continuation) inference and compare output to reference videos using SSIM.
Parameters derived from examples/inference/basic/basic_longcat_vc.py
"""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
model_id = "LongCat-Video-VC"
output_dir = os.path.join(script_dir, "generated_videos", model_id, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
# Get video path for this prompt
prompt_idx = VC_TEST_PROMPTS.index(prompt)
video_path = _resolve_asset_path(VC_VIDEO_PATHS[prompt_idx])
if not os.path.exists(video_path):
pytest.skip(f"Input video not found at {video_path}")
init_kwargs = {
"num_gpus": LONGCAT_VC_PARAMS["num_gpus"],
"use_fsdp_inference": False,
"dit_cpu_offload": False,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": False,
"enable_bsa": False,
}
generation_kwargs = {
"output_path": output_dir,
"video_path": video_path,
"num_cond_frames": LONGCAT_VC_PARAMS["num_cond_frames"],
"height": LONGCAT_VC_PARAMS["height"],
"width": LONGCAT_VC_PARAMS["width"],
"num_frames": LONGCAT_VC_PARAMS["num_frames"],
"num_inference_steps": LONGCAT_VC_PARAMS["num_inference_steps"],
"guidance_scale": LONGCAT_VC_PARAMS["guidance_scale"],
"fps": LONGCAT_VC_PARAMS["fps"],
"seed": LONGCAT_VC_PARAMS["seed"],
"negative_prompt": LONGCAT_VC_PARAMS["negative_prompt"],
}
generator = VideoGenerator.from_pretrained(
model_path=LONGCAT_VC_PARAMS["model_path"], **init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
generator.shutdown()
generated_video_path = os.path.join(output_dir, output_video_name)
assert os.path.exists(generated_video_path), (
f"Output video was not generated at {generated_video_path}"
)
# Find reference video
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
raise FileNotFoundError(
f"Reference video not found for prompt: {prompt[:50]}... with backend: {ATTENTION_BACKEND}"
)
reference_video_path = os.path.join(reference_folder, reference_video_name)
logger.info(f"Computing SSIM between {reference_video_path} and {generated_video_path}")
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
write_ssim_results(
output_dir, ssim_values, reference_video_path, generated_video_path,
LONGCAT_VC_PARAMS["num_inference_steps"], prompt
)
min_acceptable_ssim = 0.90
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}"
)
@@ -88,6 +88,7 @@ def test_matrixgame_similarity(prompt, ATTENTION_BACKEND, model_id):
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"use_fsdp_inference": True,
"dit_layerwise_offload": False,
"dit_cpu_offload": False,
"vae_cpu_offload": False,
"text_encoder_cpu_offload": True,
@@ -108,9 +109,7 @@ def test_matrixgame_similarity(prompt, ATTENTION_BACKEND, model_id):
"save_video": True,
}
generator = VideoGenerator.from_pretrained(
model_path=BASE_PARAMS["model_path"], **init_kwargs
)
generator = VideoGenerator.from_pretrained(model_path=BASE_PARAMS["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
if isinstance(generator.executor, MultiprocExecutor):
@@ -32,7 +32,7 @@ else:
# TurboDiffusion parameters (1-4 step generation with RCM scheduler + SLA attention)
TURBODIFFUSION_PARAMS = {
"num_gpus": 2,
"num_gpus": 4,
"model_path": "loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
"height": 480,
"width": 832,
@@ -40,7 +40,7 @@ TURBODIFFUSION_PARAMS = {
"num_inference_steps": 4, # TurboDiffusion uses 1-4 steps
"guidance_scale": 1.0, # No CFG for TurboDiffusion
"seed": 42,
"sp_size": 2,
"sp_size": 4,
"tp_size": 1,
"fps": 24,
}
@@ -94,10 +94,7 @@ def test_turbodiffusion_inference_similarity(prompt, model_id):
"fps": BASE_PARAMS["fps"],
}
generator = VideoGenerator.from_pretrained(
model_path=BASE_PARAMS["model_path"],
**init_kwargs
)
generator = VideoGenerator.from_pretrained(model_path=BASE_PARAMS["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
if isinstance(generator.executor, MultiprocExecutor):
@@ -220,6 +217,8 @@ def test_turbodiffusion_i2v_inference_similarity(prompt, model_id):
"override_pipeline_cls_name": "TurboDiffusionI2VPipeline",
# Keep both transformers in VRAM - avoids CPU RAM bottleneck
"dit_cpu_offload": False,
"use_fsdp_inference": True,
"dit_layerwise_offload": False,
}
generation_kwargs = {
@@ -6,6 +6,11 @@ import json
from huggingface_hub import snapshot_download
import torch
# Ensure backend selection happens during import-time initialization
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
# Force VSA to use Triton implementation even on H100 / when CUDA extension is available
# os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
from fastvideo.training.wan_training_pipeline import main
@@ -19,8 +24,6 @@ h200_reference_wandb_summary_file = "fastvideo/tests/training/VSA/h200_reference
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
def run_worker():
"""Worker function that will be run on each GPU"""
# Create and populate args
@@ -118,10 +121,10 @@ def test_distributed_training():
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 1.0,
'train_loss': 0.02
'avg_step_time': 3,
'grad_norm': 0.2,
'step_time': 2.5,
'train_loss': 0.04
}
failures = []
+23 -2
View File
@@ -17,10 +17,11 @@ from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
os.environ["MASTER_PORT"] = "29701"
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
@@ -121,4 +122,24 @@ def test_wan_transformer():
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
if __name__ == "__main__":
from fastvideo.distributed import (
cleanup_dist_env_and_memory,
maybe_init_distributed_environment_and_model_parallel,
)
# Allow running this test file directly without pytest.
maybe_init_distributed_environment_and_model_parallel(1, 1)
try:
test_wan_transformer()
logger.info("test_wan_transformer finished successfully.")
finally:
cleanup_dist_env_and_memory()
+8 -1
View File
@@ -1,5 +1,12 @@
from .distillation_pipeline import DistillationPipeline
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
from fastvideo.training.rl import RLPipeline, create_rl_pipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
__all__ = [
"TrainingPipeline",
"WanTrainingPipeline",
"DistillationPipeline",
"RLPipeline",
"create_rl_pipeline",
]
+6
View File
@@ -0,0 +1,6 @@
from .rl_pipeline import RLPipeline, create_rl_pipeline
__all__ = [
"RLPipeline",
"create_rl_pipeline",
]
+11
View File
@@ -0,0 +1,11 @@
from .rewards import (
create_reward_models,
MultiRewardAggregator,
ValueModel
)
__all__ = [
"create_reward_models",
"MultiRewardAggregator",
"ValueModel",
]
+63
View File
@@ -0,0 +1,63 @@
# SPDX-License-Identifier: Apache-2.0
"""
Abstract base class for VIDEO reward models.
All VIDEO reward models should inherit from this class and implement
the compute_reward() method.
IMPORTANT: Reward models must process FULL VIDEO SEQUENCES, not individual frames.
Input shape is [B, T, C, H, W] where T is the temporal (frame) dimension.
For video-specific rewards, consider:
- Temporal coherence across frames
- Motion quality and smoothness
- Video-text alignment (not just frame-text)
- Multi-frame aesthetic quality
"""
from typing import Any
from abc import ABC, abstractmethod
import torch
import torch.nn as nn
class BaseRewardModel(ABC, nn.Module):
def __init__(self, model_path: str | None = None, device: str = "cuda"):
super().__init__()
self.model_path = model_path
self.device = device
@abstractmethod
def compute_reward(
self,
videos: torch.Tensor, # [B, T, C, H, W] decoded video sequences
prompts: list[str] | None, # Text prompts
**kwargs: Any
) -> torch.Tensor:
"""
Compute rewards for generated VIDEO sequences.
IMPORTANT: This method must process the FULL temporal sequence [B, T, C, H, W].
Do NOT evaluate individual frames independently and average.
Args:
videos: Decoded video tensors [B, T, C, H, W] in range [0, 1]
B = batch size
T = number of frames (temporal dimension)
C = channels (typically 3 for RGB)
H, W = height, width
prompts: List of text prompts (length B) describing each video
**kwargs: Additional model-specific arguments
Returns:
rewards: Tensor of shape [B] with reward scores for each video sequence
Example:
>>> videos = torch.rand(4, 17, 3, 256, 256) # 4 videos, 17 frames each
>>> prompts = ["A cat jumping", "A dog running", ...]
>>> rewards = model.compute_reward(videos, prompts)
"""
raise NotImplementedError("Subclasses must implement compute_reward()")
def __repr__(self) -> str:
return f"{self.__class__.__name__}(model_path={self.model_path})"
+206
View File
@@ -0,0 +1,206 @@
from paddleocr import PaddleOCR
import torch
import numpy as np
from Levenshtein import distance
from typing import Any
from PIL import Image
from fastvideo.training.rl.rewards.base import BaseRewardModel
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class OcrScorerVideo(BaseRewardModel):
"""
OCR reward model for multi-frame video OCR evaluation.
This model evaluates multiple frames across the video sequence,
sampling frames at a specified interval and averaging the OCR scores.
"""
def __init__(self,
model_path: str | None = None,
device: str = "cpu",
frame_interval: int = 4):
"""
OCR reward calculator for videos
Args:
model_path: Not used for PaddleOCR (kept for BaseRewardModel compatibility)
device: Device string (used to determine use_gpu if not explicitly set)
frame_interval: Sample every Nth frame (default: 4)
"""
super().__init__(model_path=model_path, device=device)
self.frame_interval = frame_interval
self.ocr = PaddleOCR(
use_angle_cls=False,
lang="en",
use_gpu=False,
show_log=False # Disable unnecessary log output
)
logger.info("Initialized OcrScorerVideo (device=%s, frame_interval=%d)",
device, frame_interval)
def _process_single_video(self, video_tensor: torch.Tensor,
prompt: str) -> float:
"""
Process a single video tensor and return its OCR reward.
Args:
video_tensor: Video tensor of shape [C, T, H, W]
prompt: Text prompt containing target OCR text in quotes
Returns:
Average reward across positive-scoring frames
"""
# Extract target text from prompt
try:
target_text = prompt.split('"')[1].replace(' ', '').lower()
except IndexError:
logger.warning("Failed to extract quoted text from prompt: %s",
prompt)
target_text = prompt.replace(' ', '').lower()
if not target_text:
return 0.0
# video_tensor is [C, T, H, W]
C, T, H, W = video_tensor.shape
# Convert to numpy and move to CPU if needed
video_np = video_tensor.detach().cpu().numpy()
# Convert from [C, T, H, W] to [T, H, W, C] for easier frame extraction
video_np = np.transpose(video_np, (1, 2, 3, 0)) # [T, H, W, C]
logger.info(f"in ocr 1.5, video_np[0][0]: {video_np[0][0]}")
# Normalize to [0, 255] uint8 if needed
if video_np.max() <= 1.0:
video_np = (video_np * 255).astype(np.uint8)
else:
video_np = video_np.astype(np.uint8)
frame_rewards = []
# Sample frames at specified interval
for frame_idx in range(0, T, self.frame_interval):
frame = video_np[frame_idx] # [H, W, C]
logger.info(f"in ocr 2, frame.shape: {frame.shape}")
# Run OCR
try:
result = self.ocr.ocr(frame, cls=False)
logger.info(f"in ocr 3, result: {result}")
if result and result[0]:
recognized_text = "".join(
[line[1][0] for line in result[0] if line[1][1] > 0])
else:
recognized_text = ""
except Exception as e:
logger.info("OCR failed on frame %d: %s", frame_idx, str(e))
recognized_text = ''
logger.info(f"in ocr 4, recognized_text: {recognized_text}")
recognized_text = recognized_text.replace(' ', '').lower()
if target_text in recognized_text:
dist = 0
else:
dist = distance(recognized_text, target_text)
dist = min(dist, len(target_text))
reward = 1.0 - dist / len(target_text)
logger.info(f"in ocr 5, reward: {reward}")
if reward > 0:
frame_rewards.append(reward)
logger.info(f"in ocr 6, frame_rewards: {frame_rewards}")
return sum([reward / len(frame_rewards)
for reward in frame_rewards]) if frame_rewards else 0.0
@torch.no_grad()
def compute_reward(self, videos: torch.Tensor, prompts: list[str],
**kwargs: Any) -> torch.Tensor:
"""
Calculate OCR reward by evaluating sampled frames across the video.
Args:
videos: Video tensor of shape [B, C, T, H, W]
B = batch size
C = channels (typically 3 for RGB)
T = number of frames (temporal dimension)
H, W = height, width
prompts: List of text prompts containing target OCR text in quotes (length B)
**kwargs: Additional arguments
Returns:
Reward tensor [B] with averaged OCR similarity scores across frames
"""
# Ensure videos is a torch tensor with correct shape
assert isinstance(
videos,
torch.Tensor), f"videos must be torch.Tensor, got {type(videos)}"
assert videos.ndim == 5, f"videos must have 5 dimensions [B, C, T, H, W], got shape {videos.shape}"
logger.info(f"in ocr 1, videos.shape: {videos.shape}")
B, C, T, H, W = videos.shape
assert len(
prompts
) == B, f"Number of prompts ({len(prompts)}) must match batch size ({B})"
rewards = []
for b in range(B):
# Extract single video: [C, T, H, W]
video = videos[b]
reward = self._process_single_video(video, prompts[b])
rewards.append(reward)
logger.info(f"in ocr 7, rewards: {rewards}")
rewards = torch.tensor(rewards, dtype=torch.float32, device=self.device)
logger.info(f"in ocr 8, rewards: {rewards}")
# Check for NaN or Inf values
if torch.isnan(rewards).any() or torch.isinf(rewards).any():
logger.warning(
"NaN or Inf detected in OCR rewards, returning zero tensor")
return torch.zeros_like(rewards)
return rewards
if __name__ == "__main__":
example_image_path = "flowgrpo_cmd.png"
example_image = Image.open(example_image_path)
example_prompt = '/f1ow_grpo$'
# Convert image to RGB if needed
if example_image.mode != 'RGB':
example_image = example_image.convert('RGB')
# Convert PIL Image to numpy array [H, W, C]
image_np = np.array(example_image)
# Normalize to [0, 1] range and convert to float32
image_np = image_np.astype(np.float32) / 255.0
# Convert to torch tensor and reshape: [H, W, C] -> [C, H, W]
image_tensor = torch.from_numpy(image_np).permute(2, 0, 1)
# Add temporal dimension: [C, H, W] -> [C, T, H, W] where T=1
video_tensor = image_tensor.unsqueeze(1) # [C, 1, H, W]
# Add batch dimension: [C, T, H, W] -> [B, C, T, H, W] where B=1
video_tensor = video_tensor.unsqueeze(0) # [1, C, 1, H, W]
# Instantiate scorer
scorer = OcrScorerVideo(device="cpu")
# Call compute_reward method with video tensor
reward = scorer.compute_reward(video_tensor, [example_prompt])
print(f"OCR Reward: {reward.item()}")
+338
View File
@@ -0,0 +1,338 @@
# SPDX-License-Identifier: Apache-2.0
"""
Base infrastructure for VIDEO reward models in RL/GRPO training.
IMPORTANT: This module is designed exclusively for VIDEO generation models.
All reward models must operate on video sequences [B, T, C, H, W], not single frames.
This module provides:
1. Multi-reward aggregation for video
2. Value model wrapper
3. Integration with FastVideo video generation infrastructure
"""
from typing import Any
import torch
import torch.nn as nn
from fastvideo.logger import init_logger
from fastvideo.training.rl.rewards.ocr import OcrScorerVideo
from fastvideo.training.rl.rewards.base import BaseRewardModel
logger = init_logger(__name__)
class MultiRewardAggregator(nn.Module):
"""
Aggregates multiple reward models with configurable weights.
This implements the multi-reward aggregation strategy from flow_grpo,
allowing combination of different reward signals (aesthetic quality,
text-video alignment, compositional understanding, etc.)
"""
def __init__(
self,
reward_models: list[BaseRewardModel],
reward_weights: list[float] | None = None,
normalize_rewards: bool = True
):
"""
Initialize multi-reward aggregator.
Args:
reward_models: List of reward model instances
reward_weights: Weights for each reward model (default: uniform)
normalize_rewards: Whether to normalize rewards before aggregation
"""
super().__init__()
self.reward_models = nn.ModuleList(reward_models)
if reward_weights is None:
reward_weights = [1.0 / len(reward_models)] * len(reward_models)
assert len(reward_weights) == len(reward_models), \
f"Number of weights ({len(reward_weights)}) must match number of models ({len(reward_models)})"
assert abs(sum(reward_weights) - 1.0) < 1e-6, \
f"Reward weights must sum to 1.0, got {sum(reward_weights)}"
self.reward_weights = reward_weights
self.normalize_rewards = normalize_rewards
logger.info(
"Initialized MultiRewardAggregator with %d models: %s",
len(reward_models),
[(type(m).__name__, w) for m, w in zip(reward_models, reward_weights, strict=False)]
)
def compute_reward(
self,
videos: torch.Tensor,
prompts: list[str],
return_individual: bool = False,
**kwargs: Any
) -> torch.Tensor | dict[str, torch.Tensor]:
"""
Compute aggregated reward from multiple models.
Args:
videos: Decoded video tensors [B, C, T, H, W]
prompts: List of text prompts
return_individual: If True, return dict with individual rewards
**kwargs: Additional arguments passed to reward models
Returns:
If return_individual=False: aggregated_rewards [B]
If return_individual=True: dict with "aggregated" and individual model rewards
"""
batch_size = videos.shape[0]
individual_rewards: dict[str, torch.Tensor] = {}
# Collect rewards from all models
all_rewards = []
for i, (model, weight) in enumerate(zip(self.reward_models, self.reward_weights, strict=False)):
reward = model.compute_reward(videos, prompts, **kwargs)
assert reward.shape == (batch_size,), \
f"Reward model {i} returned shape {reward.shape}, expected ({batch_size},)"
# Optionally normalize individual rewards
if self.normalize_rewards:
reward = (reward - reward.mean()) / (reward.std() + 1e-8)
individual_rewards[f"reward_{type(model).__name__}"] = reward
all_rewards.append(weight * reward)
# Aggregate with weights
aggregated = sum(all_rewards)
if return_individual:
individual_rewards["aggregated"] = aggregated
return individual_rewards
return aggregated
def __repr__(self) -> str:
models_str = ", ".join([
f"{type(m).__name__}(w={w:.3f})"
for m, w in zip(self.reward_models, self.reward_weights, strict=False)
])
return f"MultiRewardAggregator({models_str})"
class ValueModel(nn.Module):
"""
Value function model wrapper for RL training.
The value model can either:
1. Share the transformer backbone with the policy (memory efficient)
2. Use a separate transformer (more flexible)
For now, this is a placeholder that will be expanded based on
the chosen architecture strategy.
"""
def __init__(
self,
transformer: nn.Module,
share_backbone: bool = False,
hidden_size: int | None = None
):
"""
Initialize value model.
Args:
transformer: Transformer model (policy or separate)
share_backbone: Whether to share backbone with policy
hidden_size: Hidden size for value head (inferred if None)
"""
super().__init__()
self.transformer = transformer
self.share_backbone = share_backbone
# Value head will be added later based on transformer architecture
# For now, just store the transformer reference
logger.info(
"Initialized ValueModel (share_backbone=%s)",
share_backbone
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.Tensor,
**kwargs: Any
) -> torch.Tensor:
"""
Forward pass to compute value predictions.
Args:
hidden_states: Latent states [B, C, T, H, W]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Timesteps [B]
**kwargs: Additional transformer arguments
Returns:
values: Value predictions [B]
"""
# TODO: Implement value prediction
# For now, return dummy values
batch_size = hidden_states.shape[0]
return torch.zeros(batch_size, device=hidden_states.device)
class DummyRewardModel(BaseRewardModel):
"""
Dummy VIDEO reward model for testing and development.
Returns random rewards in the range [0, 1] for VIDEO inputs.
This is a placeholder for testing the RL pipeline before real video reward models
are implemented.
NOTE: This does NOT actually evaluate video quality - it's just for testing!
"""
def __init__(self, mean: float = 0.5, std: float = 0.1):
super().__init__(model_path=None)
self.mean = mean
self.std = std
logger.info("Initialized DummyRewardModel (VIDEO) - mean=%.2f, std=%.2f", mean, std)
logger.warning(
"DummyRewardModel is for TESTING ONLY - does not evaluate actual video quality!"
)
def compute_reward(
self,
videos: torch.Tensor, # [B, T, C, H, W]
prompts: list[str],
**kwargs: Any
) -> torch.Tensor:
"""
Return random rewards for testing.
Args:
videos: Video sequences [B, T, C, H, W]
prompts: Text prompts
Returns:
Random rewards [B] in range [0, 1]
"""
batch_size = videos.shape[0]
num_frames = videos.shape[1]
logger.debug(
"DummyRewardModel processing %d videos with %d frames each",
batch_size,
num_frames
)
# Generate random rewards (not based on actual video content!)
rewards = torch.randn(batch_size, device=videos.device) * self.std + self.mean
return rewards.clamp(0.0, 1.0)
def load_model(self) -> None:
"""No model to load for dummy."""
pass
def create_reward_models(
reward_models: dict,
device: str = "cuda"
) -> MultiRewardAggregator:
"""
Factory function to create VIDEO reward models from configuration strings.
IMPORTANT: Only creates VIDEO reward models. Image-only reward models
(PickScore, ImageReward, GenEval, etc.) are NOT supported.
Args:
reward_models: dictionary of reward model names to weights
Example: {"paddle_ocr": 0.5, "video_score": 0.5}
device: Device to load models on
Returns:
MultiRewardAggregator with loaded VIDEO reward models
Supported VIDEO Reward Types:
- "paddle_ocr": PaddleOCR multi-frame video text recognition
- "video_score": Video aesthetic quality (multi-frame) - TODO
- "video_text_alignment": CLIP-based video-text similarity - TODO
- "temporal_coherence": Frame-to-frame consistency - TODO
- "motion_quality": Motion smoothness and realism - TODO
- "dummy": Random rewards for testing (VIDEO-aware)
NOT Supported (Image-Only):
- "pickscore": Image aesthetic (use "video_score" instead)
- "imagereward": Image quality (use "video_score" instead)
- "geneval": Image compositional (no video equivalent yet)
- Any single-frame reward models
Example:
>>> models = create_reward_models(
... reward_models={
... "paddle_ocr": 0.5,
... "video_text_alignment": 0.5
... },
... device="cuda"
... )
"""
assert reward_models, "No reward models specified. Please select at least 1 reward model"
types = [t.strip() for t in reward_models.keys()]
weights = list(reward_models.values())
assert len(types) == len(weights), \
f"Number of models ({len(types)}) must match number of weights ({len(weights)})"
# Create reward models based on types
models_list: list[BaseRewardModel] = []
for reward_type in types:
if reward_type == "dummy":
model = DummyRewardModel()
elif reward_type == "paddle_ocr":
logger.info("Creating PaddleOCR reward model")
model = OcrScorerVideo(device=device)
elif reward_type == "video_score":
# TODO: Implement VideoScore reward model (Phase 2)
logger.warning(
"VideoScore reward not implemented yet, using DummyRewardModel"
)
model = DummyRewardModel()
elif reward_type == "video_text_alignment":
# TODO: Implement VideoTextAlignment reward model (Phase 2)
logger.warning(
"VideoTextAlignment reward not implemented yet, using DummyRewardModel"
)
model = DummyRewardModel()
elif reward_type == "temporal_coherence":
# TODO: Implement TemporalCoherence reward model (Phase 2)
logger.warning(
"TemporalCoherence reward not implemented yet, using DummyRewardModel"
)
model = DummyRewardModel()
elif reward_type == "motion_quality":
# TODO: Implement MotionQuality reward model (Phase 2)
logger.warning(
"MotionQuality reward not implemented yet, using DummyRewardModel"
)
model = DummyRewardModel()
else:
logger.warning(
"Unknown VIDEO reward type '%s', using DummyRewardModel",
reward_type
)
model = DummyRewardModel()
models_list.append(model)
logger.info(
"Created MultiRewardAggregator with %d VIDEO reward models",
len(models_list)
)
return MultiRewardAggregator(models_list, weights, normalize_rewards=True)
File diff suppressed because it is too large Load Diff
+385
View File
@@ -0,0 +1,385 @@
# SPDX-License-Identifier: Apache-2.0
"""
Utility functions for RL/GRPO training.
"""
from typing import Any
import torch
import torch.nn.functional as F
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def compute_gae(
rewards: torch.Tensor,
values: torch.Tensor,
next_values: torch.Tensor,
dones: torch.Tensor | None = None,
gamma: float = 0.99,
lambda_: float = 0.95
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Compute Generalized Advantage Estimation (GAE-lambda).
GAE reduces variance in advantage estimation while allowing some bias.
This is a key component of modern policy gradient methods like PPO and GRPO.
Args:
rewards: Rewards at each step [B, T] or [B]
values: Value predictions at each step [B, T] or [B]
next_values: Value predictions at next step [B, T] or [B]
dones: Episode termination flags [B, T] or [B] (1 if done, 0 otherwise)
gamma: Discount factor
lambda_: GAE lambda parameter (0=TD(0), 1=Monte Carlo)
Returns:
advantages: GAE advantages [B, T] or [B]
returns: TD(lambda) returns [B, T] or [B]
Reference:
Schulman et al. "High-Dimensional Continuous Control Using Generalized Advantage Estimation"
https://arxiv.org/abs/1506.02438
"""
if dones is None:
dones = torch.zeros_like(rewards)
# Compute TD residuals: delta_t = r_t + gamma * V(s_{t+1}) - V(s_t)
deltas = rewards + gamma * next_values * (1.0 - dones) - values
# If single step (no time dimension), return directly
if deltas.dim() == 1:
advantages = deltas
returns = advantages + values
return advantages, returns
# Multi-step: compute GAE recursively
batch_size, num_steps = deltas.shape
advantages = torch.zeros_like(deltas)
gae = torch.zeros(batch_size, device=deltas.device)
# Backward pass to compute GAE
for t in reversed(range(num_steps)):
gae = deltas[:, t] + gamma * lambda_ * (1.0 - dones[:, t]) * gae
advantages[:, t] = gae
# Returns are advantages + values
returns = advantages + values
return advantages, returns
def normalize_advantages(
advantages: torch.Tensor,
epsilon: float = 1e-8
) -> torch.Tensor:
"""
Normalize advantages to have zero mean and unit variance.
This is a common practice in PPO and GRPO to stabilize training.
Args:
advantages: Raw advantages [B, ...]
epsilon: Small constant for numerical stability
Returns:
normalized_advantages: Normalized advantages [B, ...]
"""
mean = advantages.mean()
std = advantages.std()
return (advantages - mean) / (std + epsilon)
#TODO(jiali): refactor into algorithm
def compute_grpo_policy_loss(
log_probs: torch.Tensor,
old_log_probs: torch.Tensor,
advantages: torch.Tensor,
clip_range: float = 0.2,
use_ratio_norm: bool = True,
max_importance_ratio: float = 10.0
) -> tuple[torch.Tensor, dict[str, Any]]:
"""
Compute GRPO policy loss with importance sampling and clipping.
This implements the core GRPO objective with safety mechanisms from GRPO-Guard:
- Importance ratio clipping (PPO-style)
- RatioNorm correction (GRPO-Guard)
- Ratio clamping for extreme values
Args:
log_probs: Log probabilities from current policy [B]
old_log_probs: Log probabilities from old policy [B]
advantages: Advantages [B]
clip_range: Clipping range for importance ratios
use_ratio_norm: Apply RatioNorm correction (GRPO-Guard)
max_importance_ratio: Maximum importance ratio before clamping
Returns:
loss: Policy loss (scalar)
info: Dictionary with diagnostic information
Reference:
- PPO: Schulman et al. "Proximal Policy Optimization Algorithms"
- GRPO-Guard: RatioNorm and gradient reweighting
"""
# Compute importance ratio: r_t = pi_new(a|s) / pi_old(a|s)
log_ratio = log_probs - old_log_probs
ratio = torch.exp(log_ratio)
# Clamp extreme ratios for numerical stability
ratio = torch.clamp(ratio, 1.0 / max_importance_ratio, max_importance_ratio)
# RatioNorm correction (GRPO-Guard)
# Corrects bias in importance sampling when ratio >> 1
if use_ratio_norm:
ratio_mean = ratio.mean()
ratio = ratio / (ratio_mean + 1e-8)
# Clipped surrogate objective
ratio_clipped = torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range)
surrogate1 = ratio * advantages
surrogate2 = ratio_clipped * advantages
policy_loss = -torch.min(surrogate1, surrogate2).mean()
# Compute diagnostics
with torch.no_grad():
# Clip fraction: how often ratios were clipped
clip_fraction = ((ratio < 1.0 - clip_range) | (ratio > 1.0 + clip_range)).float().mean()
# KL divergence (approximate)
kl_div = log_ratio.mean()
# Importance ratio stats
importance_ratio_mean = ratio.mean()
importance_ratio_std = ratio.std()
info = {
"policy_loss": policy_loss.item(),
"clip_fraction": clip_fraction.item(),
"kl_divergence": kl_div.item(),
"importance_ratio_mean": importance_ratio_mean.item(),
"importance_ratio_std": importance_ratio_std.item(),
}
return policy_loss, info
def compute_value_loss(
values: torch.Tensor,
returns: torch.Tensor,
old_values: torch.Tensor | None = None,
clip_range: float = 0.2,
use_clipping: bool = True
) -> tuple[torch.Tensor, dict[str, Any]]:
"""
Compute value function loss with optional clipping.
Args:
values: Value predictions from current model [B]
returns: Target returns (from GAE) [B]
old_values: Value predictions from old model [B] (for clipping)
clip_range: Clipping range for value updates
use_clipping: Whether to use clipped value loss (PPO-style)
Returns:
loss: Value loss (scalar)
info: Dictionary with diagnostic information
"""
# Standard MSE loss
value_loss_unclipped = F.mse_loss(values, returns, reduction="none")
# Clipped value loss (PPO-style)
if use_clipping and old_values is not None:
values_clipped = old_values + torch.clamp(
values - old_values,
-clip_range,
clip_range
)
value_loss_clipped = F.mse_loss(values_clipped, returns, reduction="none")
value_loss = torch.max(value_loss_unclipped, value_loss_clipped).mean()
else:
value_loss = value_loss_unclipped.mean()
# Compute diagnostics
with torch.no_grad():
explained_variance = 1.0 - (returns - values).var() / (returns.var() + 1e-8)
info = {
"value_loss": value_loss.item(),
"explained_variance": explained_variance.item(),
"value_mean": values.mean().item(),
"value_std": values.std().item(),
}
return value_loss, info
def compute_policy_entropy(log_probs: torch.Tensor) -> torch.Tensor:
"""
Compute policy entropy for exploration bonus.
Args:
log_probs: Log probabilities [B]
Returns:
entropy: Mean entropy across batch (scalar)
"""
# For continuous actions: H = -log_prob (assuming Gaussian)
# For discrete: H = -sum(p * log(p))
# Here we use a simple approximation
entropy = -log_probs.mean()
return entropy
def apply_gradient_reweighting(
gradients: torch.Tensor,
timesteps: torch.Tensor,
num_train_timesteps: int = 1000
) -> torch.Tensor:
"""
Apply GRPO-Guard gradient reweighting across denoising steps.
This reweights gradients based on the timestep to balance learning
across different noise levels.
Args:
gradients: Gradients to reweight [B, ...]
timesteps: Timesteps at which gradients were computed [B]
num_train_timesteps: Total number of training timesteps
Returns:
reweighted_gradients: Reweighted gradients [B, ...]
"""
# Compute timestep weights (higher weight for later timesteps)
# This is a simple linear weighting, can be made more sophisticated
timestep_weights = 1.0 + (timesteps.float() / num_train_timesteps)
timestep_weights = timestep_weights.view(-1, *([1] * (gradients.dim() - 1)))
return gradients * timestep_weights
def sample_random_timesteps(
batch_size: int,
min_timestep: int,
max_timestep: int,
device: torch.device,
generator: torch.Generator | None = None
) -> torch.Tensor:
"""
Sample random timesteps for noise injection (Flow-GRPO-Fast).
Args:
batch_size: Number of samples
min_timestep: Minimum timestep
max_timestep: Maximum timestep
device: Device for tensor
generator: Random generator for reproducibility
Returns:
timesteps: Random timesteps [B]
"""
if generator is not None:
timesteps = torch.randint(
min_timestep,
max_timestep + 1,
(batch_size,),
device=device,
generator=generator
)
else:
timesteps = torch.randint(
min_timestep,
max_timestep + 1,
(batch_size,),
device=device
)
return timesteps
def compute_reward_statistics(
rewards: torch.Tensor
) -> dict[str, float]:
"""
Compute statistics for reward distribution.
Args:
rewards: Reward values [B]
Returns:
stats: Dictionary with mean, std, min, max
"""
return {
"reward_mean": rewards.mean().item(),
"reward_std": rewards.std().item(),
"reward_min": rewards.min().item(),
"reward_max": rewards.max().item(),
}
def check_early_stopping(
kl_divergence: float,
target_kl: float
) -> bool:
"""
Check if training should stop early based on KL divergence.
Args:
kl_divergence: Current KL divergence
target_kl: Target KL threshold
Returns:
should_stop: True if KL exceeds target
"""
if kl_divergence > target_kl:
logger.warning(
"Early stopping triggered: KL divergence %.4f > target %.4f",
kl_divergence,
target_kl
)
return True
return False
def compute_log_probs_from_model_output(
model_output: torch.Tensor,
target: torch.Tensor,
noise_level: float = 0.1
) -> torch.Tensor:
"""
Compute log probabilities from model predictions.
For diffusion models, we approximate log probabilities using the
negative squared error (assuming Gaussian likelihood).
Args:
model_output: Model predictions [B, C, T, H, W]
target: Target values [B, C, T, H, W]
noise_level: Assumed noise level (std) for Gaussian likelihood
Returns:
log_probs: Log probabilities [B]
"""
# Compute mean squared error per sample
mse = ((model_output - target) ** 2).flatten(1).mean(dim=1)
# Log probability under Gaussian: log p(x) = -0.5 * (x - mu)^2 / sigma^2 + const
log_probs = -0.5 * mse / (noise_level ** 2)
return log_probs
def check_for_nan_inf(tensor: torch.Tensor, name: str) -> None:
"""
Check tensor for NaN or Inf values and raise error if found.
Args:
tensor: Tensor to check
name: Name for error message
"""
if torch.isnan(tensor).any():
raise ValueError(f"{name} contains NaN values")
if torch.isinf(tensor).any():
raise ValueError(f"{name} contains Inf values")
+189
View File
@@ -0,0 +1,189 @@
# SPDX-License-Identifier: Apache-2.0
"""
Per-prompt statistics tracking for GRPO training.
This module ports the PerPromptStatTracker from FlowGRPO to FastVideo.
It tracks reward statistics per unique prompt and computes normalized advantages.
Ported from:
- flow_grpo/flow_grpo/stat_tracking.py
Key adaptations:
1. Uses FastVideo's logging instead of print statements
2. Works with single GPU (no distributed logic)
3. Supports numpy arrays and torch tensors
"""
import numpy as np
from typing import Union
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class PerPromptStatTracker:
"""
Tracks reward statistics per unique prompt for advantage normalization.
This class maintains running statistics (mean, std) for each unique prompt
and computes normalized advantages using either per-prompt or global statistics.
Used in GRPO training to normalize advantages within groups of samples
generated from the same prompt, which helps stabilize training when different
prompts have different reward scales.
"""
def __init__(self, global_std: bool = False):
"""
Initialize the per-prompt stat tracker.
Args:
global_std: If True, use global std across all rewards for normalization.
If False, use per-prompt std (default, recommended for GRPO).
"""
self.global_std = global_std
self.stats: dict[str, list] = {} # Maps prompt -> list of rewards
self.history_prompts: set[int] = set() # Set of hashed prompts seen
def update(
self,
prompts: Union[list[str], np.ndarray],
rewards: Union[list[float], np.ndarray, torch.Tensor],
type: str = 'grpo'
) -> np.ndarray:
"""
Update statistics and compute normalized advantages.
Args:
prompts: List or array of prompt strings (one per sample)
rewards: Array or tensor of reward values (one per sample)
type: Advantage computation type:
- 'grpo': Normalize by (reward - mean) / std (default)
- 'rwr': Return rewards as-is (reward-weighted regression)
- 'sft': Binary advantages (1 for max, 0 otherwise)
- 'dpo': DPO-style advantages (1 for max, -1 for min)
Returns:
advantages: Normalized advantages array [num_samples] or [num_samples, ...]
Shape matches rewards shape
"""
# Convert to numpy arrays
prompts = np.array(prompts)
if isinstance(rewards, torch.Tensor):
rewards = rewards.detach().cpu().numpy()
rewards = np.array(rewards, dtype=np.float64)
# Ensure rewards are 1D (one reward per sample)
# FlowGRPO expects rewards to be aggregated per sample
if rewards.ndim > 1:
# If multi-dimensional, flatten or take mean
# For [B, num_steps] shape, we typically want one reward per sample
# So we take the mean across timesteps
if rewards.ndim == 2:
# Assume shape is [B, num_steps] - take mean across timesteps
rewards = rewards.mean(axis=1)
else:
# Flatten and take mean for higher dimensions
rewards = rewards.reshape(len(prompts), -1).mean(axis=1)
# Ensure prompts and rewards have matching lengths
assert len(prompts) == len(rewards), \
f"Prompts ({len(prompts)}) and rewards ({len(rewards)}) must have same length"
unique_prompts = np.unique(prompts)
advantages = np.zeros_like(rewards, dtype=np.float64)
# First pass: collect rewards for each prompt
for prompt in unique_prompts:
prompt_mask = prompts == prompt
prompt_rewards = rewards[prompt_mask]
# Store rewards in stats
if prompt not in self.stats:
self.stats[prompt] = []
self.stats[prompt].extend(prompt_rewards.tolist())
self.history_prompts.add(hash(prompt))
# Second pass: compute statistics and advantages
for prompt in unique_prompts:
prompt_mask = prompts == prompt
prompt_rewards = rewards[prompt_mask]
# Stack all historical rewards for this prompt
if len(self.stats[prompt]) > 0:
all_prompt_rewards = np.array(self.stats[prompt])
else:
all_prompt_rewards = prompt_rewards
# Compute mean and std
mean = np.mean(all_prompt_rewards, axis=0, keepdims=True)
if self.global_std:
# Use global std across all rewards
std = np.std(rewards, axis=0, keepdims=True) + 1e-4
else:
# Use per-prompt std
std = np.std(all_prompt_rewards, axis=0, keepdims=True) + 1e-4
# Compute advantages based on type
if type == 'grpo':
# GRPO: normalize by (reward - mean) / std
advantages[prompt_mask] = (prompt_rewards - mean) / std
elif type == 'rwr':
# Reward-weighted regression: use rewards as-is
advantages[prompt_mask] = prompt_rewards
elif type == 'sft':
# Supervised fine-tuning: binary (1 for max, 0 otherwise)
max_reward = np.max(prompt_rewards)
advantages[prompt_mask] = (prompt_rewards == max_reward).astype(np.float64)
elif type == 'dpo':
# DPO-style: 1 for max, -1 for min
prompt_rewards_tensor = torch.tensor(prompt_rewards)
max_idx = torch.argmax(prompt_rewards_tensor)
min_idx = torch.argmin(prompt_rewards_tensor)
# If all rewards are the same, use first two indices
if max_idx == min_idx:
min_idx = torch.tensor(0)
max_idx = torch.tensor(1) if len(prompt_rewards_tensor) > 1 else torch.tensor(0)
result = torch.zeros_like(prompt_rewards_tensor, dtype=torch.float64)
result[max_idx] = 1.0
result[min_idx] = -1.0
advantages[prompt_mask] = result.numpy()
else:
raise ValueError(f"Unknown advantage type: {type}. Must be one of: 'grpo', 'rwr', 'sft', 'dpo'")
return advantages
def get_stats(self) -> tuple[float, int]:
"""
Get statistics about tracked prompts.
Returns:
avg_group_size: Average number of samples per unique prompt
history_prompts: Number of unique prompts seen (across all updates)
"""
if not self.stats:
avg_group_size = 0.0
else:
total_samples = sum(len(v) for v in self.stats.values())
avg_group_size = total_samples / len(self.stats)
history_prompts = len(self.history_prompts)
return avg_group_size, history_prompts
def clear(self) -> None:
"""
Clear all statistics (but keep history_prompts for tracking).
This is typically called after each epoch to reset per-epoch statistics
while maintaining a record of all prompts seen during training.
"""
self.stats = {}
logger.debug("Cleared per-prompt statistics (kept %d unique prompts in history)",
len(self.history_prompts))
+877
View File
@@ -0,0 +1,877 @@
# SPDX-License-Identifier: Apache-2.0
"""
GRPO utilities for Wan model in FastVideo.
This module ports the SDE step and pipeline functions from FlowGRPO to work with
FastVideo's scheduler and pipeline interfaces.
Ported from:
- flow_grpo/flow_grpo/diffusers_patch/wan_pipeline_with_logprob.py
Key adaptations:
1. Uses FastVideo's FlowUniPCMultistepScheduler instead of diffusers' UniPCMultistepScheduler
2. Works with FastVideo's WanPipeline (ComposedPipelineBase) instead of diffusers' WanPipeline
3. Direct module access via pipeline.get_module() instead of pipeline attributes
4. Simplified prompt encoding (direct text encoder usage instead of pipeline stages)
"""
import math
import time
from typing import Any
import torch
from tqdm import tqdm
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.utils import get_compute_dtype
# for test_wan_transformer2
import os
from diffusers import WanTransformer3DModel
from fastvideo.utils import maybe_download_model
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.models.loader.component_loader import TransformerLoader
logger = init_logger(__name__)
def test_wan_transformer():
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(
dit_config=WanVideoConfig(),
dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = WanTransformer3DModel.from_pretrained(
TRANSFORMER_PATH, device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
seq_len = 30
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(data_type="dummy", )
with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
return_dict=False,
)[0]
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
'''
INFO 01-19 22:53:46 [wan_grpo_utils.py:74] Model 1 weight sum: 395834.3506456231████ | 1/2 [00:00<00:00, 7.84it/s]
INFO 01-19 22:53:46 [wan_grpo_utils.py:75] Model 1 weight mean: 0.0002789536598289884
INFO 01-19 22:53:47 [wan_grpo_utils.py:83] Model 2 weight sum: 395834.3506456231
INFO 01-19 22:53:47 [wan_grpo_utils.py:84] Model 2 weight mean: 0.0002789536598289884
INFO 01-19 22:53:47 [wan_grpo_utils.py:87] Weight sum difference: 0.0
INFO 01-19 22:53:47 [wan_grpo_utils.py:89] Weight mean difference: 0.0
INFO 01-19 22:53:54 [wan_grpo_utils.py:145] Max Diff: 0.08203125
INFO 01-19 22:53:54 [wan_grpo_utils.py:146] Mean Diff: 0.01129150390625
'''
def test_wan_transformer2(model2):
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
logger.info("loading model1 transformer weight")
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
model1 = WanTransformer3DModel.from_pretrained(
TRANSFORMER_PATH,
device=device,
torch_dtype=precision,
).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
seq_len = 30
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(
batch_size,
16,
21,
160,
90,
device=device,
dtype=precision,
)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(
batch_size,
seq_len + 1,
4096,
device=device,
dtype=precision,
)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(data_type="dummy", )
with torch.amp.autocast("cuda", dtype=precision):
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
return_dict=False,
)[0]
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
)
# Print basic stats for debugging (cast to float32 for stability)
out1 = output1.detach().float()
out2 = output2.detach().float()
logger.info(
"output1 stats: min=%s max=%s mean=%s std=%s",
out1.min().item(),
out1.max().item(),
out1.mean().item(),
out1.std(unbiased=False).item(),
)
logger.info(
"output2 stats: min=%s max=%s mean=%s std=%s",
out2.min().item(),
out2.max().item(),
out2.mean().item(),
out2.std(unbiased=False).item(),
)
# Check if outputs have the same shape
assert (output1.shape == output2.shape
), f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert (output1.dtype == output2.dtype
), f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
'''
when --dit_precision "bf16", use_fsdp hardcoded to False:
INFO 01-19 22:01:24 [wan_grpo_utils.py:65] Model 1 weight sum: 395834.3506456231████████████████████████████████████████████████████████| 2/2 [00:00<00:00, 3.25it/s]
INFO 01-19 22:01:24 [wan_grpo_utils.py:66] Model 1 weight mean: 0.0002789536598289884
INFO 01-19 22:01:24 [wan_grpo_utils.py:75] Model 2 weight sum: 395125.463677882
INFO 01-19 22:01:24 [wan_grpo_utils.py:76] Model 2 weight mean: 0.0002739000890162162
INFO 01-19 22:01:24 [wan_grpo_utils.py:79] Weight sum difference: 708.8869677411276
INFO 01-19 22:01:24 [wan_grpo_utils.py:81] Weight mean difference: 5.053570812772192e-06
INFO 01-19 22:01:32 [wan_grpo_utils.py:139] output1 stats: min=-2.28125 max=1.921875 mean=-0.16638492047786713 std=0.458170622587204
INFO 01-19 22:01:32 [wan_grpo_utils.py:146] output2 stats: min=-2.296875 max=1.90625 mean=-0.166452556848526 std=0.4579130709171295
INFO 01-19 22:01:32 [wan_grpo_utils.py:165] Max Diff: 0.08984375
INFO 01-19 22:01:32 [wan_grpo_utils.py:166] Mean Diff: 0.0120849609375
when --dit_precision "fp32", use_fsdp not changed:
'''
def sde_step_with_logprob(
scheduler: FlowUniPCMultistepScheduler,
model_output: torch.FloatTensor,
timestep: float | torch.FloatTensor,
sample: torch.FloatTensor,
prev_sample: torch.FloatTensor | None = None,
generator: torch.Generator | None = None,
deterministic: bool = False,
return_pixel_log_prob: bool = False,
return_dt_and_std_dev_t: bool = False
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ...]:
"""
Predict the sample from the previous timestep by reversing the SDE.
This function propagates the flow process from the learned model outputs
(most often the predicted velocity) and computes log probabilities.
Ported from FlowGRPO's sde_step_with_logprob to work with FastVideo's
FlowUniPCMultistepScheduler.
Args:
scheduler: FastVideo FlowUniPCMultistepScheduler instance
model_output: The direct output from learned flow model
timestep: The current discrete timestep in the diffusion chain
sample: A current instance of a sample created by the diffusion process
prev_sample: Optional previous sample (if provided, used instead of sampling)
generator: Optional random number generator
deterministic: If True, no noise is added (deterministic sampling)
return_pixel_log_prob: If True, return pixel-level log probabilities (not used)
return_dt_and_std_dev_t: If True, return dt and std_dev_t separately
Returns:
If return_dt_and_std_dev_t=True:
(prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt)
Otherwise:
(prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt)
"""
# # Convert all variables to fp32 for numerical stability
# model_output = model_output.float()
# sample = sample.float()
# if prev_sample is not None:
# prev_sample = prev_sample.float()
# Get step indices for current and previous timesteps
# Handle both single timestep and batch of timesteps
if isinstance(timestep, torch.Tensor):
if timestep.ndim == 0:
timestep = timestep.unsqueeze(0)
step_indices = [
scheduler.index_for_timestep(t.item()) for t in timestep
]
else:
step_indices = [scheduler.index_for_timestep(timestep)]
prev_step_indices = [step + 1 for step in step_indices]
# Move sigmas to sample device
sigmas = scheduler.sigmas.to(sample.device)
# myregion debug: hardcode sigmas to flow_grpo's
sigmas = torch.Tensor([
0.9997, 0.9824, 0.9639, 0.9441, 0.9227, 0.8996, 0.8746, 0.8475, 0.8178,
0.7853, 0.7496, 0.7102, 0.6663, 0.6173, 0.5621, 0.4997, 0.4283, 0.3459,
0.2498, 0.1362, 0.0000
]).to(sample.device, sample.dtype)
# end region
# Get sigma values for current and previous steps
sigma = sigmas[step_indices].view(-1, 1, 1, 1, 1)
sigma_prev = sigmas[prev_step_indices].view(-1, 1, 1, 1, 1)
sigma_max = sigmas[0].item() # First sigma (highest)
sigma_min = sigmas[-1].item() # Last sigma (lowest)
dt = sigma_prev - sigma
# myregion debug
print(f"[DEBUG]: sigma_max: {sigma_max}, sigma_min: {sigma_min}, dt: {dt}")
print(f"[DEBUG]: in sde_step_with_logprob(), timestep: {timestep}")
print(f"[DEBUG]: in sde_step_with_logprob(), sigmas: {sigmas}")
print(f"[DEBUG]: in sde_step_with_logprob(), step_indices: {step_indices}")
print(
f"[DEBUG]: in sde_step_with_logprob(), prev_step_indices: {prev_step_indices}"
)
'''
[DEBUG]: in sde_step_with_logprob(), timestep: tensor([428, 428, 428, 428], device='cuda:0')
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
0.2499, 0.1363, 0.0000], device='cuda:0')
[DEBUG]: in sde_step_with_logprob(), step_indices: [16, 16, 16, 16]
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [17, 17, 17, 17]
DEBUG]: in sde_step_with_logprob(), timestep: tensor([249], device='cuda:0')███████████████▎ | 18/20 [00:04<00:00, 3.87step/s, step_time=0.26s, timestep=346.0]
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
0.2499, 0.1363, 0.0000], device='cuda:0')
[DEBUG]: in sde_step_with_logprob(), step_indices: [18]
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [19]
[DEBUG]: in sde_step_with_logprob(), timestep: tensor([617, 617, 617, 617], device='cuda:0')
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
0.2499, 0.1363, 0.0000], device='cuda:0')
[DEBUG]: in sde_step_with_logprob(), step_indices: [13, 13, 13, 13]
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [14, 14, 14, 14]
'''
# endregion
# Compute std_dev_t and prev_sample_mean using SDE formulation
std_dev_t = sigma_min + (sigma_max - sigma_min) * sigma
prev_sample_mean = (sample * (1 + std_dev_t**2 / (2 * sigma) * dt) +
model_output * (1 + std_dev_t**2 * (1 - sigma) /
(2 * sigma)) * dt)
if prev_sample is not None and generator is not None:
raise ValueError(
"Cannot pass both generator and prev_sample. Please make sure that either `generator` or"
" `prev_sample` stays `None`.")
# Sample prev_sample if not provided
if prev_sample is None:
variance_noise = randn_tensor(
model_output.shape,
generator=generator,
device=model_output.device,
dtype=model_output.dtype,
)
sqrt_dt = torch.sqrt(-1 * dt) # dt is negative (going backwards)
prev_sample = prev_sample_mean + std_dev_t * sqrt_dt * variance_noise
else:
sqrt_dt = torch.sqrt(-1 * dt)
# No noise is added during evaluation (deterministic)
if deterministic:
prev_sample = sample + dt * model_output
sqrt_dt = torch.sqrt(-1 * dt)
# Compute log probability: log p(prev_sample | sample, model_output)
# Assuming Gaussian distribution: N(prev_sample_mean, (std_dev_t * sqrt_dt)^2)
std_dev_sqrt_dt = std_dev_t * sqrt_dt
log_prob = (
-((prev_sample.detach() - prev_sample_mean)**2) /
(2 * (std_dev_sqrt_dt**2)) - torch.log(
std_dev_sqrt_dt + 1e-8) # Add small epsilon for numerical stability
- torch.log(
torch.sqrt(2 * torch.as_tensor(math.pi, device=sample.device))))
# Mean along all but batch dimension
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
if return_dt_and_std_dev_t:
return prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt
return prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt
def wan_pipeline_with_logprob(
pipeline,
prompt: str | list[str] = None,
negative_prompt: str | list[str] = None,
height: int = 480,
width: int = 832,
num_frames: int = 81,
num_inference_steps: int = 50,
guidance_scale: float = 5.0,
num_videos_per_prompt: int | None = 1,
generator: torch.Generator | list[torch.Generator] | None = None,
latents: torch.Tensor | None = None,
prompt_embeds: torch.Tensor | None = None,
negative_prompt_embeds: torch.Tensor | None = None,
output_type: str | None = "pt",
return_dict: bool = False,
attention_kwargs: dict[str, Any] | None = None,
max_sequence_length: int = 512,
deterministic: bool = False,
kl_reward: float = 0.0,
return_pixel_log_prob: bool = False,
) -> tuple[torch.Tensor, list[torch.Tensor], list[torch.Tensor],
list[torch.Tensor], torch.Tensor | None]:
"""
Wan pipeline with log probability computation for GRPO training.
Ported from FlowGRPO's wan_pipeline_with_logprob to work with FastVideo's WanPipeline.
This function generates videos and computes log probabilities at each denoising step.
Args:
pipeline: FastVideo WanPipeline instance
prompt: Text prompt(s) for generation
negative_prompt: Negative prompt(s) for classifier-free guidance
height: Height of generated video
width: Width of generated video
num_frames: Number of frames in generated video
num_inference_steps: Number of denoising steps
guidance_scale: Classifier-free guidance scale
num_videos_per_prompt: Number of videos to generate per prompt
generator: Random generator for reproducibility
latents: Optional initial latents
prompt_embeds: Optional pre-computed prompt embeddings
negative_prompt_embeds: Optional pre-computed negative prompt embeddings
output_type: Output type ("pt" for PyTorch tensor, "np" for numpy, "latent" for latents only)
return_dict: Whether to return dict (not used, always returns tuple)
attention_kwargs: Optional attention kwargs
max_sequence_length: Maximum sequence length for text encoding
deterministic: If True, use deterministic sampling (no noise)
kl_reward: KL reward coefficient (if > 0, computes KL divergence)
return_pixel_log_prob: If True, return pixel-level log probabilities (not used)
Returns:
Tuple of:
- video: Generated video tensor [B, C, T, H, W] or latents if output_type="latent"
- all_latents: List of latents at each step [num_steps+1] of shape [B, C, T, H, W]
- all_log_probs: List of log probabilities at each step [num_steps] of shape [B]
- all_kl: List of KL divergences at each step [num_steps] of shape [B] (if kl_reward > 0)
- prompt_ids: Tokenized prompt IDs [B, seq_len] (None if prompt_embeds were provided)
"""
# Get device from transformer
transformer = pipeline.get_module("transformer")
# myregion debug: test transformer output
logger.info("testing transformer, running test_wan_transformer2")
test_wan_transformer()
# test_wan_transformer2(transformer)
# endregion
# hardcode dtype for debug
# transformer_dtype = torch.float32
# use get_compute_dtype() to get dtype based on mixed precision
transformer_dtype = get_compute_dtype()
logger.info(f"[DEBUG]: transformer_dtype: {transformer_dtype}")
# Get scheduler and other modules
scheduler = pipeline.get_module("scheduler")
vae = pipeline.get_module("vae")
# Determine batch size
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
elif prompt_embeds is not None:
batch_size = prompt_embeds.shape[0]
else:
raise ValueError("Either prompt or prompt_embeds must be provided")
# Encode prompts if not provided
prompt_ids = None
if prompt_embeds is None:
# Encode prompts directly using text encoder and tokenizer
# This is a simplified encoding - for full pipeline encoding, use TextEncodingStage
text_encoder = pipeline.get_module("text_encoder")
tokenizer = pipeline.get_module("tokenizer")
# Normalize to list
if isinstance(prompt, str):
prompts_list = [prompt]
else:
prompts_list = prompt
# Tokenize prompts
text_inputs = tokenizer(prompts_list,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
return_tensors="pt").to(pipeline.device)
# Store prompt_ids for return
prompt_ids = text_inputs["input_ids"]
# Encode with text encoder
with torch.no_grad():
outputs = text_encoder(
text_inputs["input_ids"],
attention_mask=text_inputs["attention_mask"],
output_hidden_states=True,
)
# Get last hidden state (Wan typically uses last hidden state)
prompt_embeds = outputs.last_hidden_state
# Encode negative prompts if CFG is enabled
if guidance_scale > 1.0:
if negative_prompt is None:
negative_prompt = [""] * len(prompts_list)
elif isinstance(negative_prompt, str):
negative_prompt = [negative_prompt]
neg_text_inputs = tokenizer(negative_prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
return_tensors="pt").to(pipeline.device)
with torch.no_grad():
neg_outputs = text_encoder(
neg_text_inputs["input_ids"],
attention_mask=neg_text_inputs["attention_mask"],
output_hidden_states=True,
)
negative_prompt_embeds = neg_outputs.last_hidden_state
else:
negative_prompt_embeds = None
# myregion Debug: Print shapes of prompt embeddings
logger.info(
f"After encoding - prompt_embeds shape: {prompt_embeds.shape if prompt_embeds is not None else None}"
)
logger.info(
f"After encoding - negative_prompt_embeds shape: {negative_prompt_embeds.shape if negative_prompt_embeds is not None else None}"
)
logger.info(
f"After encoding - prompt_embeds dtype: {prompt_embeds.dtype if prompt_embeds is not None else None}"
)
logger.info(
f"After encoding - negative_prompt_embeds dtype: {negative_prompt_embeds.dtype if negative_prompt_embeds is not None else None}"
)
'''
INFO 01-17 05:31:13 [wan_grpo_utils.py:290] After encoding - prompt_embeds shape: torch.Size([4, 512, 4096])
INFO 01-17 05:31:13 [wan_grpo_utils.py:291] After encoding - negative_prompt_embeds shape: None
INFO 01-17 05:31:13 [wan_grpo_utils.py:292] After encoding - prompt_embeds dtype: torch.float32
INFO 01-17 05:31:13 [wan_grpo_utils.py:293] After encoding - negative_prompt_embeds dtype: None
'''
# endregion
# logger.info("wan_pipeline_with_logprob's transformer class type: %s", type(transformer))
# logger.info("Variables in transformer: %s", str(dir(transformer)))
prompt_embeds = prompt_embeds.to(transformer_dtype)
if negative_prompt_embeds is not None:
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
# Prepare timesteps
scheduler.set_timesteps(num_inference_steps, device=pipeline.device)
timesteps = scheduler.timesteps
# Prepare latent variables
num_channels_latents = transformer.config.in_channels
vae = pipeline.get_module("vae")
# Get VAE scale factors
vae_scale_factor_spatial = vae.spatial_compression_ratio
vae_scale_factor_temporal = vae.temporal_compression_ratio
if latents is None:
# Generate random latents
# Note: num_frames in latents accounts for temporal compression
num_latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
latents_shape = (
batch_size * num_videos_per_prompt,
num_channels_latents,
num_latent_frames,
height // vae_scale_factor_spatial,
width // vae_scale_factor_spatial,
)
if generator is not None:
if isinstance(generator, list):
latents = [
torch.randn(
latents_shape[1:],
generator=gen,
device=pipeline.device,
dtype=transformer_dtype,
) for gen in generator
]
latents = torch.stack(latents, dim=0)
else:
latents = torch.randn(
latents_shape,
generator=generator,
device=pipeline.device,
dtype=transformer_dtype,
)
else:
latents = torch.randn(latents_shape,
device=pipeline.device,
dtype=transformer_dtype)
else:
latents = latents.to(device=pipeline.device, dtype=transformer_dtype)
# myregion Debug: Print latents shape, dtype, and value range
logger.info("=" * 80)
logger.info("Latents Debug Information:")
logger.info(f" Shape: {latents.shape}")
logger.info(f" Dtype: {latents.dtype}")
logger.info(f" Min value: {latents.min().item():.6f}")
logger.info(f" Max value: {latents.max().item():.6f}")
logger.info(f" Mean value: {latents.mean().item():.6f}")
logger.info(f" Std value: {latents.std().item():.6f}")
logger.info(f" Device: {latents.device}")
logger.info("=" * 80)
'''
INFO 01-17 07:41:33 [wan_grpo_utils.py:355] ================================================================================
INFO 01-17 07:41:33 [wan_grpo_utils.py:356] Latents Debug Information:
INFO 01-17 07:41:33 [wan_grpo_utils.py:357] Shape: torch.Size([4, 16, 9, 30, 52])
INFO 01-17 07:41:33 [wan_grpo_utils.py:358] Dtype: torch.bfloat16
INFO 01-17 07:41:33 [wan_grpo_utils.py:359] Min value: -4.500000
INFO 01-17 07:41:33 [wan_grpo_utils.py:360] Max value: 4.656250
INFO 01-17 07:41:33 [wan_grpo_utils.py:361] Mean value: 0.000111
INFO 01-17 07:41:33 [wan_grpo_utils.py:362] Std value: 1.000000
INFO 01-17 07:41:33 [wan_grpo_utils.py:363] Device: cuda:0
INFO 01-17 07:41:33 [wan_grpo_utils.py:364] ================================================================================
'''
# endregion
all_latents = [latents]
all_log_probs = []
all_kl = []
# myregion Debug
logger.info("Tensor type issue debugging:")
logger.info(f"latents: {type(latents)}")
logger.info(f"prompt_embeds: {type(prompt_embeds)}")
logger.info(
f"[DEBUG]: before denoising loop: type(timesteps): {type(timesteps)}")
logger.info(
f"[DEBUG]: before denoising loop: timesteps.shape: {timesteps.shape}")
# endregion
# Progress bar for denoising loop
progress_bar = tqdm(enumerate(timesteps),
total=len(timesteps),
desc="Denoising steps",
unit="step")
for i, t in progress_bar:
step_start_time = time.time()
latents_ori = latents.clone()
timestep = t.expand(latents.shape[0]) if isinstance(
t, torch.Tensor) else torch.tensor([t] * latents.shape[0],
device=pipeline.device)
logger.info(
f"[DEBUG]: before set_forward_context: current_timestep=i:{i}")
# Predict noise with transformer
with set_forward_context(
current_timestep=t.item(),
attn_metadata=None,
forward_batch=None,
):
noise_pred = transformer(
hidden_states=latents,
timestep=timestep,
encoder_hidden_states=prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
noise_pred = noise_pred.to(prompt_embeds.dtype)
# Classifier-free guidance
if guidance_scale > 1.0:
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=None,
):
noise_uncond = transformer(
hidden_states=latents,
timestep=timestep,
encoder_hidden_states=negative_prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
noise_pred = noise_uncond + guidance_scale * (noise_pred -
noise_uncond)
# SDE step with log probability
latents, log_prob, prev_latents_mean, std_dev_t = sde_step_with_logprob(
scheduler,
noise_pred, #.float(),
t.unsqueeze(0) if isinstance(t, torch.Tensor) else t,
latents, #.float(),
deterministic=deterministic,
return_pixel_log_prob=return_pixel_log_prob)
# sde_step_with_logprob returns fp32
# latents = latents.to(transformer_dtype)
prev_latents = latents.clone()
all_latents.append(latents)
all_log_probs.append(log_prob)
# Compute KL divergence if kl_reward > 0 (for KL reward in sampling)
if kl_reward > 0 and not deterministic:
# Use reference model (disable adapter if using LoRA)
latent_model_input_ref = torch.cat(
[latents_ori] * 2) if guidance_scale > 1.0 else latents_ori
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=None,
):
with transformer.disable_adapter() if hasattr(
transformer, 'disable_adapter') else torch.no_grad():
noise_pred_ref = transformer(
hidden_states=latent_model_input_ref,
timestep=timestep,
encoder_hidden_states=prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
noise_pred_ref = noise_pred_ref.to(prompt_embeds.dtype)
# Perform guidance for reference model
if guidance_scale > 1.0:
noise_pred_uncond_ref, noise_pred_text_ref = noise_pred_ref.chunk(
2)
noise_pred_ref = noise_pred_uncond_ref + guidance_scale * (
noise_pred_text_ref - noise_pred_uncond_ref)
# Compute reference log prob
_, ref_log_prob, ref_prev_latents_mean, ref_std_dev_t = sde_step_with_logprob(
scheduler,
noise_pred_ref.float(),
t.unsqueeze(0) if isinstance(t, torch.Tensor) else t,
latents_ori.float(),
prev_sample=prev_latents.float(),
deterministic=deterministic,
)
# Compute KL divergence: KL = (mean_diff)^2 / (2 * std^2)
assert torch.allclose(
std_dev_t, ref_std_dev_t
), "std_dev_t should match between current and reference"
kl = (prev_latents_mean - ref_prev_latents_mean)**2 / (2 *
std_dev_t**2)
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
all_kl.append(kl)
else:
# No KL reward, set to zero
all_kl.append(torch.zeros(len(latents), device=latents.device))
# Update progress bar with timing information
step_time = time.time() - step_start_time
progress_bar.set_postfix({
"step_time":
f"{step_time:.2f}s",
"timestep":
f"{t.item() if isinstance(t, torch.Tensor) else t:.1f}"
})
# Decode latents to video if needed
if output_type != "latent":
latents = latents.to(vae.dtype)
# Apply VAE normalization (Wan VAE specific)
# Wan VAE requires denormalization before decoding
if hasattr(vae, 'config') and hasattr(vae.config,
'latents_mean') and hasattr(
vae.config, 'latents_std'):
# Get z_dim from config or VAE
z_dim = getattr(vae.config, 'z_dim', latents.shape[1])
latents_mean = (torch.tensor(vae.config.latents_mean,
device=latents.device,
dtype=latents.dtype).view(
1, z_dim, 1, 1, 1))
latents_std = (
1.0 / torch.tensor(vae.config.latents_std,
device=latents.device,
dtype=latents.dtype).view(1, z_dim, 1, 1, 1))
latents = latents / latents_std + latents_mean
elif hasattr(vae, 'latents_mean') and hasattr(vae, 'latents_std'):
# Alternative: check if latents_mean/std are direct attributes
z_dim = latents.shape[1]
latents_mean = (torch.tensor(vae.latents_mean,
device=latents.device,
dtype=latents.dtype).view(
1, z_dim, 1, 1, 1))
latents_std = (1.0 / torch.tensor(
vae.latents_std, device=latents.device,
dtype=latents.dtype).view(1, z_dim, 1, 1, 1))
latents = latents / latents_std + latents_mean
# Decode using VAE
with torch.no_grad():
video = vae.decode(latents.float(), return_dict=False)[0]
# VAE.decode returns tensor directly (not tuple)
# Postprocess video: convert from [-1, 1] to [0, 1]
# FastVideo VAE typically outputs in [-1, 1] range
video = (video / 2 + 0.5).clamp(0, 1)
else:
video = latents
return video, all_latents, all_log_probs, all_kl, prompt_ids
+27 -20
View File
@@ -174,17 +174,18 @@ class TrainingPipeline(LoRAPipeline, ABC):
last_epoch=self.init_steps - 1,
)
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=self.seed)
if not self.training_args.rl_args.rl_mode:
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=self.seed)
self.noise_scheduler = noise_scheduler
if self.training_args.boundary_ratio is not None:
@@ -192,19 +193,21 @@ class TrainingPipeline(LoRAPipeline, ABC):
else:
self.boundary_timestep = None
logger.info("train_dataloader length: %s", len(self.train_dataloader))
if not self.training_args.rl_args.rl_mode:
logger.info("train_dataloader length: %s", len(self.train_dataloader))
logger.info("train_sp_batch_size: %s",
training_args.train_sp_batch_size)
logger.info("gradient_accumulation_steps: %s",
training_args.gradient_accumulation_steps)
logger.info("sp_size: %s", training_args.sp_size)
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
training_args.train_sp_batch_size)
self.num_train_epochs = math.ceil(training_args.max_train_steps /
self.num_update_steps_per_epoch)
if not self.training_args.rl_args.rl_mode:
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
training_args.train_sp_batch_size)
self.num_train_epochs = math.ceil(training_args.max_train_steps /
self.num_update_steps_per_epoch)
# TODO(will): is there a cleaner way to track epochs?
self.current_epoch = 0
@@ -575,8 +578,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
if not self.training_args.rl_args.rl_mode:
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
else:
self.noise_random_generator = torch.Generator(device=self.device).manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed)
self.validation_random_generator = torch.Generator(
@@ -0,0 +1,100 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
from fastvideo.training.rl.rl_pipeline import RLPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanRLTrainingPipeline(RLPipeline):
"""
A training pipeline for Wan with RL/GRPO support.
This pipeline extends RLPipeline with Wan-specific initialization.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def _create_inference_pipeline(self, training_args: TrainingArgs,
dit_cpu_offload: bool):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
loaded_modules = {
"transformer": self.get_module("transformer"),
}
transformer_2 = self.get_module("transformer_2", None)
if transformer_2 is not None:
loaded_modules["transformer_2"] = transformer_2
text_encoder = self.get_module("text_encoder", None)
if text_encoder is not None:
loaded_modules["text_encoder"] = text_encoder
tokenizer = self.get_module("tokenizer", None)
if tokenizer is not None:
loaded_modules["tokenizer"] = tokenizer
vae = self.get_module("vae", None)
if vae is not None:
loaded_modules["vae"] = vae
return WanPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules=loaded_modules,
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=dit_cpu_offload)
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
self.validation_pipeline = self._create_inference_pipeline(
training_args, dit_cpu_offload=True)
def _build_sampling_pipeline(self, training_args: TrainingArgs):
return self._create_inference_pipeline(training_args,
dit_cpu_offload=False)
def main(args) -> None:
logger.info("Starting RL training pipeline...")
pipeline = WanRLTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("RL training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
# Enable RL mode
args.rl_mode = True
main(args)
+28 -1
View File
@@ -22,7 +22,7 @@ import threading
import traceback
from collections.abc import Callable
from dataclasses import dataclass, fields, is_dataclass
from functools import lru_cache, partial, wraps
from functools import lru_cache, partial, wraps, cache
from pathlib import Path
from typing import Any, TextIO, TypeVar, cast
@@ -1193,3 +1193,30 @@ def decorate_logs(process_name: str | None = None) -> None:
pid = os.getpid()
_add_prefix(sys.stdout, process_name, pid)
_add_prefix(sys.stderr, process_name, pid)
def _probe_pin_memory() -> bool:
from fastvideo.platforms import current_platform
if current_platform.is_cpu() or current_platform.is_mps(
) or current_platform.is_npu():
return False
try:
if torch.cuda.is_available():
torch.cuda.current_device()
_ = torch.empty(1024, device="cpu").pin_memory()
_ = torch.empty(1024, device="cpu", pin_memory=True)
except Exception as exc:
logger.warning("Pinned memory is unavailable: %s", exc)
return False
return True
@cache
def _cached_pin_memory_available(pid: int) -> bool:
return _probe_pin_memory()
def is_pin_memory_available() -> bool:
return _cached_pin_memory_available(os.getpid())
+2 -2
View File
@@ -63,7 +63,7 @@ dependencies = [
"remote-pdb",
# Kernel & Packaging
"fastvideo-kernel==0.2.2",
"fastvideo-kernel==0.2.4",
"wheel",
# Training Dependencies
@@ -111,7 +111,7 @@ lint = [
]
test = [
"av==14.3.0",
"av",
"pytorch-msssim==1.0.0",
"pytest",
]
+2 -2
View File
@@ -63,7 +63,7 @@ dependencies = [
"remote-pdb",
# Kernel & Packaging
"fastvideo-kernel==0.2.2",
"fastvideo-kernel==0.2.4",
"wheel",
# Training Dependencies
@@ -90,7 +90,7 @@ lint = [
]
test = [
"av==14.3.0",
"av",
"pytorch-msssim==1.0.0",
"pytest",
]
@@ -0,0 +1,106 @@
#!/usr/bin/env python3
"""Convert a PyTorch checkpoint (.pt) to a safetensors file."""
import argparse
from pathlib import Path
import torch
from safetensors.torch import save_file
def convert_pt_to_safetensors(
input_path: str,
output_path: str,
key: str | None = None,
force: bool = False,
skip_patterns: list[str] | None = None,
):
input_path = Path(input_path)
output_path = Path(output_path)
if not input_path.exists():
raise FileNotFoundError(f"Input file not found: {input_path}")
if output_path.exists() and not force:
raise FileExistsError(
f"Output file already exists: {output_path}. Use --force to overwrite."
)
checkpoint = torch.load(input_path, map_location="cpu")
state_dict: dict[str, torch.Tensor]
if isinstance(checkpoint, dict):
if key is not None:
if key not in checkpoint:
raise KeyError(f"Key {key!r} not found in checkpoint.")
state_dict = checkpoint[key]
else:
for k in ("state_dict", "model_state_dict", "model", "ema"):
if k in checkpoint:
state_dict = checkpoint[k]
break
else:
state_dict = checkpoint
else:
state_dict = checkpoint
if not isinstance(state_dict, dict):
raise TypeError(f"Expected a dict state_dict, got {type(state_dict)}")
if skip_patterns:
state_dict = {
k: v
for k, v in state_dict.items()
if not any(pat in k for pat in skip_patterns)
}
output_path.parent.mkdir(parents=True, exist_ok=True)
save_file(state_dict, str(output_path))
def main():
parser = argparse.ArgumentParser(
description="Convert a PyTorch checkpoint (.pt) to safetensors."
)
parser.add_argument(
"input",
type=str,
help="Path to input .pt checkpoint file"
)
parser.add_argument(
"output",
type=str,
help="Path to output .safetensors file"
)
parser.add_argument(
"--key",
type=str,
default=None,
help="Optional key to extract from checkpoint dict (e.g. 'model_state_dict')"
)
parser.add_argument(
"--force",
action="store_true",
help="Overwrite output file if it exists"
)
parser.add_argument(
"--skip-pattern",
action="append",
dest="skip_patterns",
help="Parameter name patterns to skip (can be used multiple times)"
)
args = parser.parse_args()
convert_pt_to_safetensors(
args.input,
args.output,
args.key,
args.force,
args.skip_patterns
)
if __name__ == "__main__":
main()