Compare commits
33
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
141a1140f6 | ||
|
|
6294015389 | ||
|
|
e31b6c9e90 | ||
|
|
bf0ff21eeb | ||
|
|
6f937102ad | ||
|
|
0164e93019 | ||
|
|
67e457aa92 | ||
|
|
d795f0c443 | ||
|
|
02452dd6e7 | ||
|
|
91ef24bc14 | ||
|
|
e76e9fda15 | ||
|
|
3b17f5a621 | ||
|
|
d758878705 | ||
|
|
689e629420 | ||
|
|
873dc9695f | ||
|
|
bfc0f46d61 | ||
|
|
39907dbe4d | ||
|
|
abdd0c9b6a | ||
|
|
f32a12200d | ||
|
|
f1d2c9e6b7 | ||
|
|
450579cb42 | ||
|
|
26d7d6cc08 | ||
|
|
44f0124eaa | ||
|
|
d3ace51394 | ||
|
|
58954c660b | ||
|
|
31f44110b5 | ||
|
|
21f3ce6577 | ||
|
|
785d123e36 | ||
|
|
d58c551c11 | ||
|
|
560628709c | ||
|
|
0f53b51e6c | ||
|
|
06093a9c4e | ||
|
|
dbddfab6d2 |
@@ -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
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
```
|
||||
|
||||
@@ -168,10 +168,6 @@ Pipelines are composed of stages, each handling a specific part of the diffusion
|
||||
- **DenoisingStage**: Performs denoising diffusion
|
||||
- **DecodingStage**: Converts latents to pixels
|
||||
|
||||
Note: `DenoisingStage` uses the unified denoising engine under the hood. You
|
||||
can inject a custom strategy via `strategy_cls` if a pipeline needs specialized
|
||||
denoising behavior.
|
||||
|
||||
### Creating Your Pipeline
|
||||
|
||||
```python
|
||||
@@ -234,7 +230,7 @@ class MyCustomPipeline(ComposedPipelineBase):
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
scheduler=self.get_module("scheduler")
|
||||
)
|
||||
)
|
||||
|
||||
@@ -249,24 +245,6 @@ class MyCustomPipeline(ComposedPipelineBase):
|
||||
EntryClass = MyCustomPipeline
|
||||
```
|
||||
|
||||
### Customizing Denoising Strategies
|
||||
|
||||
If your model requires a custom denoising loop, pass a strategy class:
|
||||
|
||||
```python
|
||||
from fastvideo.pipelines.stages.denoising_cosmos_strategy import (
|
||||
CosmosStrategy)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
strategy_cls=CosmosStrategy,
|
||||
)
|
||||
)
|
||||
```
|
||||
|
||||
### Creating Custom Stages (Optional)
|
||||
|
||||
If existing stages don't meet your needs, create custom ones:
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
Executable
+129
@@ -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[@]}"
|
||||
@@ -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()
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -38,10 +38,8 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
FASTVIDEO_STAGE_LOGGING: bool = False
|
||||
FASTVIDEO_DENOISING_PERF_LOGGING: bool = False
|
||||
FASTVIDEO_HOST_IP: str = ""
|
||||
FASTVIDEO_LOOPBACK_IP: str = ""
|
||||
FASTVIDEO_DISABLE_PIN_MEMORY: str | None = None
|
||||
|
||||
|
||||
def get_default_cache_root() -> str:
|
||||
@@ -138,10 +136,6 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_LOOPBACK_IP":
|
||||
lambda: os.getenv("FASTVIDEO_LOOPBACK_IP", ""),
|
||||
|
||||
# Disable pinned memory (e.g., on platforms that do not support it)
|
||||
"FASTVIDEO_DISABLE_PIN_MEMORY":
|
||||
lambda: os.getenv("FASTVIDEO_DISABLE_PIN_MEMORY", None),
|
||||
|
||||
# Number of GPUs per worker in Ray, if it is set to be a fraction,
|
||||
# it allows ray to schedule multiple actors on a single GPU,
|
||||
# so that users can colocate other actors on the same GPUs as FastVideo.
|
||||
@@ -281,10 +275,6 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# taken for each stage
|
||||
"FASTVIDEO_STAGE_LOGGING":
|
||||
lambda: bool(int(os.getenv("FASTVIDEO_STAGE_LOGGING", "0"))),
|
||||
|
||||
# Enable per-step denoising perf logging hooks
|
||||
"FASTVIDEO_DENOISING_PERF_LOGGING":
|
||||
lambda: bool(int(os.getenv("FASTVIDEO_DENOISING_PERF_LOGGING", "0"))),
|
||||
}
|
||||
|
||||
# end-env-vars-definition
|
||||
|
||||
+312
-8
@@ -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
|
||||
@@ -608,7 +608,6 @@ class FastVideoArgs:
|
||||
def check_fastvideo_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
from fastvideo.platforms import current_platform
|
||||
from fastvideo.pin_memory import is_pin_memory_available
|
||||
|
||||
if current_platform.is_mps():
|
||||
self.use_fsdp_inference = False
|
||||
@@ -689,11 +688,6 @@ class FastVideoArgs:
|
||||
self.pipeline_config.vae_config.load_encoder = True
|
||||
self.preprocess_config.check_preprocess_config()
|
||||
|
||||
if self.pin_cpu_memory and not is_pin_memory_available():
|
||||
logger.warning("Pinned memory is unavailable on this system; "
|
||||
"disabling pin_cpu_memory.")
|
||||
self.pin_cpu_memory = False
|
||||
|
||||
|
||||
_current_fastvideo_args = None
|
||||
|
||||
@@ -746,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):
|
||||
"""
|
||||
@@ -758,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
|
||||
@@ -868,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)
|
||||
@@ -892,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
|
||||
@@ -921,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,
|
||||
@@ -1290,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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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)]
|
||||
@@ -18,7 +18,6 @@ from fastvideo.layers.linear import (ColumnParallelLinear, LinearBase,
|
||||
QKVParallelLinear, ReplicatedLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.pin_memory import is_pin_memory_available
|
||||
from fastvideo.utils import get_mixed_precision_state
|
||||
|
||||
torch._dynamo.config.recompile_limit = 16
|
||||
@@ -156,11 +155,8 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
get_local_torch_device(),
|
||||
non_blocking=True).full_tensor().to(current_device))
|
||||
|
||||
if "cpu" in str(current_device):
|
||||
offload_policy = CPUOffloadPolicy(
|
||||
pin_memory=is_pin_memory_available())
|
||||
else:
|
||||
offload_policy = OffloadPolicy()
|
||||
offload_policy = CPUOffloadPolicy() if "cpu" in str(
|
||||
current_device) else OffloadPolicy()
|
||||
mp_policy = get_mixed_precision_state().mp_policy
|
||||
|
||||
self.base_layer = fully_shard(unsharded_base_layer,
|
||||
|
||||
@@ -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
@@ -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
|
||||
)
|
||||
|
||||
@@ -11,8 +11,6 @@ from typing import Dict, Set, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.pin_memory import is_pin_memory_available
|
||||
|
||||
|
||||
class LayerwiseOffloadManager:
|
||||
"""A lightweight layerwise CPU offload manager.
|
||||
@@ -35,8 +33,6 @@ class LayerwiseOffloadManager:
|
||||
self.module_list_attr = module_list_attr
|
||||
self.num_layers = int(num_layers)
|
||||
self.pin_cpu_memory = bool(pin_cpu_memory)
|
||||
if self.pin_cpu_memory and not is_pin_memory_available():
|
||||
self.pin_cpu_memory = False
|
||||
|
||||
self.enabled = bool(enabled and torch.cuda.is_available())
|
||||
self.device = (
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -20,11 +20,10 @@ from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
|
||||
from torch.nn.modules.module import _IncompatibleKeys
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pin_memory import is_pin_memory_available
|
||||
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__)
|
||||
|
||||
@@ -68,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,
|
||||
@@ -107,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
|
||||
@@ -142,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,
|
||||
)
|
||||
@@ -152,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
|
||||
@@ -213,11 +215,6 @@ def shard_model(
|
||||
"mp_policy": mp_policy,
|
||||
}
|
||||
if cpu_offload:
|
||||
if pin_cpu_memory and not is_pin_memory_available():
|
||||
logger.warning(
|
||||
"Pinned memory is unavailable; disabling pin_cpu_memory for "
|
||||
"FSDP offload.")
|
||||
pin_cpu_memory = False
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
|
||||
pin_memory=pin_cpu_memory)
|
||||
|
||||
@@ -299,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."
|
||||
)
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Scheduler adapter interfaces for unified denoising.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class SchedulerAdapter(Protocol):
|
||||
def scale_model_input(self, latents: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
...
|
||||
|
||||
def step(self, noise_pred: torch.Tensor, t: torch.Tensor,
|
||||
latents: torch.Tensor, **kwargs: Any) -> Any:
|
||||
...
|
||||
|
||||
def add_noise(self, latents: torch.Tensor, noise: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
...
|
||||
|
||||
def set_timesteps(self, num_steps: int, device: torch.device | None = None,
|
||||
**kwargs: Any) -> Any:
|
||||
...
|
||||
|
||||
|
||||
class DefaultSchedulerAdapter:
|
||||
def __init__(self, scheduler: Any) -> None:
|
||||
self.scheduler = scheduler
|
||||
|
||||
def scale_model_input(self, latents: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
return self.scheduler.scale_model_input(latents, t)
|
||||
|
||||
def step(self, noise_pred: torch.Tensor, t: torch.Tensor,
|
||||
latents: torch.Tensor, **kwargs: Any) -> Any:
|
||||
return self.scheduler.step(noise_pred, t, latents, **kwargs)
|
||||
|
||||
def add_noise(self, latents: torch.Tensor, noise: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
return self.scheduler.add_noise(latents, noise, t)
|
||||
|
||||
def set_timesteps(self, num_steps: int, device: torch.device | None = None,
|
||||
**kwargs: Any) -> Any:
|
||||
return self.scheduler.set_timesteps(num_steps, device=device, **kwargs)
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,42 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
from functools import cache
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _probe_pin_memory() -> bool:
|
||||
if os.getenv("FASTVIDEO_DISABLE_PIN_MEMORY",
|
||||
"") not in ("", "0", "false", "False"):
|
||||
return False
|
||||
if current_platform.is_cpu() or current_platform.is_mps():
|
||||
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())
|
||||
@@ -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
|
||||
@@ -11,12 +11,11 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, DenoisingStage,
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, CosmosDenoisingStage,
|
||||
CosmosLatentPreparationStage,
|
||||
DecodingStage, InputValidationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.pipelines.stages.denoising_cosmos_strategy import CosmosStrategy
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -74,10 +73,9 @@ class Cosmos2VideoToWorldPipeline(ComposedPipelineBase):
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=CosmosDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
strategy_cls=CosmosStrategy))
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
@@ -16,15 +16,13 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
|
||||
LongCatI2VStrategy)
|
||||
from fastvideo.pipelines.stages.longcat_image_vae_encoding import LongCatImageVAEEncodingStage
|
||||
from fastvideo.pipelines.stages.longcat_i2v_latent_preparation import LongCatI2VLatentPreparationStage
|
||||
from fastvideo.pipelines.stages.longcat_i2v_denoising import LongCatI2VDenoisingStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
|
||||
|
||||
@@ -134,13 +132,12 @@ class LongCatImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
# 8. Denoising with I2V support
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=LongCatI2VDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
strategy_cls=LongCatI2VStrategy))
|
||||
pipeline=self))
|
||||
|
||||
# 9. Decoding
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
|
||||
@@ -11,14 +11,12 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
|
||||
LongCatStrategy)
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
|
||||
|
||||
@@ -129,13 +127,12 @@ class LongCatPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=LongCatDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
strategy_cls=LongCatStrategy))
|
||||
pipeline=self))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"),
|
||||
|
||||
@@ -16,16 +16,14 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
|
||||
LongCatVCStrategy)
|
||||
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
|
||||
from fastvideo.pipelines.stages.longcat_i2v_latent_preparation import LongCatI2VLatentPreparationStage
|
||||
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
|
||||
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -133,13 +131,12 @@ class LongCatVideoContinuationPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
# 7. Denoising with VC and KV cache support
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=LongCatVCDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
strategy_cls=LongCatVCStrategy))
|
||||
pipeline=self))
|
||||
|
||||
# 8. Decoding
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
|
||||
@@ -11,11 +11,10 @@ from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, InputValidationStage,
|
||||
CausalDMDDenosingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage)
|
||||
from fastvideo.pipelines.stages.denoising_causal_strategy import (
|
||||
CausalBlockStrategy)
|
||||
# isort: on
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -48,12 +47,11 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=CausalDMDDenosingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
strategy_cls=CausalBlockStrategy))
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
@@ -14,11 +14,10 @@ from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, InputValidationStage,
|
||||
DmdDenoisingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.pipelines.stages.denoising_dmd_strategy import DmdStrategy
|
||||
# isort: on
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -64,10 +63,9 @@ class WanDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
use_btchw_layout=True))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=DmdDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=FlowMatchEulerDiscreteScheduler(shift=8.0),
|
||||
strategy_cls=DmdStrategy))
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
@@ -13,13 +13,12 @@ from fastvideo.pipelines.lora_pipeline import LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from fastvideo.pipelines.stages import (
|
||||
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
|
||||
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
|
||||
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
|
||||
TextEncodingStage, TimestepPreparationStage)
|
||||
# isort: on
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.stages.denoising_dmd_strategy import DmdStrategy
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -70,10 +69,9 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
stage=DmdDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=FlowMatchEulerDiscreteScheduler(shift=8.0),
|
||||
strategy_cls=DmdStrategy))
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -7,37 +7,50 @@ complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
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 DenoisingStage
|
||||
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
|
||||
CosmosDenoisingStage,
|
||||
DenoisingStage,
|
||||
DmdDenoisingStage)
|
||||
from fastvideo.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.pipelines.stages.image_encoding import (
|
||||
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
|
||||
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
|
||||
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
|
||||
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
|
||||
|
||||
__all__ = [
|
||||
"PipelineStage",
|
||||
"InputValidationStage",
|
||||
"TimestepPreparationStage",
|
||||
"Cosmos25TimestepPreparationStage",
|
||||
"LatentPreparationStage",
|
||||
"CosmosLatentPreparationStage",
|
||||
"Cosmos25LatentPreparationStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
"CausalDMDDenosingStage",
|
||||
"MatrixGameCausalDenoisingStage",
|
||||
"CosmosDenoisingStage",
|
||||
"Cosmos25DenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
@@ -47,8 +60,10 @@ __all__ = [
|
||||
"ImageVAEEncodingStage",
|
||||
"VideoVAEEncodingStage",
|
||||
"TextEncodingStage",
|
||||
"Cosmos25TextEncodingStage",
|
||||
"StepvideoPromptEncodingStage",
|
||||
# LongCat stages
|
||||
"LongCatVideoVAEEncodingStage",
|
||||
"LongCatKVCacheInitStage",
|
||||
"LongCatVCDenoisingStage",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,497 @@
|
||||
import torch # type: ignore
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
VideoSparseAttentionBackend = None # type: ignore
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CausalDMDDenosingStage(DenoisingStage):
|
||||
"""
|
||||
Denoising stage for causal diffusion.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
transformer,
|
||||
scheduler,
|
||||
transformer_2=None,
|
||||
vae=None) -> None:
|
||||
super().__init__(transformer, scheduler, transformer_2)
|
||||
# KV and cross-attention cache state (initialized on first forward)
|
||||
self.transformer = transformer
|
||||
self.transformer_2 = transformer_2
|
||||
self.vae = vae
|
||||
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
||||
self.num_transformer_blocks = len(self.transformer.blocks)
|
||||
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
|
||||
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
|
||||
|
||||
try:
|
||||
self.local_attn_size = getattr(self.transformer.model,
|
||||
"local_attn_size",
|
||||
-1) # type: ignore
|
||||
except Exception:
|
||||
self.local_attn_size = -1
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
|
||||
patch_ratio = self.transformer.config.arch_config.patch_size[
|
||||
-1] * self.transformer.config.arch_config.patch_size[-2]
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
# TODO(will): make this a parameter once we add i2v support
|
||||
independent_first_frame = self.transformer.independent_first_frame if hasattr(
|
||||
self.transformer, 'independent_first_frame') else False
|
||||
# Timesteps for DMD
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
|
||||
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_kwargs: dict = {}
|
||||
|
||||
pos_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
# "encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
# STA
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
# Latents and prompts
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents # [B, C, T, H, W]
|
||||
b, c, t, h, w = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
# Initialize or reset caches
|
||||
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
kv_cache2 = None
|
||||
if boundary_timestep is not None:
|
||||
# Initialize the low noise kv cache
|
||||
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
def _get_kv_cache(timestep: float) -> list[dict]:
|
||||
if boundary_timestep is not None:
|
||||
if timestep >= boundary_timestep:
|
||||
return kv_cache1
|
||||
else:
|
||||
assert kv_cache2 is not None, "kv_cache2 is not initialized"
|
||||
return kv_cache2
|
||||
return kv_cache1
|
||||
|
||||
crossattn_cache = self._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=fastvideo_args.pipeline_config.text_encoder_configs[0].
|
||||
arch_config.text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
pos_start_base = 0
|
||||
|
||||
# Determine block sizes
|
||||
if t % self.num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
|
||||
)
|
||||
num_blocks = t // self.num_frames_per_block
|
||||
block_sizes = [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
# For now hardcode the first block to be 1 frame assuming the model is Wan2.2-MoE
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
first_frame_latent = None
|
||||
if batch.pil_image is not None:
|
||||
# Causal video gen directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert self.vae is not None, "VAE is not provided for causal video gen task"
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
first_frame_latent = self.vae.encode(batch.pil_image).mean.float()
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
first_frame_latent -= self.vae.shift_factor.to(
|
||||
first_frame_latent.device, first_frame_latent.dtype)
|
||||
else:
|
||||
first_frame_latent -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
|
||||
first_frame_latent.device, first_frame_latent.dtype)
|
||||
else:
|
||||
first_frame_latent = first_frame_latent * self.vae.scaling_factor
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae = self.vae.to("cpu")
|
||||
|
||||
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
|
||||
t_zero = torch.zeros([latents.shape[0], 1],
|
||||
device=latents.device,
|
||||
dtype=torch.long)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch):
|
||||
self.transformer(
|
||||
first_frame_latent.to(target_dtype),
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
if boundary_timestep is not None:
|
||||
self.transformer_2(
|
||||
first_frame_latent.to(target_dtype),
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=kv_cache2,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
start_index += 1
|
||||
block_sizes.pop(0)
|
||||
latents[:, :, :1, :, :] = first_frame_latent
|
||||
|
||||
# DMD loop in causal blocks
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
for current_num_frames in block_sizes:
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
# use BTCHW for DMD conversion routines
|
||||
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
|
||||
video_raw_latent_shape = noise_latents_btchw.shape
|
||||
|
||||
for i, t_cur in enumerate(timesteps):
|
||||
if boundary_timestep is not None and t_cur < boundary_timestep:
|
||||
current_model = self.transformer_2
|
||||
else:
|
||||
current_model = self.transformer
|
||||
# Copy for pred conversion
|
||||
noise_latents = noise_latents_btchw.clone()
|
||||
latent_model_input = current_latents.to(target_dtype)
|
||||
|
||||
if batch.image_latent is not None and independent_first_frame and start_index == 0:
|
||||
latent_model_input = torch.cat([
|
||||
latent_model_input,
|
||||
batch.image_latent.to(target_dtype)
|
||||
],
|
||||
dim=2)
|
||||
|
||||
# Prepare inputs
|
||||
t_expand = t_cur.repeat(latent_model_input.shape[0])
|
||||
|
||||
# Attention metadata if needed
|
||||
if (vsa_available and self.attn_backend
|
||||
== VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=i, # type: ignore
|
||||
raw_latent_shape=(current_num_frames, h,
|
||||
w), # type: ignore
|
||||
patch_size=fastvideo_args.pipeline_config.
|
||||
dit_config.patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.
|
||||
VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(), # type: ignore
|
||||
) # type: ignore
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
pred_noise_btchw = current_model(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=_get_kv_cache(t_cur),
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Convert pred noise to pred video with FM Euler scheduler utilities
|
||||
if boundary_timestep is not None and t_cur >= boundary_timestep:
|
||||
pred_video_btchw = pred_noise_to_x_bound(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
boundary_timestep=torch.ones_like(t_expand) *
|
||||
boundary_timestep,
|
||||
scheduler=self.scheduler).unflatten(
|
||||
0, pred_noise_btchw.shape[:2])
|
||||
else:
|
||||
pred_video_btchw = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
scheduler=self.scheduler).unflatten(
|
||||
0, pred_noise_btchw.shape[:2])
|
||||
|
||||
if i < len(timesteps) - 1:
|
||||
next_timestep = timesteps[i + 1] * torch.ones(
|
||||
[1],
|
||||
dtype=torch.long,
|
||||
device=pred_video_btchw.device)
|
||||
noise = torch.randn(
|
||||
video_raw_latent_shape,
|
||||
dtype=pred_video_btchw.dtype,
|
||||
generator=(batch.generator[0] if isinstance(
|
||||
batch.generator, list) else
|
||||
batch.generator)).to(self.device)
|
||||
noise_btchw = noise
|
||||
if boundary_timestep is not None and i < len(
|
||||
high_noise_timesteps) - 1:
|
||||
noise_latents_btchw = self.scheduler.add_noise_high(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1), next_timestep,
|
||||
torch.ones_like(next_timestep) *
|
||||
boundary_timestep).unflatten(
|
||||
0, pred_video_btchw.shape[:2])
|
||||
elif boundary_timestep is not None and i == len(
|
||||
high_noise_timesteps) - 1:
|
||||
noise_latents_btchw = pred_video_btchw
|
||||
else:
|
||||
noise_latents_btchw = self.scheduler.add_noise(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1),
|
||||
next_timestep).unflatten(
|
||||
0, pred_video_btchw.shape[:2])
|
||||
current_latents = noise_latents_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
else:
|
||||
current_latents = pred_video_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
# Write back and advance
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = current_latents
|
||||
|
||||
# Re-run with context timestep to update KV cache using clean context
|
||||
context_noise = getattr(fastvideo_args.pipeline_config,
|
||||
"context_noise", 0)
|
||||
t_context = torch.ones([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = current_latents.to(target_dtype)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
|
||||
if boundary_timestep is not None:
|
||||
self.transformer_2(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache2,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
self.transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
start_index += current_num_frames
|
||||
|
||||
if boundary_timestep is not None:
|
||||
num_frames_to_remove = self.num_frames_per_block - 1
|
||||
latents = latents[:, :, :-num_frames_to_remove, :, :]
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
|
||||
"""
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
if self.local_attn_size != -1:
|
||||
kv_cache_size = self.local_attn_size * self.frame_seq_length
|
||||
else:
|
||||
kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames
|
||||
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
kv_cache1.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
return kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
|
||||
device) -> list[dict]:
|
||||
"""
|
||||
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
crossattn_cache = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
return crossattn_cache
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
V.none_or_tensor_with_dims(5))
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("eta", batch.eta, V.non_negative_float)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
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
|
||||
@@ -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
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,623 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Causal block denoising strategy (DMD).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
BlockContext,
|
||||
BlockDenoisingStrategy,
|
||||
BlockPlan,
|
||||
BlockPlanItem,
|
||||
ModelInputs,
|
||||
StrategyState,
|
||||
)
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
VideoSparseAttentionBackend = None # type: ignore
|
||||
|
||||
|
||||
class CausalBlockStrategy(BlockDenoisingStrategy):
|
||||
|
||||
def __init__(self, stage: Any) -> None:
|
||||
self.stage = stage
|
||||
self.num_transformer_blocks = 0
|
||||
self.num_frames_per_block = 0
|
||||
self.sliding_window_num_frames = 0
|
||||
self.local_attn_size = -1
|
||||
self.frame_seq_length = 0
|
||||
|
||||
def _ensure_model_constants(self) -> None:
|
||||
transformer = self.stage.transformer
|
||||
self.num_transformer_blocks = len(transformer.blocks)
|
||||
arch_config = transformer.config.arch_config
|
||||
self.num_frames_per_block = arch_config.num_frames_per_block
|
||||
self.sliding_window_num_frames = arch_config.sliding_window_num_frames
|
||||
try:
|
||||
self.local_attn_size = getattr(transformer.model, "local_attn_size",
|
||||
-1)
|
||||
except Exception:
|
||||
self.local_attn_size = -1
|
||||
|
||||
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
|
||||
kv_cache1 = []
|
||||
num_attention_heads = self.stage.transformer.num_attention_heads
|
||||
attention_head_dim = self.stage.transformer.attention_head_dim
|
||||
if self.local_attn_size != -1:
|
||||
kv_cache_size = self.local_attn_size * self.frame_seq_length
|
||||
else:
|
||||
kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames
|
||||
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
kv_cache1.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
return kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
|
||||
device) -> list[dict]:
|
||||
crossattn_cache = []
|
||||
num_attention_heads = self.stage.transformer.num_attention_heads
|
||||
attention_head_dim = self.stage.transformer.attention_head_dim
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
return crossattn_cache
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
self._ensure_model_constants()
|
||||
latents = batch.latents
|
||||
if latents is None:
|
||||
raise ValueError("latents must be provided")
|
||||
|
||||
latent_seq_length = latents.shape[-1] * latents.shape[-2]
|
||||
patch_ratio = (
|
||||
self.stage.transformer.config.arch_config.patch_size[-1] *
|
||||
self.stage.transformer.config.arch_config.patch_size[-2])
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
|
||||
independent_first_frame = getattr(self.stage.transformer,
|
||||
"independent_first_frame", False)
|
||||
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step",
|
||||
False):
|
||||
scheduler_timesteps = torch.cat((
|
||||
self.stage.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0], dtype=torch.float32),
|
||||
))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
|
||||
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = (boundary_ratio *
|
||||
self.stage.scheduler.num_train_timesteps)
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
image_kwargs: dict[str, Any] = {}
|
||||
pos_cond_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
if (st_attn_available
|
||||
and self.stage.attn_backend == SlidingTileAttentionBackend):
|
||||
self.stage.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
kv_cache2 = None
|
||||
if boundary_timestep is not None:
|
||||
kv_cache2 = self._initialize_kv_cache(
|
||||
batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device,
|
||||
)
|
||||
|
||||
text_len = None
|
||||
if fastvideo_args.pipeline_config.text_encoder_configs:
|
||||
text_len = getattr(
|
||||
fastvideo_args.pipeline_config.text_encoder_configs[0].
|
||||
arch_config, "text_len", None)
|
||||
if not text_len:
|
||||
if batch.prompt_attention_mask:
|
||||
text_len = batch.prompt_attention_mask[0].shape[-1]
|
||||
elif batch.prompt_embeds:
|
||||
text_len = batch.prompt_embeds[0].shape[1]
|
||||
else:
|
||||
text_len = 0
|
||||
|
||||
crossattn_cache = self._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device,
|
||||
)
|
||||
|
||||
num_frames = latents.shape[2]
|
||||
if num_frames % self.num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frames_per_block for "
|
||||
"causal DMD denoising")
|
||||
num_blocks = num_frames // self.num_frames_per_block
|
||||
block_sizes = [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
pos_start_base = 0
|
||||
if batch.pil_image is not None:
|
||||
assert self.stage.vae is not None, (
|
||||
"VAE is not provided for causal video gen task")
|
||||
self.stage.vae = self.stage.vae.to(get_local_torch_device())
|
||||
first_frame_latent = self.stage.vae.encode(
|
||||
batch.pil_image).mean.float()
|
||||
if (hasattr(self.stage.vae, "shift_factor")
|
||||
and self.stage.vae.shift_factor is not None):
|
||||
if isinstance(self.stage.vae.shift_factor, torch.Tensor):
|
||||
first_frame_latent -= self.stage.vae.shift_factor.to(
|
||||
first_frame_latent.device, first_frame_latent.dtype)
|
||||
else:
|
||||
first_frame_latent -= self.stage.vae.shift_factor
|
||||
|
||||
if isinstance(self.stage.vae.scaling_factor, torch.Tensor):
|
||||
first_frame_latent = (
|
||||
first_frame_latent * self.stage.vae.scaling_factor.to(
|
||||
first_frame_latent.device, first_frame_latent.dtype))
|
||||
else:
|
||||
first_frame_latent = (first_frame_latent *
|
||||
self.stage.vae.scaling_factor)
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.stage.vae = self.stage.vae.to("cpu")
|
||||
|
||||
t_zero = torch.zeros([latents.shape[0], 1],
|
||||
device=latents.device,
|
||||
dtype=torch.long)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch):
|
||||
self.stage.transformer(
|
||||
first_frame_latent.to(target_dtype),
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
if boundary_timestep is not None:
|
||||
self.stage.transformer_2(
|
||||
first_frame_latent.to(target_dtype),
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=kv_cache2,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
start_index += 1
|
||||
block_sizes.pop(0)
|
||||
latents[:, :, :1, :, :] = first_frame_latent
|
||||
|
||||
progress_bar = self.stage.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps))
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch":
|
||||
batch,
|
||||
"fastvideo_args":
|
||||
fastvideo_args,
|
||||
"target_dtype":
|
||||
target_dtype,
|
||||
"autocast_enabled":
|
||||
autocast_enabled,
|
||||
"boundary_timestep":
|
||||
boundary_timestep,
|
||||
"high_noise_timesteps":
|
||||
high_noise_timesteps,
|
||||
"image_kwargs":
|
||||
image_kwargs,
|
||||
"pos_cond_kwargs":
|
||||
pos_cond_kwargs,
|
||||
"kv_cache1":
|
||||
kv_cache1,
|
||||
"kv_cache2":
|
||||
kv_cache2,
|
||||
"crossattn_cache":
|
||||
crossattn_cache,
|
||||
"block_sizes":
|
||||
block_sizes,
|
||||
"start_index":
|
||||
start_index,
|
||||
"progress_bar":
|
||||
progress_bar,
|
||||
"independent_first_frame":
|
||||
independent_first_frame,
|
||||
"context_noise":
|
||||
getattr(fastvideo_args.pipeline_config, "context_noise", 0),
|
||||
"pos_start_base":
|
||||
pos_start_base,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=len(timesteps),
|
||||
prompt_embeds=batch.prompt_embeds,
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=batch.prompt_attention_mask,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=batch.image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def block_plan(self, state: StrategyState) -> BlockPlan:
|
||||
block_sizes = state.extra["block_sizes"]
|
||||
start_index = state.extra["start_index"]
|
||||
items: list[BlockPlanItem] = []
|
||||
for block_size in block_sizes:
|
||||
items.append(
|
||||
BlockPlanItem(
|
||||
start_index=start_index,
|
||||
num_frames=block_size,
|
||||
use_kv_cache=True,
|
||||
model_selector="default",
|
||||
))
|
||||
start_index += block_size
|
||||
return BlockPlan(items=items)
|
||||
|
||||
def init_block_context(self, state: StrategyState,
|
||||
block_item: BlockPlanItem,
|
||||
block_idx: int) -> BlockContext:
|
||||
return BlockContext(
|
||||
kv_cache=state.extra["kv_cache1"],
|
||||
kv_cache_2=state.extra["kv_cache2"],
|
||||
crossattn_cache=state.extra["crossattn_cache"],
|
||||
action_cache=None,
|
||||
extra={
|
||||
"block_idx": block_idx,
|
||||
"start_index": block_item.start_index,
|
||||
"num_frames": block_item.num_frames,
|
||||
},
|
||||
)
|
||||
|
||||
def process_block(self, state: StrategyState, block_ctx: BlockContext,
|
||||
block_item: BlockPlanItem) -> None:
|
||||
batch = state.extra["batch"]
|
||||
fastvideo_args = state.extra["fastvideo_args"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
boundary_timestep = state.extra["boundary_timestep"]
|
||||
high_noise_timesteps = state.extra["high_noise_timesteps"]
|
||||
image_kwargs = state.extra["image_kwargs"]
|
||||
pos_cond_kwargs = state.extra["pos_cond_kwargs"]
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
independent_first_frame = state.extra["independent_first_frame"]
|
||||
|
||||
start_index = block_item.start_index
|
||||
current_num_frames = block_item.num_frames
|
||||
|
||||
kv_cache1 = block_ctx.kv_cache
|
||||
kv_cache2 = block_ctx.kv_cache_2
|
||||
crossattn_cache = block_ctx.crossattn_cache
|
||||
|
||||
def _get_kv_cache(timestep_val: float) -> list[dict]:
|
||||
if boundary_timestep is not None:
|
||||
if timestep_val >= boundary_timestep:
|
||||
return kv_cache1
|
||||
if kv_cache2 is None:
|
||||
raise ValueError("kv_cache2 is not initialized")
|
||||
return kv_cache2
|
||||
return kv_cache1
|
||||
|
||||
current_latents = state.latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
|
||||
video_raw_latent_shape = noise_latents_btchw.shape
|
||||
h, w = current_latents.shape[-2:]
|
||||
|
||||
attn_metadata = None
|
||||
for i, t_cur in enumerate(state.timesteps):
|
||||
if boundary_timestep is not None and t_cur < boundary_timestep:
|
||||
current_model = self.stage.transformer_2
|
||||
else:
|
||||
current_model = self.stage.transformer
|
||||
noise_latents = noise_latents_btchw.clone()
|
||||
latent_model_input = current_latents.to(target_dtype)
|
||||
|
||||
if (batch.image_latent is not None and independent_first_frame
|
||||
and start_index == 0):
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input,
|
||||
batch.image_latent.to(target_dtype)],
|
||||
dim=2)
|
||||
|
||||
t_expand = t_cur.repeat(latent_model_input.shape[0])
|
||||
|
||||
if (vsa_available
|
||||
and self.stage.attn_backend == VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = (
|
||||
self.stage.attn_backend.get_builder_cls())
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = (
|
||||
self.attn_metadata_builder_cls())
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=i, # type: ignore
|
||||
raw_latent_shape=(current_num_frames, h,
|
||||
w), # type: ignore
|
||||
patch_size=fastvideo_args.pipeline_config.dit_config.
|
||||
patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.
|
||||
VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(), # type: ignore
|
||||
)
|
||||
assert attn_metadata is not None, (
|
||||
"attn_metadata cannot be None")
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
pred_noise_btchw = current_model(
|
||||
latent_model_input,
|
||||
batch.prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=_get_kv_cache(t_cur),
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=start_index * self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
if boundary_timestep is not None and t_cur >= boundary_timestep:
|
||||
pred_video_btchw = pred_noise_to_x_bound(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
boundary_timestep=torch.ones_like(t_expand) *
|
||||
boundary_timestep,
|
||||
scheduler=self.stage.scheduler,
|
||||
).unflatten(0, pred_noise_btchw.shape[:2])
|
||||
else:
|
||||
pred_video_btchw = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
scheduler=self.stage.scheduler,
|
||||
).unflatten(0, pred_noise_btchw.shape[:2])
|
||||
|
||||
if i < len(state.timesteps) - 1:
|
||||
next_timestep = state.timesteps[i + 1] * torch.ones(
|
||||
[1],
|
||||
dtype=torch.long,
|
||||
device=pred_video_btchw.device,
|
||||
)
|
||||
noise = torch.randn(
|
||||
video_raw_latent_shape,
|
||||
dtype=pred_video_btchw.dtype,
|
||||
generator=(batch.generator[0] if isinstance(
|
||||
batch.generator, list) else batch.generator),
|
||||
).to(self.stage.device)
|
||||
noise_btchw = noise
|
||||
if (boundary_timestep is not None
|
||||
and high_noise_timesteps is not None
|
||||
and i < len(high_noise_timesteps) - 1):
|
||||
noise_latents_btchw = self.stage.scheduler.add_noise_high(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1),
|
||||
next_timestep,
|
||||
torch.ones_like(next_timestep) * boundary_timestep,
|
||||
).unflatten(0, pred_video_btchw.shape[:2])
|
||||
elif (boundary_timestep is not None
|
||||
and high_noise_timesteps is not None
|
||||
and i == len(high_noise_timesteps) - 1):
|
||||
noise_latents_btchw = pred_video_btchw
|
||||
else:
|
||||
noise_latents_btchw = self.stage.scheduler.add_noise(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1),
|
||||
next_timestep,
|
||||
).unflatten(0, pred_video_btchw.shape[:2])
|
||||
current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
current_latents = pred_video_btchw.permute(0, 2, 1, 3, 4)
|
||||
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
block_ctx.extra["attn_metadata"] = attn_metadata
|
||||
state.latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = (current_latents)
|
||||
|
||||
def update_context(self, state: StrategyState, block_ctx: BlockContext,
|
||||
block_item: BlockPlanItem) -> None:
|
||||
batch = state.extra["batch"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
boundary_timestep = state.extra["boundary_timestep"]
|
||||
kv_cache1 = state.extra["kv_cache1"]
|
||||
kv_cache2 = state.extra["kv_cache2"]
|
||||
crossattn_cache = state.extra["crossattn_cache"]
|
||||
image_kwargs = state.extra["image_kwargs"]
|
||||
pos_cond_kwargs = state.extra["pos_cond_kwargs"]
|
||||
context_noise = state.extra["context_noise"]
|
||||
|
||||
start_index = block_item.start_index
|
||||
current_num_frames = block_item.num_frames
|
||||
|
||||
current_latents = state.latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
latents_device = current_latents.device
|
||||
|
||||
t_context = torch.ones([current_latents.shape[0]],
|
||||
device=latents_device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = current_latents.to(target_dtype)
|
||||
|
||||
attn_metadata = block_ctx.extra.get("attn_metadata")
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
|
||||
if boundary_timestep is not None:
|
||||
self.stage.transformer_2(
|
||||
context_bcthw,
|
||||
batch.prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache2,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=start_index * self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
self.stage.transformer(
|
||||
context_bcthw,
|
||||
batch.prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=start_index * self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
|
||||
batch = state.extra["batch"]
|
||||
boundary_timestep = state.extra["boundary_timestep"]
|
||||
|
||||
latents = state.latents
|
||||
if boundary_timestep is not None:
|
||||
num_frames_to_remove = self.num_frames_per_block - 1
|
||||
latents = latents[:, :, :-num_frames_to_remove, :, :]
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
@@ -1,325 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Cosmos denoising strategy using FlowMatchEulerDiscreteScheduler.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
DenoisingStrategy,
|
||||
ModelInputs,
|
||||
StrategyState,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CosmosStrategy(DenoisingStrategy):
|
||||
|
||||
def __init__(self, stage: Any) -> None:
|
||||
self.stage = stage
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
pipeline = self.stage.pipeline() if self.stage.pipeline else None
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.stage.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.stage.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
extra_step_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.scheduler.step,
|
||||
{
|
||||
"generator": batch.generator,
|
||||
"eta": batch.eta
|
||||
},
|
||||
)
|
||||
|
||||
if hasattr(self.stage.transformer, "module"):
|
||||
transformer_dtype = next(
|
||||
self.stage.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.stage.transformer.parameters()).dtype
|
||||
target_dtype = transformer_dtype
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
|
||||
sigma_max = 80.0
|
||||
sigma_min = 0.002
|
||||
sigma_data = 1.0
|
||||
final_sigmas_type = "sigma_min"
|
||||
|
||||
if self.stage.scheduler is not None:
|
||||
self.stage.scheduler.register_to_config(
|
||||
sigma_max=sigma_max,
|
||||
sigma_min=sigma_min,
|
||||
sigma_data=sigma_data,
|
||||
final_sigmas_type=final_sigmas_type,
|
||||
)
|
||||
|
||||
self.stage.scheduler.set_timesteps(num_inference_steps,
|
||||
device=latents.device)
|
||||
timesteps = self.stage.scheduler.timesteps
|
||||
|
||||
if (hasattr(self.stage.scheduler.config, "final_sigmas_type")
|
||||
and self.stage.scheduler.config.final_sigmas_type == "sigma_min"
|
||||
and len(self.stage.scheduler.sigmas) > 1):
|
||||
self.stage.scheduler.sigmas[-1] = self.stage.scheduler.sigmas[-2]
|
||||
|
||||
conditioning_latents = getattr(batch, "conditioning_latents", None)
|
||||
unconditioning_latents = conditioning_latents
|
||||
|
||||
progress_bar = self.stage.progress_bar(total=num_inference_steps)
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"fastvideo_args": fastvideo_args,
|
||||
"extra_step_kwargs": extra_step_kwargs,
|
||||
"target_dtype": target_dtype,
|
||||
"autocast_enabled": autocast_enabled,
|
||||
"progress_bar": progress_bar,
|
||||
"conditioning_latents": conditioning_latents,
|
||||
"unconditioning_latents": unconditioning_latents,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt_embeds=batch.prompt_embeds,
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=batch.prompt_attention_mask,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=batch.image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
state.extra["step_idx"] = step_idx
|
||||
return ModelInputs(
|
||||
latent_model_input=state.latents,
|
||||
timestep=t,
|
||||
prompt_embeds=state.prompt_embeds,
|
||||
prompt_attention_mask=state.prompt_attention_mask,
|
||||
)
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
batch = state.extra["batch"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
conditioning_latents = state.extra["conditioning_latents"]
|
||||
unconditioning_latents = state.extra["unconditioning_latents"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
|
||||
if getattr(self.stage, "interrupt", False):
|
||||
return state.latents
|
||||
|
||||
current_sigma = self.stage.scheduler.sigmas[step_idx]
|
||||
current_t = current_sigma / (current_sigma + 1)
|
||||
c_in = 1 - current_t
|
||||
c_skip = 1 - current_t
|
||||
c_out = -current_t
|
||||
|
||||
timestep = current_t.view(1, 1, 1, 1,
|
||||
1).expand(state.latents.size(0), -1,
|
||||
state.latents.size(2), -1, -1)
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
cond_latent = state.latents * c_in
|
||||
|
||||
if (hasattr(batch, "cond_indicator")
|
||||
and batch.cond_indicator is not None
|
||||
and conditioning_latents is not None):
|
||||
cond_latent = (batch.cond_indicator * conditioning_latents +
|
||||
(1 - batch.cond_indicator) * cond_latent)
|
||||
else:
|
||||
logger.warning(
|
||||
"Step %s: Missing conditioning data - "
|
||||
"cond_indicator: %s, conditioning_latents: %s", step_idx,
|
||||
hasattr(batch, "cond_indicator"), conditioning_latents
|
||||
is not None)
|
||||
|
||||
cond_latent = cond_latent.to(target_dtype)
|
||||
|
||||
cond_timestep = timestep
|
||||
if hasattr(batch,
|
||||
"cond_indicator") and batch.cond_indicator is not None:
|
||||
sigma_conditioning = 0.0001
|
||||
t_conditioning = sigma_conditioning / (sigma_conditioning + 1)
|
||||
cond_timestep = (batch.cond_indicator * t_conditioning +
|
||||
(1 - batch.cond_indicator) * timestep)
|
||||
cond_timestep = cond_timestep.to(target_dtype)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
condition_mask = (batch.cond_mask.to(target_dtype) if hasattr(
|
||||
batch, "cond_mask") else None)
|
||||
padding_mask = torch.zeros(1,
|
||||
1,
|
||||
batch.height,
|
||||
batch.width,
|
||||
device=cond_latent.device,
|
||||
dtype=target_dtype)
|
||||
|
||||
if condition_mask is None:
|
||||
batch_size, _, num_frames, height, width = cond_latent.shape
|
||||
condition_mask = torch.zeros(batch_size,
|
||||
1,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
device=cond_latent.device,
|
||||
dtype=target_dtype)
|
||||
|
||||
noise_pred = self.stage.transformer(
|
||||
hidden_states=cond_latent,
|
||||
timestep=cond_timestep.to(target_dtype),
|
||||
encoder_hidden_states=batch.prompt_embeds[0].to(
|
||||
target_dtype),
|
||||
fps=24,
|
||||
condition_mask=condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
cond_pred = (c_skip * state.latents +
|
||||
c_out * noise_pred.float()).to(target_dtype)
|
||||
|
||||
if (hasattr(batch, "cond_indicator")
|
||||
and batch.cond_indicator is not None
|
||||
and conditioning_latents is not None):
|
||||
cond_pred = (batch.cond_indicator * conditioning_latents +
|
||||
(1 - batch.cond_indicator) * cond_pred)
|
||||
|
||||
if (state.do_cfg and batch.negative_prompt_embeds is not None):
|
||||
uncond_latent = state.latents * c_in
|
||||
|
||||
if (hasattr(batch, "uncond_indicator")
|
||||
and batch.uncond_indicator is not None
|
||||
and unconditioning_latents is not None):
|
||||
uncond_latent = (
|
||||
batch.uncond_indicator * unconditioning_latents +
|
||||
(1 - batch.uncond_indicator) * uncond_latent)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
uncond_condition_mask = (
|
||||
batch.uncond_mask.to(target_dtype) if
|
||||
(hasattr(batch, "uncond_mask")
|
||||
and batch.uncond_mask is not None) else condition_mask)
|
||||
|
||||
uncond_timestep = timestep
|
||||
if (hasattr(batch, "uncond_indicator")
|
||||
and batch.uncond_indicator is not None):
|
||||
sigma_conditioning = 0.0001
|
||||
t_conditioning = sigma_conditioning / (
|
||||
sigma_conditioning + 1)
|
||||
uncond_timestep = (
|
||||
batch.uncond_indicator * t_conditioning +
|
||||
(1 - batch.uncond_indicator) * timestep)
|
||||
uncond_timestep = uncond_timestep.to(target_dtype)
|
||||
|
||||
noise_pred_uncond = self.stage.transformer(
|
||||
hidden_states=uncond_latent.to(target_dtype),
|
||||
timestep=uncond_timestep.to(target_dtype),
|
||||
encoder_hidden_states=batch.negative_prompt_embeds[0].
|
||||
to(target_dtype),
|
||||
fps=24,
|
||||
condition_mask=uncond_condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
uncond_pred = (
|
||||
c_skip * state.latents +
|
||||
c_out * noise_pred_uncond.float()).to(target_dtype)
|
||||
|
||||
if (hasattr(batch, "uncond_indicator")
|
||||
and batch.uncond_indicator is not None
|
||||
and unconditioning_latents is not None):
|
||||
uncond_pred = (
|
||||
batch.uncond_indicator * unconditioning_latents +
|
||||
(1 - batch.uncond_indicator) * uncond_pred)
|
||||
|
||||
guidance_diff = cond_pred - uncond_pred
|
||||
final_pred = cond_pred + state.guidance_scale * guidance_diff
|
||||
else:
|
||||
final_pred = cond_pred
|
||||
|
||||
if current_sigma > 1e-8:
|
||||
noise_for_scheduler = (state.latents - final_pred) / current_sigma
|
||||
else:
|
||||
logger.warning(
|
||||
"Step %s: current_sigma too small (%s), using final_pred directly",
|
||||
step_idx, current_sigma)
|
||||
noise_for_scheduler = final_pred
|
||||
|
||||
if torch.isnan(noise_for_scheduler).sum() > 0:
|
||||
logger.error(
|
||||
"Step %s: NaN detected in noise_for_scheduler, sum: %s",
|
||||
step_idx,
|
||||
noise_for_scheduler.float().sum().item())
|
||||
logger.error(
|
||||
"Step %s: latents sum: %s, final_pred sum: %s, current_sigma: %s",
|
||||
step_idx,
|
||||
state.latents.float().sum().item(),
|
||||
final_pred.float().sum().item(), current_sigma)
|
||||
|
||||
return noise_for_scheduler
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
return noise_pred
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
latents = self.stage.scheduler.step(
|
||||
noise_pred,
|
||||
t,
|
||||
state.latents,
|
||||
**state.extra["extra_step_kwargs"],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
return latents
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
batch = state.extra["batch"]
|
||||
batch.latents = state.latents
|
||||
return batch
|
||||
@@ -1,269 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
DMD denoising strategy (FlowMatch-based).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
DenoisingStrategy,
|
||||
ModelInputs,
|
||||
StrategyState,
|
||||
)
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
VideoSparseAttentionBackend = None # type: ignore
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DmdStrategy(DenoisingStrategy):
|
||||
|
||||
def __init__(self, stage: Any) -> None:
|
||||
self.stage = stage
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
timesteps = batch.timesteps
|
||||
if timesteps is None:
|
||||
raise ValueError("Timesteps must be provided")
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
num_warmup_steps = len(
|
||||
timesteps) - num_inference_steps * self.stage.scheduler.order
|
||||
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"mask_strategy": dict_to_3d_list(
|
||||
None, t_max=50, l_max=60, h_max=24)
|
||||
},
|
||||
)
|
||||
|
||||
pos_cond_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
if (st_attn_available
|
||||
and self.stage.attn_backend == SlidingTileAttentionBackend):
|
||||
self.stage.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents
|
||||
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert not torch.isnan(
|
||||
prompt_embeds[0]).any(), "prompt_embeds contains nan"
|
||||
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
|
||||
progress_bar = self.stage.progress_bar(total=len(timesteps))
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"fastvideo_args": fastvideo_args,
|
||||
"target_dtype": target_dtype,
|
||||
"autocast_enabled": autocast_enabled,
|
||||
"num_warmup_steps": num_warmup_steps,
|
||||
"image_kwargs": image_kwargs,
|
||||
"pos_cond_kwargs": pos_cond_kwargs,
|
||||
"progress_bar": progress_bar,
|
||||
"video_raw_latent_shape": latents.shape,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=batch.prompt_attention_mask,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
state.extra["step_idx"] = step_idx
|
||||
return ModelInputs(
|
||||
latent_model_input=state.latents,
|
||||
timestep=t,
|
||||
prompt_embeds=state.prompt_embeds,
|
||||
prompt_attention_mask=state.prompt_attention_mask,
|
||||
)
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
batch = state.extra["batch"]
|
||||
fastvideo_args = state.extra["fastvideo_args"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
|
||||
if getattr(self.stage, "interrupt", False):
|
||||
state.extra["pred_latents"] = state.latents
|
||||
return state.latents
|
||||
|
||||
latents = state.latents
|
||||
noise_latents = latents.clone()
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
|
||||
if batch.image_latent is not None:
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input,
|
||||
batch.image_latent.permute(0, 2, 1, 3, 4)],
|
||||
dim=2).to(target_dtype)
|
||||
|
||||
t_expand = model_inputs.timestep.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = None
|
||||
if fastvideo_args.pipeline_config.embedded_cfg_scale is not None:
|
||||
guidance_expand = (torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) * 1000.0)
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
if (vsa_available
|
||||
and self.stage.attn_backend == VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = (
|
||||
self.stage.attn_backend.get_builder_cls())
|
||||
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = (
|
||||
self.attn_metadata_builder_cls())
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=step_idx, # type: ignore
|
||||
raw_latent_shape=batch.
|
||||
raw_latent_shape[2:5], # type: ignore
|
||||
patch_size=fastvideo_args.pipeline_config.dit_config.
|
||||
patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.
|
||||
VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(), # type: ignore
|
||||
)
|
||||
assert attn_metadata is not None, (
|
||||
"attn_metadata cannot be None")
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
pred_noise = self.stage.transformer(
|
||||
latent_model_input.permute(0, 2, 1, 3, 4),
|
||||
state.prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**state.extra["image_kwargs"],
|
||||
**state.extra["pos_cond_kwargs"],
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
scheduler=self.stage.scheduler).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
if step_idx < len(state.timesteps) - 1:
|
||||
next_timestep = state.timesteps[step_idx + 1] * torch.ones(
|
||||
[1], dtype=torch.long, device=pred_video.device)
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list):
|
||||
generator = generator[0] if generator else None
|
||||
noise = torch.randn(state.extra["video_raw_latent_shape"],
|
||||
dtype=pred_video.dtype,
|
||||
generator=generator).to(self.stage.device)
|
||||
latents = self.stage.scheduler.add_noise(pred_video.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
next_timestep).unflatten(
|
||||
0,
|
||||
pred_video.shape[:2])
|
||||
else:
|
||||
latents = pred_video
|
||||
|
||||
state.extra["pred_latents"] = latents
|
||||
return latents
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
return noise_pred
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
num_warmup_steps = state.extra["num_warmup_steps"]
|
||||
|
||||
if step_idx == len(state.timesteps) - 1 or (
|
||||
(step_idx + 1) > num_warmup_steps and
|
||||
(step_idx + 1) % self.stage.scheduler.order == 0
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
return state.extra.get("pred_latents", state.latents)
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
|
||||
batch = state.extra["batch"]
|
||||
latents = state.extra["pred_latents"].permute(0, 2, 1, 3, 4)
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -1,177 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Scaffolding for a unified denoising engine with hook support.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import Protocol, TYPE_CHECKING
|
||||
from collections.abc import Sequence
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.schedulers.adapter import SchedulerAdapter
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
BlockDenoisingStrategy,
|
||||
DenoisingStrategy,
|
||||
StrategyState,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
class EngineHook(Protocol):
|
||||
|
||||
def on_init(self, engine: DenoisingEngine, batch: ForwardBatch,
|
||||
args: FastVideoArgs) -> None:
|
||||
...
|
||||
|
||||
def pre_run(self, state: StrategyState) -> None:
|
||||
...
|
||||
|
||||
def pre_step(self, state: StrategyState, step_idx: int,
|
||||
t: torch.Tensor) -> None:
|
||||
...
|
||||
|
||||
def post_step(self, state: StrategyState, step_idx: int,
|
||||
t: torch.Tensor) -> None:
|
||||
...
|
||||
|
||||
def post_run(self, state: StrategyState, batch: ForwardBatch) -> None:
|
||||
...
|
||||
|
||||
|
||||
class BaseEngineHook:
|
||||
|
||||
def on_init(self, engine: DenoisingEngine, batch: ForwardBatch,
|
||||
args: FastVideoArgs) -> None:
|
||||
return None
|
||||
|
||||
def pre_run(self, state: StrategyState) -> None:
|
||||
return None
|
||||
|
||||
def pre_step(self, state: StrategyState, step_idx: int,
|
||||
t: torch.Tensor) -> None:
|
||||
return None
|
||||
|
||||
def post_step(self, state: StrategyState, step_idx: int,
|
||||
t: torch.Tensor) -> None:
|
||||
return None
|
||||
|
||||
def post_run(self, state: StrategyState, batch: ForwardBatch) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class GuidanceCache:
|
||||
|
||||
def __init__(self, maxsize: int = 8) -> None:
|
||||
self._build = lru_cache(maxsize=maxsize)(self._build_impl)
|
||||
|
||||
def _build_impl(self, batch_size: int, dtype: torch.dtype,
|
||||
device: torch.device, guidance_val: float) -> torch.Tensor:
|
||||
return (torch.full(
|
||||
(batch_size, ),
|
||||
guidance_val,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).to(dtype) * 1000.0)
|
||||
|
||||
def get(self, batch_size: int, dtype: torch.dtype, device: torch.device,
|
||||
guidance_val: float | None) -> torch.Tensor | None:
|
||||
if guidance_val is None:
|
||||
return None
|
||||
return self._build(batch_size, dtype, device, guidance_val)
|
||||
|
||||
|
||||
class DenoisingEngine:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
strategy: DenoisingStrategy,
|
||||
*,
|
||||
scheduler_adapter: SchedulerAdapter | None = None,
|
||||
hooks: Sequence[EngineHook] | None = None,
|
||||
) -> None:
|
||||
self.strategy = strategy
|
||||
self.scheduler_adapter = scheduler_adapter
|
||||
self.hooks = list(hooks) if hooks is not None else []
|
||||
|
||||
def run(self, batch: ForwardBatch, args: FastVideoArgs) -> ForwardBatch:
|
||||
for hook in self.hooks:
|
||||
hook.on_init(self, batch, args)
|
||||
|
||||
state = self.strategy.prepare(batch, args)
|
||||
if self.scheduler_adapter is not None:
|
||||
state.extra.setdefault("scheduler_adapter", self.scheduler_adapter)
|
||||
for hook in self.hooks:
|
||||
hook.pre_run(state)
|
||||
|
||||
if isinstance(self.strategy, BlockDenoisingStrategy):
|
||||
self.run_blocks(state)
|
||||
else:
|
||||
timesteps = state.timesteps
|
||||
for i, t in enumerate(timesteps):
|
||||
for hook in self.hooks:
|
||||
hook.pre_step(state, i, t)
|
||||
|
||||
model_inputs = self.strategy.make_model_inputs(state, t, i)
|
||||
noise_pred = self.strategy.forward(state, model_inputs)
|
||||
noise_pred = self.strategy.cfg_combine(state, noise_pred)
|
||||
state.latents = self.strategy.scheduler_step(
|
||||
state,
|
||||
noise_pred,
|
||||
t,
|
||||
)
|
||||
|
||||
for hook in self.hooks:
|
||||
hook.post_step(state, i, t)
|
||||
|
||||
for hook in self.hooks:
|
||||
hook.post_run(state, batch)
|
||||
|
||||
return self.strategy.postprocess(state)
|
||||
|
||||
def run_blocks(
|
||||
self,
|
||||
state: StrategyState,
|
||||
*,
|
||||
block_plan=None,
|
||||
start_block: int = 0,
|
||||
num_blocks: int | None = None,
|
||||
) -> None:
|
||||
if not isinstance(self.strategy, BlockDenoisingStrategy):
|
||||
raise TypeError("run_blocks requires a BlockDenoisingStrategy")
|
||||
|
||||
strategy = self.strategy
|
||||
if block_plan is None:
|
||||
block_plan = strategy.block_plan(state)
|
||||
|
||||
items = block_plan.items
|
||||
end_block = len(items)
|
||||
if num_blocks is not None:
|
||||
end_block = min(end_block, start_block + num_blocks)
|
||||
|
||||
for block_idx in range(start_block, end_block):
|
||||
block_item = items[block_idx]
|
||||
t_hook = self._block_hook_t(state, block_idx)
|
||||
for hook in self.hooks:
|
||||
hook.pre_step(state, block_idx, t_hook)
|
||||
block_ctx = strategy.init_block_context(
|
||||
state,
|
||||
block_item,
|
||||
block_idx,
|
||||
)
|
||||
strategy.process_block(state, block_ctx, block_item)
|
||||
strategy.update_context(state, block_ctx, block_item)
|
||||
for hook in self.hooks:
|
||||
hook.post_step(state, block_idx, t_hook)
|
||||
|
||||
def _block_hook_t(self, state: StrategyState,
|
||||
block_idx: int) -> torch.Tensor:
|
||||
timesteps = state.timesteps
|
||||
if timesteps.numel() > 0:
|
||||
return timesteps[0]
|
||||
return torch.tensor(block_idx, device=state.latents.device)
|
||||
@@ -1,77 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Engine hook implementations for denoising runs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
|
||||
from fastvideo.pipelines.stages.denoising_strategies import StrategyState
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PerfLoggingHook:
|
||||
"""
|
||||
Record per-step and total denoising times into batch logging info.
|
||||
"""
|
||||
|
||||
def __init__(self, stage_name: str = "DenoisingEngine") -> None:
|
||||
self.stage_name = stage_name
|
||||
self._batch: ForwardBatch | None = None
|
||||
self._run_start = 0.0
|
||||
self._step_starts: dict[int, float] = {}
|
||||
self._step_times_ms: list[float] = []
|
||||
|
||||
def on_init(self, engine: DenoisingEngine, batch: ForwardBatch, args):
|
||||
self._batch = batch
|
||||
|
||||
def pre_run(self, state: StrategyState) -> None:
|
||||
self._run_start = time.perf_counter()
|
||||
self._step_starts.clear()
|
||||
self._step_times_ms.clear()
|
||||
|
||||
def pre_step(self, state: StrategyState, step_idx: int, t) -> None:
|
||||
self._step_starts[step_idx] = time.perf_counter()
|
||||
|
||||
def post_step(self, state: StrategyState, step_idx: int, t) -> None:
|
||||
start = self._step_starts.pop(step_idx, None)
|
||||
if start is None:
|
||||
return
|
||||
self._step_times_ms.append((time.perf_counter() - start) * 1000.0)
|
||||
|
||||
def post_run(self, state: StrategyState, batch: ForwardBatch) -> None:
|
||||
total_ms = (time.perf_counter() - self._run_start) * 1000.0
|
||||
target_batch = self._batch or batch
|
||||
if target_batch is None:
|
||||
return
|
||||
target_batch.logging_info.add_stage_metric(self.stage_name,
|
||||
"denoise_step_times_ms",
|
||||
list(self._step_times_ms))
|
||||
target_batch.logging_info.add_stage_metric(self.stage_name,
|
||||
"denoise_total_ms", total_ms)
|
||||
if not self._step_times_ms:
|
||||
return
|
||||
if not self._is_primary_rank():
|
||||
return
|
||||
mean_ms = sum(self._step_times_ms) / len(self._step_times_ms)
|
||||
logger.info(
|
||||
"[%s] denoise steps=%d total_ms=%.2f mean_step_ms=%.2f min=%.2f max=%.2f",
|
||||
self.stage_name,
|
||||
len(self._step_times_ms),
|
||||
total_ms,
|
||||
mean_ms,
|
||||
min(self._step_times_ms),
|
||||
max(self._step_times_ms),
|
||||
)
|
||||
|
||||
def _is_primary_rank(self) -> bool:
|
||||
try:
|
||||
from fastvideo.distributed import get_world_group
|
||||
return get_world_group().local_rank == 0
|
||||
except Exception:
|
||||
return True
|
||||
@@ -1,551 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat denoising strategies (base, I2V, VC).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
DenoisingStrategy,
|
||||
ModelInputs,
|
||||
StrategyState,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class _BaseLongCatStrategy(DenoisingStrategy):
|
||||
|
||||
def __init__(self, stage: Any) -> None:
|
||||
self.stage = stage
|
||||
|
||||
def _load_transformer(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
attach_pipeline: bool) -> None:
|
||||
if fastvideo_args.model_loaded["transformer"]:
|
||||
return
|
||||
loader = TransformerLoader()
|
||||
self.stage.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if attach_pipeline:
|
||||
pipeline = self.stage.pipeline() if self.stage.pipeline else None
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.stage.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
def _build_prompt_inputs(self, batch: ForwardBatch):
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
do_cfg = batch.do_classifier_free_guidance
|
||||
|
||||
if do_cfg:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
return prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
|
||||
prompt_attention_mask_combined
|
||||
|
||||
def optimized_scale(self, positive_flat: torch.Tensor,
|
||||
negative_flat: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Calculate optimized scale from CFG-zero paper.
|
||||
|
||||
st_star = (v_cond^T * v_uncond) / ||v_uncond||^2
|
||||
"""
|
||||
dot_product = torch.sum(positive_flat * negative_flat,
|
||||
dim=1,
|
||||
keepdim=True)
|
||||
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
|
||||
return dot_product / squared_norm
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
return noise_pred
|
||||
|
||||
|
||||
class LongCatStrategy(_BaseLongCatStrategy):
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
self._load_transformer(batch, fastvideo_args, attach_pipeline=True)
|
||||
|
||||
if hasattr(self.stage.transformer, "module"):
|
||||
transformer_dtype = next(
|
||||
self.stage.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.stage.transformer.parameters()).dtype
|
||||
|
||||
target_dtype = transformer_dtype
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
|
||||
prompt_attention_mask_combined = self._build_prompt_inputs(batch)
|
||||
|
||||
timesteps = batch.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
progress_bar = tqdm(total=num_inference_steps, desc="LongCat Denoising")
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"target_dtype": target_dtype,
|
||||
"autocast_enabled": autocast_enabled,
|
||||
"prompt_embeds_combined": prompt_embeds_combined,
|
||||
"prompt_attention_mask_combined": prompt_attention_mask_combined,
|
||||
"progress_bar": progress_bar,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=batch.latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt_embeds=[prompt_embeds],
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=[prompt_attention_mask]
|
||||
if prompt_attention_mask is not None else None,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=batch.image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
latents = state.latents
|
||||
|
||||
if state.do_cfg:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
state.extra["step_idx"] = step_idx
|
||||
return ModelInputs(
|
||||
latent_model_input=latent_model_input,
|
||||
timestep=timestep,
|
||||
prompt_embeds=state.prompt_embeds,
|
||||
prompt_attention_mask=state.prompt_attention_mask,
|
||||
)
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
batch = state.extra["batch"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
prompt_embeds_combined = state.extra["prompt_embeds_combined"]
|
||||
prompt_attention_mask_combined = state.extra[
|
||||
"prompt_attention_mask_combined"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.stage.transformer(
|
||||
hidden_states=model_inputs.latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=model_inputs.timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
)
|
||||
|
||||
if state.do_cfg:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
noise_pred = (noise_pred_uncond * st_star + state.guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
noise_pred = -noise_pred
|
||||
return noise_pred
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
latents = self.stage.scheduler.step(noise_pred,
|
||||
t,
|
||||
state.latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
return latents
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
batch = state.extra["batch"]
|
||||
batch.latents = state.latents
|
||||
return batch
|
||||
|
||||
|
||||
class LongCatI2VStrategy(_BaseLongCatStrategy):
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
self._load_transformer(batch, fastvideo_args, attach_pipeline=False)
|
||||
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
|
||||
prompt_attention_mask_combined = self._build_prompt_inputs(batch)
|
||||
|
||||
timesteps = batch.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
progress_bar = tqdm(total=num_inference_steps, desc="I2V Denoising")
|
||||
|
||||
num_cond_latents = getattr(batch, "num_cond_latents", 0)
|
||||
if num_cond_latents > 0:
|
||||
logger.info("I2V Denoising: num_cond_latents=%s, latent_shape=%s",
|
||||
num_cond_latents, batch.latents.shape)
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"target_dtype": target_dtype,
|
||||
"autocast_enabled": autocast_enabled,
|
||||
"prompt_embeds_combined": prompt_embeds_combined,
|
||||
"prompt_attention_mask_combined": prompt_attention_mask_combined,
|
||||
"num_cond_latents": num_cond_latents,
|
||||
"progress_bar": progress_bar,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=batch.latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt_embeds=[prompt_embeds],
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=[prompt_attention_mask]
|
||||
if prompt_attention_mask is not None else None,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=batch.image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
num_cond_latents = state.extra["num_cond_latents"]
|
||||
latents = state.latents
|
||||
|
||||
if state.do_cfg:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(target_dtype)
|
||||
timestep = timestep.unsqueeze(-1).repeat(1, latent_model_input.shape[2])
|
||||
if num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
state.extra["step_idx"] = step_idx
|
||||
return ModelInputs(
|
||||
latent_model_input=latent_model_input,
|
||||
timestep=timestep,
|
||||
prompt_embeds=state.prompt_embeds,
|
||||
prompt_attention_mask=state.prompt_attention_mask,
|
||||
extra_kwargs={"num_cond_latents": num_cond_latents},
|
||||
)
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
batch = state.extra["batch"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
prompt_embeds_combined = state.extra["prompt_embeds_combined"]
|
||||
prompt_attention_mask_combined = state.extra[
|
||||
"prompt_attention_mask_combined"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
num_cond_latents = state.extra["num_cond_latents"]
|
||||
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.stage.transformer(
|
||||
hidden_states=model_inputs.latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=model_inputs.timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
num_cond_latents=num_cond_latents,
|
||||
)
|
||||
|
||||
if state.do_cfg:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
noise_pred = (noise_pred_uncond * st_star + state.guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
noise_pred = -noise_pred
|
||||
return noise_pred
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
num_cond_latents = state.extra["num_cond_latents"]
|
||||
latents = state.latents
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.stage.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
latents = self.stage.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
return latents
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
batch = state.extra["batch"]
|
||||
batch.latents = state.latents
|
||||
return batch
|
||||
|
||||
|
||||
class LongCatVCStrategy(_BaseLongCatStrategy):
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
self._load_transformer(batch, fastvideo_args, attach_pipeline=False)
|
||||
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
|
||||
prompt_attention_mask_combined = self._build_prompt_inputs(batch)
|
||||
|
||||
timesteps = batch.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
progress_bar = tqdm(total=num_inference_steps, desc="VC Denoising")
|
||||
|
||||
num_cond_latents = getattr(batch, "num_cond_latents", 0)
|
||||
use_kv_cache = getattr(batch, "use_kv_cache", False)
|
||||
kv_cache_dict = getattr(batch, "kv_cache_dict", {})
|
||||
|
||||
logger.info(
|
||||
"VC Denoising: num_cond_latents=%d, use_kv_cache=%s, latent_shape=%s",
|
||||
num_cond_latents, use_kv_cache, batch.latents.shape)
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"target_dtype": target_dtype,
|
||||
"autocast_enabled": autocast_enabled,
|
||||
"prompt_embeds_combined": prompt_embeds_combined,
|
||||
"prompt_attention_mask_combined": prompt_attention_mask_combined,
|
||||
"num_cond_latents": num_cond_latents,
|
||||
"use_kv_cache": use_kv_cache,
|
||||
"kv_cache_dict": kv_cache_dict,
|
||||
"progress_bar": progress_bar,
|
||||
"step_times": [],
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=batch.latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt_embeds=[prompt_embeds],
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=[prompt_attention_mask]
|
||||
if prompt_attention_mask is not None else None,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=batch.image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
num_cond_latents = state.extra["num_cond_latents"]
|
||||
use_kv_cache = state.extra["use_kv_cache"]
|
||||
latents = state.latents
|
||||
|
||||
if state.do_cfg:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(target_dtype)
|
||||
timestep = timestep.unsqueeze(-1).repeat(1, latent_model_input.shape[2])
|
||||
if not use_kv_cache and num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
state.extra["step_idx"] = step_idx
|
||||
state.extra["step_start"] = time.time()
|
||||
|
||||
extra_kwargs = {"num_cond_latents": num_cond_latents}
|
||||
if use_kv_cache:
|
||||
extra_kwargs["kv_cache_dict"] = state.extra["kv_cache_dict"]
|
||||
|
||||
return ModelInputs(
|
||||
latent_model_input=latent_model_input,
|
||||
timestep=timestep,
|
||||
prompt_embeds=state.prompt_embeds,
|
||||
prompt_attention_mask=state.prompt_attention_mask,
|
||||
extra_kwargs=extra_kwargs,
|
||||
)
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
batch = state.extra["batch"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
prompt_embeds_combined = state.extra["prompt_embeds_combined"]
|
||||
prompt_attention_mask_combined = state.extra[
|
||||
"prompt_attention_mask_combined"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.stage.transformer(
|
||||
hidden_states=model_inputs.latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=model_inputs.timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
**model_inputs.extra_kwargs,
|
||||
)
|
||||
|
||||
if state.do_cfg:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
noise_pred = (noise_pred_uncond * st_star + state.guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
noise_pred = -noise_pred
|
||||
return noise_pred
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
num_cond_latents = state.extra["num_cond_latents"]
|
||||
use_kv_cache = state.extra["use_kv_cache"]
|
||||
latents = state.latents
|
||||
|
||||
if use_kv_cache:
|
||||
latents = self.stage.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.stage.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
else:
|
||||
latents = self.stage.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
step_time = time.time() - state.extra["step_start"]
|
||||
state.extra["step_times"].append(step_time)
|
||||
if state.extra["step_idx"] < 3:
|
||||
logger.info("Step %d: %.2fs", state.extra["step_idx"], step_time)
|
||||
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
return latents
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
|
||||
batch = state.extra["batch"]
|
||||
use_kv_cache = state.extra["use_kv_cache"]
|
||||
step_times = state.extra["step_times"]
|
||||
|
||||
latents = state.latents
|
||||
if use_kv_cache and hasattr(
|
||||
batch, "cond_latents") and batch.cond_latents is not None:
|
||||
latents = torch.cat([batch.cond_latents, latents], dim=2)
|
||||
logger.info(
|
||||
"Concatenated conditioning latents back, final shape: %s",
|
||||
latents.shape)
|
||||
|
||||
if step_times:
|
||||
avg_time = sum(step_times) / len(step_times)
|
||||
logger.info("Average step time: %.2fs (total: %.1fs)", avg_time,
|
||||
sum(step_times))
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -1,324 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
MatrixGame causal block denoising strategy.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
BlockContext,
|
||||
BlockDenoisingStrategy,
|
||||
BlockPlan,
|
||||
BlockPlanItem,
|
||||
ModelInputs,
|
||||
StrategyState,
|
||||
)
|
||||
from fastvideo.pipelines.stages.matrixgame_denoising import BlockProcessingContext
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
|
||||
class MatrixGameBlockStrategy(BlockDenoisingStrategy):
|
||||
|
||||
def __init__(self, stage: Any) -> None:
|
||||
self.stage = stage
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
target_dtype = torch.bfloat16
|
||||
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")
|
||||
|
||||
latent_seq_length = latents.shape[-1] * latents.shape[-2]
|
||||
patch_size = self.stage.transformer.patch_size
|
||||
patch_ratio = patch_size[-1] * patch_size[-2]
|
||||
self.stage.frame_seq_length = latent_seq_length // patch_ratio
|
||||
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step",
|
||||
False):
|
||||
scheduler_timesteps = torch.cat((
|
||||
self.stage.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0], dtype=torch.float32),
|
||||
))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
|
||||
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
|
||||
"boundary_ratio", None)
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = (boundary_ratio *
|
||||
self.stage.scheduler.num_train_timesteps)
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = {"encoder_hidden_states_image": image_embeds}
|
||||
pos_cond_kwargs: dict[str, Any] = {}
|
||||
|
||||
if (st_attn_available
|
||||
and self.stage.attn_backend == SlidingTileAttentionBackend):
|
||||
self.stage.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
kv_cache1 = self.stage._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
kv_cache2 = None
|
||||
if boundary_timestep is not None:
|
||||
kv_cache2 = self.stage._initialize_kv_cache(
|
||||
batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device,
|
||||
)
|
||||
|
||||
kv_cache_mouse = None
|
||||
kv_cache_keyboard = None
|
||||
if self.stage.use_action_module:
|
||||
kv_cache_mouse, kv_cache_keyboard = (
|
||||
self.stage._initialize_action_kv_cache(
|
||||
batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device,
|
||||
))
|
||||
|
||||
crossattn_cache = self.stage._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=257,
|
||||
dtype=target_dtype,
|
||||
device=latents.device,
|
||||
)
|
||||
|
||||
num_frames = latents.shape[2]
|
||||
if num_frames % self.stage.num_frame_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frame_per_block for "
|
||||
"causal denoising")
|
||||
num_blocks = num_frames // self.stage.num_frame_per_block
|
||||
block_sizes = [self.stage.num_frame_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
ctx = BlockProcessingContext(
|
||||
batch=batch,
|
||||
block_idx=0,
|
||||
start_index=0,
|
||||
kv_cache1=kv_cache1,
|
||||
kv_cache2=kv_cache2,
|
||||
kv_cache_mouse=kv_cache_mouse,
|
||||
kv_cache_keyboard=kv_cache_keyboard,
|
||||
crossattn_cache=crossattn_cache,
|
||||
timesteps=timesteps,
|
||||
block_sizes=block_sizes,
|
||||
noise_pool=None,
|
||||
fastvideo_args=fastvideo_args,
|
||||
target_dtype=target_dtype,
|
||||
autocast_enabled=autocast_enabled,
|
||||
boundary_timestep=boundary_timestep,
|
||||
high_noise_timesteps=high_noise_timesteps,
|
||||
context_noise=getattr(fastvideo_args.pipeline_config,
|
||||
"context_noise", 0),
|
||||
image_kwargs=image_kwargs,
|
||||
pos_cond_kwargs=pos_cond_kwargs,
|
||||
)
|
||||
|
||||
progress_bar = self.stage.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps))
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"fastvideo_args": fastvideo_args,
|
||||
"ctx": ctx,
|
||||
"block_sizes": block_sizes,
|
||||
"start_index": start_index,
|
||||
"progress_bar": progress_bar,
|
||||
"boundary_timestep": boundary_timestep,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=len(timesteps),
|
||||
prompt_embeds=batch.prompt_embeds,
|
||||
negative_prompt_embeds=batch.negative_prompt_embeds,
|
||||
prompt_attention_mask=batch.prompt_attention_mask,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def block_plan(self, state: StrategyState) -> BlockPlan:
|
||||
block_sizes = state.extra["block_sizes"]
|
||||
start_index = state.extra["start_index"]
|
||||
items: list[BlockPlanItem] = []
|
||||
for block_size in block_sizes:
|
||||
items.append(
|
||||
BlockPlanItem(
|
||||
start_index=start_index,
|
||||
num_frames=block_size,
|
||||
use_kv_cache=True,
|
||||
model_selector="default",
|
||||
))
|
||||
start_index += block_size
|
||||
return BlockPlan(items=items)
|
||||
|
||||
def init_block_context(self, state: StrategyState,
|
||||
block_item: BlockPlanItem,
|
||||
block_idx: int) -> BlockContext:
|
||||
ctx = state.extra["ctx"]
|
||||
ctx.block_idx = block_idx
|
||||
ctx.start_index = block_item.start_index
|
||||
|
||||
action_kwargs = self.stage._prepare_action_kwargs(
|
||||
state.extra["batch"],
|
||||
block_item.start_index,
|
||||
block_item.num_frames,
|
||||
)
|
||||
|
||||
return BlockContext(
|
||||
kv_cache=ctx.kv_cache1,
|
||||
kv_cache_2=ctx.kv_cache2,
|
||||
crossattn_cache=ctx.crossattn_cache,
|
||||
action_cache=None,
|
||||
extra={
|
||||
"ctx": ctx,
|
||||
"action_kwargs": action_kwargs,
|
||||
"start_index": block_item.start_index,
|
||||
"num_frames": block_item.num_frames,
|
||||
},
|
||||
)
|
||||
|
||||
def process_block(self, state: StrategyState, block_ctx: BlockContext,
|
||||
block_item: BlockPlanItem) -> None:
|
||||
ctx = block_ctx.extra["ctx"]
|
||||
batch = state.extra["batch"]
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
|
||||
start_index = block_item.start_index
|
||||
current_num_frames = block_item.num_frames
|
||||
action_kwargs = block_ctx.extra["action_kwargs"]
|
||||
|
||||
current_latents = state.latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
noise_generator = None
|
||||
if ctx.noise_pool is not None:
|
||||
latents_device = state.latents.device
|
||||
|
||||
def noise_generator(shape: tuple, dtype: torch.dtype,
|
||||
step_idx: int) -> torch.Tensor:
|
||||
if step_idx < len(ctx.noise_pool):
|
||||
noise = ctx.noise_pool[step_idx]
|
||||
if noise.shape != shape:
|
||||
noise = noise[:, :shape[1], :, :, :]
|
||||
return noise.to(device=latents_device, dtype=dtype)
|
||||
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list):
|
||||
generator = generator[0] if generator else None
|
||||
return torch.randn(shape, dtype=dtype,
|
||||
generator=generator).to(latents_device)
|
||||
|
||||
current_latents = self.stage._process_single_block(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
timesteps=state.timesteps,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
progress_bar=progress_bar,
|
||||
noise_generator=noise_generator,
|
||||
)
|
||||
|
||||
state.latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = (current_latents)
|
||||
|
||||
def update_context(self, state: StrategyState, block_ctx: BlockContext,
|
||||
block_item: BlockPlanItem) -> None:
|
||||
ctx = block_ctx.extra["ctx"]
|
||||
batch = state.extra["batch"]
|
||||
action_kwargs = block_ctx.extra["action_kwargs"]
|
||||
start_index = block_item.start_index
|
||||
current_num_frames = block_item.num_frames
|
||||
|
||||
current_latents = state.latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
self.stage._update_context_cache(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
context_noise=ctx.context_noise,
|
||||
)
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
|
||||
batch = state.extra["batch"]
|
||||
boundary_timestep = state.extra["boundary_timestep"]
|
||||
|
||||
latents = state.latents
|
||||
if boundary_timestep is not None:
|
||||
num_frames_to_remove = self.stage.num_frame_per_block - 1
|
||||
if num_frames_to_remove > 0:
|
||||
latents = latents[:, :, :-num_frames_to_remove, :, :]
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
@@ -1,552 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Standard denoising strategy backed by the legacy DenoisingStage utilities.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines.base import STA_Mode
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising_strategies import (
|
||||
DenoisingStrategy,
|
||||
ModelInputs,
|
||||
StrategyState,
|
||||
)
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
|
||||
from fastvideo.utils import is_vmoba_available
|
||||
vmoba_attn_available = is_vmoba_available()
|
||||
except ImportError:
|
||||
vmoba_attn_available = False
|
||||
VMOBAAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
VideoSparseAttentionBackend = None # type: ignore
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class StandardStrategy(DenoisingStrategy):
|
||||
|
||||
def __init__(self, stage: Any) -> None:
|
||||
self.stage = stage
|
||||
|
||||
def prepare(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> StrategyState:
|
||||
pipeline = self.stage.pipeline() if self.stage.pipeline else None
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.stage.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.stage.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
extra_step_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.scheduler.step,
|
||||
{
|
||||
"generator": batch.generator,
|
||||
"eta": batch.eta
|
||||
},
|
||||
)
|
||||
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
timesteps = batch.timesteps
|
||||
if timesteps is None:
|
||||
raise ValueError("Timesteps must be provided")
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
num_warmup_steps = len(
|
||||
timesteps) - num_inference_steps * self.stage.scheduler.order
|
||||
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert not torch.isnan(
|
||||
image_embeds[0]).any(), "image_embeds contains nan"
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"mask_strategy": dict_to_3d_list(
|
||||
None, t_max=50, l_max=60, h_max=24)
|
||||
},
|
||||
)
|
||||
|
||||
pos_cond_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
neg_cond_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_2": batch.clip_embedding_neg,
|
||||
"encoder_attention_mask": batch.negative_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
action_kwargs = self.stage.prepare_extra_func_kwargs(
|
||||
self.stage.transformer.forward,
|
||||
{
|
||||
"mouse_cond": batch.mouse_cond,
|
||||
"keyboard_cond": batch.keyboard_cond,
|
||||
},
|
||||
)
|
||||
|
||||
if (st_attn_available
|
||||
and self.stage.attn_backend == SlidingTileAttentionBackend):
|
||||
self.stage.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert not torch.isnan(
|
||||
prompt_embeds[0]).any(), "prompt_embeds contains nan"
|
||||
|
||||
neg_prompt_embeds = None
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds
|
||||
assert neg_prompt_embeds is not None
|
||||
assert not torch.isnan(
|
||||
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
boundary_ratio = (
|
||||
fastvideo_args.pipeline_config.dit_config.boundary_ratio)
|
||||
if batch.boundary_ratio is not None:
|
||||
logger.info("Overriding boundary ratio from %s to %s",
|
||||
boundary_ratio, batch.boundary_ratio)
|
||||
boundary_ratio = batch.boundary_ratio
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = (boundary_ratio *
|
||||
self.stage.scheduler.num_train_timesteps)
|
||||
else:
|
||||
boundary_timestep = None
|
||||
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if latent_model_input.ndim == 5:
|
||||
assert latent_model_input.shape[0] == 1, (
|
||||
"only support batch size 1")
|
||||
|
||||
ti2v_mask = None
|
||||
ti2v_z = None
|
||||
ti2v_seq_len = None
|
||||
if (fastvideo_args.pipeline_config.ti2v_task
|
||||
and batch.pil_image is not None):
|
||||
assert batch.image_latent is None, (
|
||||
"TI2V task should not have image latents")
|
||||
assert self.stage.vae is not None, (
|
||||
"VAE is not provided for TI2V task")
|
||||
z = self.stage.vae.encode(batch.pil_image).mean.float()
|
||||
if (hasattr(self.stage.vae, "shift_factor")
|
||||
and self.stage.vae.shift_factor is not None):
|
||||
if isinstance(self.stage.vae.shift_factor, torch.Tensor):
|
||||
z -= self.stage.vae.shift_factor.to(z.device, z.dtype)
|
||||
else:
|
||||
z -= self.stage.vae.shift_factor
|
||||
|
||||
if isinstance(self.stage.vae.scaling_factor, torch.Tensor):
|
||||
z = z * self.stage.vae.scaling_factor.to(z.device, z.dtype)
|
||||
else:
|
||||
z = z * self.stage.vae.scaling_factor
|
||||
|
||||
latent_model_input = latents.to(target_dtype).squeeze(0)
|
||||
_, mask2 = masks_like([latent_model_input], zero=True)
|
||||
latent_model_input = ((1. - mask2[0]) * z +
|
||||
mask2[0] * latent_model_input)
|
||||
latent_model_input = latent_model_input.to(get_local_torch_device())
|
||||
latents = latent_model_input
|
||||
F = batch.num_frames
|
||||
temporal_scale = (fastvideo_args.pipeline_config.vae_config.
|
||||
arch_config.scale_factor_temporal)
|
||||
spatial_scale = (fastvideo_args.pipeline_config.vae_config.
|
||||
arch_config.scale_factor_spatial)
|
||||
patch_size = (fastvideo_args.pipeline_config.dit_config.arch_config.
|
||||
patch_size)
|
||||
seq_len = ((F - 1) // temporal_scale +
|
||||
1) * (batch.height // spatial_scale) * (
|
||||
batch.width // spatial_scale) // (patch_size[1] *
|
||||
patch_size[2])
|
||||
ti2v_mask = mask2[0]
|
||||
ti2v_z = z
|
||||
ti2v_seq_len = seq_len
|
||||
|
||||
trajectory_timesteps: list[torch.Tensor] | None = None
|
||||
trajectory_latents: list[torch.Tensor] | None = None
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps = []
|
||||
trajectory_latents = []
|
||||
|
||||
progress_bar = self.stage.progress_bar(total=num_inference_steps)
|
||||
|
||||
extra: dict[str, Any] = {
|
||||
"batch": batch,
|
||||
"fastvideo_args": fastvideo_args,
|
||||
"extra_step_kwargs": extra_step_kwargs,
|
||||
"target_dtype": target_dtype,
|
||||
"autocast_enabled": autocast_enabled,
|
||||
"num_warmup_steps": num_warmup_steps,
|
||||
"image_kwargs": image_kwargs,
|
||||
"pos_cond_kwargs": pos_cond_kwargs,
|
||||
"neg_cond_kwargs": neg_cond_kwargs,
|
||||
"action_kwargs": action_kwargs,
|
||||
"boundary_timestep": boundary_timestep,
|
||||
"progress_bar": progress_bar,
|
||||
"trajectory_timesteps": trajectory_timesteps,
|
||||
"trajectory_latents": trajectory_latents,
|
||||
"ti2v_mask": ti2v_mask,
|
||||
"ti2v_z": ti2v_z,
|
||||
"ti2v_seq_len": ti2v_seq_len,
|
||||
}
|
||||
|
||||
return StrategyState(
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=neg_prompt_embeds,
|
||||
prompt_attention_mask=batch.prompt_attention_mask,
|
||||
negative_attention_mask=batch.negative_attention_mask,
|
||||
image_embeds=image_embeds,
|
||||
guidance_scale=batch.guidance_scale,
|
||||
guidance_scale_2=batch.guidance_scale_2,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
do_cfg=batch.do_classifier_free_guidance,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
batch = state.extra["batch"]
|
||||
fastvideo_args = state.extra["fastvideo_args"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
boundary_timestep = state.extra["boundary_timestep"]
|
||||
|
||||
if getattr(self.stage, "interrupt", False):
|
||||
state.extra["skip_step"] = True
|
||||
else:
|
||||
state.extra["skip_step"] = False
|
||||
|
||||
if boundary_timestep is None or t >= boundary_timestep:
|
||||
if (fastvideo_args.dit_cpu_offload
|
||||
and not fastvideo_args.dit_layerwise_offload
|
||||
and self.stage.transformer_2 is not None
|
||||
and next(self.stage.transformer_2.parameters()).device.type
|
||||
== 'cuda'):
|
||||
self.stage.transformer_2.to('cpu')
|
||||
current_model = self.stage.transformer
|
||||
if (fastvideo_args.dit_cpu_offload
|
||||
and not fastvideo_args.dit_layerwise_offload
|
||||
and not fastvideo_args.use_fsdp_inference
|
||||
and current_model is not None):
|
||||
transformer_device = next(
|
||||
current_model.parameters()).device.type
|
||||
if transformer_device == 'cpu':
|
||||
current_model.to(get_local_torch_device())
|
||||
current_guidance_scale = batch.guidance_scale
|
||||
else:
|
||||
if (fastvideo_args.dit_cpu_offload
|
||||
and not fastvideo_args.dit_layerwise_offload
|
||||
and next(self.stage.transformer.parameters()).device.type
|
||||
== 'cuda'):
|
||||
self.stage.transformer.to('cpu')
|
||||
current_model = self.stage.transformer_2
|
||||
if (fastvideo_args.dit_cpu_offload
|
||||
and not fastvideo_args.dit_layerwise_offload
|
||||
and not fastvideo_args.use_fsdp_inference
|
||||
and current_model is not None):
|
||||
transformer_2_device = next(
|
||||
current_model.parameters()).device.type
|
||||
if transformer_2_device == 'cpu':
|
||||
current_model.to(get_local_torch_device())
|
||||
current_guidance_scale = batch.guidance_scale_2
|
||||
|
||||
assert current_model is not None, "current_model is None"
|
||||
state.extra["current_model"] = current_model
|
||||
state.extra["current_guidance_scale"] = current_guidance_scale
|
||||
state.extra["step_idx"] = step_idx
|
||||
|
||||
latent_model_input = state.latents.to(target_dtype)
|
||||
if batch.video_latent is not None:
|
||||
latent_model_input = torch.cat([
|
||||
latent_model_input, batch.video_latent,
|
||||
torch.zeros_like(state.latents)
|
||||
],
|
||||
dim=1).to(target_dtype)
|
||||
elif batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, (
|
||||
"image latents should not be provided for TI2V task")
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
|
||||
if (fastvideo_args.pipeline_config.ti2v_task
|
||||
and batch.pil_image is not None):
|
||||
timestep = torch.stack([t]).to(get_local_torch_device())
|
||||
mask2 = state.extra["ti2v_mask"]
|
||||
seq_len = state.extra["ti2v_seq_len"]
|
||||
temp_ts = (mask2[0][:, ::2, ::2] * timestep).flatten()
|
||||
temp_ts = torch.cat([
|
||||
temp_ts,
|
||||
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
|
||||
])
|
||||
timestep = temp_ts.unsqueeze(0)
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
|
||||
latent_model_input = self.stage.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
guidance_expand = None
|
||||
if fastvideo_args.pipeline_config.embedded_cfg_scale is not None:
|
||||
guidance_expand = (torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) * 1000.0)
|
||||
|
||||
state.extra["guidance_expand"] = guidance_expand
|
||||
return ModelInputs(
|
||||
latent_model_input=latent_model_input,
|
||||
timestep=t_expand,
|
||||
prompt_embeds=state.prompt_embeds,
|
||||
prompt_attention_mask=state.prompt_attention_mask,
|
||||
)
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
if state.extra.get("skip_step", False):
|
||||
return state.latents
|
||||
|
||||
batch = state.extra["batch"]
|
||||
fastvideo_args = state.extra["fastvideo_args"]
|
||||
target_dtype = state.extra["target_dtype"]
|
||||
autocast_enabled = state.extra["autocast_enabled"]
|
||||
current_model = state.extra["current_model"]
|
||||
current_guidance_scale = state.extra["current_guidance_scale"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
guidance_expand = state.extra["guidance_expand"]
|
||||
|
||||
image_kwargs = state.extra["image_kwargs"]
|
||||
pos_cond_kwargs = state.extra["pos_cond_kwargs"]
|
||||
neg_cond_kwargs = state.extra["neg_cond_kwargs"]
|
||||
action_kwargs = state.extra["action_kwargs"]
|
||||
|
||||
if ((st_attn_available
|
||||
and self.stage.attn_backend == SlidingTileAttentionBackend) or
|
||||
(vsa_available
|
||||
and self.stage.attn_backend == VideoSparseAttentionBackend)):
|
||||
self.attn_metadata_builder_cls = (
|
||||
self.stage.attn_backend.get_builder_cls())
|
||||
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls()
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=step_idx, # type: ignore
|
||||
raw_latent_shape=batch.
|
||||
raw_latent_shape[2:5], # type: ignore
|
||||
patch_size=fastvideo_args.pipeline_config.dit_config.
|
||||
patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(),
|
||||
)
|
||||
assert attn_metadata is not None, (
|
||||
"attn_metadata cannot be None")
|
||||
else:
|
||||
attn_metadata = None
|
||||
elif (vmoba_attn_available
|
||||
and self.stage.attn_backend == VMOBAAttentionBackend):
|
||||
self.attn_metadata_builder_cls = (
|
||||
self.stage.attn_backend.get_builder_cls())
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls()
|
||||
moba_params = fastvideo_args.moba_config.copy()
|
||||
moba_params.update({
|
||||
"current_timestep":
|
||||
step_idx,
|
||||
"raw_latent_shape":
|
||||
batch.raw_latent_shape[2:5],
|
||||
"patch_size":
|
||||
fastvideo_args.pipeline_config.dit_config.patch_size,
|
||||
"device":
|
||||
get_local_torch_device(),
|
||||
})
|
||||
attn_metadata = self.attn_metadata_builder.build(**moba_params)
|
||||
assert attn_metadata is not None, (
|
||||
"attn_metadata cannot be None")
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
noise_pred = current_model(
|
||||
model_inputs.latent_model_input,
|
||||
state.prompt_embeds,
|
||||
model_inputs.timestep,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
|
||||
if state.do_cfg:
|
||||
batch.is_cfg_negative = True
|
||||
with set_forward_context(
|
||||
current_timestep=step_idx,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
noise_pred_uncond = current_model(
|
||||
model_inputs.latent_model_input,
|
||||
state.negative_prompt_embeds,
|
||||
model_inputs.timestep,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
|
||||
noise_pred_text = noise_pred
|
||||
noise_pred = noise_pred_uncond + current_guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond)
|
||||
|
||||
if state.guidance_rescale > 0.0:
|
||||
noise_pred = self.stage.rescale_noise_cfg(
|
||||
noise_pred,
|
||||
noise_pred_text,
|
||||
guidance_rescale=state.guidance_rescale,
|
||||
)
|
||||
|
||||
return noise_pred
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
return noise_pred
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
if state.extra.get("skip_step", False):
|
||||
return state.latents
|
||||
|
||||
batch = state.extra["batch"]
|
||||
extra_step_kwargs = state.extra["extra_step_kwargs"]
|
||||
latents = self.stage.scheduler.step(
|
||||
noise_pred,
|
||||
t,
|
||||
state.latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if (state.extra["ti2v_mask"] is not None
|
||||
and batch.pil_image is not None):
|
||||
mask2 = state.extra["ti2v_mask"]
|
||||
z = state.extra["ti2v_z"]
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2) * z + mask2 * latents
|
||||
|
||||
if state.extra["trajectory_latents"] is not None:
|
||||
state.extra["trajectory_timesteps"].append(t)
|
||||
state.extra["trajectory_latents"].append(latents)
|
||||
|
||||
progress_bar = state.extra["progress_bar"]
|
||||
step_idx = state.extra["step_idx"]
|
||||
num_warmup_steps = state.extra["num_warmup_steps"]
|
||||
timesteps = state.timesteps
|
||||
if step_idx == len(timesteps) - 1 or (
|
||||
(step_idx + 1) > num_warmup_steps and
|
||||
(step_idx + 1) % self.stage.scheduler.order == 0
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
return latents
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
batch = state.extra["batch"]
|
||||
fastvideo_args = state.extra["fastvideo_args"]
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
|
||||
if state.extra["trajectory_latents"]:
|
||||
trajectory_tensor = torch.stack(state.extra["trajectory_latents"],
|
||||
dim=1)
|
||||
trajectory_timesteps_tensor = torch.stack(
|
||||
state.extra["trajectory_timesteps"], dim=0)
|
||||
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
|
||||
batch.latents = state.latents
|
||||
|
||||
if fastvideo_args.dit_layerwise_offload:
|
||||
mgr = getattr(self.stage.transformer, "_layerwise_offload_manager",
|
||||
None)
|
||||
if mgr is not None and getattr(mgr, "enabled", False):
|
||||
mgr.release_all()
|
||||
if self.stage.transformer_2 is not None:
|
||||
mgr2 = getattr(self.stage.transformer_2,
|
||||
"_layerwise_offload_manager", None)
|
||||
if mgr2 is not None and getattr(mgr2, "enabled", False):
|
||||
mgr2.release_all()
|
||||
|
||||
if (st_attn_available
|
||||
and self.stage.attn_backend == SlidingTileAttentionBackend
|
||||
and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING):
|
||||
self.stage.save_sta_search_results(batch)
|
||||
|
||||
pipeline = self.stage.pipeline() if self.stage.pipeline else None
|
||||
if torch.backends.mps.is_available():
|
||||
logger.info("Memory before deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory())
|
||||
del self.stage.transformer
|
||||
if pipeline is not None and "transformer" in pipeline.modules:
|
||||
del pipeline.modules["transformer"]
|
||||
fastvideo_args.model_loaded["transformer"] = False
|
||||
logger.info("Memory after deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory())
|
||||
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
|
||||
return batch
|
||||
@@ -1,106 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Strategy interfaces and shared types for unified denoising.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
@dataclass
|
||||
class StrategyState:
|
||||
latents: torch.Tensor
|
||||
timesteps: torch.Tensor
|
||||
num_inference_steps: int
|
||||
prompt_embeds: list[torch.Tensor]
|
||||
negative_prompt_embeds: list[torch.Tensor] | None
|
||||
prompt_attention_mask: list[torch.Tensor] | None
|
||||
negative_attention_mask: list[torch.Tensor] | None
|
||||
image_embeds: list[torch.Tensor]
|
||||
guidance_scale: float
|
||||
guidance_scale_2: float | None
|
||||
guidance_rescale: float
|
||||
do_cfg: bool
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelInputs:
|
||||
latent_model_input: torch.Tensor
|
||||
timestep: torch.Tensor
|
||||
prompt_embeds: torch.Tensor | list[torch.Tensor]
|
||||
prompt_attention_mask: torch.Tensor | list[torch.Tensor] | None
|
||||
extra_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BlockPlanItem:
|
||||
start_index: int
|
||||
num_frames: int
|
||||
use_kv_cache: bool
|
||||
model_selector: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class BlockPlan:
|
||||
items: list[BlockPlanItem] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BlockContext:
|
||||
kv_cache: list[dict] | None
|
||||
kv_cache_2: list[dict] | None
|
||||
crossattn_cache: list[dict] | None
|
||||
action_cache: dict[str, list[dict]] | None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class DenoisingStrategy(Protocol):
|
||||
|
||||
def prepare(self, batch: ForwardBatch, args) -> StrategyState:
|
||||
...
|
||||
|
||||
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
|
||||
step_idx: int) -> ModelInputs:
|
||||
...
|
||||
|
||||
def forward(self, state: StrategyState,
|
||||
model_inputs: ModelInputs) -> torch.Tensor:
|
||||
...
|
||||
|
||||
def cfg_combine(self, state: StrategyState,
|
||||
noise_pred: torch.Tensor) -> torch.Tensor:
|
||||
...
|
||||
|
||||
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
|
||||
t: torch.Tensor) -> torch.Tensor:
|
||||
...
|
||||
|
||||
def postprocess(self, state: StrategyState) -> ForwardBatch:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class BlockDenoisingStrategy(DenoisingStrategy, Protocol):
|
||||
|
||||
def block_plan(self, state: StrategyState) -> BlockPlan:
|
||||
...
|
||||
|
||||
def init_block_context(self, state: StrategyState,
|
||||
block_item: BlockPlanItem,
|
||||
block_idx: int) -> BlockContext:
|
||||
...
|
||||
|
||||
def process_block(self, state: StrategyState, block_ctx: BlockContext,
|
||||
block_item: BlockPlanItem) -> None:
|
||||
...
|
||||
|
||||
def update_context(self, state: StrategyState, block_ctx: BlockContext,
|
||||
block_item: BlockPlanItem) -> None:
|
||||
...
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat-specific denoising stage implementing CFG-zero optimized guidance.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatDenoisingStage(DenoisingStage):
|
||||
"""
|
||||
LongCat denoising stage with CFG-zero optimized guidance scale.
|
||||
|
||||
Implements:
|
||||
1. Optimized CFG scale from CFG-zero paper
|
||||
2. Negation of noise prediction before scheduler step (flow matching convention)
|
||||
3. Batched CFG computation (unlike standard FastVideo separate passes)
|
||||
"""
|
||||
|
||||
def optimized_scale(self, positive_flat, negative_flat) -> torch.Tensor:
|
||||
"""
|
||||
Calculate optimized scale from CFG-zero paper.
|
||||
|
||||
st_star = (v_cond^T * v_uncond) / ||v_uncond||^2
|
||||
|
||||
Args:
|
||||
positive_flat: Conditional prediction, flattened [B, -1]
|
||||
negative_flat: Unconditional prediction, flattened [B, -1]
|
||||
|
||||
Returns:
|
||||
st_star: Optimized scale [B, 1]
|
||||
"""
|
||||
# Calculate dot product
|
||||
dot_product = torch.sum(positive_flat * negative_flat,
|
||||
dim=1,
|
||||
keepdim=True)
|
||||
# Squared norm of uncondition
|
||||
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
|
||||
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
|
||||
st_star = dot_product / squared_norm
|
||||
return st_star
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run LongCat denoising loop with optimized CFG.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Get transformer dtype
|
||||
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
|
||||
|
||||
# Extract batch parameters
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0] # LongCat uses single encoder
|
||||
prompt_attention_mask = batch.prompt_attention_mask[
|
||||
0] if batch.prompt_attention_mask else None
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get negative prompts if doing CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
# Concatenate for batched processing
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="LongCat Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# Run transformer with context
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
)
|
||||
|
||||
# Apply CFG with optimized scale
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
# Calculate optimized scale (CFG-zero)
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
|
||||
# Reshape for broadcasting
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
# Apply optimized CFG formula
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# CRITICAL: Negate noise prediction for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# Compute previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat I2V Denoising Stage with conditioning support.
|
||||
|
||||
This stage implements Tier 3 I2V denoising:
|
||||
1. Per-frame timestep masking (timestep[:, :num_cond_latents] = 0)
|
||||
2. Passes num_cond_latents to transformer (for RoPE skipping)
|
||||
3. Selective denoising (only updates non-conditioned frames)
|
||||
4. CFG-zero optimized guidance
|
||||
"""
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatI2VDenoisingStage(LongCatDenoisingStage):
|
||||
"""
|
||||
LongCat denoising with I2V conditioning support.
|
||||
|
||||
Key modifications from base LongCat denoising:
|
||||
1. Sets timestep=0 for conditioning frames
|
||||
2. Passes num_cond_latents to transformer
|
||||
3. Only applies scheduler step to non-conditioned frames
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run denoising loop with I2V conditioning."""
|
||||
|
||||
# Load transformer if needed
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Setup
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get num_cond_latents from batch
|
||||
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
|
||||
|
||||
if num_cond_latents > 0:
|
||||
logger.info("I2V Denoising: num_cond_latents=%s, latent_shape=%s",
|
||||
num_cond_latents, latents.shape)
|
||||
|
||||
# Prepare negative prompts for CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="I2V Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
|
||||
# 1. Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# 2. Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# 3. CRITICAL: Expand timestep to temporal dimension
|
||||
# and set conditioning frames to timestep=0
|
||||
timestep = timestep.unsqueeze(-1).repeat(
|
||||
1, latent_model_input.shape[2])
|
||||
|
||||
# Mark conditioning frames as clean (timestep=0)
|
||||
if num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
# 4. Run transformer with num_cond_latents
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
num_cond_latents=num_cond_latents,
|
||||
)
|
||||
|
||||
# 5. Apply CFG with optimized scale
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
# CFG-zero optimized scale
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# 6. CRITICAL: Negate for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# 7. CRITICAL: Only update non-conditioned frames
|
||||
# The conditioning frames stay FIXED throughout denoising
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
# No conditioning, update all frames
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,217 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat VC Denoising Stage with KV cache support.
|
||||
|
||||
This stage extends the I2V denoising stage to support:
|
||||
1. KV cache for conditioning frames
|
||||
2. Video continuation with multiple conditioning frames
|
||||
"""
|
||||
|
||||
import time
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatVCDenoisingStage(LongCatDenoisingStage):
|
||||
"""
|
||||
LongCat denoising with Video Continuation and KV cache support.
|
||||
|
||||
Key differences from I2V denoising:
|
||||
- Supports KV cache (reuses cached K/V from conditioning frames)
|
||||
- Handles larger num_cond_latents
|
||||
- Concatenates conditioning latents back after denoising
|
||||
|
||||
When use_kv_cache=True:
|
||||
- batch.latents contains ONLY noise frames (cond removed by KV cache init)
|
||||
- batch.kv_cache_dict contains cached K/V
|
||||
- batch.cond_latents contains conditioning latents for post-concat
|
||||
|
||||
When use_kv_cache=False:
|
||||
- batch.latents contains ALL frames (cond + noise)
|
||||
- Timestep masking: timestep[:, :num_cond_latents] = 0
|
||||
- Selective denoising: only update noise frames
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run denoising loop with VC conditioning and optional KV cache."""
|
||||
|
||||
# Load transformer if needed
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Setup
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get VC-specific parameters
|
||||
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
|
||||
use_kv_cache = getattr(batch, 'use_kv_cache', False)
|
||||
kv_cache_dict = getattr(batch, 'kv_cache_dict', {})
|
||||
|
||||
logger.info(
|
||||
"VC Denoising: num_cond_latents=%d, use_kv_cache=%s, latent_shape=%s",
|
||||
num_cond_latents, use_kv_cache, latents.shape)
|
||||
|
||||
# Prepare negative prompts for CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
step_times = []
|
||||
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="VC Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
step_start = time.time()
|
||||
|
||||
# 1. Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# 2. Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# 3. Expand timestep to temporal dimension
|
||||
timestep = timestep.unsqueeze(-1).repeat(
|
||||
1, latent_model_input.shape[2])
|
||||
|
||||
# 4. Timestep masking (only when NOT using KV cache)
|
||||
if not use_kv_cache and num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
# 5. Prepare transformer kwargs
|
||||
# IMPORTANT: num_cond_latents is ALWAYS passed - needed for RoPE position offset
|
||||
transformer_kwargs = {
|
||||
'num_cond_latents': num_cond_latents,
|
||||
}
|
||||
if use_kv_cache:
|
||||
transformer_kwargs['kv_cache_dict'] = kv_cache_dict
|
||||
|
||||
# 6. Run transformer
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
**transformer_kwargs,
|
||||
)
|
||||
|
||||
# 7. Apply CFG with optimized scale (CFG-zero)
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# 8. Negate for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# 9. Scheduler step
|
||||
if use_kv_cache:
|
||||
# All latents are noise frames (conditioning is in cache)
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
# Only update noise frames (skip conditioning)
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
else:
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
step_time = time.time() - step_start
|
||||
step_times.append(step_time)
|
||||
|
||||
# Log timing for first few steps
|
||||
if i < 3:
|
||||
logger.info("Step %d: %.2fs", i, step_time)
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# 10. If using KV cache, concatenate conditioning latents back
|
||||
if use_kv_cache and hasattr(
|
||||
batch, 'cond_latents') and batch.cond_latents is not None:
|
||||
latents = torch.cat([batch.cond_latents, latents], dim=2)
|
||||
logger.info(
|
||||
"Concatenated conditioning latents back, final shape: %s",
|
||||
latents.shape)
|
||||
|
||||
# Log average timing
|
||||
avg_time = sum(step_times) / len(step_times)
|
||||
logger.info("Average step time: %.2fs (total: %.1fs)", avg_time,
|
||||
sum(step_times))
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -15,6 +15,14 @@ from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
@@ -115,22 +123,169 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
|
||||
|
||||
self._streaming_initialized: bool = False
|
||||
self._streaming_ctx: BlockProcessingContext | None = None
|
||||
self._streaming_engine = None
|
||||
self._streaming_state = None
|
||||
self._streaming_block_plan = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
|
||||
from fastvideo.pipelines.stages.denoising_matrixgame_strategy import (
|
||||
MatrixGameBlockStrategy)
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
engine = DenoisingEngine(MatrixGameBlockStrategy(self),
|
||||
hooks=self._build_engine_hooks())
|
||||
return engine.run(batch, fastvideo_args)
|
||||
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
|
||||
patch_size = self.transformer.patch_size
|
||||
patch_ratio = patch_size[-1] * patch_size[-2]
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
|
||||
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
|
||||
'boundary_ratio', None)
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
# directly set the kwarg.
|
||||
image_kwargs = {"encoder_hidden_states_image": image_embeds}
|
||||
pos_cond_kwargs: dict[str, Any] = {}
|
||||
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents
|
||||
b, c, t, h, w = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
kv_cache2 = None
|
||||
if boundary_timestep is not None:
|
||||
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
kv_cache_mouse = None
|
||||
kv_cache_keyboard = None
|
||||
if self.use_action_module:
|
||||
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
|
||||
batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
crossattn_cache = self._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=257, # 1 CLS + 256 patch tokens
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
if t % self.num_frame_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frame_per_block for causal denoising"
|
||||
)
|
||||
num_blocks = t // self.num_frame_per_block
|
||||
block_sizes = [self.num_frame_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
# NOTE: MatrixGame does NOT process the first frame separately.
|
||||
# The first frame information is already encoded in batch.image_latent (cond_concat)
|
||||
# and will be used by the model via channel concatenation: torch.cat([x, cond_concat], dim=1)
|
||||
|
||||
ctx = BlockProcessingContext(
|
||||
batch=batch,
|
||||
block_idx=0,
|
||||
start_index=0,
|
||||
kv_cache1=kv_cache1,
|
||||
kv_cache2=kv_cache2,
|
||||
kv_cache_mouse=kv_cache_mouse,
|
||||
kv_cache_keyboard=kv_cache_keyboard,
|
||||
crossattn_cache=crossattn_cache,
|
||||
timesteps=timesteps,
|
||||
block_sizes=block_sizes,
|
||||
noise_pool=None,
|
||||
fastvideo_args=fastvideo_args,
|
||||
target_dtype=target_dtype,
|
||||
autocast_enabled=autocast_enabled,
|
||||
boundary_timestep=boundary_timestep,
|
||||
high_noise_timesteps=high_noise_timesteps,
|
||||
context_noise=getattr(fastvideo_args.pipeline_config,
|
||||
"context_noise", 0),
|
||||
image_kwargs=image_kwargs,
|
||||
pos_cond_kwargs=pos_cond_kwargs,
|
||||
)
|
||||
|
||||
context_noise = getattr(fastvideo_args.pipeline_config, "context_noise",
|
||||
0)
|
||||
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
for block_idx, current_num_frames in enumerate(block_sizes):
|
||||
ctx.block_idx = block_idx
|
||||
ctx.start_index = start_index
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
action_kwargs = self._prepare_action_kwargs(
|
||||
batch, start_index, current_num_frames)
|
||||
|
||||
current_latents = self._process_single_block(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
timesteps=timesteps,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
progress_bar=progress_bar,
|
||||
)
|
||||
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = current_latents
|
||||
|
||||
# Update KV caches with clean context
|
||||
self._update_context_cache(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
context_noise=context_noise,
|
||||
)
|
||||
|
||||
start_index += current_num_frames
|
||||
|
||||
if boundary_timestep is not None:
|
||||
num_frames_to_remove = self.num_frame_per_block - 1
|
||||
if num_frames_to_remove > 0:
|
||||
latents = latents[:, :, :-num_frames_to_remove, :, :]
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
def _prepare_action_kwargs(self, batch: ForwardBatch, start_index: int,
|
||||
num_frames: int) -> dict[str, Any]:
|
||||
@@ -497,44 +652,123 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
|
||||
|
||||
def streaming_reset(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
|
||||
from fastvideo.pipelines.stages.denoising_matrixgame_strategy import (
|
||||
MatrixGameBlockStrategy)
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
strategy = MatrixGameBlockStrategy(self)
|
||||
engine = DenoisingEngine(strategy, hooks=self._build_engine_hooks())
|
||||
state = strategy.prepare(batch, fastvideo_args)
|
||||
for hook in engine.hooks:
|
||||
hook.on_init(engine, batch, fastvideo_args)
|
||||
for hook in engine.hooks:
|
||||
hook.pre_run(state)
|
||||
block_plan = strategy.block_plan(state)
|
||||
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
|
||||
patch_size = self.transformer.patch_size
|
||||
patch_ratio = patch_size[-1] * patch_size[-2]
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
|
||||
progress_bar = state.extra.get("progress_bar")
|
||||
if progress_bar is not None:
|
||||
progress_bar.close()
|
||||
state.extra["progress_bar"] = None
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
|
||||
ctx = state.extra["ctx"]
|
||||
latents = state.latents
|
||||
assert latents is not None, "latents must be provided"
|
||||
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
|
||||
'boundary_ratio', None)
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
# directly set the kwarg.
|
||||
image_kwargs = {"encoder_hidden_states_image": image_embeds}
|
||||
pos_cond_kwargs: dict[str, Any] = {}
|
||||
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents
|
||||
b, c, t, h, w = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
num_denoising_steps = len(state.timesteps)
|
||||
# Initialize caches
|
||||
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
kv_cache2 = None
|
||||
if boundary_timestep is not None:
|
||||
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
kv_cache_mouse = None
|
||||
kv_cache_keyboard = None
|
||||
if self.use_action_module:
|
||||
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
|
||||
batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
crossattn_cache = self._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=257, # 1 CLS + 256 patch tokens
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
# Calculate block sizes
|
||||
if t % self.num_frame_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frame_per_block for causal denoising"
|
||||
)
|
||||
num_blocks = t // self.num_frame_per_block
|
||||
block_sizes = [self.num_frame_per_block] * num_blocks
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
# Pre-allocate noise pool
|
||||
num_denoising_steps = len(timesteps)
|
||||
noise_shape = (b, self.num_frame_per_block, c, h, w)
|
||||
noise_pool = [
|
||||
torch.randn(
|
||||
noise_shape,
|
||||
dtype=ctx.target_dtype,
|
||||
dtype=target_dtype,
|
||||
device=latents.device,
|
||||
) for _ in range(max(num_denoising_steps - 1, 0))
|
||||
) for _ in range(num_denoising_steps - 1)
|
||||
]
|
||||
ctx.noise_pool = noise_pool
|
||||
|
||||
self._streaming_engine = engine
|
||||
self._streaming_state = state
|
||||
self._streaming_block_plan = block_plan
|
||||
self._streaming_ctx = ctx
|
||||
# Create and store context
|
||||
self._streaming_ctx = BlockProcessingContext(
|
||||
batch=batch,
|
||||
block_idx=0,
|
||||
start_index=0,
|
||||
kv_cache1=kv_cache1,
|
||||
kv_cache2=kv_cache2,
|
||||
kv_cache_mouse=kv_cache_mouse,
|
||||
kv_cache_keyboard=kv_cache_keyboard,
|
||||
crossattn_cache=crossattn_cache,
|
||||
timesteps=timesteps,
|
||||
block_sizes=block_sizes,
|
||||
noise_pool=noise_pool,
|
||||
fastvideo_args=fastvideo_args,
|
||||
target_dtype=target_dtype,
|
||||
autocast_enabled=autocast_enabled,
|
||||
boundary_timestep=boundary_timestep,
|
||||
high_noise_timesteps=high_noise_timesteps,
|
||||
context_noise=getattr(fastvideo_args.pipeline_config,
|
||||
"context_noise", 0),
|
||||
image_kwargs=image_kwargs,
|
||||
pos_cond_kwargs=pos_cond_kwargs,
|
||||
)
|
||||
|
||||
self._streaming_initialized = True
|
||||
return batch
|
||||
|
||||
@@ -542,25 +776,23 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
|
||||
self,
|
||||
keyboard_action: torch.Tensor | None = None,
|
||||
mouse_action: torch.Tensor | None = None) -> ForwardBatch:
|
||||
if (not self._streaming_initialized or self._streaming_ctx is None
|
||||
or self._streaming_engine is None
|
||||
or self._streaming_state is None
|
||||
or self._streaming_block_plan is None):
|
||||
if not self._streaming_initialized or self._streaming_ctx is None:
|
||||
raise RuntimeError(
|
||||
"Streaming not initialized! Call streaming_reset first.")
|
||||
|
||||
ctx = self._streaming_ctx
|
||||
block_plan = self._streaming_block_plan
|
||||
if ctx.block_idx >= len(block_plan.items):
|
||||
if ctx.block_idx >= len(ctx.block_sizes):
|
||||
return ctx.batch
|
||||
|
||||
batch = ctx.batch
|
||||
latents = batch.latents
|
||||
assert latents is not None, "latents must be set in batch"
|
||||
|
||||
block_item = block_plan.items[ctx.block_idx]
|
||||
start_index = block_item.start_index
|
||||
current_num_frames = block_item.num_frames
|
||||
current_num_frames = ctx.block_sizes[ctx.block_idx]
|
||||
start_index = ctx.start_index
|
||||
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
# Update batch with new actions for this block
|
||||
if keyboard_action is not None or mouse_action is not None:
|
||||
@@ -578,32 +810,58 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
|
||||
batch.mouse_cond[:, start_frame:start_frame +
|
||||
n] = mouse_action.to(batch.mouse_cond.device)
|
||||
|
||||
self._streaming_engine.run_blocks(
|
||||
self._streaming_state,
|
||||
block_plan=block_plan,
|
||||
start_block=ctx.block_idx,
|
||||
num_blocks=1,
|
||||
action_kwargs = self._prepare_action_kwargs(batch, start_index,
|
||||
current_num_frames)
|
||||
|
||||
# Create noise generator that uses pre-allocated noise pool
|
||||
def streaming_noise_generator(shape: tuple, dtype: torch.dtype,
|
||||
step_idx: int) -> torch.Tensor:
|
||||
if ctx.noise_pool is not None and step_idx < len(ctx.noise_pool):
|
||||
return ctx.noise_pool[step_idx][:, :shape[1], :, :, :].to(
|
||||
latents.device)
|
||||
else:
|
||||
# Fallback to dynamic allocation if pool not available
|
||||
return torch.randn(
|
||||
shape,
|
||||
dtype=dtype,
|
||||
generator=(batch.generator[0] if isinstance(
|
||||
batch.generator, list) else batch.generator)).to(
|
||||
latents.device)
|
||||
|
||||
current_latents = self._process_single_block(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
timesteps=ctx.timesteps,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
noise_generator=streaming_noise_generator,
|
||||
)
|
||||
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = current_latents
|
||||
|
||||
# Update KV caches with clean context
|
||||
self._update_context_cache(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
context_noise=ctx.context_noise,
|
||||
)
|
||||
|
||||
# Advance streaming state
|
||||
ctx.start_index = start_index + current_num_frames
|
||||
ctx.start_index += current_num_frames
|
||||
ctx.block_idx += 1
|
||||
|
||||
return batch
|
||||
|
||||
def streaming_clear(self) -> None:
|
||||
if (self._streaming_engine is not None
|
||||
and self._streaming_state is not None):
|
||||
batch = (self._streaming_ctx.batch if self._streaming_ctx
|
||||
is not None else self._streaming_state.extra.get("batch"))
|
||||
if batch is not None:
|
||||
for hook in self._streaming_engine.hooks:
|
||||
hook.post_run(self._streaming_state, batch)
|
||||
self._streaming_initialized = False
|
||||
self._streaming_ctx = None
|
||||
self._streaming_engine = None
|
||||
self._streaming_state = None
|
||||
self._streaming_block_plan = None
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -21,8 +21,3 @@ def distributed_setup():
|
||||
yield
|
||||
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
config.addinivalue_line("markers",
|
||||
"gpu: requires CUDA with BF16 support")
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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():
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user