Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cec2bc5ff6 | ||
|
|
b7f69c2c1d | ||
|
|
23a4531491 | ||
|
|
7d52ad0118 | ||
|
|
4d7bf35fa3 | ||
|
|
a6a9c9ca07 | ||
|
|
f4704847c2 | ||
|
|
d9c996310b | ||
|
|
d6651afd2e | ||
|
|
cf67618cad | ||
|
|
2f0a2b3c57 | ||
|
|
e7748d9952 | ||
|
|
8eb3140b2f | ||
|
|
d6ddcea682 | ||
|
|
3559ba2377 | ||
|
|
61e63ea0d7 |
@@ -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/c7g1qdD" 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/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 490 KiB |
Binary file not shown.
@@ -41,7 +41,7 @@ Clone the repository and build the kernel:
|
||||
|
||||
```bash
|
||||
# Clone recursively to get ThunderKittens submodule
|
||||
git clone --recursive https://github.com/hao-ai-lab/FastVideo.git
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo/fastvideo-kernel
|
||||
|
||||
# Build and install
|
||||
|
||||
@@ -11,15 +11,13 @@ 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
|
||||
@@ -28,6 +26,12 @@ 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+
|
||||
|
||||
@@ -15,6 +15,12 @@ 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
|
||||
|
||||
@@ -106,6 +106,8 @@ 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
|
||||
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# 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
|
||||
@@ -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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
|
||||
@@ -14,7 +14,7 @@ def main():
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
# 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"
|
||||
|
||||
@@ -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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
|
||||
@@ -0,0 +1,210 @@
|
||||
"""
|
||||
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()
|
||||
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
"""
|
||||
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()
|
||||
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
"""
|
||||
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()
|
||||
|
||||
|
||||
@@ -43,7 +43,7 @@ def main():
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
|
||||
@@ -15,12 +15,14 @@ def main() -> None:
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
# TurboDiffusion uses a custom pipeline with RCM scheduler
|
||||
override_pipeline_cls_name="TurboDiffusionPipeline",
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
|
||||
# set to false if using RTX 4090
|
||||
# pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
# TurboDiffusion uses guidance_scale=1.0 (no CFG) and only 4 steps
|
||||
# TurboDiffusion defaults: guidance_scale=1.0 and num_inference_steps=4 (from config)
|
||||
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 "
|
||||
@@ -30,9 +32,7 @@ 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,9 +47,7 @@ def main() -> None:
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_inference_steps=4,
|
||||
seed=42,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -15,8 +15,6 @@ 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 = (
|
||||
@@ -28,9 +26,7 @@ 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!
|
||||
@@ -45,9 +41,7 @@ def main() -> None:
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_inference_steps=4,
|
||||
seed=42,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
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()
|
||||
@@ -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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
|
||||
@@ -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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
|
||||
@@ -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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
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=True,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
|
||||
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.2.1"
|
||||
version = "0.2.2"
|
||||
description = "Unified CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.2.1"
|
||||
__version__ = "0.2.2"
|
||||
|
||||
@@ -26,7 +26,6 @@ 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
|
||||
|
||||
@@ -17,11 +17,7 @@ from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""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.
|
||||
"""
|
||||
"""Extended DiTArchConfig with LongCat-specific fields."""
|
||||
# LongCat-specific architecture parameters
|
||||
adaln_tembed_dim: int = 512
|
||||
caption_channels: int = 4096
|
||||
@@ -88,20 +84,16 @@ def umt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
|
||||
@dataclass
|
||||
class LongCatT2V480PConfig(PipelineConfig):
|
||||
"""Configuration for LongCat pipeline (480p) aligned to LongCat-Video modules.
|
||||
"""Configuration for LongCat pipeline (480p).
|
||||
|
||||
Components expected by loaders:
|
||||
- tokenizer: AutoTokenizer
|
||||
- text_encoder: UMT5EncoderModel
|
||||
- transformer: LongCatVideoTransformer3DModel (Phase 1 wrapper)
|
||||
OR LongCatTransformer3DModel (Phase 2 native)
|
||||
- transformer: LongCatTransformer3DModel
|
||||
- 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()))
|
||||
|
||||
|
||||
@@ -10,6 +10,9 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
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 (
|
||||
@@ -55,11 +58,25 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"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,
|
||||
# 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":
|
||||
@@ -78,13 +95,15 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
lambda id: "stepvideo" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
"turbodiffusion":
|
||||
lambda id: "turbodiffusion" in id.lower() or "turbowan" 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,
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
@@ -96,7 +115,8 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
"stepvideo": StepVideoT2VConfig,
|
||||
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
# 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
|
||||
@@ -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",
|
||||
|
||||
@@ -26,6 +26,11 @@ 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,
|
||||
@@ -34,36 +39,48 @@ from fastvideo.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"FastVideo/FastHunyuan-diffusers":
|
||||
FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo":
|
||||
HunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
|
||||
Hunyuan15_480P_SamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
|
||||
Hunyuan15_720P_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers":
|
||||
StepVideoT2VSamplingParam,
|
||||
|
||||
# Wan2.1
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
|
||||
WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
|
||||
# Wan2.2
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
|
||||
# FastWan2.1
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480P_SamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
FastWanT2V480P_SamplingParam,
|
||||
|
||||
# FastWan2.2
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
|
||||
@@ -80,9 +97,20 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
Cosmos_Predict2_2B_Video2World_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,
|
||||
"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,
|
||||
|
||||
# Add other specific weight variants
|
||||
}
|
||||
@@ -105,6 +133,8 @@ 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(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -121,6 +151,8 @@ 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
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# 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
|
||||
@@ -132,7 +132,7 @@ class FastVideoArgs:
|
||||
|
||||
# CPU offload parameters
|
||||
dit_cpu_offload: bool = True
|
||||
use_fsdp_inference: bool = True
|
||||
use_fsdp_inference: bool = False
|
||||
dit_layerwise_offload: bool = False
|
||||
text_encoder_cpu_offload: bool = True
|
||||
image_encoder_cpu_offload: bool = True
|
||||
@@ -431,7 +431,9 @@ class FastVideoArgs:
|
||||
"--use-fsdp-inference",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
|
||||
"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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-cpu-offload",
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
# 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
|
||||
@@ -129,7 +126,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
|
||||
# Cast to model dtype before MLP (matching original LongCat)
|
||||
# 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
|
||||
@@ -166,13 +163,14 @@ 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.SiLU()
|
||||
self.act = nn.GELU(approximate="tanh") # Match original LongCat
|
||||
self.linear_2 = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
@@ -268,10 +266,19 @@ 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:
|
||||
) -> torch.Tensor | tuple:
|
||||
"""
|
||||
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
|
||||
@@ -290,6 +297,12 @@ 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)
|
||||
@@ -301,6 +314,50 @@ 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
|
||||
@@ -348,6 +405,96 @@ 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
|
||||
|
||||
|
||||
@@ -394,6 +541,8 @@ 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:
|
||||
"""
|
||||
@@ -402,9 +551,57 @@ 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)
|
||||
@@ -475,19 +672,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 (matching original LongCat).
|
||||
Apply modulation in FP32 for numerical stability.
|
||||
|
||||
shift and scale should already be FP32 from torch.amp.autocast context.
|
||||
Converts inputs to FP32 for the modulation operation, then casts back.
|
||||
"""
|
||||
# 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}"
|
||||
|
||||
orig_dtype = x.dtype
|
||||
|
||||
# Convert to FP32 for numerical stability
|
||||
shift_fp32 = shift.float()
|
||||
scale_fp32 = scale.float()
|
||||
|
||||
# Normalize and modulate in FP32
|
||||
x_norm = norm(x.to(torch.float32))
|
||||
x_mod = x_norm * (scale + 1) + shift
|
||||
x_mod = x_norm * (scale_fp32 + 1) + shift_fp32
|
||||
|
||||
return x_mod.to(orig_dtype)
|
||||
|
||||
@@ -568,10 +765,21 @@ 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:
|
||||
) -> torch.Tensor | tuple:
|
||||
"""
|
||||
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
|
||||
@@ -592,17 +800,47 @@ 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)
|
||||
|
||||
attn_out = self.self_attn(x_norm, latent_shape=latent_shape)
|
||||
# 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
|
||||
|
||||
# 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 ===
|
||||
x_norm_cross = self.norm_cross(x)
|
||||
cross_out = self.cross_attn(x_norm_cross, context)
|
||||
x = x + cross_out
|
||||
# === 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
|
||||
|
||||
# === FFN ===
|
||||
x_norm_ffn = modulate_fp32(self.norm_ffn, x.view(B, T, -1, C), shift_mlp, scale_mlp)
|
||||
@@ -615,6 +853,8 @@ 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
|
||||
|
||||
|
||||
@@ -670,16 +910,12 @@ class FinalLayer(nn.Module):
|
||||
B, N, C = x.shape
|
||||
T, _, _ = latent_shape
|
||||
|
||||
# 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)
|
||||
# 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)
|
||||
|
||||
# Modulate
|
||||
# Modulate (converts to FP32 internally for stability)
|
||||
x = modulate_fp32(self.norm, x.view(B, T, -1, C), shift, scale)
|
||||
x = x.reshape(B, N, C)
|
||||
|
||||
@@ -696,8 +932,6 @@ 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
|
||||
@@ -789,13 +1023,28 @@ 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:
|
||||
) -> torch.Tensor | tuple[torch.Tensor, dict]:
|
||||
"""
|
||||
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
|
||||
|
||||
@@ -825,12 +1074,31 @@ class LongCatTransformer3DModel(CachableDiT):
|
||||
encoder_attention_mask=encoder_attention_mask
|
||||
) # [B, N_text, C]
|
||||
|
||||
# 4. Transformer blocks
|
||||
# 4. Transformer blocks with optional KV cache
|
||||
kv_cache_dict_ret = {} if return_kv else None
|
||||
|
||||
for i, block in enumerate(self.blocks):
|
||||
x = block(
|
||||
# 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, context, t,
|
||||
latent_shape=(N_t, N_h, N_w)
|
||||
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,
|
||||
)
|
||||
|
||||
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))
|
||||
@@ -841,6 +1109,8 @@ 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:
|
||||
|
||||
@@ -30,8 +30,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"),
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
|
||||
@@ -2,5 +2,10 @@
|
||||
"""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"]
|
||||
__all__ = [
|
||||
"LongCatPipeline", "LongCatImageToVideoPipeline",
|
||||
"LongCatVideoContinuationPipeline"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
# 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 (Phase 1: Wrapper).
|
||||
LongCat video diffusion pipeline implementation.
|
||||
|
||||
This module contains a wrapper implementation of the LongCat video diffusion pipeline
|
||||
using FastVideo's modular pipeline architecture with the original LongCat modules.
|
||||
This module implements the LongCat video diffusion pipeline using FastVideo's
|
||||
modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -26,9 +26,6 @@ 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 = [
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
# 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
|
||||
@@ -0,0 +1,6 @@
|
||||
from fastvideo.pipelines.basic.turbodiffusion.turbodiffusion_pipeline import (
|
||||
TurboDiffusionPipeline, )
|
||||
from fastvideo.pipelines.basic.turbodiffusion.turbodiffusion_i2v_pipeline import (
|
||||
TurboDiffusionI2VPipeline, )
|
||||
|
||||
__all__ = ["TurboDiffusionPipeline", "TurboDiffusionI2VPipeline"]
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
# 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,11 +36,6 @@ 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."""
|
||||
|
||||
|
||||
@@ -24,6 +24,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanVideoToVideoPipeline": "wan",
|
||||
"WanCausalDMDPipeline": "wan",
|
||||
"TurboDiffusionPipeline": "turbodiffusion",
|
||||
"TurboDiffusionI2VPipeline": "turbodiffusion",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
"HunyuanVideo15Pipeline": "hunyuan15",
|
||||
@@ -31,6 +32,8 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"MatrixGamePipeline": "matrixgame",
|
||||
"MatrixGameCausalDMDPipeline": "matrixgame",
|
||||
"LongCatPipeline": "longcat",
|
||||
"LongCatImageToVideoPipeline": "longcat",
|
||||
"LongCatVideoContinuationPipeline": "longcat",
|
||||
}
|
||||
|
||||
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
|
||||
|
||||
@@ -28,6 +28,11 @@ from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
from fastvideo.pipelines.stages.timestep_preparation import (
|
||||
TimestepPreparationStage)
|
||||
|
||||
# LongCat stages
|
||||
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
|
||||
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
|
||||
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
|
||||
|
||||
__all__ = [
|
||||
"PipelineStage",
|
||||
"InputValidationStage",
|
||||
@@ -50,4 +55,8 @@ __all__ = [
|
||||
"VideoVAEEncodingStage",
|
||||
"TextEncodingStage",
|
||||
"StepvideoPromptEncodingStage",
|
||||
# LongCat stages
|
||||
"LongCatVideoVAEEncodingStage",
|
||||
"LongCatKVCacheInitStage",
|
||||
"LongCatVCDenoisingStage",
|
||||
]
|
||||
|
||||
@@ -64,7 +64,7 @@ class LongCatDenoisingStage(DenoisingStage):
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
from fastvideo.models.model_loader import TransformerLoader
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat I2V Denoising Stage with conditioning support.
|
||||
|
||||
This stage implements Tier 3 I2V denoising:
|
||||
1. Per-frame timestep masking (timestep[:, :num_cond_latents] = 0)
|
||||
2. Passes num_cond_latents to transformer (for RoPE skipping)
|
||||
3. Selective denoising (only updates non-conditioned frames)
|
||||
4. CFG-zero optimized guidance
|
||||
"""
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatI2VDenoisingStage(LongCatDenoisingStage):
|
||||
"""
|
||||
LongCat denoising with I2V conditioning support.
|
||||
|
||||
Key modifications from base LongCat denoising:
|
||||
1. Sets timestep=0 for conditioning frames
|
||||
2. Passes num_cond_latents to transformer
|
||||
3. Only applies scheduler step to non-conditioned frames
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run denoising loop with I2V conditioning."""
|
||||
|
||||
# Load transformer if needed
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Setup
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get num_cond_latents from batch
|
||||
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
|
||||
|
||||
if num_cond_latents > 0:
|
||||
logger.info("I2V Denoising: num_cond_latents=%s, latent_shape=%s",
|
||||
num_cond_latents, latents.shape)
|
||||
|
||||
# Prepare negative prompts for CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="I2V Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
|
||||
# 1. Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# 2. Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# 3. CRITICAL: Expand timestep to temporal dimension
|
||||
# and set conditioning frames to timestep=0
|
||||
timestep = timestep.unsqueeze(-1).repeat(
|
||||
1, latent_model_input.shape[2])
|
||||
|
||||
# Mark conditioning frames as clean (timestep=0)
|
||||
if num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
# 4. Run transformer with num_cond_latents
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
num_cond_latents=num_cond_latents,
|
||||
)
|
||||
|
||||
# 5. Apply CFG with optimized scale
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
# CFG-zero optimized scale
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# 6. CRITICAL: Negate for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# 7. CRITICAL: Only update non-conditioned frames
|
||||
# The conditioning frames stay FIXED throughout denoising
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
# No conditioning, update all frames
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,105 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat I2V Latent Preparation Stage.
|
||||
|
||||
This stage prepares latents with image conditioning for the first frame.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatI2VLatentPreparationStage(LatentPreparationStage):
|
||||
"""
|
||||
Prepare latents with image conditioning for first frame.
|
||||
|
||||
This stage:
|
||||
1. Generates random noise for all frames
|
||||
2. Replaces first latent frame with encoded image latent
|
||||
3. Marks conditioning information in batch
|
||||
"""
|
||||
|
||||
# Uses parent __init__ - no need for additional constructor
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Prepare latents with I2V conditioning."""
|
||||
|
||||
# IMPORTANT: Skip if latents already prepared (e.g., by refinement init stage)
|
||||
# The refine_init stage encodes stage1 video and mixes with noise - don't overwrite!
|
||||
if batch.latents is not None:
|
||||
logger.info(
|
||||
"I2V Latent Prep: Skipping - latents already prepared "
|
||||
"(shape=%s), likely from refinement stage", batch.latents.shape)
|
||||
return batch
|
||||
|
||||
# 1. Calculate dimensions
|
||||
num_frames = batch.num_frames
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
# Get VAE compression factors
|
||||
# IMPORTANT: Use VAE's temporal compression (4), NOT transformer's patch_size[0] (1)
|
||||
vae_temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
|
||||
vae_spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
|
||||
num_latent_frames = (num_frames - 1) // vae_temporal_scale + 1
|
||||
latent_height = height // vae_spatial_scale
|
||||
latent_width = width // vae_spatial_scale
|
||||
|
||||
num_channels = self.transformer.config.in_channels
|
||||
|
||||
logger.info(
|
||||
"I2V Latent Prep: num_frames=%s, num_latent_frames=%s "
|
||||
"(vae_temporal_scale=%s), latent_shape=(%s, %s)", num_frames,
|
||||
num_latent_frames, vae_temporal_scale, latent_height, latent_width)
|
||||
|
||||
# 2. Generate random noise for all frames
|
||||
# batch_size might not be set, default to 1
|
||||
batch_size = batch.batch_size if batch.batch_size is not None else 1
|
||||
shape = (batch_size, num_channels, num_latent_frames, latent_height,
|
||||
latent_width)
|
||||
|
||||
# Handle generator - may be a list for batch handling
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list):
|
||||
generator = generator[0] if generator else None
|
||||
|
||||
# torch.randn requires specific argument order: size, generator, dtype
|
||||
latents = torch.randn(*shape,
|
||||
generator=generator).to(get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Replace first frame with conditioned image latent
|
||||
if batch.image_latent is not None:
|
||||
num_cond_latents = batch.num_cond_latents
|
||||
latents[:, :, :
|
||||
num_cond_latents] = batch.image_latent[:, :, :
|
||||
num_cond_latents]
|
||||
|
||||
logger.info(
|
||||
"I2V: Replaced first %s latent frame(s) with image conditioning",
|
||||
num_cond_latents)
|
||||
else:
|
||||
logger.warning(
|
||||
"No image_latent found in batch, proceeding without conditioning"
|
||||
)
|
||||
|
||||
# 4. Store in batch
|
||||
batch.latents = latents
|
||||
|
||||
# Required by base class output validator
|
||||
batch.raw_latent_shape = (num_latent_frames, latent_height,
|
||||
latent_width)
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,162 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Image VAE Encoding Stage for I2V generation.
|
||||
|
||||
This stage handles encoding a single input image to latent space with
|
||||
LongCat-specific normalization for I2V conditioning.
|
||||
"""
|
||||
|
||||
import PIL
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vision_utils import (normalize, numpy_to_pt, pil_to_numpy,
|
||||
resize)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatImageVAEEncodingStage(PipelineStage):
|
||||
"""
|
||||
Encode input image to latent space for I2V conditioning.
|
||||
|
||||
This stage:
|
||||
1. Preprocesses image to match target dimensions
|
||||
2. Encodes via VAE to latent space
|
||||
3. Applies LongCat-specific normalization
|
||||
4. Stores latent and calculates num_cond_latents
|
||||
"""
|
||||
|
||||
def __init__(self, vae):
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Encode image to latent for I2V conditioning."""
|
||||
|
||||
# Skip image encoding for refinement tasks - we're refining an existing video
|
||||
if getattr(batch, 'stage1_video', None) is not None or getattr(
|
||||
batch, 'refine_from', None) is not None:
|
||||
logger.info(
|
||||
"Skipping image encoding - refinement mode (using stage1_video)"
|
||||
)
|
||||
return batch
|
||||
|
||||
# 1. Get image from batch
|
||||
image = batch.pil_image # PIL.Image
|
||||
if image is None:
|
||||
raise ValueError("pil_image must be provided for I2V")
|
||||
|
||||
if not isinstance(image, PIL.Image.Image):
|
||||
raise TypeError(f"pil_image must be PIL.Image, got {type(image)}")
|
||||
|
||||
# 2. Get target dimensions
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
if height is None or width is None:
|
||||
raise ValueError("height and width must be set for I2V")
|
||||
|
||||
# 3. Preprocess image
|
||||
image = resize(image, height, width, resize_mode="default")
|
||||
image = pil_to_numpy(image)
|
||||
image = numpy_to_pt(image)
|
||||
image = normalize(image) # to [-1, 1]
|
||||
|
||||
# 4. Add temporal dimension
|
||||
# After numpy_to_pt: [1, C, H, W] (batch already added by pil_to_numpy)
|
||||
# Add T dimension: [1, C, H, W] -> [1, C, 1, H, W] = [B, C, T, H, W]
|
||||
image = image.unsqueeze(2)
|
||||
image = image.to(get_local_torch_device(), dtype=torch.float32)
|
||||
|
||||
# 5. Encode via VAE
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
|
||||
if not vae_autocast_enabled:
|
||||
image = image.to(vae_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
encoder_output = self.vae.encode(image)
|
||||
latent = self.retrieve_latents(encoder_output, batch.generator)
|
||||
|
||||
# 6. Apply LongCat-specific normalization
|
||||
# Formula: (latents - mean) / std
|
||||
latent = self.normalize_latents(latent)
|
||||
|
||||
# 7. Calculate num_cond_latents
|
||||
# Formula: 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
# For single image (num_cond_frames=1): always 1 latent frame
|
||||
num_cond_frames = 1 # Single image
|
||||
vae_temporal_scale = self.vae.config.scale_factor_temporal
|
||||
batch.num_cond_latents = 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
|
||||
# 8. Store in batch
|
||||
batch.image_latent = latent
|
||||
batch.num_cond_frames = 1
|
||||
|
||||
logger.info(
|
||||
"I2V: Encoded image to latent shape %s, num_cond_latents=%s",
|
||||
latent.shape, batch.num_cond_latents)
|
||||
|
||||
# Offload VAE if needed
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self, encoder_output: object,
|
||||
generator: torch.Generator | None) -> torch.Tensor:
|
||||
"""Sample from VAE posterior."""
|
||||
# WAN VAE returns an object with .sample() method
|
||||
if hasattr(encoder_output, 'sample'):
|
||||
return encoder_output.sample(generator)
|
||||
elif hasattr(encoder_output, 'latent_dist'):
|
||||
return encoder_output.latent_dist.sample(generator)
|
||||
elif hasattr(encoder_output, 'latents'):
|
||||
return encoder_output.latents
|
||||
else:
|
||||
raise AttributeError("Could not access latents from encoder output")
|
||||
|
||||
def normalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply LongCat-specific latent normalization.
|
||||
|
||||
Formula: (latents - mean) / std
|
||||
|
||||
This matches the original LongCat implementation and is DIFFERENT
|
||||
from standard VAE scaling (which uses scaling_factor).
|
||||
"""
|
||||
if not hasattr(self.vae.config, 'latents_mean') or not hasattr(
|
||||
self.vae.config, 'latents_std'):
|
||||
raise ValueError(
|
||||
"VAE config must have 'latents_mean' and 'latents_std' "
|
||||
"for LongCat normalization")
|
||||
|
||||
latents_mean = torch.tensor(self.vae.config.latents_mean).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
latents_std = torch.tensor(self.vae.config.latents_std).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
return (latents - latents_mean) / latents_std
|
||||
@@ -0,0 +1,123 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat KV Cache Initialization Stage for Video Continuation (VC).
|
||||
|
||||
This stage pre-computes K/V cache for conditioning frames.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatKVCacheInitStage(PipelineStage):
|
||||
"""
|
||||
Pre-compute KV cache for conditioning frames.
|
||||
|
||||
After this stage:
|
||||
- batch.kv_cache_dict contains {block_idx: (k, v)}
|
||||
- batch.cond_latents contains the conditioning latents
|
||||
- batch.latents contains ONLY noise latents
|
||||
"""
|
||||
|
||||
def __init__(self, transformer):
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Initialize KV cache from conditioning latents."""
|
||||
|
||||
# Check if KV cache is enabled
|
||||
use_kv_cache = getattr(fastvideo_args.pipeline_config, 'use_kv_cache',
|
||||
True)
|
||||
if not use_kv_cache:
|
||||
batch.kv_cache_dict = {}
|
||||
batch.use_kv_cache = False
|
||||
logger.info("KV cache disabled, skipping initialization")
|
||||
return batch
|
||||
|
||||
batch.use_kv_cache = True
|
||||
offload_kv_cache = getattr(fastvideo_args.pipeline_config,
|
||||
'offload_kv_cache', False)
|
||||
|
||||
# Get conditioning latents
|
||||
num_cond_latents = batch.num_cond_latents
|
||||
if num_cond_latents <= 0:
|
||||
batch.kv_cache_dict = {}
|
||||
logger.warning("num_cond_latents <= 0, skipping KV cache init")
|
||||
return batch
|
||||
|
||||
# Extract conditioning latents
|
||||
cond_latents = batch.latents[:, :, :num_cond_latents].clone()
|
||||
|
||||
logger.info(
|
||||
"Initializing KV cache for %d conditioning latents, shape: %s",
|
||||
num_cond_latents, cond_latents.shape)
|
||||
|
||||
# Timestep = 0 for conditioning (they are "clean")
|
||||
B = cond_latents.shape[0]
|
||||
T_cond = cond_latents.shape[2]
|
||||
timestep = torch.zeros(B,
|
||||
T_cond,
|
||||
device=cond_latents.device,
|
||||
dtype=cond_latents.dtype)
|
||||
|
||||
# Empty prompt embeddings (cross-attn will be skipped)
|
||||
max_seq_len = 512
|
||||
# Get caption dimension from transformer config
|
||||
caption_dim = self.transformer.config.caption_channels
|
||||
empty_embeds = torch.zeros(B,
|
||||
max_seq_len,
|
||||
caption_dim,
|
||||
device=cond_latents.device,
|
||||
dtype=cond_latents.dtype)
|
||||
|
||||
# Get transformer dtype
|
||||
if hasattr(self.transformer, 'module'):
|
||||
transformer_dtype = next(self.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
|
||||
# Run transformer with return_kv=True, skip_crs_attn=True
|
||||
with (
|
||||
torch.no_grad(),
|
||||
set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
),
|
||||
torch.autocast(device_type='cuda', dtype=transformer_dtype),
|
||||
):
|
||||
_, kv_cache_dict = self.transformer(
|
||||
hidden_states=cond_latents.to(transformer_dtype),
|
||||
encoder_hidden_states=empty_embeds.to(transformer_dtype),
|
||||
timestep=timestep.to(transformer_dtype),
|
||||
return_kv=True,
|
||||
skip_crs_attn=True,
|
||||
offload_kv_cache=offload_kv_cache,
|
||||
)
|
||||
|
||||
# Store cache and save cond_latents for later concatenation
|
||||
batch.kv_cache_dict = kv_cache_dict
|
||||
batch.cond_latents = cond_latents
|
||||
|
||||
# Remove conditioning latents from main latents
|
||||
# After this, batch.latents contains ONLY noise frames
|
||||
batch.latents = batch.latents[:, :, num_cond_latents:]
|
||||
|
||||
logger.info(
|
||||
"KV cache initialized: %d blocks, offload=%s, remaining latents shape: %s",
|
||||
len(kv_cache_dict), offload_kv_cache, batch.latents.shape)
|
||||
|
||||
return batch
|
||||
@@ -257,12 +257,16 @@ class LongCatRefineInitStage(PipelineStage):
|
||||
num_cond_frames_added, num_noise_frames_added,
|
||||
new_num_frames)
|
||||
|
||||
# VAE encode
|
||||
logger.info("Encoding stage1 video with VAE...")
|
||||
# VAE encode with tiling for memory efficiency
|
||||
logger.info("Encoding stage1 video with VAE (tiling enabled)...")
|
||||
vae_dtype = next(self.vae.parameters()).dtype
|
||||
vae_device = next(self.vae.parameters()).device
|
||||
video_up = video_up.to(dtype=vae_dtype, device=vae_device)
|
||||
|
||||
# Enable tiling for large video encoding
|
||||
if hasattr(self.vae, 'enable_tiling'):
|
||||
self.vae.enable_tiling()
|
||||
|
||||
with torch.no_grad():
|
||||
latent_dist = self.vae.encode(video_up)
|
||||
# Extract tensor from latent distribution
|
||||
@@ -301,10 +305,14 @@ class LongCatRefineInitStage(PipelineStage):
|
||||
|
||||
logger.info("Applied t_thresh=%s noise mixing", t_thresh)
|
||||
|
||||
# Store in batch
|
||||
batch.latents = latent_up.to(dtype)
|
||||
# Store in batch - ensure correct dtype and device
|
||||
# The latents need to be on the same device as the transformer (CUDA)
|
||||
target_device = batch.prompt_embeds[0].device
|
||||
batch.latents = latent_up.to(device=target_device, dtype=dtype)
|
||||
batch.raw_latent_shape = latent_up.shape
|
||||
|
||||
logger.info("Latents device: %s, dtype: %s", batch.latents.device,
|
||||
batch.latents.dtype)
|
||||
logger.info("LongCat refinement initialization complete")
|
||||
|
||||
return batch
|
||||
|
||||
@@ -0,0 +1,217 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat VC Denoising Stage with KV cache support.
|
||||
|
||||
This stage extends the I2V denoising stage to support:
|
||||
1. KV cache for conditioning frames
|
||||
2. Video continuation with multiple conditioning frames
|
||||
"""
|
||||
|
||||
import time
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatVCDenoisingStage(LongCatDenoisingStage):
|
||||
"""
|
||||
LongCat denoising with Video Continuation and KV cache support.
|
||||
|
||||
Key differences from I2V denoising:
|
||||
- Supports KV cache (reuses cached K/V from conditioning frames)
|
||||
- Handles larger num_cond_latents
|
||||
- Concatenates conditioning latents back after denoising
|
||||
|
||||
When use_kv_cache=True:
|
||||
- batch.latents contains ONLY noise frames (cond removed by KV cache init)
|
||||
- batch.kv_cache_dict contains cached K/V
|
||||
- batch.cond_latents contains conditioning latents for post-concat
|
||||
|
||||
When use_kv_cache=False:
|
||||
- batch.latents contains ALL frames (cond + noise)
|
||||
- Timestep masking: timestep[:, :num_cond_latents] = 0
|
||||
- Selective denoising: only update noise frames
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run denoising loop with VC conditioning and optional KV cache."""
|
||||
|
||||
# Load transformer if needed
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Setup
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get VC-specific parameters
|
||||
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
|
||||
use_kv_cache = getattr(batch, 'use_kv_cache', False)
|
||||
kv_cache_dict = getattr(batch, 'kv_cache_dict', {})
|
||||
|
||||
logger.info(
|
||||
"VC Denoising: num_cond_latents=%d, use_kv_cache=%s, latent_shape=%s",
|
||||
num_cond_latents, use_kv_cache, latents.shape)
|
||||
|
||||
# Prepare negative prompts for CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
step_times = []
|
||||
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="VC Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
step_start = time.time()
|
||||
|
||||
# 1. Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# 2. Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# 3. Expand timestep to temporal dimension
|
||||
timestep = timestep.unsqueeze(-1).repeat(
|
||||
1, latent_model_input.shape[2])
|
||||
|
||||
# 4. Timestep masking (only when NOT using KV cache)
|
||||
if not use_kv_cache and num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
# 5. Prepare transformer kwargs
|
||||
# IMPORTANT: num_cond_latents is ALWAYS passed - needed for RoPE position offset
|
||||
transformer_kwargs = {
|
||||
'num_cond_latents': num_cond_latents,
|
||||
}
|
||||
if use_kv_cache:
|
||||
transformer_kwargs['kv_cache_dict'] = kv_cache_dict
|
||||
|
||||
# 6. Run transformer
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
**transformer_kwargs,
|
||||
)
|
||||
|
||||
# 7. Apply CFG with optimized scale (CFG-zero)
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# 8. Negate for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# 9. Scheduler step
|
||||
if use_kv_cache:
|
||||
# All latents are noise frames (conditioning is in cache)
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
# Only update noise frames (skip conditioning)
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
else:
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
step_time = time.time() - step_start
|
||||
step_times.append(step_time)
|
||||
|
||||
# Log timing for first few steps
|
||||
if i < 3:
|
||||
logger.info("Step %d: %.2fs", i, step_time)
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# 10. If using KV cache, concatenate conditioning latents back
|
||||
if use_kv_cache and hasattr(
|
||||
batch, 'cond_latents') and batch.cond_latents is not None:
|
||||
latents = torch.cat([batch.cond_latents, latents], dim=2)
|
||||
logger.info(
|
||||
"Concatenated conditioning latents back, final shape: %s",
|
||||
latents.shape)
|
||||
|
||||
# Log average timing
|
||||
avg_time = sum(step_times) / len(step_times)
|
||||
logger.info("Average step time: %.2fs (total: %.1fs)", avg_time,
|
||||
sum(step_times))
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Video VAE Encoding Stage for Video Continuation (VC) generation.
|
||||
|
||||
This stage handles encoding multiple video frames to latent space with
|
||||
LongCat-specific normalization for VC conditioning.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vision_utils import normalize, numpy_to_pt, pil_to_numpy, resize
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatVideoVAEEncodingStage(PipelineStage):
|
||||
"""
|
||||
Encode video frames to latent space for VC conditioning.
|
||||
|
||||
This stage:
|
||||
1. Loads video frames from path or uses provided frames
|
||||
2. Takes the last num_cond_frames from the video
|
||||
3. Preprocesses and stacks frames
|
||||
4. Encodes via VAE to latent space
|
||||
5. Applies LongCat-specific normalization
|
||||
6. Calculates num_cond_latents
|
||||
"""
|
||||
|
||||
def __init__(self, vae):
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Encode video frames to latent for VC conditioning."""
|
||||
|
||||
# Get video from batch - can be path, list of PIL images, or already loaded
|
||||
video = getattr(batch, 'video_frames', None) or getattr(
|
||||
batch, 'video_path', None)
|
||||
num_cond_frames = getattr(batch, 'num_cond_frames',
|
||||
13) # Default 13 for VC
|
||||
|
||||
if video is None:
|
||||
raise ValueError(
|
||||
"video_frames or video_path must be provided for VC")
|
||||
|
||||
# Load video if path
|
||||
if isinstance(video, str):
|
||||
from diffusers.utils import load_video
|
||||
video = load_video(video)
|
||||
logger.info("Loaded video from path: %d frames", len(video))
|
||||
|
||||
# Take last num_cond_frames
|
||||
if len(video) > num_cond_frames:
|
||||
video = video[-num_cond_frames:]
|
||||
logger.info("Using last %d frames for conditioning",
|
||||
num_cond_frames)
|
||||
elif len(video) < num_cond_frames:
|
||||
logger.warning(
|
||||
"Video has only %d frames, less than num_cond_frames=%d",
|
||||
len(video), num_cond_frames)
|
||||
num_cond_frames = len(video)
|
||||
|
||||
# Get target dimensions
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
if height is None or width is None:
|
||||
raise ValueError("height and width must be set for VC")
|
||||
|
||||
# Preprocess and stack frames
|
||||
processed_frames = []
|
||||
for frame in video:
|
||||
if not isinstance(frame, PIL.Image.Image):
|
||||
raise TypeError(f"Frame must be PIL.Image, got {type(frame)}")
|
||||
|
||||
frame = resize(frame, height, width, resize_mode="default")
|
||||
frame = pil_to_numpy(frame) # Returns [1, H, W, C] then converted
|
||||
frame = numpy_to_pt(frame) # Returns [1, C, H, W]
|
||||
frame = normalize(frame) # to [-1, 1]
|
||||
processed_frames.append(frame)
|
||||
|
||||
# Stack frames: [num_frames, C, H, W] -> [1, C, T, H, W]
|
||||
video_tensor = torch.cat(processed_frames, dim=0) # [T, C, H, W]
|
||||
video_tensor = video_tensor.permute(1, 0, 2,
|
||||
3).unsqueeze(0) # [1, C, T, H, W]
|
||||
video_tensor = video_tensor.to(get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
logger.info("VC: Preprocessed video tensor shape: %s",
|
||||
video_tensor.shape)
|
||||
|
||||
# Encode via VAE
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
|
||||
if not vae_autocast_enabled:
|
||||
video_tensor = video_tensor.to(vae_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
encoder_output = self.vae.encode(video_tensor)
|
||||
latent = self.retrieve_latents(encoder_output, batch.generator)
|
||||
|
||||
# Apply LongCat-specific normalization
|
||||
latent = self.normalize_latents(latent)
|
||||
|
||||
# Calculate num_cond_latents
|
||||
# Formula: 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
vae_temporal_scale = self.vae.config.scale_factor_temporal
|
||||
num_cond_latents = 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
|
||||
# Store in batch
|
||||
batch.video_latent = latent
|
||||
batch.num_cond_frames = num_cond_frames
|
||||
batch.num_cond_latents = num_cond_latents
|
||||
|
||||
logger.info(
|
||||
"VC: Encoded %d frames to latent shape %s, num_cond_latents=%d",
|
||||
num_cond_frames, latent.shape, num_cond_latents)
|
||||
|
||||
# Offload VAE if needed
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self, encoder_output: Any,
|
||||
generator: torch.Generator | None) -> torch.Tensor:
|
||||
"""Sample from VAE posterior."""
|
||||
if hasattr(encoder_output, 'sample'):
|
||||
return encoder_output.sample(generator)
|
||||
elif hasattr(encoder_output, 'latent_dist'):
|
||||
return encoder_output.latent_dist.sample(generator)
|
||||
elif hasattr(encoder_output, 'latents'):
|
||||
return encoder_output.latents
|
||||
else:
|
||||
raise AttributeError("Could not access latents from encoder output")
|
||||
|
||||
def normalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply LongCat-specific latent normalization.
|
||||
|
||||
Formula: (latents - mean) / std
|
||||
"""
|
||||
if not hasattr(self.vae.config, 'latents_mean') or not hasattr(
|
||||
self.vae.config, 'latents_std'):
|
||||
raise ValueError(
|
||||
"VAE config must have 'latents_mean' and 'latents_std' "
|
||||
"for LongCat normalization")
|
||||
|
||||
latents_mean = torch.tensor(self.vae.config.latents_mean).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
latents_std = torch.tensor(self.vae.config.latents_std).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
return (latents - latents_mean) / latents_std
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"mean_ssim": 0.9805062346988254,
|
||||
"min_ssim": 0.9636241793632507,
|
||||
"max_ssim": 0.9951534271240234,
|
||||
"reference_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.2-I2V-A14B-Diffusers/SLA_ATTN/An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space reali.mp4",
|
||||
"generated_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.2-I2V-A14B-Diffusers/SLA_ATTN/An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space reali.mp4",
|
||||
"parameters": {
|
||||
"num_inference_steps": 4,
|
||||
"prompt": "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
}
|
||||
}
|
||||
+5
-5
@@ -1,9 +1,9 @@
|
||||
{
|
||||
"mean_ssim": 0.9917645101194028,
|
||||
"min_ssim": 0.9908460974693298,
|
||||
"max_ssim": 0.9925968050956726,
|
||||
"reference_video": "/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
|
||||
"generated_video": "/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
|
||||
"mean_ssim": 0.7606927804004999,
|
||||
"min_ssim": 0.7035917639732361,
|
||||
"max_ssim": 0.7920367121696472,
|
||||
"reference_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
|
||||
"generated_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
|
||||
"parameters": {
|
||||
"num_inference_steps": 4,
|
||||
"prompt": "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
|
||||
@@ -152,7 +152,147 @@ def test_turbodiffusion_inference_similarity(prompt, model_id):
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
# TurboDiffusion uses fewer steps, may have slightly lower SSIM
|
||||
min_acceptable_ssim = 0.90
|
||||
min_acceptable_ssim = 0.95
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
|
||||
f"for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
)
|
||||
|
||||
|
||||
# TurboDiffusion I2V parameters (dual-model with RCM scheduler + SLA attention)
|
||||
TURBODIFFUSION_I2V_PARAMS = {
|
||||
"num_gpus": 4,
|
||||
"model_path": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 4, # TurboDiffusion uses 1-4 steps
|
||||
"guidance_scale": 1.0, # No CFG for TurboDiffusion
|
||||
"seed": 42,
|
||||
"sp_size": 4,
|
||||
"tp_size": 1,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
TURBODIFFUSION_I2V_MODEL_TO_PARAMS = {
|
||||
"TurboWan2.2-I2V-A14B-Diffusers": TURBODIFFUSION_I2V_PARAMS,
|
||||
}
|
||||
|
||||
TURBODIFFUSION_I2V_TEST_PROMPTS = [
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot.",
|
||||
]
|
||||
|
||||
TURBODIFFUSION_I2V_IMAGE_PATHS = [
|
||||
"https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", TURBODIFFUSION_I2V_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("model_id", list(TURBODIFFUSION_I2V_MODEL_TO_PARAMS.keys()))
|
||||
def test_turbodiffusion_i2v_inference_similarity(prompt, model_id):
|
||||
"""
|
||||
Test that runs TurboDiffusion I2V inference with dual-model switching,
|
||||
then compares the output to reference videos using SSIM.
|
||||
"""
|
||||
# TurboDiffusion requires SLA attention backend
|
||||
ATTENTION_BACKEND = "SLA_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
assert len(TURBODIFFUSION_I2V_TEST_PROMPTS) == len(TURBODIFFUSION_I2V_IMAGE_PATHS), \
|
||||
"Expect number of prompts equal to number of images"
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
|
||||
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
BASE_PARAMS = TURBODIFFUSION_I2V_MODEL_TO_PARAMS[model_id]
|
||||
num_inference_steps = BASE_PARAMS["num_inference_steps"]
|
||||
image_path = TURBODIFFUSION_I2V_IMAGE_PATHS[TURBODIFFUSION_I2V_TEST_PROMPTS.index(prompt)]
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"override_pipeline_cls_name": "TurboDiffusionI2VPipeline",
|
||||
# Keep both transformers in VRAM - avoids CPU RAM bottleneck
|
||||
"dit_cpu_offload": False,
|
||||
}
|
||||
|
||||
generation_kwargs = {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"output_path": output_dir,
|
||||
"image_path": image_path,
|
||||
"height": BASE_PARAMS["height"],
|
||||
"width": BASE_PARAMS["width"],
|
||||
"num_frames": BASE_PARAMS["num_frames"],
|
||||
"guidance_scale": BASE_PARAMS["guidance_scale"],
|
||||
"seed": BASE_PARAMS["seed"],
|
||||
"fps": BASE_PARAMS["fps"],
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=BASE_PARAMS["model_path"],
|
||||
**init_kwargs
|
||||
)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
if isinstance(generator.executor, MultiprocExecutor):
|
||||
generator.executor.shutdown()
|
||||
|
||||
assert os.path.exists(output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(
|
||||
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
|
||||
)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}"
|
||||
)
|
||||
|
||||
# Find the matching reference video based on the prompt
|
||||
reference_video_name = None
|
||||
|
||||
for filename in os.listdir(reference_folder):
|
||||
if filename.endswith('.mp4') and prompt[:100].strip() in filename:
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
if not reference_video_name:
|
||||
logger.error(
|
||||
f"Reference video not found for prompt: {prompt} with backend: {ATTENTION_BACKEND}"
|
||||
)
|
||||
raise FileNotFoundError(f"Reference video missing")
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
|
||||
logger.info(
|
||||
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
|
||||
)
|
||||
ssim_values = compute_video_ssim_torchvision(
|
||||
reference_video_path, generated_video_path, use_ms_ssim=True
|
||||
)
|
||||
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info(f"SSIM mean value: {mean_ssim}")
|
||||
logger.info(f"Writing SSIM results to directory: {output_dir}")
|
||||
|
||||
success = write_ssim_results(
|
||||
output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path, num_inference_steps, prompt
|
||||
)
|
||||
|
||||
if not success:
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
# TurboDiffusion I2V uses fewer steps, may have slightly lower SSIM
|
||||
min_acceptable_ssim = 0.95
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
|
||||
f"for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.1.6"
|
||||
__version__ = "0.1.7"
|
||||
|
||||
@@ -153,6 +153,8 @@ nav:
|
||||
- Video Sparse Attention: attention/vsa/index.md
|
||||
- Sliding Tile Attention: attention/sta/index.md
|
||||
- Adding a New Attention Backend: attention/developer/index.md
|
||||
- Utilities:
|
||||
- LoRA: utilities/lora.md
|
||||
- Design:
|
||||
- Overview: design/overview.md
|
||||
- Developer Guide:
|
||||
|
||||
+3
-2
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.6"
|
||||
version = "0.1.7"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
@@ -63,6 +63,7 @@ dependencies = [
|
||||
"remote-pdb",
|
||||
|
||||
# Kernel & Packaging
|
||||
"fastvideo-kernel==0.2.2",
|
||||
"wheel",
|
||||
|
||||
# Training Dependencies
|
||||
@@ -110,7 +111,7 @@ lint = [
|
||||
]
|
||||
|
||||
test = [
|
||||
"av==14.3.0",
|
||||
"av",
|
||||
"pytorch-msssim==1.0.0",
|
||||
"pytest",
|
||||
]
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.6"
|
||||
version = "0.1.7"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
@@ -63,6 +63,7 @@ dependencies = [
|
||||
"remote-pdb",
|
||||
|
||||
# Kernel & Packaging
|
||||
"fastvideo-kernel==0.2.2",
|
||||
"wheel",
|
||||
|
||||
# Training Dependencies
|
||||
@@ -89,7 +90,7 @@ lint = [
|
||||
]
|
||||
|
||||
test = [
|
||||
"av==14.3.0",
|
||||
"av",
|
||||
"pytorch-msssim==1.0.0",
|
||||
"pytest",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Convert TurboDiffusion I2V .pth checkpoint to Diffusers safetensors format.
|
||||
|
||||
TurboDiffusion I2V uses two models: high-noise and low-noise.
|
||||
This script converts both checkpoints to Diffusers format.
|
||||
|
||||
Usage:
|
||||
python convert_turbodiffusion_i2v_to_diffusers.py \
|
||||
--high_noise_path /path/to/TurboWan2.2-I2V-A14B-high-720P.pth \
|
||||
--low_noise_path /path/to/TurboWan2.2-I2V-A14B-low-720P.pth \
|
||||
--output_dir /path/to/output \
|
||||
--reference_repo Wan-AI/Wan2.1-I2V-14B-720P-Diffusers
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
import torch
|
||||
import shutil
|
||||
import glob
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import save_file
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
# Weight mapping from TurboDiffusion -> Diffusers/FastVideo format
|
||||
# Same as T2V but may need additional I2V-specific mappings
|
||||
TURBODIFFUSION_WEIGHT_MAPPING = {
|
||||
# Self attention
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
|
||||
# Cross attention
|
||||
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
|
||||
# I2V-specific cross attention (add_k/v_proj for image context)
|
||||
r"^blocks\.(\d+)\.cross_attn\.add_k\.(.*)$": r"blocks.\1.attn2.add_k_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.add_v\.(.*)$": r"blocks.\1.attn2.add_v_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_add_k\.(.*)$": r"blocks.\1.attn2.norm_added_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_add_q\.(.*)$": r"blocks.\1.attn2.norm_added_q.\2",
|
||||
# Norms and FFN
|
||||
r"^blocks\.(\d+)\.norm1\.(.*)$": r"blocks.\1.norm1.\2",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.norm3.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
|
||||
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
|
||||
# Embeddings
|
||||
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
# Head
|
||||
r"^head\.head\.(.*)$": r"proj_out.\1",
|
||||
r"^head\.norm\.(.*)$": r"norm_out.\1",
|
||||
r"^head\.modulation$": r"scale_shift_table",
|
||||
# SLA proj_l weights
|
||||
r"^blocks\.(\d+)\.self_attn\.attn_op\.local_attn\.proj_l\.(.*)$": r"blocks.\1.attn1.attn_impl.proj_l.\2",
|
||||
}
|
||||
|
||||
SKIP_PATTERNS = []
|
||||
|
||||
|
||||
def should_skip_key(key: str) -> bool:
|
||||
for pattern in SKIP_PATTERNS:
|
||||
if re.match(pattern, key):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def convert_key(turbo_key: str) -> str:
|
||||
for pattern, replacement in TURBODIFFUSION_WEIGHT_MAPPING.items():
|
||||
if re.match(pattern, turbo_key):
|
||||
return re.sub(pattern, replacement, turbo_key)
|
||||
return turbo_key
|
||||
|
||||
|
||||
def reshape_patch_embedding(tensor: torch.Tensor, target_shape: tuple) -> torch.Tensor:
|
||||
if len(tensor.shape) == 2 and len(target_shape) == 5:
|
||||
return tensor.view(target_shape)
|
||||
return tensor
|
||||
|
||||
|
||||
def get_reference_shapes(reference_repo: str) -> dict:
|
||||
print(f"Downloading reference model shapes from {reference_repo}...")
|
||||
|
||||
local_dir = snapshot_download(
|
||||
repo_id=reference_repo,
|
||||
allow_patterns=["transformer/config.json", "transformer/diffusion_pytorch_model*.safetensors"],
|
||||
local_dir_use_symlinks=False
|
||||
)
|
||||
|
||||
weight_files = glob.glob(os.path.join(local_dir, "transformer", "*.safetensors"))
|
||||
shapes = {}
|
||||
|
||||
for wf in weight_files:
|
||||
with safe_open(wf, framework="pt") as f:
|
||||
for key in f.keys():
|
||||
shapes[key] = f.get_tensor(key).shape
|
||||
|
||||
return shapes, local_dir
|
||||
|
||||
|
||||
def convert_checkpoint(input_path: str, output_dir: str, ref_shapes: dict, model_name: str) -> None:
|
||||
"""Convert a single TurboDiffusion checkpoint to Diffusers format."""
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Converting {model_name}: {input_path}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
turbo_state_dict = torch.load(input_path, map_location="cpu", weights_only=True)
|
||||
print(f"Loaded {len(turbo_state_dict)} keys")
|
||||
|
||||
converted_state_dict = {}
|
||||
skipped_keys = []
|
||||
|
||||
for turbo_key, tensor in turbo_state_dict.items():
|
||||
if should_skip_key(turbo_key):
|
||||
skipped_keys.append(turbo_key)
|
||||
continue
|
||||
|
||||
new_key = convert_key(turbo_key)
|
||||
|
||||
if "patch_embedding" in new_key and new_key in ref_shapes:
|
||||
target_shape = ref_shapes[new_key]
|
||||
if tensor.shape != target_shape:
|
||||
print(f"Reshaping {new_key}: {tensor.shape} -> {target_shape}")
|
||||
tensor = reshape_patch_embedding(tensor, target_shape)
|
||||
|
||||
if new_key in ref_shapes:
|
||||
if tensor.shape != ref_shapes[new_key]:
|
||||
print(f"WARNING: Shape mismatch for {new_key}: got {tensor.shape}, expected {ref_shapes[new_key]}")
|
||||
|
||||
converted_state_dict[new_key] = tensor
|
||||
|
||||
print(f"Converted: {len(converted_state_dict)} keys, Skipped: {len(skipped_keys)} keys")
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
output_path = os.path.join(output_dir, "diffusion_pytorch_model.safetensors")
|
||||
print(f"Saving to {output_path}...")
|
||||
save_file(converted_state_dict, output_path)
|
||||
|
||||
return converted_state_dict
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Convert TurboDiffusion I2V checkpoints to Diffusers format")
|
||||
parser.add_argument("--high_noise_path", type=str, required=True,
|
||||
help="Path to high-noise TurboDiffusion .pth checkpoint")
|
||||
parser.add_argument("--low_noise_path", type=str, required=True,
|
||||
help="Path to low-noise TurboDiffusion .pth checkpoint")
|
||||
parser.add_argument("--output_dir", type=str, required=True,
|
||||
help="Output directory for converted safetensors")
|
||||
parser.add_argument("--reference_repo", type=str, default="Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
help="Reference HF repo to get expected tensor shapes")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Get reference shapes
|
||||
ref_shapes, ref_local_dir = get_reference_shapes(args.reference_repo)
|
||||
print(f"Got {len(ref_shapes)} reference shapes")
|
||||
|
||||
# Convert high-noise model
|
||||
high_noise_output = os.path.join(args.output_dir, "transformer_high")
|
||||
convert_checkpoint(args.high_noise_path, high_noise_output, ref_shapes, "high-noise")
|
||||
|
||||
# Convert low-noise model
|
||||
low_noise_output = os.path.join(args.output_dir, "transformer_low")
|
||||
convert_checkpoint(args.low_noise_path, low_noise_output, ref_shapes, "low-noise")
|
||||
|
||||
# Copy config.json to both
|
||||
src_config = os.path.join(ref_local_dir, "transformer", "config.json")
|
||||
shutil.copy(src_config, os.path.join(high_noise_output, "config.json"))
|
||||
shutil.copy(src_config, os.path.join(low_noise_output, "config.json"))
|
||||
print("Copied config.json to both transformer directories")
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Conversion complete!")
|
||||
print(f"High-noise model: {high_noise_output}")
|
||||
print(f"Low-noise model: {low_noise_output}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,14 +1,30 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=2
|
||||
# LongCat Text-to-Video (T2V) Inference Script
|
||||
#
|
||||
# This script runs LongCat T2V inference using the fastvideo CLI.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# For longcat, we must first convert the official weights to FastVideo native format
|
||||
# Model path options:
|
||||
# Option 1: HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-T2V-Diffusers
|
||||
|
||||
# Option 2: Local weights (uncomment if you have local weights)
|
||||
# For local weights, convert the official weights to FastVideo native format
|
||||
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
|
||||
# --source /path/to/LongCat-Video/weights/LongCat-Video
|
||||
# --output weights/longcat-native
|
||||
export MODEL_BASE=weights/longcat-native
|
||||
# export MODEL_BASE=weights/longcat-native
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
@@ -26,7 +42,7 @@ fastvideo generate \
|
||||
--num-inference-steps 50 \
|
||||
--fps 15 \
|
||||
--guidance-scale 4.0 \
|
||||
--prompt-txt assets/prompt.txt \
|
||||
--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 \
|
||||
--output-path outputs_video/longcat_480p
|
||||
--output-path outputs_video/longcat_t2v
|
||||
|
||||
@@ -1,12 +1,23 @@
|
||||
#!/bin/bash
|
||||
|
||||
# LongCat T2V Distilled Inference Script
|
||||
#
|
||||
# This script runs LongCat T2V with distillation LoRA (16 steps instead of 50).
|
||||
# Uses the distilled LoRA for faster generation.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_distill.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
# For longcat, we must first convert the official weights to FastVideo native format
|
||||
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
|
||||
# --source /path/to/LongCat-Video/weights/LongCat-Video
|
||||
# --output weights/longcat-native
|
||||
export MODEL_BASE=weights/longcat-native
|
||||
|
||||
# Model path - HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-T2V-Diffusers
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
@@ -14,11 +25,11 @@ fastvideo generate \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload False \
|
||||
--vae-cpu-offload False \
|
||||
--text-encoder-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa False \
|
||||
--lora-path "$MODEL_BASE/lora/distilled" \
|
||||
--lora-path "FastVideo/LongCat-Video-T2V-Distilled-LoRA" \
|
||||
--lora-nickname "distilled" \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
@@ -26,6 +37,7 @@ fastvideo generate \
|
||||
--num-inference-steps 16 \
|
||||
--fps 15 \
|
||||
--guidance-scale 1.0 \
|
||||
--prompt "In a realistic photography style, an asian boy around seven or eight years old sits on a park bench, wearing a light yellow 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." \
|
||||
--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 \
|
||||
--output-path outputs_video/longcat_distill
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
#!/bin/bash
|
||||
|
||||
# LongCat Image-to-Video (I2V) Inference Script
|
||||
#
|
||||
# This script runs LongCat I2V inference using the fastvideo CLI.
|
||||
# LongCat I2V takes an input image and generates a video from it.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_i2v.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
# - Or use local weights if you have them
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# Model path options:
|
||||
# Option 1: HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-I2V-Diffusers
|
||||
|
||||
# Option 2: Local weights (uncomment if you have local weights)
|
||||
# export MODEL_BASE=weights/longcat-for-i2v
|
||||
|
||||
# Input image path (must be square for LongCat I2V)
|
||||
IMAGE_PATH="assets/girl.png"
|
||||
|
||||
# Check if image exists
|
||||
if [ ! -f "$IMAGE_PATH" ]; then
|
||||
echo "Error: Image not found at $IMAGE_PATH"
|
||||
echo "Please provide a valid image path"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa False \
|
||||
--image-path "$IMAGE_PATH" \
|
||||
--height 480 \
|
||||
--width 480 \
|
||||
--num-frames 93 \
|
||||
--num-inference-steps 50 \
|
||||
--fps 15 \
|
||||
--guidance-scale 4.0 \
|
||||
--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" \
|
||||
--seed 42 \
|
||||
--output-path outputs_video/longcat_i2v
|
||||
|
||||
|
||||
|
||||
@@ -1,18 +1,30 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
# For longcat, we must first convert the official weights to FastVideo native format
|
||||
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
|
||||
# --source /path/to/LongCat-Video/weights/LongCat-Video
|
||||
# --output weights/longcat-native
|
||||
export MODEL_BASE=weights/longcat-native
|
||||
# LongCat T2V Refinement Script (480p -> 720p)
|
||||
#
|
||||
# This script refines a 480p distilled video to 720p using the refinement LoRA.
|
||||
# Run v1_inference_longcat_distill.sh first to generate the 480p video.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_refine_fromvideo.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
# - Run v1_inference_longcat_distill.sh first to generate input video
|
||||
|
||||
INPUT_VIDEO="outputs_video/longcat_distill/In a realistic photography style, an asian boy around seven or eight years old sits on a park bench,.mp4"
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# Model path - HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-T2V-Diffusers
|
||||
|
||||
INPUT_VIDEO="outputs_video/longcat_distill/In a realistic photography style, a white boy around seven or eight years old sits on a park bench,.mp4"
|
||||
REFINE_OUTPUT="outputs_video/longcat_refine_720p"
|
||||
|
||||
# Prompt used for base generation
|
||||
PROMPT="In a realistic photography style, an asian boy around seven or eight years old sits on a park bench, wearing a light yellow 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."
|
||||
# Prompt used for base generation (must match distill script)
|
||||
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."
|
||||
|
||||
echo "=========================================="
|
||||
echo "LongCat 480p -> 720p Refinement"
|
||||
@@ -25,14 +37,14 @@ echo ""
|
||||
# Check if input video exists
|
||||
if [ ! -f "$INPUT_VIDEO" ]; then
|
||||
echo "Error: Input video not found: $INPUT_VIDEO"
|
||||
echo "Please set INPUT_VIDEO to your 480p video path"
|
||||
echo "Please run v1_inference_longcat_distill.sh first to generate the 480p video"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "🔧 Configuring refinement (BSA enabled, refinement LoRA)..."
|
||||
echo "✅ Input video: $INPUT_VIDEO"
|
||||
echo "✅ BSA enabled with sparsity=0.875"
|
||||
echo "✅ Refinement LoRA loaded"
|
||||
echo "Configuring refinement (BSA enabled, refinement LoRA)..."
|
||||
echo "Input video: $INPUT_VIDEO"
|
||||
echo "BSA enabled with sparsity=0.875"
|
||||
echo "Refinement LoRA loaded"
|
||||
echo ""
|
||||
|
||||
fastvideo generate \
|
||||
@@ -41,14 +53,14 @@ fastvideo generate \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload True \
|
||||
--vae-cpu-offload False \
|
||||
--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 "$MODEL_BASE/lora/refinement" \
|
||||
--lora-path "FastVideo/LongCat-Video-T2V-Refinement-LoRA" \
|
||||
--lora-nickname "refinement" \
|
||||
--refine-from "$INPUT_VIDEO" \
|
||||
--t-thresh 0.5 \
|
||||
@@ -60,12 +72,13 @@ fastvideo generate \
|
||||
--fps 30 \
|
||||
--guidance-scale 1.0 \
|
||||
--prompt "$PROMPT" \
|
||||
--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 \
|
||||
--output-path "$REFINE_OUTPUT"
|
||||
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "✓ Refinement Complete!"
|
||||
echo "Refinement Complete!"
|
||||
echo "=========================================="
|
||||
echo ""
|
||||
echo "Output directory: $REFINE_OUTPUT"
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
#!/bin/bash
|
||||
|
||||
# LongCat Video Continuation (VC) Inference Script
|
||||
#
|
||||
# This script runs LongCat VC inference using the fastvideo CLI.
|
||||
# LongCat VC takes an input video and generates a continuation of it.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_vc.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
# - Or use local weights if you have them
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# Model path options:
|
||||
# Option 1: HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-VC-Diffusers
|
||||
|
||||
# Option 2: Local weights (uncomment if you have local weights)
|
||||
# export MODEL_BASE=weights/longcat-vc-upload
|
||||
|
||||
# Input video path
|
||||
VIDEO_PATH="assets/motorcycle.mp4"
|
||||
|
||||
# Check if video exists
|
||||
if [ ! -f "$VIDEO_PATH" ]; then
|
||||
echo "Error: Video not found at $VIDEO_PATH"
|
||||
echo "Please provide a valid video path"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa False \
|
||||
--video-path "$VIDEO_PATH" \
|
||||
--num-cond-frames 13 \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
--num-frames 93 \
|
||||
--num-inference-steps 50 \
|
||||
--fps 15 \
|
||||
--guidance-scale 4.0 \
|
||||
--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" \
|
||||
--seed 42 \
|
||||
--output-path outputs_video/longcat_vc
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user