Compare commits

..
Author SHA1 Message Date
SolitaryThinker 14e334d0b4 fix 2026-01-05 08:30:16 +00:00
178 changed files with 839 additions and 19537 deletions
+1 -1
View File
@@ -61,7 +61,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- TEST_TYPE=ssim
+1 -16
View File
@@ -156,23 +156,8 @@ 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 \
--exclude libtorch_cuda.so \
--exclude libtorch_cpu.so \
--exclude libtorch.so \
--exclude libc10.so \
--exclude libc10_cuda.so \
--exclude libtorch_python.so
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist
# Move fixed wheels back to dist for upload consistency
rm dist/*.whl
mv fixed_dist/*.whl dist/
+1 -1
View File
@@ -68,7 +68,7 @@ repos:
entry: bash
args:
- -c
- 'git ls-files | grep -v "^\"*fastvideo/tests/ssim/" | grep -v "^\"*fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
- 'git ls-files | grep -v "^fastvideo/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
language: system
always_run: true
pass_filenames: false
+1 -1
View File
@@ -3,7 +3,7 @@
</div>
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/c7g1qdD" target="_blank"> <b> WeChat </b> </a> |
</p>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 490 KiB

Binary file not shown.
+1 -1
View File
@@ -41,7 +41,7 @@ Clone the repository and build the kernel:
```bash
# Clone recursively to get ThunderKittens submodule
git clone https://github.com/hao-ai-lab/FastVideo.git
git clone --recursive https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo/fastvideo-kernel
# Build and install
+1 -3
View File
@@ -13,13 +13,11 @@ 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,
block_sizes,
block_sizes,
variable_block_sizes=block_sizes,
topk=32
)
```
+6 -10
View File
@@ -11,13 +11,15 @@ FastVideo supports the following hardware platforms:
### Using pip
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
pip install fastvideo
```
### Using conda
```bash
conda install -c conda-forge fastvideo
```
### From source
```bash
@@ -26,12 +28,6 @@ cd FastVideo
pip install -e .
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
## Hardware Requirements
- **NVIDIA GPUs**: CUDA 11.8+ with compute capability 7.0+
-6
View File
@@ -15,12 +15,6 @@ conda activate fastvideo
pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
## Basic Usage
### Text-to-Video Generation
-2
View File
@@ -106,8 +106,6 @@ If you encounter CUDA out of memory errors:
- Enable memory optimization with `enable_model_cpu_offload`
- Try a smaller model or use distilled versions
- Use `num_gpus` > 1 if multiple GPUs are available
- Try enabling FSDP inference with `use_fsdp_inference=True` (may slow down generation)
- Try enabling DiT layerwise offload with `dit_layerwise_offload=True` (now only a few models support this, but may introduce less overhead than FSDP)
### Slow Generation
-71
View File
@@ -1,71 +0,0 @@
# LoRA Extraction and Merging
Tools for extracting and merging LoRA adapters for FastVideo models.
## Extract LoRA Adapter
```bash
python scripts/lora_extraction/extract_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--out adapter_r32.safetensors \
--rank 32
```
**Options:**
- `--base`: Base model (HuggingFace ID or local path)
- `--ft`: Fine-tuned model (HuggingFace ID or local path)
- `--out`: Output adapter file
- `--rank`: LoRA rank (16, 32, 64, 128)
- `--full-rank`: Extract full-rank adapter (optional)
## Merge Adapter
```bash
python scripts/lora_extraction/merge_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--adapter adapter_r32.safetensors \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--output merged_model
```
**Options:**
- `--base`: Base model (HuggingFace ID or local path)
- `--adapter`: LoRA adapter file (.safetensors)
- `--ft`: Fine-tuned model (for configuration)
- `--output`: Output directory
## Validate Quality (Optional)
```bash
python scripts/lora_extraction/lora_inference_comparison.py \
--base merged_model \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--adapter NONE \
--output-dir results \
--prompt "A cat sitting on a windowsill" \
--seed 42 \
--height 480 \
--width 480 \
--num-frames 49 \
--num-inference-steps 32 \
--compute-ssim \
--compute-lpips
```
**Options:**
- `--base`: Merged model or base model path
- `--ft`: Fine-tuned model (reference)
- `--adapter`: Path to adapter or NONE
- `--output-dir`: Output directory
- `--prompt`: Text prompt (default: "A cat sitting on a windowsill")
- `--seed`: Random seed (default: 42)
- `--height`: Video height (default: 480)
- `--width`: Video width (default: 832)
- `--num-frames`: Number of frames (default: 49)
- `--num-inference-steps`: Inference steps (default: 32)
- `--compute-ssim`: Compute SSIM metric
- `--compute-lpips`: Compute LPIPS metric
+1 -1
View File
@@ -12,7 +12,7 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
@@ -1,42 +0,0 @@
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()
+2 -1
View File
@@ -14,13 +14,14 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
# Adjust these offload parameters if you have < 32GB of VRAM
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
init_weights_from_safetensors="/mnt/weka/home/hao.zhang/wl/release/dmd_distill_1.3_4n_syn/checkpoint-900_weight_only/generator_inference_transformer"
)
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
+1 -1
View File
@@ -12,7 +12,7 @@ def main():
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
@@ -1,210 +0,0 @@
"""
LongCat Image-to-Video (I2V) Example Script
This script demonstrates LongCat I2V inference using the FastVideo Python API.
LongCat I2V takes an input image and generates a video from it.
It runs both basic generation (50 steps) and distill+refine generation
(16 steps distill + 50 steps refinement to 720p with BSA).
Usage:
python examples/inference/basic/basic_longcat_i2v.py
Note:
Refinement uses 768x768 dimensions where latent (48x48) is divisible by 8,
compatible with BSA chunks [4, 4, 8].
"""
import glob
import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = (
"A woman sits at a wooden table by the window in a cozy café. She reaches out "
"with her right hand, picks up the white coffee cup from the saucer, and gently "
"brings it to her lips to take a sip. After drinking, she places the cup back on "
"the table and looks out the window, enjoying the peaceful atmosphere."
)
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
# Input image path
IMAGE_PATH = "assets/girl.png"
SEED = 42
def basic_generation():
"""
Run basic LongCat I2V generation (50 steps at 480p).
This uses the full 50-step denoising process for highest quality.
"""
print("=" * 60)
print("LongCat I2V: Basic Generation (50 steps, 480p)")
print("=" * 60)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-I2V-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_i2v_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
image_path=IMAGE_PATH,
output_path=output_path,
save_video=True,
height=480,
width=480, # Square
num_frames=93,
num_inference_steps=50,
fps=15,
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
def distill_refine_generation():
"""
Run LongCat I2V with distill+refine pipeline (16 steps + refinement to 768p).
This uses the distilled LoRA for fast 480p generation (16 steps),
then refines to 768p using the refinement LoRA with BSA enabled.
"""
print("\n" + "=" * 60)
print("LongCat I2V: Distill + Refine Pipeline")
print("=" * 60)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-I2V-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_i2v_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
image_path=IMAGE_PATH,
output_path=distill_output_path,
save_video=True,
height=480,
width=480, # Square
num_frames=93,
num_inference_steps=16,
fps=15,
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 768p)
print("\n[Stage 2] Refinement (480p -> 768p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
raise FileNotFoundError(f"No video file found in {distill_output_path}")
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
# Note: Refinement uses the T2V model (not I2V) since it's upscaling the generated video
# For BSA [4, 4, 8]: latent must be divisible by 8
# 768x768: latent 48x48, 48%8=0 ✓
refine_generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=True,
bsa_sparsity=0.875,
bsa_chunk_q=[4, 4, 4],
bsa_chunk_k=[4, 4, 4],
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_i2v_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
output_path=refine_output_path,
save_video=True,
refine_from=distill_video_path,
t_thresh=0.5,
spatial_refine_only=False,
num_cond_frames=0,
height=720,
width=720,
num_inference_steps=50,
fps=30,
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
def main():
"""Run both basic and distill+refine generation pipelines."""
print("\n" + "=" * 60)
print("LongCat Image-to-Video Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
if __name__ == "__main__":
main()
@@ -1,198 +0,0 @@
"""
LongCat Text-to-Video (T2V) Example Script
This script demonstrates LongCat T2V inference using the FastVideo Python API.
It runs both basic generation (50 steps) and distill+refine generation
(16 steps distill + 50 steps refinement to 720p).
Usage:
python examples/inference/basic/basic_longcat_t2v.py
"""
import glob
import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = (
"In a realistic photography style, a white boy around seven or eight years old "
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
"features a green lawn and several tall trees, creating a warm and loving scene."
)
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
SEED = 42
def basic_generation():
"""
Run basic LongCat T2V generation (50 steps at 480p).
This uses the full 50-step denoising process for highest quality.
"""
print("=" * 60)
print("LongCat T2V: Basic Generation (50 steps, 480p)")
print("=" * 60)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_t2v_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
output_path=output_path,
save_video=True,
height=480,
width=832,
num_frames=93,
num_inference_steps=50,
fps=15,
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
def distill_refine_generation():
"""
Run LongCat T2V with distill+refine pipeline (16 steps + refinement to 720p).
This uses the distilled LoRA for fast 480p generation (16 steps),
then refines to 720p using the refinement LoRA with BSA enabled.
"""
print("\n" + "=" * 60)
print("LongCat T2V: Distill + Refine Pipeline")
print("=" * 60)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_t2v_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
output_path=distill_output_path,
save_video=True,
height=480,
width=832,
num_frames=93,
num_inference_steps=16,
fps=15,
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 720p)
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
raise FileNotFoundError(f"No video file found in {distill_output_path}")
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
refine_generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=True,
bsa_sparsity=0.875,
bsa_chunk_q=[4, 4, 8],
bsa_chunk_k=[4, 4, 8],
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_t2v_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
output_path=refine_output_path,
save_video=True,
refine_from=distill_video_path,
t_thresh=0.5,
spatial_refine_only=False,
num_cond_frames=0,
height=720,
width=1280,
num_inference_steps=50,
fps=30,
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
def main():
"""Run both basic and distill+refine generation pipelines."""
print("\n" + "=" * 60)
print("LongCat Text-to-Video Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
if __name__ == "__main__":
main()
@@ -1,228 +0,0 @@
"""
LongCat Video Continuation (VC) Example Script
This script demonstrates LongCat VC inference using the FastVideo Python API.
LongCat VC takes an input video and generates a continuation of it.
It runs both basic generation (50 steps) and distill+refine generation
(16 steps distill + 50 steps refinement to 720p).
Usage:
python examples/inference/basic/basic_longcat_vc.py
Prerequisites:
- Ensure the input video exists at assets/motorcycle.mp4
(or provide your own video)
"""
import glob
import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = (
"A person rides a motorcycle along a long, straight road that stretches between "
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
"the motorcycle centered between the guardrails, while the scenery passes by on "
"both sides. The video captures the journey from the rider's perspective, emphasizing "
"the sense of motion and adventure."
)
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
# Input video path
VIDEO_PATH = "assets/motorcycle.mp4"
# Number of conditioning frames from the input video
NUM_COND_FRAMES = 13
SEED = 42
def basic_generation():
"""
Run basic LongCat VC generation (50 steps at 480p).
This uses the full 50-step denoising process for highest quality.
"""
print("=" * 60)
print("LongCat VC: Basic Generation (50 steps, 480p)")
print("=" * 60)
# Check if video exists
if not os.path.exists(VIDEO_PATH):
raise FileNotFoundError(
f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path."
)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-VC-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_vc_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
video_path=VIDEO_PATH,
num_cond_frames=NUM_COND_FRAMES,
output_path=output_path,
save_video=True,
height=480,
width=832,
num_frames=93,
num_inference_steps=50,
fps=15,
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
def distill_refine_generation():
"""
Run LongCat VC with distill+refine pipeline (16 steps + refinement to 720p).
This uses the distilled LoRA for fast 480p generation (16 steps),
then refines to 720p using the refinement LoRA with BSA enabled.
"""
print("\n" + "=" * 60)
print("LongCat VC: Distill + Refine Pipeline")
print("=" * 60)
# Check if video exists
if not os.path.exists(VIDEO_PATH):
raise FileNotFoundError(
f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path."
)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-VC-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_vc_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
video_path=VIDEO_PATH,
num_cond_frames=NUM_COND_FRAMES,
output_path=distill_output_path,
save_video=True,
height=480,
width=832,
num_frames=93,
num_inference_steps=16,
fps=15,
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 720p)
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
raise FileNotFoundError(f"No video file found in {distill_output_path}")
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
# Note: Refinement uses the T2V model (not VC) since it's upscaling the generated video
refine_generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=True,
bsa_sparsity=0.875,
bsa_chunk_q=[4, 4, 8],
bsa_chunk_k=[4, 4, 8],
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_vc_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
output_path=refine_output_path,
save_video=True,
refine_from=distill_video_path,
t_thresh=0.5,
spatial_refine_only=False,
num_cond_frames=0, # For refinement, no conditioning frames
height=720,
width=1280,
num_inference_steps=50,
fps=30,
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
def main():
"""Run both basic and distill+refine generation pipelines."""
print("\n" + "=" * 60)
print("LongCat Video Continuation Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
if __name__ == "__main__":
main()
-34
View File
@@ -1,34 +0,0 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
+1 -1
View File
@@ -43,7 +43,7 @@ def main():
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
@@ -46,7 +46,7 @@ async def main():
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
@@ -13,7 +13,7 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
text_encoder_cpu_offload=False,
dit_cpu_offload=False,
)
@@ -14,7 +14,7 @@ def main():
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
dit_precision="fp32",
vae_cpu_offload=False,
@@ -14,7 +14,7 @@ def main():
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
@@ -15,14 +15,12 @@ def main() -> None:
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
# set to false if using RTX 4090
# pin_cpu_memory=False,
# TurboDiffusion uses a custom pipeline with RCM scheduler
override_pipeline_cls_name="TurboDiffusionPipeline",
)
# Generate videos with the same simple API, regardless of GPU count
# TurboDiffusion defaults: guidance_scale=1.0 and num_inference_steps=4 (from config)
# TurboDiffusion uses guidance_scale=1.0 (no CFG) and only 4 steps
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
@@ -32,7 +30,9 @@ def main() -> None:
prompt,
output_path=OUTPUT_PATH,
save_video=True,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
# Generate another video with a different prompt, without reloading the model!
@@ -47,7 +47,9 @@ def main() -> None:
prompt2,
output_path=OUTPUT_PATH,
save_video=True,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
@@ -15,6 +15,8 @@ def main() -> None:
"loayrashid/TurboWan2.1-T2V-14B-Diffusers",
# 14B model needs more GPUs
num_gpus=2,
# TurboDiffusion uses a custom pipeline with RCM scheduler
override_pipeline_cls_name="TurboDiffusionPipeline",
)
prompt = (
@@ -26,7 +28,9 @@ def main() -> None:
prompt,
output_path=OUTPUT_PATH,
save_video=True,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
# Generate another video with a different prompt, without reloading the model!
@@ -41,7 +45,9 @@ def main() -> None:
prompt2,
output_path=OUTPUT_PATH,
save_video=True,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
@@ -1,37 +0,0 @@
import os
# Set SLA attention backend BEFORE fastvideo imports
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
from fastvideo import VideoGenerator
# Use local model path
MODEL_PATH = "loayrashid/TurboWan2.2-I2V-A14B-Diffusers"
OUTPUT_PATH = "video_samples_turbodiffusion_i2v"
def main() -> None:
# TurboDiffusion I2V: 1-4 step image-to-video generation
generator = VideoGenerator.from_pretrained(
MODEL_PATH,
num_gpus=2,
)
# Example prompt and image for I2V
prompt = ("Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside.")
# Use an example image path
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
video = generator.generate_video(
prompt,
image_path=image_path,
output_path=OUTPUT_PATH,
save_video=True,
seed=42,
)
if __name__ == "__main__":
main()
+1 -1
View File
@@ -12,7 +12,7 @@ def main():
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
+1 -1
View File
@@ -14,7 +14,7 @@ def main():
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
+1 -1
View File
@@ -12,7 +12,7 @@ def main():
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
@@ -11,7 +11,7 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
-26
View File
@@ -10,8 +10,6 @@ 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
@@ -155,30 +153,6 @@ if(BUILD_CXX_KERNELS)
$<$<COMPILE_LANGUAGE:CUDA>:${CUDA_FLAGS}>
)
# Link against Torch libraries to avoid undefined symbols at import time
# (e.g., torch::autograd vtables) when loading the extension module.
target_link_libraries(fastvideo_kernel_ops PRIVATE ${TORCH_LIBRARIES})
# Also link against libtorch_python to satisfy Python-binding symbols
# (e.g., torch::PyWarningHandler) required by torch/extension.h.
execute_process(
COMMAND "${Python_EXECUTABLE}" -c "import torch; from pathlib import Path; p=Path(torch.__file__).parent/'lib'; m=sorted(p.glob('libtorch_python*')); print(str(m[0]) if m else '')"
OUTPUT_VARIABLE TORCH_PYTHON_LIBRARY_PATH
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(TORCH_PYTHON_LIBRARY_PATH)
message(STATUS "TORCH_PYTHON_LIBRARY_PATH: ${TORCH_PYTHON_LIBRARY_PATH}")
target_link_libraries(fastvideo_kernel_ops PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
else()
message(WARNING "Could not locate libtorch_python; fastvideo_kernel_ops may fail to import.")
endif()
# Link CUDA runtime + driver explicitly (fixes missing symbols like cuGetErrorString at import time)
if(NOT GPU_BACKEND STREQUAL "ROCM")
target_link_libraries(fastvideo_kernel_ops PRIVATE CUDA::cudart CUDA::cuda_driver)
endif()
# We install it to fastvideo_kernel/_C so we can load it to register the ops
install(TARGETS fastvideo_kernel_ops LIBRARY DESTINATION fastvideo_kernel/_C)
endif()
+1 -1
View File
@@ -34,7 +34,7 @@ from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_att
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
# Example: Video Sparse Attention (with Triton fallback)
out = video_sparse_attn(q, k, v, block_sizes, block_sizes, topk=5)
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
# Example: VMoBA
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
@@ -639,8 +639,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
// store kq and vq
// ! the following two line seems unnecessary.
// tma::store_async_wait(); // ensure qg is finished
// ensuring all writes are finished
__syncthreads();
warpgroup::store(kg_smem[0], kg_reg);
@@ -661,6 +660,145 @@ 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>
@@ -672,32 +810,23 @@ block_sparse_attention_forward(
torch::Tensor v,
torch::Tensor q2k_block_sparse_index,
torch::Tensor q2k_block_sparse_num,
torch::Tensor kv_block_size
torch::Tensor 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 q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto seq_len = q.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 = q2k_block_sparse_index.size(2);
auto num_kv_blocks = kv_block_size.size(0);
auto num_q_blocks = 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 * 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.");
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");
// 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");
@@ -705,9 +834,11 @@ 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(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(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(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");
@@ -733,12 +864,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>(q_seq_len),
static_cast<const uint>(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>(q_seq_len),
static_cast<const uint>(seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
@@ -749,110 +880,32 @@ 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) {
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()),
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()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
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
reinterpret_cast<int32_t*>(block_size.data_ptr()),
stream
);
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()),
} 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()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
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
reinterpret_cast<int32_t*>(block_size.data_ptr()),
stream
);
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
CHECK_CUDA_ERROR(cudaGetLastError());
// cudaStreamSynchronize(stream);
} else {
TORCH_CHECK(false, "Unsupported head_dim: ", head_dim, ". Only 64 and 128 are supported.");
}
return {o, l_vec};
@@ -868,7 +921,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 kv_block_size)
torch::Tensor block_size)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -877,23 +930,11 @@ 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 q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto seq_len = q.size(2);
auto head_dim = q.size(3);
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
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))");
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match 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");
@@ -904,18 +945,23 @@ 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(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(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(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);
@@ -942,20 +988,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>(q_seq_len),
static_cast<const uint>(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>(kv_seq_len),
static_cast<const uint>(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>(kv_seq_len),
static_cast<const uint>(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>(q_seq_len),
static_cast<const uint>(seq_len),
static_cast<const uint>(1)}, l_vec.options());
float* qg_ptr = qg.data_ptr<float>();
@@ -984,7 +1030,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(q_seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
dim3 grid_bwd(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>;
@@ -997,9 +1043,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>(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)};
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)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1036,15 +1082,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>(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_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_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1055,14 +1101,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(kv_seq_len), // N is not used in the kernel
static_cast<int>(seq_len),
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*>(kv_block_size.data_ptr())};
reinterpret_cast<int32_t*>(block_size.data_ptr())};
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
@@ -1101,9 +1147,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>(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)};
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)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1140,15 +1186,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>(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_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_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1159,14 +1205,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(kv_seq_len), // N is not used in the kernel
static_cast<int>(seq_len),
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*>(kv_block_size.data_ptr())};
reinterpret_cast<int32_t*>(block_size.data_ptr())};
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
@@ -1187,4 +1233,4 @@ block_sparse_attention_backward(torch::Tensor q,
return {qg, kg, vg};
//cudadevicesynchronize();
}
}
@@ -4,7 +4,6 @@
#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"
@@ -15,6 +14,10 @@ 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();
@@ -23,70 +26,31 @@ auto layer_norm(
Output.emplace(
torch::empty(
{m, n},
torch::TensorOptions().device(input_device).dtype(Input.scalar_type())
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
)
);
}
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();
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);
});
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()
);
});
});
} 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,22 +68,9 @@ 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) {
int idx = thr_n_offset + i;
if (idx < params.n) {
x[i] -= u;
} else {
x[i] = 0.f;
}
}
for (int i = 0; i < NumElementPerThread; ++i)
x[i] -= u;
__syncthreads();
// var reduction
@@ -4,7 +4,6 @@
#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"
@@ -17,6 +16,10 @@ 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();
@@ -25,51 +28,27 @@ auto rms_norm(
Output.emplace(
torch::empty(
{m, n},
torch::TensorOptions().device(input_device).dtype(Input.scalar_type())
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
)
);
}
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();
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());
}
CONFIG_SWITCH(n, [&]{
rmsnorm<
ElementIn, ElementOut, ElementWeight,
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA
> (
Iptr, Wptr,
Optr,
eps, m, n,
at::cuda::getCurrentCUDAStream().stream()
);
});
return Output;
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.4"
version = "0.2.1"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -1,298 +0,0 @@
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)
+15 -58
View File
@@ -1,6 +1,5 @@
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
@@ -46,22 +45,14 @@ 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_h, k_h,
v_h, o_h,
q[:, head_idx:head_idx + 1], k[:, head_idx:head_idx + 1],
v[:, head_idx:head_idx + 1], output[:, head_idx:head_idx + 1],
t, h, w, text_length, False, has_text, flag
)
output[:, head_idx:head_idx + 1] = o_h
if has_text:
sta_fwd(q.contiguous(), k.contiguous(), v.contiguous(), output, 3, 3, 3, text_length, True, True, flag)
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
return output[:, :, :seq_length]
@@ -71,7 +62,6 @@ 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,
@@ -80,42 +70,14 @@ 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, 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()}"
)
batch, heads, seq_len, dim = q.shape
# Compression branch
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.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_c.float().sum(dim=3) / q_variable_block_sizes.view(1, 1, -1, 1)).to(
q_c = (q_c.float().sum(dim=3) / 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)
@@ -126,9 +88,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, q_num_blocks, 1, dim)
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, q_seq_len, dim)
1).view(batch, heads, seq_len, dim)
# Sparse branch
topk_idx = torch.topk(scores, topk, dim=-1).indices
@@ -138,17 +100,12 @@ def video_sparse_attn(
idx, num = map_to_index(mask)
if block_sparse_fwd is not None:
# 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]
out_s = block_sparse_fwd(
q, k, v, idx, num, variable_block_sizes.int()
)[0] # block_sparse_fwd returns vector<Tensor>
else:
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)
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.4"
__version__ = "0.2.1"
+31 -6
View File
@@ -42,13 +42,37 @@ 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 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
)
# 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
output = output_padded[:, :, q_non_pad_index, :]
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.backward(dO)
return output, Q.grad, K.grad, V.grad
@@ -240,6 +264,7 @@ 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
+20 -7
View File
@@ -3,7 +3,6 @@ import sys
from typing import Tuple
import torch
import pytest
from .utils import (
generate_block_sparse_mask_for_function,
@@ -58,14 +57,23 @@ 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 autograd-enabled wrapper (internally dispatches SM90 C++ vs Triton)
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
# Use raw kernel or triton
try:
out_padded, _aux = block_sparse_attn(
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
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
)
except RuntimeError as e:
pytest.skip(str(e))
# Remove padding on the query side
out = out_padded[:, :, q_non_pad_index, :]
@@ -148,6 +156,11 @@ 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,9 +276,8 @@ class VideoSparseAttentionImpl(AttentionImpl):
query,
key,
value,
attn_metadata.variable_block_sizes,
attn_metadata.variable_block_sizes,
cur_topk,
variable_block_sizes=attn_metadata.variable_block_sizes,
topk=cur_topk,
block_size=VSA_TILE_SIZE,
compress_attn_weight=gate_compress).transpose(1, 2)
+1 -12
View File
@@ -2,16 +2,5 @@ from fastvideo.configs.models.base import ModelConfig
from fastvideo.configs.models.dits.base import DiTConfig
from fastvideo.configs.models.encoders.base import EncoderConfig
from fastvideo.configs.models.vaes.base import VAEConfig
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig)
__all__ = [
"ModelConfig",
"VAEConfig",
"DiTConfig",
"EncoderConfig",
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
@@ -1,13 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.configs.models.audio.ltx2_audio_vae import (
LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
)
__all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -1,31 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 audio VAE and vocoder configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig, ModelConfig
@dataclass
class LTX2AudioArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=list)
@dataclass
class LTX2AudioEncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioEncoder"]))
@dataclass
class LTX2AudioDecoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioDecoder"]))
@dataclass
class LTX2VocoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2Vocoder"]))
+1 -2
View File
@@ -3,12 +3,11 @@ from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig"
"LongCatVideoConfig"
]
+1
View File
@@ -26,6 +26,7 @@ class LongCatVideoArchConfig(DiTArchConfig):
default_factory=lambda: [is_longcat_blocks])
# Parameter name mapping for weight conversion
# Maps original LongCat third_party names -> native FastVideo names
param_names_mapping: dict = field(
default_factory=lambda: {
# Embedders
-84
View File
@@ -1,84 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 Transformer configuration for native FastVideo integration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_ltx2_blocks(name: str, _module) -> bool:
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
@dataclass
class LTX2VideoArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 video transformer."""
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_ltx2_blocks])
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo)
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.diffusion_model\.(.*)$": r"model.\1",
r"^diffusion_model\.(.*)$": r"model.\1",
r"^model\.(.*)$": r"model.\1",
r"^(.*)$": r"model.\1",
})
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Core transformer settings (defaults from LTX-2 metadata)
num_attention_heads: int = 32
attention_head_dim: int = 128
num_layers: int = 48
cross_attention_dim: int = 4096
caption_channels: int = 3840
norm_eps: float = 1e-6
attention_type: str = "default"
rope_type: str = "split"
double_precision_rope: bool = True
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20, 2048, 2048])
timestep_scale_multiplier: int = 1000
use_middle_indices_grid: bool = True
# Patchification (video-only path)
patch_size: tuple[int, int, int] = (1, 1, 1)
num_channels_latents: int = 128
in_channels: int | None = None
out_channels: int | None = None
# Audio defaults (reserved for joint AV ports)
audio_num_attention_heads: int = 32
audio_attention_head_dim: int = 64
audio_in_channels: int = 128
audio_out_channels: int = 128
audio_cross_attention_dim: int = 2048
audio_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20])
av_ca_timestep_scale_multiplier: int = 1
def __post_init__(self):
super().__post_init__()
patch_volume = self.patch_size[0] * self.patch_size[
1] * self.patch_size[2]
if self.in_channels is None:
self.in_channels = self.num_channels_latents * patch_volume
if self.out_channels is None:
self.out_channels = self.in_channels
@dataclass
class LTX2VideoConfig(DiTConfig):
"""Main configuration for LTX-2 transformer."""
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
prefix: str = "ltx2"
@@ -7,12 +7,10 @@ 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
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig"
"Qwen2_5_VLConfig"
]
@@ -1,48 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class LTX2GemmaArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(
default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
hidden_size: int = 3840
num_hidden_layers: int = 48
num_attention_heads: int = 30
text_len: int = 1024
pad_token_id: int = 0
eos_token_id: int = 2
gemma_model_path: str = ""
gemma_dtype: str = "bfloat16"
padding_side: str = "left"
feature_extractor_in_features: int = 3840 * 49
feature_extractor_out_features: int = 3840
connector_num_attention_heads: int = 30
connector_attention_head_dim: int = 128
connector_num_layers: int = 2
connector_positional_embedding_theta: float = 10000.0
connector_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [4096])
connector_rope_type: str = "split"
connector_double_precision_rope: bool = False
connector_num_learnable_registers: int | None = 128
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LTX2GemmaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=LTX2GemmaArchConfig)
prefix: str = "ltx2_gemma"
@@ -1,72 +0,0 @@
# 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,8 +1,6 @@
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.ltx2vae import LTX2VAEConfig
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
@@ -11,7 +9,5 @@ __all__ = [
"WanVAEConfig",
"StepVideoVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
]
@@ -1,223 +0,0 @@
"""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
-45
View File
@@ -1,45 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 VAE configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class LTX2VAEArchConfig(VAEArchConfig):
# Mirrors LTX-2 safetensors metadata config under "vae"
_class_name: str = "CausalVideoAutoencoder"
dims: int = 3
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 128
encoder_blocks: list = field(default_factory=list)
decoder_blocks: list = field(default_factory=list)
patch_size: int = 4
norm_layer: str = "pixel_norm"
latent_log_var: str = "uniform"
encoder_spatial_padding_mode: str = "zeros"
decoder_spatial_padding_mode: str = "reflect"
causal_decoder: bool = False
timestep_conditioning: bool = True
use_quant_conv: bool = False
scaling_factor: float = 1.0
normalize_latent_channels: bool = False
# Match FastVideo naming for compression ratios (LTX-2 default)
temporal_compression_ratio: int = 8
spatial_compression_ratio: int = 32
@dataclass
class LTX2VAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
ltx2_spatial_tile_size_in_pixels: int = 512
ltx2_spatial_tile_overlap_in_pixels: int = 64
ltx2_temporal_tile_size_in_frames: int = 64
ltx2_temporal_tile_overlap_in_frames: int = 24
+1 -4
View File
@@ -1,10 +1,8 @@
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.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -17,6 +15,5 @@ __all__ = [
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"get_pipeline_config_cls_from_name"
"CosmosConfig", "get_pipeline_config_cls_from_name"
]
-89
View File
@@ -1,89 +0,0 @@
# 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
+11 -3
View File
@@ -17,7 +17,11 @@ from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class LongCatDiTArchConfig(DiTArchConfig):
"""Extended DiTArchConfig with LongCat-specific fields."""
"""Extended DiTArchConfig with LongCat-specific fields.
NOTE: This is for Phase 1 wrapper compatibility. For native model (Phase 2),
use LongCatVideoConfig from fastvideo.configs.models.dits.longcat instead.
"""
# LongCat-specific architecture parameters
adaln_tembed_dim: int = 512
caption_channels: int = 4096
@@ -84,16 +88,20 @@ def umt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
@dataclass
class LongCatT2V480PConfig(PipelineConfig):
"""Configuration for LongCat pipeline (480p).
"""Configuration for LongCat pipeline (480p) aligned to LongCat-Video modules.
Components expected by loaders:
- tokenizer: AutoTokenizer
- text_encoder: UMT5EncoderModel
- transformer: LongCatTransformer3DModel
- transformer: LongCatVideoTransformer3DModel (Phase 1 wrapper)
OR LongCatTransformer3DModel (Phase 2 native)
- vae: AutoencoderKLWan (Wan VAE, 4x8 compression)
- scheduler: FlowMatchEulerDiscreteScheduler
"""
# DiT config with LongCat-specific arch_config
# NOTE: For Phase 1 wrapper, uses LongCatDiTArchConfig
# For Phase 2 native model, can use LongCatVideoConfig directly
dit_config: DiTConfig = field(
default_factory=lambda: DiTConfig(arch_config=LongCatDiTArchConfig()))
-50
View File
@@ -1,50 +0,0 @@
# 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, ModelConfig,
LTX2AudioDecoderConfig, LTX2VocoderConfig,
VAEConfig)
from fastvideo.configs.models.dits import LTX2VideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, LTX2GemmaConfig
from fastvideo.configs.models.vaes import LTX2VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def ltx2_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
return outputs.last_hidden_state
@dataclass
class LTX2T2VConfig(PipelineConfig):
"""Configuration for LTX-2 T2V pipeline."""
dit_config: DiTConfig = field(default_factory=LTX2VideoConfig)
vae_config: VAEConfig = field(default_factory=LTX2VAEConfig)
vae_tiling: bool = True
vae_sp: bool = False
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LTX2GemmaConfig(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(ltx2_postprocess_text, ))
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", ))
audio_decoder_config: ModelConfig = field(
default_factory=LTX2AudioDecoderConfig)
vocoder_config: ModelConfig = field(default_factory=LTX2VocoderConfig)
audio_decoder_precision: str = "bf16"
vocoder_precision: str = "bf16"
def __post_init__(self) -> None:
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
+4 -37
View File
@@ -6,15 +6,10 @@ 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.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionT2V_1_3B_Config, TurboDiffusionT2V_14B_Config,
TurboDiffusionI2V_A14B_Config)
# isort: off
from fastvideo.configs.pipelines.wan import (
@@ -57,32 +52,14 @@ 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,
# LongCat Video models
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
# LTX-2 models
"Lightricks/LTX-2": LTX2T2VConfig,
"converted/ltx2_diffusers": LTX2T2VConfig,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers": TurboDiffusionT2V_1_3B_Config,
"loayrashid/TurboWan2.1-T2V-14B-Diffusers": TurboDiffusionT2V_14B_Config,
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers": TurboDiffusionI2V_A14B_Config,
# Add other specific weight variants
}
# For determining pipeline type from model ID
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"longcatimagetovideo":
lambda id: "longcatimagetovideo" in id.lower(),
"longcatvideocontinuation":
lambda id: "longcatvideocontinuation" in id.lower(),
"longcat":
lambda id: "longcat" in id.lower(),
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
@@ -100,23 +77,15 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"stepvideo":
lambda id: "stepvideo" in id.lower(),
"cosmos":
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(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
lambda id: "cosmos" in id.lower(),
"longcat":
lambda id: "longcat" in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
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,
@@ -127,9 +96,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"wanimagetovideo": WanI2V480PConfig,
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig,
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
"ltx2": LTX2T2VConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
@@ -1,128 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
TurboDiffusion pipeline configurations.
TurboDiffusion uses RCM (recurrent Consistency Model) scheduler with
SLA (Sparse-Linear Attention) for fast 1-4 step video generation.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.models.encoders import CLIPVisionConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.wan import t5_postprocess_text, T5Config, BaseEncoderOutput
import torch
from collections.abc import Callable
@dataclass
class TurboDiffusionT2VConfig(PipelineConfig):
"""Base configuration for TurboDiffusion T2V pipeline.
Uses RCM scheduler with sigma_max=80 for 1-4 step generation.
No boundary_ratio (single model, no switching).
"""
# DiT
dit_config: DiTConfig = field(default_factory=WanVideoConfig)
# VAE
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# Denoising stage
flow_shift: float | None = 3.0
# No boundary_ratio for T2V (single model)
boundary_ratio: float | None = None
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp32"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# self-forcing params
warp_denoising_step: bool = True
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
# Ensure no boundary_ratio is set in dit_config
self.dit_config.boundary_ratio = None
@dataclass
class TurboDiffusionT2V_1_3B_Config(TurboDiffusionT2VConfig):
"""Configuration for TurboDiffusion T2V 1.3B model."""
pass
@dataclass
class TurboDiffusionT2V_14B_Config(TurboDiffusionT2VConfig):
"""Configuration for TurboDiffusion T2V 14B model.
Uses same config as 1.3B but with higher flow_shift for 14B model.
"""
flow_shift: float | None = 5.0
@dataclass
class TurboDiffusionI2VConfig(PipelineConfig):
"""Base configuration for TurboDiffusion I2V pipeline.
Uses RCM scheduler with sigma_max=200 for 1-4 step generation.
Uses boundary_ratio=0.9 for high-noise to low-noise model switching.
"""
# DiT
dit_config: DiTConfig = field(default_factory=WanVideoConfig)
# VAE
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# Denoising stage
flow_shift: float | None = 5.0
boundary_ratio: float | None = 0.9
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
# Image encoder for I2V
image_encoder_config: EncoderConfig = field(
default_factory=CLIPVisionConfig)
image_encoder_precision: str = "fp32"
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp32"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# self-forcing params
warp_denoising_step: bool = True
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self.dit_config.boundary_ratio = self.boundary_ratio
@dataclass
class TurboDiffusionI2V_A14B_Config(TurboDiffusionI2VConfig):
"""Configuration for TurboDiffusion I2V A14B model."""
pass
+1 -1
View File
@@ -223,7 +223,7 @@ class SamplingParam:
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--video-path",
"--video_path",
type=str,
default=SamplingParam.video_path,
help="Path to input video for video-to-video generation",
-19
View File
@@ -1,19 +0,0 @@
# 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
-20
View File
@@ -1,20 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
"""
seed: int = 10
num_frames: int = 121
height: int = 1024
width: int = 1536
fps: int = 24
num_inference_steps: int = 8
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
+3 -36
View File
@@ -9,8 +9,6 @@ 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
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
@@ -28,11 +26,6 @@ from fastvideo.configs.sample.wan import (
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
MatrixGame2_SamplingParam,
)
from fastvideo.configs.sample.turbodiffusion import (
TurboDiffusionT2V_1_3B_SamplingParam,
TurboDiffusionT2V_14B_SamplingParam,
TurboDiffusionI2V_A14B_SamplingParam,
)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -86,27 +79,11 @@ 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,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers":
TurboDiffusionT2V_1_3B_SamplingParam,
"loayrashid/TurboWan2.1-T2V-14B-Diffusers":
TurboDiffusionT2V_14B_SamplingParam,
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers":
TurboDiffusionI2V_A14B_SamplingParam,
# LTX-2 models
"Lightricks/LTX-2": LTX2SamplingParam,
"FastVideo/LTX2-Distilled-Diffusers": LTX2SamplingParam,
# Add other specific weight variants
}
@@ -128,14 +105,6 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "wancausaldmdpipeline" in id.lower(),
"matrixgame":
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(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -152,11 +121,6 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam,
"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,
"ltx2": LTX2SamplingParam,
# Other fallbacks by architecture
}
@@ -180,6 +144,9 @@ 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)
@@ -1,73 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
TurboDiffusion sampling parameters.
TurboDiffusion uses RCM (recurrent Consistency Model) scheduler for
1-4 step video generation with no classifier-free guidance.
"""
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class TurboDiffusionT2V_1_3B_SamplingParam(SamplingParam):
"""Sampling parameters for TurboDiffusion T2V 1.3B model.
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
"""
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
guidance_scale: float = 1.0
num_inference_steps: int = 4
# No negative prompt needed for TurboDiffusion (no CFG)
negative_prompt: str | None = None
@dataclass
class TurboDiffusionT2V_14B_SamplingParam(SamplingParam):
"""Sampling parameters for TurboDiffusion T2V 14B model.
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
"""
# Video parameters (720p for 14B)
height: int = 720
width: int = 1280
num_frames: int = 81
fps: int = 16
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
guidance_scale: float = 1.0
num_inference_steps: int = 4
# No negative prompt needed for TurboDiffusion (no CFG)
negative_prompt: str | None = None
@dataclass
class TurboDiffusionI2V_A14B_SamplingParam(SamplingParam):
"""Sampling parameters for TurboDiffusion I2V A14B model.
Uses 4-step RCM sampling with dual-model switching (high/low noise).
"""
# Video parameters (720p for A14B I2V)
height: int = 720
width: int = 1280
num_frames: int = 81
fps: int = 16
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
guidance_scale: float = 1.0
num_inference_steps: int = 4
# Note: boundary_ratio is set in the pipeline config (TurboDiffusionI2VConfig),
# not here. This keeps sampling params and pipeline config separate.
# No negative prompt needed for TurboDiffusion (no CFG)
negative_prompt: str | None = None
-100
View File
@@ -18,8 +18,6 @@ import numpy as np
import torch
import torchvision
from einops import rearrange
import shutil
import tempfile
from fastvideo.configs.sample import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
@@ -391,11 +389,6 @@ class VideoGenerator:
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if (audio is not None and audio_sample_rate is not None and
not self._mux_audio(output_path, audio, audio_sample_rate)):
logger.warning("Audio mux failed; saved video without audio.")
if batch.return_frames:
return frames
@@ -403,7 +396,6 @@ class VideoGenerator:
return {
"samples": samples,
"frames": frames,
"audio": output_batch.extra.get("audio"),
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
@@ -413,98 +405,6 @@ class VideoGenerator:
"trajectory_decoded": output_batch.trajectory_decoded,
}
@staticmethod
def _mux_audio(
video_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
) -> bool:
"""Mux audio into video using PyAV."""
try:
import av
except ImportError:
logger.warning("PyAV not installed; cannot mux audio. "
"Install with: pip install av")
return False
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
logger.warning("Unexpected audio shape %s; skipping mux.",
audio_np.shape)
return False
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
num_channels = audio_int16.shape[1]
layout = "stereo" if num_channels == 2 else "mono"
try:
import wave
with tempfile.TemporaryDirectory() as tmpdir:
out_path = os.path.join(tmpdir, "muxed.mp4")
wav_path = os.path.join(tmpdir, "audio.wav")
# Write audio to WAV file
with wave.open(wav_path, "wb") as wav_file:
wav_file.setnchannels(num_channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(audio_int16.tobytes())
# Open input video and audio
input_video = av.open(video_path)
input_audio = av.open(wav_path)
# Create output with both streams
output = av.open(out_path, mode="w")
# Add video stream (copy codec from input)
in_video_stream = input_video.streams.video[0]
out_video_stream = output.add_stream(
codec_name=in_video_stream.codec_context.name,
rate=in_video_stream.average_rate,
)
out_video_stream.width = in_video_stream.width
out_video_stream.height = in_video_stream.height
out_video_stream.pix_fmt = in_video_stream.pix_fmt
# Add audio stream (AAC)
out_audio_stream = output.add_stream("aac", rate=sample_rate)
out_audio_stream.layout = layout
# Remux video (decode and re-encode to be safe)
for frame in input_video.decode(video=0):
for packet in out_video_stream.encode(frame):
output.mux(packet)
for packet in out_video_stream.encode():
output.mux(packet)
# Encode audio
for frame in input_audio.decode(audio=0):
frame.pts = None # Let encoder assign PTS
for packet in out_audio_stream.encode(frame):
output.mux(packet)
for packet in out_audio_stream.encode():
output.mux(packet)
input_video.close()
input_audio.close()
output.close()
shutil.move(out_path, video_path)
return True
except Exception as e:
logger.warning("Audio mux failed: %s", e)
return False
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
+3 -87
View File
@@ -132,8 +132,8 @@ class FastVideoArgs:
# CPU offload parameters
dit_cpu_offload: bool = True
use_fsdp_inference: bool = False
dit_layerwise_offload: bool = True
use_fsdp_inference: bool = True
dit_layerwise_offload: bool = False
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
vae_cpu_offload: bool = True
@@ -166,14 +166,6 @@ class FastVideoArgs:
# Prompt text file for batch processing
prompt_txt: str | None = None
# LTX-2 VAE tiling overrides
ltx2_vae_tiling: bool | None = None
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
ltx2_vae_temporal_tile_size_in_frames: int | None = None
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
ltx2_initial_latent_path: str | None = None
# model paths for correct deallocation
model_paths: dict[str, str] = field(default_factory=dict)
model_loaded: dict[str, bool] = field(default_factory=lambda: {
@@ -211,44 +203,8 @@ class FastVideoArgs:
logger.error("Failed to load V-MoBA config from %s: %s",
self.moba_config_path, e)
raise
self._apply_ltx2_vae_overrides()
self.check_fastvideo_args()
def _apply_ltx2_vae_overrides(self) -> None:
if self.pipeline_config is None:
return
vae_config = self.pipeline_config.vae_config
has_any = any(value is not None for value in (
self.ltx2_vae_spatial_tile_size_in_pixels,
self.ltx2_vae_spatial_tile_overlap_in_pixels,
self.ltx2_vae_temporal_tile_size_in_frames,
self.ltx2_vae_temporal_tile_overlap_in_frames,
))
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
"vae_tiling"):
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
self.pipeline_config.vae_tiling = True
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
vae_config.ltx2_spatial_tile_size_in_pixels = (
self.ltx2_vae_spatial_tile_size_in_pixels)
if hasattr(
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
self.ltx2_vae_spatial_tile_overlap_in_pixels)
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
vae_config.ltx2_temporal_tile_size_in_frames = (
self.ltx2_vae_temporal_tile_size_in_frames)
if hasattr(
vae_config, "ltx2_temporal_tile_overlap_in_frames"
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
vae_config.ltx2_temporal_tile_overlap_in_frames = (
self.ltx2_vae_temporal_tile_overlap_in_frames)
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
# Model and path configuration
@@ -369,44 +325,6 @@ class FastVideoArgs:
"Path to a text file containing prompts (one per line) for batch processing",
)
# LTX-2 VAE tiling overrides
parser.add_argument(
"--ltx2-vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.ltx2_vae_tiling,
help="Enable LTX-2 VAE tiling overrides.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-size-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
help="LTX-2 VAE spatial tile size in pixels.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-overlap-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
help="LTX-2 VAE spatial tile overlap in pixels.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-size-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
help="LTX-2 VAE temporal tile size in frames.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-overlap-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
help="LTX-2 VAE temporal tile overlap in frames.",
)
parser.add_argument(
"--ltx2-initial-latent-path",
type=str,
default=FastVideoArgs.ltx2_initial_latent_path,
help="Path to load/save a precomputed LTX-2 initial latent.",
)
# LoRA parameters (inference-time adapter loading)
parser.add_argument(
"--lora-path",
@@ -513,9 +431,7 @@ class FastVideoArgs:
"--use-fsdp-inference",
action=StoreBoolean,
help=
"Use FSDP for inference by sharding the model weights. FSDP helps reduce GPU memory usage but may introduce"
+
" weight transfer overhead depending on the specific setup. Enable if run out of memory.",
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
)
parser.add_argument(
"--text-encoder-cpu-offload",
View File
-122
View File
@@ -1,122 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import functools
from typing import Any
from torch import nn
class ForwardHook:
"""
Base class for forward hooks.
Hooks are used in the way:
modified_args, modified_kwargs = hook.pre_forward(module, *args, **kwargs)
output = module.forward(*modified_args, **modified_kwargs)
modified_output = hook.post_forward(module, output)
"""
@classmethod
def name(cls) -> str:
raise NotImplementedError
def on_attach(self, module: nn.Module): # noqa: B027
"""Called once when the hook is attached to the module."""
pass
def on_detach(self, module: nn.Module): # noqa: B027
"""
Called once when the hook is detached from the module.
Note: this function is not guaranteed to be called if the module is
deleted before the hook is detached.
"""
pass
def pre_forward(self, module: nn.Module, *args,
**kwargs) -> tuple[tuple[Any, ...], dict[str, Any]]:
"""Called before the module's forward method is executed."""
return args, kwargs
def post_forward(self, module: nn.Module, output: Any) -> Any:
"""Called after the module's forward method is executed."""
return output
class ModuleHookManager:
module_hook_attribute = "_hook_manager"
def __init__(self, module: nn.Module):
self.module = module
self.forward_hooks: dict[str, ForwardHook] = {}
self.original_forward = module.forward
@classmethod
def get_from(cls, module: nn.Module) -> "ModuleHookManager | None":
if hasattr(module, cls.module_hook_attribute):
return getattr(module, cls.module_hook_attribute)
return None
@classmethod
def get_from_or_default(cls, module: nn.Module) -> "ModuleHookManager":
if not hasattr(module, cls.module_hook_attribute):
setattr(module, cls.module_hook_attribute, cls(module))
def forward_hook_wrapper(mod: nn.Module, *args, **kwargs):
manager: ModuleHookManager = getattr(mod,
cls.module_hook_attribute)
for hook in manager.forward_hooks.values():
args, kwargs = hook.pre_forward(mod, *args, **kwargs)
output = manager.original_forward(*args, **kwargs)
for hook in reversed(manager.forward_hooks.values()):
output = hook.post_forward(mod, output)
return output
module.forward = functools.partial(forward_hook_wrapper, module)
return getattr(module, cls.module_hook_attribute)
@staticmethod
def remove_from_manager(module: nn.Module) -> None:
if hasattr(module, ModuleHookManager.module_hook_attribute):
manager: ModuleHookManager = getattr(
module, ModuleHookManager.module_hook_attribute)
module.forward = manager.original_forward
delattr(module, ModuleHookManager.module_hook_attribute)
def _check_manager_attached(self) -> None:
if not hasattr(self.module, self.module_hook_attribute):
raise ValueError("ModuleHookManager is not attached to the module.")
if getattr(self.module, self.module_hook_attribute) is not self:
raise ValueError(
"ModuleHookManager attached to the module is different.")
def append_forward_hook(self, hook: ForwardHook):
self._check_manager_attached()
if hook.name() in self.forward_hooks:
raise ValueError(
f"Hook with name {hook.name()} is already registered.")
# after python 3.7, dicts maintain insertion order
self.forward_hooks[hook.name()] = hook
hook.on_attach(self.module)
def replace_forward_hook(self,
hook_name: str,
new_hook: ForwardHook,
run_on_attach: bool = True):
self._check_manager_attached()
if hook_name not in self.forward_hooks:
raise ValueError(f"No hook with name {hook_name} found.")
old_hook = self.forward_hooks[hook_name]
if run_on_attach:
old_hook.on_detach(self.module)
self.forward_hooks[hook_name] = new_hook
new_hook.on_attach(self.module)
def remove_forward_hook(self, hook_name: str, run_detach: bool = True):
self._check_manager_attached()
if hook_name not in self.forward_hooks:
raise ValueError(f"No hook with name {hook_name} found.")
if run_detach:
self.forward_hooks[hook_name].on_detach(self.module)
del self.forward_hooks[hook_name]
def get_forward_hook(self, hook_name: str) -> ForwardHook | None:
return self.forward_hooks.get(hook_name, None)
-164
View File
@@ -1,164 +0,0 @@
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)]
-9
View File
@@ -1,9 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.models.audio.ltx2_audio_vae import (
LTX2AudioDecoder,
LTX2AudioEncoder,
LTX2Vocoder,
)
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
File diff suppressed because it is too large Load Diff
+34 -304
View File
@@ -1,6 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
"""
Native LongCat Video DiT implementation using FastVideo conventions.
This is a Phase 2 reimplementation that replaces the third_party wrapper
with native FastVideo layers for better performance and integration.
"""
from typing import Any
@@ -126,7 +129,7 @@ class TimestepEmbedder(nn.Module):
# Sinusoidal embedding in FP32
t_freq = self.timestep_embedding(t.flatten(), self.frequency_embedding_size)
# Cast to model dtype before MLP (matching original LongCat)
# Cast to model dtype before MLP
# Handle LoRA wrapper if present
linear_layer = self.linear_1.base_layer if hasattr(self.linear_1, 'base_layer') else self.linear_1
target_dtype = linear_layer.weight.dtype
@@ -163,14 +166,13 @@ class CaptionEmbedder(nn.Module):
self.text_tokens_zero_pad = text_tokens_zero_pad
# Two-layer MLP using ReplicatedLinear
# CRITICAL: Original LongCat uses GELU(approximate="tanh"), NOT SiLU!
self.linear_1 = ReplicatedLinear(
caption_channels,
hidden_size,
bias=True,
params_dtype=dtype,
)
self.act = nn.GELU(approximate="tanh") # Match original LongCat
self.act = nn.SiLU()
self.linear_2 = ReplicatedLinear(
hidden_size,
hidden_size,
@@ -266,19 +268,10 @@ class LongCatSelfAttention(nn.Module):
self,
x: torch.Tensor, # [B, N, C]
latent_shape: tuple, # (T, H, W)
num_cond_latents: int = 0, # Number of conditioning latent frames (for I2V)
return_kv: bool = False, # Return K/V for caching
**kwargs
) -> torch.Tensor | tuple:
) -> torch.Tensor:
"""
Forward pass with 3D RoPE and optional BSA.
For I2V mode (num_cond_latents > 0):
- Conditioned tokens only attend to themselves
- Noise tokens attend to ALL tokens (cond + noise)
Args:
return_kv: If True, return (output, (k_cache, v_cache)) for KV caching
"""
B, N, C = x.shape
T, H, W = latent_shape
@@ -297,12 +290,6 @@ class LongCatSelfAttention(nn.Module):
q = self.q_norm(q)
k = self.k_norm(k)
# Save pre-RoPE K/V for cache if requested (before RoPE is applied)
if return_kv:
# [B, N, num_heads, head_dim] -> [B, num_heads, N, head_dim]
k_cache = k.transpose(1, 2).clone()
v_cache = v.transpose(1, 2).clone()
# For RoPE: need [B, num_heads, N, head_dim]
q_rope = q.transpose(1, 2)
k_rope = k.transpose(1, 2)
@@ -314,50 +301,6 @@ class LongCatSelfAttention(nn.Module):
q = q_rope.transpose(1, 2)
k = k_rope.transpose(1, 2)
# === I2V Split Attention ===
# For I2V, conditioned tokens and noise tokens are processed separately
if num_cond_latents > 0:
# Calculate number of conditioned tokens (cond_latents * spatial_tokens_per_frame)
num_cond_tokens = num_cond_latents * (N // T)
# Conditioned tokens: only attend to themselves (same seq length, use self.attn)
q_cond = q[:, :num_cond_tokens].contiguous()
k_cond = k[:, :num_cond_tokens].contiguous()
v_cond = v[:, :num_cond_tokens].contiguous()
out_cond, _ = self.attn(q_cond, k_cond, v_cond)
# Noise tokens: attend to ALL tokens (different seq lengths!)
# Need to use flash attention directly since q has different length than k/v
q_noise = q[:, num_cond_tokens:].contiguous() # [B, N_noise, num_heads, head_dim]
# k, v are full: [B, N, num_heads, head_dim]
# Transpose for flash attention: [B, num_heads, seq, head_dim]
q_noise_t = q_noise.transpose(1, 2)
k_t = k.transpose(1, 2)
v_t = v.transpose(1, 2)
# Use scaled dot product attention (handles different q/kv lengths)
out_noise_t = torch.nn.functional.scaled_dot_product_attention(
q_noise_t, k_t, v_t,
attn_mask=None,
dropout_p=0.0,
is_causal=False
) # [B, num_heads, N_noise, head_dim]
# Transpose back: [B, N_noise, num_heads, head_dim]
out_noise = out_noise_t.transpose(1, 2)
# Merge conditioned and noise outputs
out = torch.cat([out_cond, out_noise], dim=1)
# Reshape and project out
out = out.reshape(B, N, C)
out, _ = self.to_out(out)
if return_kv:
return out, (k_cache, v_cache)
return out
# === Attention: BSA or standard ===
if self.enable_bsa and T > 1: # Only use BSA for multi-frame videos
# BSA expects [B, H, S, D] format
@@ -405,96 +348,6 @@ class LongCatSelfAttention(nn.Module):
out = out.reshape(B, N, C)
out, _ = self.to_out(out)
if return_kv:
return out, (k_cache, v_cache)
return out
def forward_with_kv_cache(
self,
x: torch.Tensor, # [B, N_noise, C] - only noise tokens
latent_shape: tuple, # (T_noise, H, W) - shape for noise only
num_cond_latents: int, # Number of conditioning latent frames
kv_cache: tuple, # (k_cond, v_cond) - [B, heads, N_cond, head_dim]
) -> torch.Tensor:
"""
Forward using cached K/V from conditioning frames.
x contains only NOISE tokens.
kv_cache contains pre-computed K/V for CONDITIONING tokens.
CRITICAL: RoPE positions for noise tokens must start AFTER conditioning.
We achieve this by padding Q with dummy tokens for conditioning positions,
applying RoPE to the full sequence, then extracting only noise token Q.
"""
B, N, C = x.shape
T, H, W = latent_shape
k_cache, v_cache = kv_cache
# Handle batch size mismatch (cache might be smaller for CFG)
# When using CFG, latent_model_input is doubled [neg, pos], but cache is for original batch
if k_cache.shape[0] != B:
# Expand cache to match input batch size
# For CFG: repeat the cache for both negative and positive branches
repeat_factor = B // k_cache.shape[0]
k_cache = k_cache.repeat(repeat_factor, 1, 1, 1)
v_cache = v_cache.repeat(repeat_factor, 1, 1, 1)
# Project to Q/K/V for noise tokens
q, _ = self.to_q(x)
k, _ = self.to_k(x)
v, _ = self.to_v(x)
# Reshape to heads: [B, N, num_heads, head_dim]
q = q.view(B, N, self.num_heads, self.head_dim)
k = k.view(B, N, self.num_heads, self.head_dim)
v = v.view(B, N, self.num_heads, self.head_dim)
# Per-head RMS normalization
q = self.q_norm(q)
k = self.k_norm(k)
# Transpose for RoPE: [B, heads, N, head_dim]
q_rope = q.transpose(1, 2)
k_rope = k.transpose(1, 2)
v = v.transpose(1, 2)
# CRITICAL: Apply RoPE with correct positional offset
# Noise frame queries need positions starting from num_cond_latents
# Following the original LongCat approach:
# 1. Pad Q with dummy tokens matching k_cache shape
# 2. Apply RoPE to full sequence (T_cond + T_noise)
# 3. Extract only the noise portion of Q
# Create dummy Q padding to fill conditioning positions
# k_cache shape: [B, heads, N_cond, head_dim]
q_padding = torch.cat([torch.empty_like(k_cache), q_rope], dim=2).contiguous()
# Concatenate cached K with noise K for RoPE
k_full = torch.cat([k_cache, k_rope], dim=2)
v_full = torch.cat([v_cache, v], dim=2)
# Apply RoPE to full sequence (includes both cond and noise positions)
# Grid size: (T_cond + T_noise, H, W)
full_T = num_cond_latents + T
q_padding, k_full = self.rope_3d(q_padding, k_full, grid_size=(full_T, H, W))
# Extract only the noise portion of Q (last N tokens)
q_rope = q_padding[:, :, -N:].contiguous()
# Run attention: Q_noise attends to full K/V (cond + noise)
out = torch.nn.functional.scaled_dot_product_attention(
q_rope, k_full, v_full,
attn_mask=None,
dropout_p=0.0,
is_causal=False
) # [B, heads, N_noise, head_dim]
# Transpose back: [B, N_noise, heads, head_dim]
out = out.transpose(1, 2)
out = out.reshape(B, N, C)
out, _ = self.to_out(out)
return out
@@ -541,8 +394,6 @@ class LongCatCrossAttention(nn.Module):
self,
x: torch.Tensor, # [B, N_img, C]
context: torch.Tensor, # [B, N_text, C]
latent_shape: tuple = None, # (T, H, W) - needed for I2V
num_cond_latents: int = 0, # Number of conditioning latent frames (for I2V)
**kwargs
) -> torch.Tensor:
"""
@@ -551,57 +402,9 @@ class LongCatCrossAttention(nn.Module):
Args:
x: Image tokens [B, N_img, C]
context: Text tokens [B, N_text, C] (standard padded format)
latent_shape: (T, H, W) - needed for calculating num_cond_tokens
num_cond_latents: Number of conditioning latent frames (for I2V)
For I2V mode (num_cond_latents > 0):
- Conditioned tokens get ZERO cross-attention output
- Only noise tokens get cross-attention with text
"""
B, N_img, C = x.shape
# === I2V: Only noise tokens get cross-attention ===
if num_cond_latents > 0 and latent_shape is not None:
T, H, W = latent_shape
num_cond_tokens = num_cond_latents * (N_img // T)
# Only process noise tokens
x_noise = x[:, num_cond_tokens:] # [B, N_noise, C]
# Project Q, K, V for noise tokens only
q, _ = self.to_q(x_noise)
k, _ = self.to_k(context)
v, _ = self.to_v(context)
N_text = context.shape[1]
N_noise = x_noise.shape[1]
# Reshape to heads
q = q.view(B, N_noise, self.num_heads, self.head_dim)
k = k.view(B, N_text, self.num_heads, self.head_dim)
v = v.view(B, N_text, self.num_heads, self.head_dim)
# Per-head RMS normalization
q = self.q_norm(q)
k = self.k_norm(k)
# Run cross-attention
out_noise = self.attn(q, k, v) # [B, N_noise, num_heads, head_dim]
out_noise = out_noise.reshape(B, N_noise, C)
out_noise, _ = self.to_out(out_noise)
# Conditioned tokens get zero output
out_cond = torch.zeros(
(B, num_cond_tokens, C),
dtype=out_noise.dtype,
device=out_noise.device
)
# Merge
out = torch.cat([out_cond, out_noise], dim=1)
return out
# === Standard cross-attention ===
# Project Q, K, V (standard cross-attention like WanVideo/StepVideo/Cosmos)
q, _ = self.to_q(x)
k, _ = self.to_k(context)
@@ -672,19 +475,19 @@ class LongCatSwiGLUFFN(nn.Module):
def modulate_fp32(norm: nn.Module, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
"""
Apply modulation in FP32 for numerical stability.
Apply modulation in FP32 for numerical stability (matching original LongCat).
Converts inputs to FP32 for the modulation operation, then casts back.
shift and scale should already be FP32 from torch.amp.autocast context.
"""
orig_dtype = x.dtype
# Ensure modulation params are FP32 (should be from autocast)
assert shift.dtype == torch.float32 and scale.dtype == torch.float32, \
f"shift and scale must be FP32, got {shift.dtype} and {scale.dtype}"
# Convert to FP32 for numerical stability
shift_fp32 = shift.float()
scale_fp32 = scale.float()
orig_dtype = x.dtype
# Normalize and modulate in FP32
x_norm = norm(x.to(torch.float32))
x_mod = x_norm * (scale_fp32 + 1) + shift_fp32
x_mod = x_norm * (scale + 1) + shift
return x_mod.to(orig_dtype)
@@ -765,21 +568,10 @@ class LongCatTransformerBlock(nn.Module):
context: torch.Tensor, # [B, N_text, C]
t: torch.Tensor, # [B, T, C_t]
latent_shape: tuple, # (T, H, W)
num_cond_latents: int = 0, # Number of conditioning latent frames (for I2V)
return_kv: bool = False, # Return K/V for caching
kv_cache: tuple | None = None, # Pre-computed K/V cache
skip_crs_attn: bool = False, # Skip cross-attention (for cache init)
**kwargs
) -> torch.Tensor | tuple:
) -> torch.Tensor:
"""
Forward pass with AdaLN modulation.
Args:
num_cond_latents: For I2V, number of conditioning latent frames.
These frames use split attention behavior.
return_kv: If True, return (x, (k_cache, v_cache))
kv_cache: Pre-computed K/V from conditioning frames
skip_crs_attn: If True, skip cross-attention (used during cache init)
"""
B, N, C = x.shape
T, H, W = latent_shape
@@ -800,47 +592,17 @@ class LongCatTransformerBlock(nn.Module):
x_norm = modulate_fp32(self.norm_attn, x.view(B, T, -1, C), shift_msa, scale_msa)
x_norm = x_norm.view(B, N, C)
# Handle KV cache
if kv_cache is not None:
# Move cache to device if offloaded
kv_cache = (kv_cache[0].to(x.device), kv_cache[1].to(x.device))
attn_out = self.self_attn.forward_with_kv_cache(
x_norm,
latent_shape=latent_shape,
num_cond_latents=num_cond_latents,
kv_cache=kv_cache,
)
kv_cache_new = None # Don't return cache when using cache
else:
attn_result = self.self_attn(
x_norm,
latent_shape=latent_shape,
num_cond_latents=num_cond_latents,
return_kv=return_kv,
)
if return_kv:
attn_out, kv_cache_new = attn_result
else:
attn_out = attn_result
kv_cache_new = None
attn_out = self.self_attn(x_norm, latent_shape=latent_shape)
# Residual with gating (CRITICAL: FP32 like original, then cast back)
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
x = x + (gate_msa * attn_out.view(B, T, -1, C)).view(B, N, C)
x = x.to(x_orig_dtype)
# === Cross-Attention (skip if requested) ===
if not skip_crs_attn:
x_norm_cross = self.norm_cross(x)
# When using KV cache, no need for num_cond_latents in cross-attn
cross_num_cond = 0 if kv_cache is not None else num_cond_latents
cross_out = self.cross_attn(
x_norm_cross,
context,
latent_shape=latent_shape,
num_cond_latents=cross_num_cond
)
x = x + cross_out
# === Cross-Attention ===
x_norm_cross = self.norm_cross(x)
cross_out = self.cross_attn(x_norm_cross, context)
x = x + cross_out
# === FFN ===
x_norm_ffn = modulate_fp32(self.norm_ffn, x.view(B, T, -1, C), shift_mlp, scale_mlp)
@@ -853,8 +615,6 @@ class LongCatTransformerBlock(nn.Module):
x = x + (gate_mlp * ffn_out.view(B, T, -1, C)).view(B, N, C)
x = x.to(x_orig_dtype)
if return_kv:
return x, kv_cache_new
return x
@@ -910,12 +670,16 @@ class FinalLayer(nn.Module):
B, N, C = x.shape
T, _, _ = latent_shape
# AdaLN modulation
t_mod = self.adaln_act(t)
mod_params, _ = self.adaln_linear(t_mod)
shift, scale = mod_params.unsqueeze(2).chunk(2, dim=-1)
# AdaLN modulation (FP32 for stability like original)
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
t_mod = self.adaln_act(t)
mod_params, _ = self.adaln_linear(t_mod)
# Ensure FP32 output (needed when LoRA is applied)
if mod_params.dtype != torch.float32:
mod_params = mod_params.float()
shift, scale = mod_params.unsqueeze(2).chunk(2, dim=-1)
# Modulate (converts to FP32 internally for stability)
# Modulate
x = modulate_fp32(self.norm, x.view(B, T, -1, C), shift, scale)
x = x.reshape(B, N, C)
@@ -932,6 +696,8 @@ class FinalLayer(nn.Module):
class LongCatTransformer3DModel(CachableDiT):
"""
Native LongCat Video Transformer using FastVideo layers.
This is a Phase 2 implementation that replaces third_party dependencies.
"""
# FSDP sharding: shard at each transformer block
@@ -1023,28 +789,13 @@ class LongCatTransformer3DModel(CachableDiT):
encoder_attention_mask: torch.Tensor | None = None, # [B, N_text]
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance: float | None = None, # Unused, for API compatibility
num_cond_latents: int = 0, # For I2V: number of conditioning latent frames
# === KV Cache Parameters ===
return_kv: bool = False, # If True, return (output, kv_cache_dict)
kv_cache_dict: dict | None = None, # Pre-computed {block_idx: (k, v)}
skip_crs_attn: bool = False, # Skip cross-attention (for cache init)
offload_kv_cache: bool = False, # Move cache to CPU after compute
**kwargs
) -> torch.Tensor | tuple[torch.Tensor, dict]:
) -> torch.Tensor:
"""
Forward pass with FastVideo parameter ordering.
NOTE: This follows FastVideo convention:
(hidden_states, encoder_hidden_states, timestep)
Args:
num_cond_latents: For I2V, number of conditioning latent frames.
These frames are treated as "clean" (timestep=0)
and use split attention behavior.
return_kv: If True, return (output, kv_cache_dict)
kv_cache_dict: Pre-computed K/V cache {block_idx: (k, v)}
skip_crs_attn: If True, skip cross-attention (for cache init)
offload_kv_cache: If True, move cache to CPU after compute
"""
B, _, T, H, W = hidden_states.shape
@@ -1074,31 +825,12 @@ class LongCatTransformer3DModel(CachableDiT):
encoder_attention_mask=encoder_attention_mask
) # [B, N_text, C]
# 4. Transformer blocks with optional KV cache
kv_cache_dict_ret = {} if return_kv else None
# 4. Transformer blocks
for i, block in enumerate(self.blocks):
# Get cache for this block if available
block_kv_cache = kv_cache_dict.get(i, None) if kv_cache_dict else None
block_out = block(
x = block(
x, context, t,
latent_shape=(N_t, N_h, N_w),
num_cond_latents=num_cond_latents,
return_kv=return_kv,
kv_cache=block_kv_cache,
skip_crs_attn=skip_crs_attn,
latent_shape=(N_t, N_h, N_w)
)
if return_kv:
x, kv_cache = block_out
# Store cache
if offload_kv_cache:
kv_cache_dict_ret[i] = (kv_cache[0].cpu(), kv_cache[1].cpu())
else:
kv_cache_dict_ret[i] = (kv_cache[0].contiguous(), kv_cache[1].contiguous())
else:
x = block_out
# 5. Output projection
output = self.final_layer(x, t, latent_shape=(N_t, N_h, N_w))
@@ -1109,8 +841,6 @@ class LongCatTransformer3DModel(CachableDiT):
# Cast to float32 for better accuracy (as per original)
output = output.to(torch.float32)
if return_kv:
return output, kv_cache_dict_ret
return output
def unpatchify(self, x: torch.Tensor, N_t: int, N_h: int, N_w: int) -> torch.Tensor:
File diff suppressed because it is too large Load Diff
+12 -2
View File
@@ -734,8 +734,18 @@ class WanTransformer3DModel(CachableDiT):
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
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,
timestep_proj, freqs_cis, attention_mask)
# if teacache is enabled, we need to cache the original hidden states
-563
View File
@@ -1,563 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
import os
from typing import Iterable
import torch
from torch import nn
from transformers import Gemma3ForConditionalGeneration
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.dits.ltx2 import (
FeedForward,
LTXRopeType,
apply_ltx_rotary_emb,
generate_ltx_freq_grid_np,
generate_ltx_freq_grid_pytorch,
precompute_ltx_freqs_cis,
)
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.platforms import AttentionBackendEnum
def _debug_log_line(message: str) -> None:
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") != "1":
return
log_path = os.getenv("LTX2_PIPELINE_DEBUG_PATH", "")
if not log_path:
return
log_dir = os.path.dirname(log_path)
if log_dir:
os.makedirs(log_dir, exist_ok=True)
with open(log_path, "a", encoding="utf-8") as f:
f.write(message + "\n")
def _debug_gemma_log_line(message: str) -> None:
log_path = os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", "")
if not log_path:
return
log_dir = os.path.dirname(log_path)
if log_dir:
os.makedirs(log_dir, exist_ok=True)
with open(log_path, "a", encoding="utf-8") as f:
f.write(message + "\n")
@dataclass(frozen=True)
class GemmaConnectorConfig:
num_attention_heads: int
attention_head_dim: int
num_layers: int
positional_embedding_theta: float
positional_embedding_max_pos: list[int]
rope_type: LTXRopeType
double_precision_rope: bool
num_learnable_registers: int | None
class GemmaFeaturesExtractorProjLinear(nn.Module):
"""Linear projection that aggregates stacked Gemma hidden states."""
def __init__(self, in_features: int, out_features: int) -> None:
super().__init__()
self.aggregate_embed = nn.Linear(in_features, out_features, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.aggregate_embed(x)
class _BasicTransformerBlock1D(nn.Module):
"""1D transformer block for connector processing."""
def __init__(
self,
dim: int,
heads: int,
dim_head: int,
rope_type: LTXRopeType,
norm_eps: float = 1e-6,
) -> None:
super().__init__()
self.attn1 = _GemmaAttention(
query_dim=dim,
context_dim=None,
heads=heads,
dim_head=dim_head,
norm_eps=norm_eps,
rope_type=rope_type,
)
self.ff = FeedForward(dim, dim_out=dim)
self.norm_eps = norm_eps
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
norm_hidden_states = torch.nn.functional.rms_norm(
hidden_states, (hidden_states.shape[-1],), eps=self.norm_eps
)
if norm_hidden_states.ndim == 4:
norm_hidden_states = norm_hidden_states.squeeze(1)
attn_output = self.attn1(
norm_hidden_states,
mask=attention_mask,
pe=pe,
)
hidden_states = attn_output + hidden_states
if hidden_states.ndim == 4:
hidden_states = hidden_states.squeeze(1)
norm_hidden_states = torch.nn.functional.rms_norm(
hidden_states, (hidden_states.shape[-1],), eps=self.norm_eps
)
ff_output = self.ff(norm_hidden_states)
hidden_states = ff_output + hidden_states
if hidden_states.ndim == 4:
hidden_states = hidden_states.squeeze(1)
return hidden_states
class _GemmaAttention(nn.Module):
"""Attention implementation aligned with LTX-2 text encoder."""
def __init__(
self,
query_dim: int,
context_dim: int | None,
heads: int,
dim_head: int,
norm_eps: float,
rope_type: LTXRopeType,
) -> None:
super().__init__()
inner_dim = dim_head * heads
context_dim = query_dim if context_dim is None else context_dim
self.heads = heads
self.dim_head = dim_head
self.rope_type = rope_type
self.q_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
self.k_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
self.to_q = nn.Linear(query_dim, inner_dim, bias=True)
self.to_k = nn.Linear(context_dim, inner_dim, bias=True)
self.to_v = nn.Linear(context_dim, inner_dim, bias=True)
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim, bias=True), nn.Identity())
def forward(
self,
x: torch.Tensor,
context: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
k_pe: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
q = self.to_q(x)
context = x if context is None else context
k = self.to_k(context)
v = self.to_v(context)
q = self.q_norm(q)
k = self.k_norm(k)
if pe is not None:
q = apply_ltx_rotary_emb(q, pe, self.rope_type)
k = apply_ltx_rotary_emb(k, pe if k_pe is None else k_pe, self.rope_type)
b, q_len, _ = q.shape
k_len = k.shape[1]
q = q.view(b, q_len, self.heads, self.dim_head).transpose(1, 2)
k = k.view(b, k_len, self.heads, self.dim_head).transpose(1, 2)
v = v.view(b, k_len, self.heads, self.dim_head).transpose(1, 2)
if mask is not None:
if mask.ndim == 2:
mask = mask.unsqueeze(0)
if mask.ndim == 3:
mask = mask.unsqueeze(1)
out = torch.nn.functional.scaled_dot_product_attention(
q,
k,
v,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
)
out = out.transpose(1, 2).reshape(b, q_len, -1)
return self.to_out(out)
class Embeddings1DConnector(nn.Module):
"""Transformer connector that refines Gemma embeddings for LTX-2."""
_supports_gradient_checkpointing = True
def __init__(self, config: GemmaConnectorConfig) -> None:
super().__init__()
self.num_attention_heads = config.num_attention_heads
self.inner_dim = config.num_attention_heads * config.attention_head_dim
self.positional_embedding_theta = config.positional_embedding_theta
self.positional_embedding_max_pos = config.positional_embedding_max_pos
self.rope_type = config.rope_type
self.double_precision_rope = config.double_precision_rope
self.transformer_1d_blocks = nn.ModuleList(
[
_BasicTransformerBlock1D(
dim=self.inner_dim,
heads=config.num_attention_heads,
dim_head=config.attention_head_dim,
rope_type=config.rope_type,
)
for _ in range(config.num_layers)
]
)
self.num_learnable_registers = config.num_learnable_registers
if self.num_learnable_registers:
self.learnable_registers = nn.Parameter(
torch.rand(
self.num_learnable_registers,
self.inner_dim,
dtype=torch.bfloat16,
)
* 2.0
- 1.0
)
def _replace_padded_with_learnable_registers(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
assert hidden_states.shape[1] % self.num_learnable_registers == 0, (
f"Hidden states sequence length {hidden_states.shape[1]} must be divisible by "
f"num_learnable_registers {self.num_learnable_registers}."
)
num_registers_duplications = (
hidden_states.shape[1] // self.num_learnable_registers
)
learnable_registers = torch.tile(
self.learnable_registers, (num_registers_duplications, 1)
)
attention_mask_binary = (
attention_mask.squeeze(1).squeeze(1).unsqueeze(-1) >= -9000.0
).int()
non_zero_hidden_states = hidden_states[
:, attention_mask_binary.squeeze().bool(), :
]
non_zero_nums = non_zero_hidden_states.shape[1]
pad_length = hidden_states.shape[1] - non_zero_nums
adjusted_hidden_states = torch.nn.functional.pad(
non_zero_hidden_states, pad=(0, 0, 0, pad_length), value=0
)
flipped_mask = torch.flip(attention_mask_binary, dims=[1])
hidden_states = flipped_mask * adjusted_hidden_states + (
1 - flipped_mask
) * learnable_registers
attention_mask = torch.full_like(
attention_mask,
0.0,
dtype=attention_mask.dtype,
device=attention_mask.device,
)
return hidden_states, attention_mask
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
if self.num_learnable_registers:
hidden_states, attention_mask = (
self._replace_padded_with_learnable_registers(
hidden_states, attention_mask
)
)
indices_grid = torch.arange(
hidden_states.shape[1],
dtype=torch.float32,
device=hidden_states.device,
)
indices_grid = indices_grid[None, None, :]
freq_grid_generator = (
generate_ltx_freq_grid_np
if self.double_precision_rope
else generate_ltx_freq_grid_pytorch
)
freqs_cis = precompute_ltx_freqs_cis(
indices_grid=indices_grid,
dim=self.inner_dim,
out_dtype=hidden_states.dtype,
theta=self.positional_embedding_theta,
max_pos=self.positional_embedding_max_pos,
num_attention_heads=self.num_attention_heads,
rope_type=self.rope_type,
freq_grid_generator=freq_grid_generator,
)
for block in self.transformer_1d_blocks:
hidden_states = block(
hidden_states, attention_mask=attention_mask, pe=freqs_cis
)
hidden_states = torch.nn.functional.rms_norm(
hidden_states, (hidden_states.shape[-1],), eps=1e-6
)
return hidden_states, attention_mask
class LTX2GemmaTextEncoderModel(TextEncoder):
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__(config)
arch = config.arch_config
self.feature_extractor_linear = GemmaFeaturesExtractorProjLinear(
in_features=arch.feature_extractor_in_features,
out_features=arch.feature_extractor_out_features,
)
connector_config = GemmaConnectorConfig(
num_attention_heads=arch.connector_num_attention_heads,
attention_head_dim=arch.connector_attention_head_dim,
num_layers=arch.connector_num_layers,
positional_embedding_theta=arch.connector_positional_embedding_theta,
positional_embedding_max_pos=arch.connector_positional_embedding_max_pos,
rope_type=LTXRopeType(arch.connector_rope_type),
double_precision_rope=arch.connector_double_precision_rope,
num_learnable_registers=arch.connector_num_learnable_registers,
)
self.embeddings_connector = Embeddings1DConnector(connector_config)
self.audio_embeddings_connector = Embeddings1DConnector(connector_config)
self.gemma_model_path = arch.gemma_model_path
self.gemma_dtype = arch.gemma_dtype
self.padding_side = arch.padding_side
self._gemma_model: Gemma3ForConditionalGeneration | None = None
def named_parameters(self, prefix: str = "", recurse: bool = True):
for name, param in super().named_parameters(
prefix=prefix, recurse=recurse
):
if name.startswith("gemma_model."):
continue
yield name, param
@property
def gemma_model(self) -> Gemma3ForConditionalGeneration:
if self._gemma_model is None:
gemma_path = self.gemma_model_path
if not gemma_path:
raise ValueError(
"gemma_model_path must be set (expected text_encoder/gemma)."
)
dtype = getattr(torch, self.gemma_dtype, torch.bfloat16)
self._gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
gemma_path,
local_files_only=True,
torch_dtype=dtype,
)
# Configure model-level attention implementation when using TORCH_SDPA.
# Note: torch.backends.cuda.enable_*_sdp() settings should be configured
# at application/pipeline initialization level, not here, to avoid
# unexpected side effects across the application.
if os.getenv("FASTVIDEO_ATTENTION_BACKEND") == "TORCH_SDPA":
if hasattr(self._gemma_model.config, "attn_implementation"):
self._gemma_model.config.attn_implementation = "sdpa"
if hasattr(self._gemma_model.config, "_attn_implementation"):
self._gemma_model.config._attn_implementation = "sdpa"
device = next(self.feature_extractor_linear.parameters()).device
self._gemma_model.to(device=device)
self._gemma_model.eval()
return self._gemma_model
def _run_feature_extractor(
self,
hidden_states: tuple[torch.Tensor, ...],
attention_mask: torch.Tensor,
padding_side: str,
) -> torch.Tensor:
encoded_text_features = torch.stack(hidden_states, dim=-1)
if os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", ""):
for idx, layer in enumerate(hidden_states):
_debug_gemma_log_line(
f"fastvideo:gemma_hidden_state_{idx}"
f":sum={layer.float().sum().item():.6f}"
)
_debug_gemma_log_line(
"fastvideo:gemma_hidden_states_stack"
f":sum={encoded_text_features.float().sum().item():.6f}"
)
encoded_text_features_dtype = encoded_text_features.dtype
sequence_lengths = attention_mask.sum(dim=-1)
normed_text_features = _norm_and_concat_padded_batch(
encoded_text_features, sequence_lengths, padding_side=padding_side
)
return self.feature_extractor_linear(
normed_text_features.to(encoded_text_features_dtype)
)
def _convert_to_additive_mask(
self, attention_mask: torch.Tensor, dtype: torch.dtype
) -> torch.Tensor:
return (attention_mask - 1).to(dtype).reshape(
(attention_mask.shape[0], 1, -1, attention_mask.shape[-1])
) * torch.finfo(dtype).max
def _run_connectors(
self,
encoded_input: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
connector_attention_mask = self._convert_to_additive_mask(
attention_mask, encoded_input.dtype
)
encoded, encoded_connector_attention_mask = self.embeddings_connector(
encoded_input, connector_attention_mask
)
attention_mask = (encoded_connector_attention_mask < 0.000001).to(
torch.int64
)
attention_mask = attention_mask.reshape(
[encoded.shape[0], encoded.shape[1], 1]
)
encoded = encoded * attention_mask
encoded_for_audio, _ = self.audio_embeddings_connector(
encoded_input, connector_attention_mask
)
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
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:
if input_ids is None:
raise ValueError("input_ids is required for Gemma text encoding.")
if attention_mask is None:
attention_mask = torch.ones_like(input_ids)
model = self.gemma_model
input_ids = input_ids.to(device=model.device)
attention_mask = attention_mask.to(device=model.device)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
encoded_inputs = self._run_feature_extractor(
outputs.hidden_states,
attention_mask,
padding_side=self.padding_side,
)
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
_debug_log_line(
"fastvideo:gemma_feature"
f":sum={encoded_inputs.float().sum().item():.6f} "
f"shape={tuple(encoded_inputs.shape)}"
)
video_encoding, audio_encoding, attention_mask = self._run_connectors(
encoded_inputs, attention_mask
)
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
_debug_log_line(
"fastvideo:gemma_video_encoding"
f":sum={video_encoding.float().sum().item():.6f} "
f"shape={tuple(video_encoding.shape)}"
)
_debug_log_line(
"fastvideo:gemma_audio_encoding"
f":sum={audio_encoding.float().sum().item():.6f} "
f"shape={tuple(audio_encoding.shape)}"
)
hidden_states = (audio_encoding, ) if output_hidden_states else None
return BaseEncoderOutput(
last_hidden_state=video_encoding,
hidden_states=hidden_states,
attention_mask=attention_mask,
)
def load_weights(
self, weights: Iterable[tuple[str, torch.Tensor]]
) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if name == "aggregate_embed.weight":
name = "feature_extractor_linear.aggregate_embed.weight"
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)
loaded_params.add(name)
return loaded_params
def _norm_and_concat_padded_batch(
encoded_text: torch.Tensor,
sequence_lengths: torch.Tensor,
padding_side: str = "right",
) -> torch.Tensor:
b, t, d, l = encoded_text.shape
device = encoded_text.device
token_indices = torch.arange(t, device=device)[None, :]
if padding_side == "right":
mask = token_indices < sequence_lengths[:, None]
elif padding_side == "left":
start_indices = t - sequence_lengths[:, None]
mask = token_indices >= start_indices
else:
raise ValueError(
f"padding_side must be 'left' or 'right', got {padding_side}"
)
mask = mask.reshape(b, t, 1, 1)
eps = 1e-6
masked = encoded_text.masked_fill(~mask, 0.0)
denom = (sequence_lengths * d).view(b, 1, 1, 1)
mean = masked.sum(dim=(1, 2), keepdim=True) / (denom + eps)
x_min = encoded_text.masked_fill(~mask, float("inf")).amin(
dim=(1, 2), keepdim=True
)
x_max = encoded_text.masked_fill(~mask, float("-inf")).amax(
dim=(1, 2), keepdim=True
)
range_ = x_max - x_min
normed = 8 * (encoded_text - mean) / (range_ + eps)
normed = normed.reshape(b, t, -1)
mask_flattened = mask.reshape(b, t, 1).expand(-1, -1, d * l)
normed = normed.masked_fill(~mask_flattened, 0.0)
return normed
File diff suppressed because it is too large Load Diff
-353
View File
@@ -1,353 +0,0 @@
# 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
)
+51 -260
View File
@@ -15,7 +15,8 @@ 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, AutoProcessor, AutoTokenizer
from transformers import AutoImageProcessor, AutoTokenizer
from transformers import UMT5EncoderModel
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.configs.models import EncoderConfig
@@ -34,8 +35,8 @@ from fastvideo.models.loader.weight_utils import (
safetensors_weights_iterator,
)
from fastvideo.models.registry import ModelRegistry
from fastvideo.utils import PRECISION_TO_TYPE, is_pin_memory_available
from fastvideo.hooks.layerwise_offload import enable_layerwise_offload
from fastvideo.utils import PRECISION_TO_TYPE
from fastvideo.models.layerwise_offload import LayerwiseOffloadManager
logger = init_logger(__name__)
@@ -80,9 +81,6 @@ class ComponentLoader(ABC):
"transformer": (TransformerLoader, "diffusers"),
"transformer_2": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"audio_vae": (AudioDecoderLoader, "diffusers"),
"audio_decoder": (AudioDecoderLoader, "diffusers"),
"vocoder": (VocoderLoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
"tokenizer": (TokenizerLoader, "transformers"),
@@ -93,12 +91,10 @@ class ComponentLoader(ABC):
if module_type in module_loaders:
loader_cls, expected_library = module_loaders[module_type]
# 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}"
# 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
@@ -245,47 +241,6 @@ class TextEncoderLoader(ComponentLoader):
model_config.pop("model_type", None)
model_config.pop("tokenizer_class", None)
model_config.pop("torch_dtype", None)
repo_root = os.path.dirname(model_path)
index_path = os.path.join(repo_root, "model_index.json")
gemma_path = ""
gemma_path_from_candidate = False
if os.path.isfile(index_path):
try:
with open(index_path, encoding="utf-8") as f:
model_index = json.load(f)
gemma_path = model_index.get("gemma_model_path", "")
except json.JSONDecodeError:
gemma_path = ""
if not gemma_path:
candidate = os.path.normpath(os.path.join(model_path, "gemma"))
if os.path.isdir(candidate):
gemma_path = candidate
gemma_path_from_candidate = True
model_config["gemma_model_path"] = gemma_path
if gemma_path and not gemma_path_from_candidate:
if not os.path.isabs(gemma_path):
model_config["gemma_model_path"] = os.path.normpath(
os.path.join(repo_root, gemma_path)
)
transformer_config_path = os.path.join(
repo_root, "transformer", "config.json"
)
if os.path.isfile(transformer_config_path):
try:
with open(transformer_config_path, encoding="utf-8") as f:
transformer_config = json.load(f)
if (
"connector_double_precision_rope" not in model_config
or not model_config["connector_double_precision_rope"]
):
if transformer_config.get("double_precision_rope") is True:
model_config["connector_double_precision_rope"] = True
if "connector_rope_type" not in model_config:
rope_type = transformer_config.get("rope_type")
if rope_type is not None:
model_config["connector_rope_type"] = rope_type
except json.JSONDecodeError:
pass
logger.info("HF Model config: %s", model_config)
# @TODO(Wei): Better way to handle this?
@@ -324,7 +279,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
@@ -341,10 +296,7 @@ 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"
@@ -361,10 +313,7 @@ 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],
@@ -391,7 +340,6 @@ 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(
@@ -409,7 +357,7 @@ class TextEncoderLoader(ComponentLoader):
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
)
else:
mesh = init_device_mesh(
@@ -423,7 +371,7 @@ class TextEncoderLoader(ComponentLoader):
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
)
# We only enable strict check for non-quantized models
# that have loaded weights tracking currently.
@@ -501,52 +449,13 @@ 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"
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
# other method of config?
padding_size="right",
)
padding_side = None
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
try:
arch_config = fastvideo_args.pipeline_config.text_encoder_configs[
0
].arch_config
padding_side = getattr(arch_config, "padding_side", None)
except Exception:
padding_side = None
if padding_side:
tokenizer.padding_side = padding_side
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
tokenizer.pad_token = tokenizer.eos_token
logger.info("Loaded tokenizer: %s", tokenizer.__class__.__name__)
return tokenizer
@@ -557,12 +466,15 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
class_name = config.get("_class_name")
class_name = config.pop("_class_name")
assert class_name is not None, (
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
)
fastvideo_args.model_paths["vae"] = model_path
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
from fastvideo.platforms import current_platform
if fastvideo_args.vae_cpu_offload:
@@ -579,149 +491,23 @@ 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()
# LTX-2 uses CausalVideoAutoencoder with nested "vae" config
if class_name == "CausalVideoAutoencoder" and "vae" in config:
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(config).to(target_device)
if hasattr(vae, "set_tiling_config"):
vae_config = fastvideo_args.pipeline_config.vae_config
vae.set_tiling_config(
spatial_tile_size_in_pixels=getattr(
vae_config, "ltx2_spatial_tile_size_in_pixels", 512),
spatial_tile_overlap_in_pixels=getattr(
vae_config, "ltx2_spatial_tile_overlap_in_pixels", 64),
temporal_tile_size_in_frames=getattr(
vae_config, "ltx2_temporal_tile_size_in_frames", 64),
temporal_tile_overlap_in_frames=getattr(
vae_config,
"ltx2_temporal_tile_overlap_in_frames", 24),
)
else:
config.pop("_class_name", None)
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
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"))
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.
os.path.join(str(model_path), "*.safetensors")
)
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
# LTX-2 CausalVideoAutoencoder needs per_channel_statistics remapping
if class_name == "CausalVideoAutoencoder" and "vae" in config:
per_channel_prefixes = (
"per_channel_statistics.",
"vae.per_channel_statistics.",
)
remapped = {}
for key, tensor in loaded.items():
remapped[key] = tensor
for prefix in per_channel_prefixes:
if key.startswith(prefix):
suffix = key[len(prefix):]
remapped.setdefault(
f"encoder.per_channel_statistics.{suffix}",
tensor,
)
remapped.setdefault(
f"decoder.per_channel_statistics.{suffix}",
tensor,
)
break
loaded = remapped
vae.load_state_dict(loaded, strict=False)
vae.load_state_dict(
loaded, strict=False
) # We might only load encoder or decoder
return vae.eval()
class AudioDecoderLoader(ComponentLoader):
"""Loader for LTX-2 audio decoder (audio_vae component)."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name", None) or "LTX2AudioDecoder"
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
target_device = get_local_torch_device()
precision = getattr(
fastvideo_args.pipeline_config, "audio_decoder_precision", "bf16"
)
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
audio_decoder = model_cls(config).to(target_device)
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
decoder_state = {}
for name, tensor in loaded.items():
if name.startswith("decoder."):
decoder_state[name.replace("decoder.", "")] = tensor
elif name.startswith("per_channel_statistics."):
decoder_state[name] = tensor
target_module = getattr(audio_decoder, "model", audio_decoder)
target_module.load_state_dict(decoder_state, strict=False)
return audio_decoder.eval()
class VocoderLoader(ComponentLoader):
"""Loader for LTX-2 vocoder."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name", None) or "LTX2Vocoder"
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
target_device = get_local_torch_device()
precision = getattr(
fastvideo_args.pipeline_config, "vocoder_precision", "bf16"
)
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
vocoder = model_cls(config).to(target_device)
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
target_module = getattr(vocoder, "model", vocoder)
target_module.load_state_dict(loaded, strict=False)
return vocoder.eval()
class TransformerLoader(ComponentLoader):
"""Loader for transformer."""
@@ -795,17 +581,7 @@ 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},
@@ -813,7 +589,6 @@ 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,
@@ -836,19 +611,35 @@ class TransformerLoader(ComponentLoader):
model = model.eval()
if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload:
# Check if model has nn.ModuleList for layerwise offload compatibility
has_module_list = any(
isinstance(m, nn.ModuleList) for m in model.children()
)
if has_module_list:
enable_layerwise_offload(model)
else:
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 requested but model %s does not have "
"nn.ModuleList structure. Skipping layerwise offload.",
"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)
return model
+2 -24
View File
@@ -23,7 +23,7 @@ from fastvideo.logger import init_logger
from fastvideo.models.loader.utils import (get_param_names_mapping,
hf_to_custom_state_dict)
from fastvideo.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available
from fastvideo.utils import set_mixed_precision_policy
logger = init_logger(__name__)
@@ -67,7 +67,6 @@ 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,7 +106,6 @@ 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
@@ -143,7 +141,7 @@ def maybe_load_fsdp_model(
weight_iterator,
device,
default_dtype,
strict=strict,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
@@ -295,26 +293,6 @@ 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."
)
+1 -17
View File
@@ -30,10 +30,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -52,10 +50,6 @@ _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"),
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
@@ -69,14 +63,7 @@ _VAE_MODELS = {
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
}
_AUDIO_MODELS = {
"LTX2AudioEncoder": ("audio", "ltx2_audio_vae", "LTX2AudioEncoder"),
"LTX2AudioDecoder": ("audio", "ltx2_audio_vae", "LTX2AudioDecoder"),
"LTX2Vocoder": ("audio", "ltx2_audio_vae", "LTX2Vocoder"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
}
_SCHEDULERS = {
@@ -85,8 +72,6 @@ _SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
"FlowUniPCMultistepScheduler":
("schedulers", "scheduling_flow_unipc_multistep", "FlowUniPCMultistepScheduler"),
"SelfForcingFlowMatchScheduler":
("schedulers", "scheduling_self_forcing_flow_match",
"SelfForcingFlowMatchScheduler"),
@@ -100,7 +85,6 @@ _FAST_VIDEO_MODELS = {
**_TEXT_ENCODER_MODELS,
**_IMAGE_ENCODER_MODELS,
**_VAE_MODELS,
**_AUDIO_MODELS,
**_SCHEDULERS,
}
@@ -109,9 +109,6 @@ 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
@@ -174,8 +171,6 @@ 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).
@@ -191,44 +186,21 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
" you have to pass a value for `mu` when `use_dynamic_shifting` is set to be `True`"
)
# 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 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 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 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]) /
@@ -446,12 +418,8 @@ 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)
# 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))
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = sample.device
@@ -462,8 +430,7 @@ 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(torch.clamp(alpha_si, min=eps)) - torch.log(
torch.clamp(sigma_si, min=eps))
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
@@ -596,11 +563,8 @@ 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)
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))
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = this_sample.device
@@ -611,8 +575,7 @@ 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(torch.clamp(alpha_si, min=eps)) - torch.log(
torch.clamp(sigma_si, min=eps))
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
-735
View File
@@ -1,735 +0,0 @@
#!/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
File diff suppressed because it is too large Load Diff
@@ -1,61 +0,0 @@
# 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
@@ -2,10 +2,5 @@
"""LongCat pipeline module."""
from fastvideo.pipelines.basic.longcat.longcat_pipeline import LongCatPipeline
from fastvideo.pipelines.basic.longcat.longcat_i2v_pipeline import LongCatImageToVideoPipeline
from fastvideo.pipelines.basic.longcat.longcat_vc_pipeline import LongCatVideoContinuationPipeline
__all__ = [
"LongCatPipeline", "LongCatImageToVideoPipeline",
"LongCatVideoContinuationPipeline"
]
__all__ = ["LongCatPipeline"]
@@ -1,148 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat Image-to-Video pipeline implementation.
This module implements I2V (Image-to-Video) generation for LongCat using Tier 3
conditioning with timestep masking, num_cond_latents support, and RoPE skipping.
Supports:
- Basic I2V (50 steps, guidance_scale=4.0)
- Distilled I2V with LoRA (16 steps, guidance_scale=1.0)
- Refinement I2V for 720p upscaling (with refinement LoRA + BSA)
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
DecodingStage,
InputValidationStage,
TextEncodingStage,
TimestepPreparationStage,
)
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
logger = init_logger(__name__)
class LongCatImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
"""
LongCat Image-to-Video pipeline.
Generates video from a single input image using Tier 3 I2V conditioning:
- Per-frame timestep masking (timestep[:, 0] = 0)
- num_cond_latents parameter to transformer
- RoPE skipping for conditioning frames
- Selective denoising (skip first frame in scheduler)
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize LongCat-specific components."""
# Same BSA initialization as base LongCat pipeline
pipeline_config = fastvideo_args.pipeline_config
transformer = self.get_module("transformer", None)
if transformer is None:
return
# Enable BSA if configured
if pipeline_config.enable_bsa:
bsa_params_cfg = getattr(pipeline_config, 'bsa_params', None) or {}
sparsity = getattr(pipeline_config, 'bsa_sparsity', None)
cdf_threshold = getattr(pipeline_config, 'bsa_cdf_threshold', None)
chunk_q = getattr(pipeline_config, 'bsa_chunk_q', None)
chunk_k = getattr(pipeline_config, 'bsa_chunk_k', None)
effective_bsa_params = dict(bsa_params_cfg) if isinstance(
bsa_params_cfg, dict) else {}
if sparsity is not None:
effective_bsa_params['sparsity'] = sparsity
if cdf_threshold is not None:
effective_bsa_params['cdf_threshold'] = cdf_threshold
if chunk_q is not None:
effective_bsa_params['chunk_3d_shape_q'] = chunk_q
if chunk_k is not None:
effective_bsa_params['chunk_3d_shape_k'] = chunk_k
# Provide defaults
effective_bsa_params.setdefault('sparsity', 0.9375)
effective_bsa_params.setdefault('chunk_3d_shape_q', [4, 4, 4])
effective_bsa_params.setdefault('chunk_3d_shape_k', [4, 4, 4])
if hasattr(transformer, 'enable_bsa'):
logger.info("Enabling BSA for LongCat I2V transformer")
transformer.enable_bsa()
if hasattr(transformer, 'blocks'):
try:
for blk in transformer.blocks:
if hasattr(blk, 'self_attn'):
blk.self_attn.bsa_params = effective_bsa_params
except Exception as e:
logger.warning("Failed to set BSA params: %s", e)
logger.info("BSA parameters: %s", effective_bsa_params)
else:
if hasattr(transformer, 'disable_bsa'):
transformer.disable_bsa()
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up I2V-specific pipeline stages."""
# 1. Input validation
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
# 2. Text encoding (same as T2V)
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
# 3. Image VAE encoding (for I2V - skipped in refinement mode)
self.add_stage(
stage_name="image_vae_encoding_stage",
stage=LongCatImageVAEEncodingStage(vae=self.get_module("vae")))
# 4. Refinement initialization (skipped if not refining)
self.add_stage(stage_name="longcat_refine_init_stage",
stage=LongCatRefineInitStage(vae=self.get_module("vae")))
# 5. Timestep preparation (generic)
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
# 6. Refinement timestep override (skipped if not refining)
self.add_stage(stage_name="longcat_refine_timestep_stage",
stage=LongCatRefineTimestepStage(
scheduler=self.get_module("scheduler")))
# 7. Latent preparation with I2V conditioning
self.add_stage(stage_name="latent_preparation_stage",
stage=LongCatI2VLatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
# 8. Denoising with I2V support
self.add_stage(stage_name="denoising_stage",
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))
# 9. Decoding
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae"),
pipeline=self))
EntryClass = LongCatImageToVideoPipeline
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat video diffusion pipeline implementation.
LongCat video diffusion pipeline implementation (Phase 1: Wrapper).
This module implements the LongCat video diffusion pipeline using FastVideo's
modular pipeline architecture.
This module contains a wrapper implementation of the LongCat video diffusion pipeline
using FastVideo's modular pipeline architecture with the original LongCat modules.
"""
from fastvideo.fastvideo_args import FastVideoArgs
@@ -26,6 +26,9 @@ logger = init_logger(__name__)
class LongCatPipeline(LoRAPipeline, ComposedPipelineBase):
"""
LongCat video diffusion pipeline with LoRA support.
Phase 1 implementation using wrapper modules from third_party/longcat_video.
This validates the pipeline infrastructure before full FastVideo integration.
"""
_required_config_modules = [
@@ -1,168 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat Video Continuation (VC) pipeline implementation.
This module implements VC (Video Continuation) generation for LongCat with
KV cache optimization for 2-3x speedup.
Supports:
- Basic VC (50 steps, guidance_scale=4.0)
- Distilled VC with LoRA (16 steps, guidance_scale=1.0)
- KV cache for conditioning frames
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
DecodingStage,
InputValidationStage,
TextEncodingStage,
TimestepPreparationStage,
)
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__)
class LongCatVideoContinuationPipeline(LoRAPipeline, ComposedPipelineBase):
"""
LongCat Video Continuation pipeline.
Generates video continuation from multiple conditioning frames using
optional KV cache for 2-3x speedup.
Key features:
- Takes video input (13+ frames typically)
- Encodes conditioning frames via VAE
- Optionally pre-computes KV cache for conditioning
- Uses cached K/V during denoising for speedup
- Concatenates conditioning back after denoising
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize LongCat-specific components."""
pipeline_config = fastvideo_args.pipeline_config
transformer = self.get_module("transformer", None)
if transformer is None:
return
# Enable BSA if configured (for VC, BSA may not be needed)
if getattr(pipeline_config, 'enable_bsa', False):
bsa_params_cfg = getattr(pipeline_config, 'bsa_params', None) or {}
sparsity = getattr(pipeline_config, 'bsa_sparsity', None)
cdf_threshold = getattr(pipeline_config, 'bsa_cdf_threshold', None)
chunk_q = getattr(pipeline_config, 'bsa_chunk_q', None)
chunk_k = getattr(pipeline_config, 'bsa_chunk_k', None)
effective_bsa_params = dict(bsa_params_cfg) if isinstance(
bsa_params_cfg, dict) else {}
if sparsity is not None:
effective_bsa_params['sparsity'] = sparsity
if cdf_threshold is not None:
effective_bsa_params['cdf_threshold'] = cdf_threshold
if chunk_q is not None:
effective_bsa_params['chunk_3d_shape_q'] = chunk_q
if chunk_k is not None:
effective_bsa_params['chunk_3d_shape_k'] = chunk_k
# Provide defaults
effective_bsa_params.setdefault('sparsity', 0.9375)
effective_bsa_params.setdefault('chunk_3d_shape_q', [4, 4, 4])
effective_bsa_params.setdefault('chunk_3d_shape_k', [4, 4, 4])
if hasattr(transformer, 'enable_bsa'):
logger.info("Enabling BSA for LongCat VC transformer")
transformer.enable_bsa()
if hasattr(transformer, 'blocks'):
try:
for blk in transformer.blocks:
if hasattr(blk, 'self_attn'):
blk.self_attn.bsa_params = effective_bsa_params
except Exception as e:
logger.warning("Failed to set BSA params: %s", e)
logger.info("BSA parameters: %s", effective_bsa_params)
else:
if hasattr(transformer, 'disable_bsa'):
transformer.disable_bsa()
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up VC-specific pipeline stages."""
# 1. Input validation
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
# 2. Text encoding
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
# 3. Video VAE encoding (encodes conditioning frames)
self.add_stage(
stage_name="video_vae_encoding_stage",
stage=LongCatVideoVAEEncodingStage(vae=self.get_module("vae")))
# 4. Timestep preparation
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
# 5. Latent preparation (reuse I2V stage - it handles video_latent too)
self.add_stage(stage_name="latent_preparation_stage",
stage=LongCatVCLatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
# 6. KV cache initialization (optional, based on config)
# This is always added but will skip if use_kv_cache=False
self.add_stage(stage_name="kv_cache_init_stage",
stage=LongCatKVCacheInitStage(
transformer=self.get_module("transformer")))
# 7. Denoising with VC and KV cache support
self.add_stage(stage_name="denoising_stage",
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))
# 8. Decoding
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae"),
pipeline=self))
class LongCatVCLatentPreparationStage(LongCatI2VLatentPreparationStage):
"""
Prepare latents with video conditioning for first N frames.
Extends I2V latent preparation to handle video_latent (multiple frames)
instead of image_latent (single frame).
"""
def forward(self, batch, fastvideo_args):
"""Prepare latents with VC conditioning."""
# Check if we have video_latent (from VC encoding stage)
video_latent = getattr(batch, 'video_latent', None)
if video_latent is not None:
# Set image_latent to video_latent for parent class compatibility
batch.image_latent = video_latent
# Call parent class forward
return super().forward(batch, fastvideo_args)
EntryClass = LongCatVideoContinuationPipeline
@@ -1,150 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 text-to-video pipeline.
"""
import os
from typing import Any
from transformers import AutoTokenizer
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import PipelineComponentLoader
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (DecodingStage, InputValidationStage,
LTX2AudioDecodingStage,
LTX2DenoisingStage,
LTX2LatentPreparationStage,
TextEncodingStage)
logger = init_logger(__name__)
class LTX2Pipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder",
"tokenizer",
"transformer",
"vae",
"audio_vae",
"vocoder",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(
stage_name="input_validation_stage",
stage=InputValidationStage(),
)
self.add_stage(
stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=LTX2LatentPreparationStage(
transformer=self.get_module("transformer"), ),
)
self.add_stage(
stage_name="denoising_stage",
stage=LTX2DenoisingStage(
transformer=self.get_module("transformer"), ),
)
self.add_stage(
stage_name="audio_decoding_stage",
stage=LTX2AudioDecodingStage(
audio_decoder=self.get_module("audio_vae"),
vocoder=self.get_module("vocoder"),
),
)
self.add_stage(
stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")),
)
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
tokenizer = self.get_module("tokenizer")
if tokenizer is not None:
tokenizer.padding_side = "left"
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
def load_modules(
self,
fastvideo_args: FastVideoArgs,
loaded_modules: dict[str, Any] | None = None,
) -> dict[str, Any]:
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
model_index.pop("workload_type", None)
if len(model_index) <= 1:
raise ValueError(
"model_index.json must contain at least one pipeline module")
required_modules = self.required_config_modules
modules: dict[str, Any] = {}
for module_name, module_spec in model_index.items():
if not isinstance(module_spec, list) or len(module_spec) < 1:
continue
transformers_or_diffusers = module_spec[0]
if transformers_or_diffusers is None:
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
if module_name not in required_modules:
continue
if loaded_modules is not None and module_name in loaded_modules:
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
if module_name == "tokenizer" and not os.path.isdir(
component_model_path):
gemma_path = os.path.join(self.model_path, "text_encoder",
"gemma")
if os.path.isdir(gemma_path):
component_model_path = gemma_path
else:
raise ValueError(
"Tokenizer directory missing and Gemma weights were not found."
)
module = PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=fastvideo_args,
)
logger.info("Loaded module %s from %s", module_name,
component_model_path)
modules[module_name] = module
if "tokenizer" in required_modules and "tokenizer" not in modules:
gemma_path = os.path.join(self.model_path, "text_encoder", "gemma")
if os.path.isdir(gemma_path):
modules["tokenizer"] = AutoTokenizer.from_pretrained(
gemma_path, local_files_only=True)
for module_name in required_modules:
if module_name not in modules or modules[module_name] is None:
raise ValueError(
f"Required module {module_name} was not loaded properly")
return modules
EntryClass = LTX2Pipeline
@@ -1,6 +0,0 @@
from fastvideo.pipelines.basic.turbodiffusion.turbodiffusion_pipeline import (
TurboDiffusionPipeline, )
from fastvideo.pipelines.basic.turbodiffusion.turbodiffusion_i2v_pipeline import (
TurboDiffusionI2VPipeline, )
__all__ = ["TurboDiffusionPipeline", "TurboDiffusionI2VPipeline"]
@@ -1,89 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
TurboDiffusion I2V (Image-to-Video) Pipeline Implementation.
This module contains an implementation of the TurboDiffusion I2V pipeline
for 1-4 step image-to-video generation using rCM (recurrent Consistency Model)
sampling with SLA (Sparse-Linear Attention).
Key differences from T2V:
- Uses dual models (high/low noise) with boundary switching
- sigma_max=200 (vs 80 for T2V)
- Mask conditioning with encoded first frame
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_rcm import RCMScheduler
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, ImageVAEEncodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class TurboDiffusionI2VPipeline(LoRAPipeline, ComposedPipelineBase):
"""
TurboDiffusion I2V pipeline for 1-4 step image-to-video generation.
Uses RCM scheduler, SLA attention, and dual model switching for
high-quality I2V generation.
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "transformer_2",
"scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Use RCM scheduler with higher sigma_max for I2V
logger.info(
"Initializing RCM scheduler for TurboDiffusion I2V (sigma_max=200)")
self.modules["scheduler"] = RCMScheduler(sigma_max=200.0)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
# I2V: Encode initial image to latent space
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
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))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae"),
pipeline=self))
EntryClass = TurboDiffusionI2VPipeline
@@ -36,6 +36,11 @@ class TurboDiffusionPipeline(LoRAPipeline, ComposedPipelineBase):
logger.info("Initializing RCM scheduler for TurboDiffusion")
self.modules["scheduler"] = RCMScheduler(sigma_max=80.0)
# Store checkpoint path for later loading
self._turbodiffusion_checkpoint = getattr(fastvideo_args,
'turbodiffusion_checkpoint',
None)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
+92 -273
View File
@@ -1,9 +1,7 @@
# 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
@@ -14,13 +12,8 @@ from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
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.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
@@ -29,89 +22,17 @@ 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 = ""
# model_name -> layers
lora_layers: dict[str, LoRAModelLayers] = {}
lora_layers: dict[str, dict[str, BaseLayerWithLoRA]] = {}
fastvideo_args: FastVideoArgs | TrainingArgs
exclude_lora_layers: dict[str, list[str]] = {}
device: torch.device = get_local_torch_device()
@@ -127,10 +48,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(
@@ -138,23 +59,19 @@ 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
@@ -166,33 +83,20 @@ 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
@@ -200,8 +104,7 @@ 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:
@@ -211,9 +114,9 @@ class LoRAPipeline(ComposedPipelineBase):
def set_trainable(self) -> None:
def set_lora_grads(lora_layers: LoRAModelLayers,
def set_lora_grads(lora_layers: dict[str, BaseLayerWithLoRA],
device_mesh: DeviceMesh):
for name, layer in lora_layers.all_lora_layers():
for name, layer in lora_layers.items():
layer.lora_A.requires_grad_(True)
layer.lora_B.requires_grad_(True)
layer.base_layer.requires_grad_(False)
@@ -228,15 +131,10 @@ 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:
@@ -253,71 +151,32 @@ 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:
# 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
self.lora_layers[transformer_name] = LoRAModelLayers(block_list)
self.lora_layers[transformer_name] = {}
logger.info("Converting %s to LoRA Transformer", transformer_name)
# scan every module and convert to LoRA layer if applicable
for name, layer in transformer_module.named_modules():
if not self.is_target_layer(name):
continue
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
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
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
logger.info("Converted %d layers to LoRA layers", converted_count)
def set_lora_adapter(self,
@@ -350,7 +209,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.", "")
@@ -364,14 +223,13 @@ 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:
@@ -403,84 +261,45 @@ class LoRAPipeline(ComposedPipelineBase):
# Merge the new adapter
adapted_count = 0
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,
)
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)
def merge_lora_weights(self) -> None:
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()
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
layer.merge_lora_weights()
def unmerge_lora_weights(self) -> None:
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()
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
layer.unmerge_lora_weights()
-5
View File
@@ -24,18 +24,13 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanVideoToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"TurboDiffusionPipeline": "turbodiffusion",
"TurboDiffusionI2VPipeline": "turbodiffusion",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"Cosmos2_5Pipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
"MatrixGameCausalDMDPipeline": "matrixgame",
"LongCatPipeline": "longcat",
"LongCatImageToVideoPipeline": "longcat",
"LongCatVideoContinuationPipeline": "longcat",
"LTX2Pipeline": "ltx2",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
+4 -27
View File
@@ -10,8 +10,7 @@ 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 (Cosmos25DenoisingStage,
CosmosDenoisingStage,
from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
@@ -20,44 +19,27 @@ from fastvideo.pipelines.stages.image_encoding import (
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import (
Cosmos25LatentPreparationStage, CosmosLatentPreparationStage,
LatentPreparationStage)
from fastvideo.pipelines.stages.ltx2_audio_decoding import LTX2AudioDecodingStage
from fastvideo.pipelines.stages.ltx2_denoising import LTX2DenoisingStage
from fastvideo.pipelines.stages.ltx2_latent_preparation import (
LTX2LatentPreparationStage)
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 (Cosmos25TextEncodingStage,
TextEncodingStage)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.timestep_preparation import (
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
TimestepPreparationStage)
__all__ = [
"PipelineStage",
"InputValidationStage",
"TimestepPreparationStage",
"Cosmos25TimestepPreparationStage",
"LatentPreparationStage",
"CosmosLatentPreparationStage",
"Cosmos25LatentPreparationStage",
"LTX2LatentPreparationStage",
"LTX2AudioDecodingStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"MatrixGameCausalDenoisingStage",
"CosmosDenoisingStage",
"Cosmos25DenoisingStage",
"LTX2DenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
@@ -67,10 +49,5 @@ __all__ = [
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
"TextEncodingStage",
"Cosmos25TextEncodingStage",
"StepvideoPromptEncodingStage",
# LongCat stages
"LongCatVideoVAEEncodingStage",
"LongCatKVCacheInitStage",
"LongCatVCDenoisingStage",
]

Some files were not shown because too many files have changed in this diff Show More