Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3572c6821e | ||
|
|
deb901f6fc | ||
|
|
beae2943cf | ||
|
|
fadb71bb64 | ||
|
|
65c6fbaa46 | ||
|
|
e32a8a9504 | ||
|
|
6907d87871 | ||
|
|
9775fed31a | ||
|
|
5c6f635d73 | ||
|
|
c19708fd58 | ||
|
|
b2e4fb0743 | ||
|
|
a3ad4852b0 | ||
|
|
add2be21b5 | ||
|
|
41203d92b8 | ||
|
|
5fa8415c0b | ||
|
|
3a182925f3 | ||
|
|
c1e4787775 | ||
|
|
2c6bf47b9f | ||
|
|
548cc08817 | ||
|
|
7521b06693 | ||
|
|
fb9ad77086 | ||
|
|
31f44110b5 | ||
|
|
21f3ce6577 | ||
|
|
785d123e36 | ||
|
|
d58c551c11 | ||
|
|
560628709c | ||
|
|
0f53b51e6c | ||
|
|
06093a9c4e | ||
|
|
dbddfab6d2 | ||
|
|
7188170277 | ||
|
|
b7f69c2c1d | ||
|
|
23a4531491 | ||
|
|
7d52ad0118 | ||
|
|
4d7bf35fa3 | ||
|
|
a6a9c9ca07 | ||
|
|
f4704847c2 | ||
|
|
d9c996310b | ||
|
|
d6651afd2e | ||
|
|
cf67618cad | ||
|
|
2f0a2b3c57 | ||
|
|
e7748d9952 | ||
|
|
8eb3140b2f | ||
|
|
d6ddcea682 | ||
|
|
3559ba2377 | ||
|
|
61e63ea0d7 | ||
|
|
4ce4ac4734 | ||
|
|
e7f6db9bd1 | ||
|
|
d83f45a6a0 | ||
|
|
dd91542cd1 | ||
|
|
581e8115fe | ||
|
|
dea69cf651 | ||
|
|
60ac6537df | ||
|
|
5285116e73 |
@@ -61,7 +61,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 45m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
@@ -76,7 +76,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
label: "LoRA Inference Tests"
|
||||
env:
|
||||
- TEST_TYPE=inference_lora
|
||||
|
||||
@@ -156,8 +156,23 @@ jobs:
|
||||
|
||||
# Fix the wheel to be manylinux compliant
|
||||
pip install auditwheel
|
||||
# Point auditwheel at torch libs, but do not vendor them into the wheel.
|
||||
TORCH_LIB_DIR=$(python - <<'PY'
|
||||
import os
|
||||
import torch
|
||||
|
||||
print(os.path.join(os.path.dirname(torch.__file__), "lib"))
|
||||
PY
|
||||
)
|
||||
export LD_LIBRARY_PATH="${TORCH_LIB_DIR}:${LD_LIBRARY_PATH}"
|
||||
# Target manylinux_2_35 (Ubuntu 22.04 native)
|
||||
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist
|
||||
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist \
|
||||
--exclude libtorch_cuda.so \
|
||||
--exclude libtorch_cpu.so \
|
||||
--exclude libtorch.so \
|
||||
--exclude libc10.so \
|
||||
--exclude libc10_cuda.so \
|
||||
--exclude libtorch_python.so
|
||||
# Move fixed wheels back to dist for upload consistency
|
||||
rm dist/*.whl
|
||||
mv fixed_dist/*.whl dist/
|
||||
|
||||
@@ -68,7 +68,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
- 'git ls-files | grep -v "^\"*fastvideo/tests/ssim/" | grep -v "^\"*fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
|
||||
@@ -1,41 +1,47 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-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>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/fastwan.png width="90%"/>
|
||||
</div>
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
<details>
|
||||
<summary>More</summary>
|
||||
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
</details>
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
- End-to-end post-training support:
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
|
||||
- Data preprocessing pipeline for video data
|
||||
- End-to-end post-training support for bidirectional and autoregressive models:
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
|
||||
- Data preprocessing pipeline for video, image, and text data
|
||||
- Distribution Matching Distillation (DMD2) stepwise distillation.
|
||||
- Sparse attention with [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achineve >50x denoising speedup
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
|
||||
- Causal distillation through Self-Forcing
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Sequence Parallelism for distributed inference
|
||||
- Multiple state-of-the-art attention backends
|
||||
- User-friendly CLI and Python API
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/optimizations/) for full list of supported optimizations.
|
||||
- Diverse hardware and OS support
|
||||
- Support H100, A100, 4090
|
||||
- Support Linux, Windows, MacOS
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/hardware_support/) for full list of supported hardware and OS.
|
||||
|
||||
## Getting Started
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 490 KiB |
Binary file not shown.
@@ -50,21 +50,25 @@ class MyNewAttnBackend(AttentionBackend):
|
||||
FastVideo uses a `ForwardContext` to pass global metadata (like current timestep, batch info, or custom attention configurations) to attention backends without changing the `forward` signature of every layer. **This is optional and only required if your backend needs dynamic per-step information.**
|
||||
|
||||
To use this:
|
||||
|
||||
1. **Set Context**: In your pipeline or generation loop, use the `set_forward_context` context manager.
|
||||
2. **Access Context**: Inside your attention backend, use `get_forward_context()`.
|
||||
|
||||
See `docs/attention/sta/index.md` (Sliding Tile Attention) for an example of how complex configuration (window sizes) is passed this way.
|
||||
See [`docs/attention/sta/index.md`](../sta/index.md) (Sliding Tile Attention) for an example of how complex configuration (window sizes) is passed this way.
|
||||
|
||||
## 3. Adding Compiled Kernels (C++/CUDA)
|
||||
|
||||
If your backend requires custom CUDA kernels, you need to add them to the `fastvideo-kernel` package.
|
||||
|
||||
### A. Add Source Files
|
||||
|
||||
Place your kernel implementation files in `fastvideo-kernel/csrc/attention/`.
|
||||
|
||||
* `mynew_attn.cu` (CUDA implementation)
|
||||
* `mynew_attn.h` (Optional headers)
|
||||
|
||||
### B. Register in Extension
|
||||
|
||||
Update `fastvideo-kernel/csrc/common_extension.cpp` to expose your function to Python.
|
||||
|
||||
```cpp
|
||||
@@ -84,6 +88,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
```
|
||||
|
||||
### C. Update CMakeLists.txt
|
||||
|
||||
Update `fastvideo-kernel/CMakeLists.txt` to compile your new files.
|
||||
|
||||
**Case 1: General CUDA Kernel (Runs on all GPUs)**
|
||||
@@ -111,6 +116,7 @@ endif()
|
||||
```
|
||||
|
||||
### D. Expose in Python Ops
|
||||
|
||||
Update `fastvideo-kernel/python/fastvideo_kernel/ops.py` to make the function importable and handle fallbacks gracefully.
|
||||
|
||||
```python
|
||||
@@ -133,6 +139,7 @@ def my_compiled_attn_func(q, k, v):
|
||||
```
|
||||
|
||||
### E. Expose in Package Init
|
||||
|
||||
Update `fastvideo-kernel/python/fastvideo_kernel/__init__.py` to export the function.
|
||||
|
||||
```python
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -13,11 +13,13 @@ from fastvideo_kernel import video_sparse_attn
|
||||
|
||||
# q, k, v: [batch_size, num_heads, seq_len, head_dim]
|
||||
# variable_block_sizes: Number of valid tokens per block
|
||||
# q_variable_block_sizes: Number of valid tokens per q block (can differ from KV for q/k of different lengths)
|
||||
# topk: Number of blocks to attend
|
||||
|
||||
output = video_sparse_attn(
|
||||
q, k, v,
|
||||
variable_block_sizes=block_sizes,
|
||||
block_sizes,
|
||||
block_sizes,
|
||||
topk=32
|
||||
)
|
||||
```
|
||||
|
||||
+35
-14
@@ -4,25 +4,26 @@ This document outlines FastVideo's architecture for developers interested in fra
|
||||
|
||||
## Table of Contents - Directory Structure and Files
|
||||
|
||||
- [`fastvideo/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/models/`](#design-model-components) - Model implementations
|
||||
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
|
||||
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
|
||||
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
|
||||
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
|
||||
- [`fastvideo/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- [`fastvideo/pipelines/`](#pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/models/`](#model-components) - Model implementations
|
||||
- [`dits/`](#transformer-models) - Transformer-based diffusion models
|
||||
- [`vaes/`](#vae-variational-auto-encoder) - Variational autoencoders
|
||||
- [`encoders/`](#text-and-image-encoders) - Text and image encoders
|
||||
- [`schedulers/`](#schedulers) - Diffusion schedulers
|
||||
- [`fastvideo/attention/`](#optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/distributed/`](#distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/layers/`](#tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/platforms/`](#platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/worker/`](#executor-and-worker-system) - Multi-GPU process management
|
||||
- [`fastvideo/fastvideo_args.py`](#fastvideoargs) - Argument handling
|
||||
- [`fastvideo/forward_context.py`](#forward-context-management) - Forward pass context management
|
||||
- `fastvideo/utils.py` - Utility functions
|
||||
- [`fastvideo/logger.py`](#design-logger) - Logging infrastructure
|
||||
- [`fastvideo/logger.py`](#logger) - Logging infrastructure
|
||||
|
||||
## Core Architecture
|
||||
|
||||
FastVideo separates model components from execution logic with these principles:
|
||||
|
||||
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
|
||||
- **Modular Design**: Components can be independently replaced
|
||||
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
|
||||
@@ -34,12 +35,14 @@ FastVideo separates model components from execution logic with these principles:
|
||||
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
|
||||
Key features include:
|
||||
|
||||
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
|
||||
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
|
||||
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
|
||||
- **Parameter Validation**: Ensures valid combinations of settings
|
||||
|
||||
Common configuration areas:
|
||||
|
||||
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
|
||||
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
|
||||
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
|
||||
@@ -90,7 +93,9 @@ class MyCustomPipeline(ComposedPipelineBase):
|
||||
```
|
||||
|
||||
### Pipeline Stages
|
||||
|
||||
Each stage handles a specific diffusion process component:
|
||||
|
||||
- **Input Validation**: Parameter verification
|
||||
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
|
||||
- **Image Encoding**: Image input processing
|
||||
@@ -133,6 +138,7 @@ Transformer networks perform the actual denoising during diffusion:
|
||||
- `HunyuanVideoTransformer3DModel`
|
||||
|
||||
Features include:
|
||||
|
||||
- Text/image conditioning
|
||||
- Standardized interface for model-specific optimizations
|
||||
|
||||
@@ -161,6 +167,7 @@ VAEs handle conversion between pixel space and latent space:
|
||||
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
|
||||
|
||||
FastVideo's VAE implementations include:
|
||||
|
||||
- Efficient video batch processing
|
||||
- Memory optimization
|
||||
- Optional tiling for large frames
|
||||
@@ -179,6 +186,7 @@ Encoders process conditioning inputs into embeddings:
|
||||
- `CLIPVisionModel`
|
||||
|
||||
FastVideo implements optimizations such as:
|
||||
|
||||
- Vocab parallelism for distributed processing
|
||||
- Caching for common prompts
|
||||
- Precision-tuned computation
|
||||
@@ -193,6 +201,7 @@ Schedulers manage the diffusion sampling process:
|
||||
- `FlowMatchEulerDiscreteScheduler`
|
||||
|
||||
These components control:
|
||||
|
||||
- Diffusion timestep sequences
|
||||
- Noise prediction to latent update conversions
|
||||
- Quality/speed trade-offs
|
||||
@@ -219,7 +228,9 @@ This diagram shows how models are discovered, validated, and loaded across entry
|
||||
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
|
||||
### Attention Backends
|
||||
|
||||
Multiple implementations with automatic selection:
|
||||
|
||||
- **FLASH_ATTN**: Optimized for supporting hardware
|
||||
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
|
||||
- **SLIDING_TILE_ATTN**: For very long sequences
|
||||
@@ -240,7 +251,9 @@ self.attn = LocalAttention(
|
||||

|
||||
|
||||
### Attention Patterns
|
||||
|
||||
Supports various patterns with memory optimization techniques:
|
||||
|
||||
- **Cross/Self/Temporal/Global-Local Attention**
|
||||
- Chunking, progressive computation, optimized masking
|
||||
|
||||
@@ -296,6 +309,7 @@ self.attn = DistributedAttention(
|
||||
```
|
||||
|
||||
### Communication Primitives
|
||||
|
||||
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
|
||||
|
||||
Efficient communication primitives minimize distributed overhead:
|
||||
@@ -314,6 +328,7 @@ Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-sp
|
||||
- **Profiling Data**: Potential hooks for performance metrics collection
|
||||
|
||||
This context-based approach enables:
|
||||
|
||||
- Dynamic optimization based on execution state (e.g., attention backend selection)
|
||||
- Step-specific customizations within model components
|
||||
|
||||
@@ -339,12 +354,14 @@ FastVideo implements a flexible execution model for distributed processing:
|
||||
- **GPU Workers**: Handle actual model execution on individual GPUs
|
||||
|
||||
The MultiProcExecutor implementation:
|
||||
|
||||
1. Spawns worker processes for each GPU
|
||||
2. Establishes communication channels via pipes
|
||||
3. Coordinates distributed operations across workers
|
||||
4. Handles graceful startup and shutdown of the process group
|
||||
|
||||
Each GPU worker:
|
||||
|
||||
1. Initializes the distributed environment
|
||||
2. Builds the pipeline for the specified model
|
||||
3. Executes requested operations on its assigned GPU
|
||||
@@ -359,11 +376,13 @@ The `fastvideo/platforms/` directory provides hardware platform abstractions tha
|
||||
### Platform Abstraction
|
||||
|
||||
FastVideo's platform abstraction layer enables:
|
||||
|
||||
- **Hardware Detection**: Automatic detection of available hardware
|
||||
- **Backend Selection**: Appropriate selection of compute kernels
|
||||
- **Memory Management**: Efficient utilization of hardware-specific memory features
|
||||
|
||||
The primary components include:
|
||||
|
||||
- **Platform Interface**: Defines the common API for all platform implementations
|
||||
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
|
||||
- **Backend Enum**: Used throughout the codebase for feature selection
|
||||
@@ -383,6 +402,7 @@ else:
|
||||
The platform system is designed to be extensible for future hardware targets.
|
||||
|
||||
## Logger
|
||||
|
||||
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
|
||||
@@ -397,6 +417,7 @@ If you're a new contributor, here are some common areas to explore:
|
||||
4. **Hardware support**: Extend the `platforms` module for new hardware targets
|
||||
|
||||
When adding code, follow these practices:
|
||||
|
||||
- Use type hints for better code readability
|
||||
- Add appropriate docstrings
|
||||
- Maintain the separation between model components and execution logic
|
||||
|
||||
@@ -13,7 +13,7 @@ We provide two distilled models:
|
||||
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
|
||||
|
||||
## ⚙️ Inference
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
|
||||
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
|
||||
@@ -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+
|
||||
@@ -38,4 +42,4 @@ pip install -e .
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/) - Explore example scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
|
||||
@@ -84,12 +84,12 @@ pip install flash-attn --no-build-isolation
|
||||
|
||||
## Set up using Docker
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
[Docker Images](../../contributing/developer_env/docker.md)
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
[Contributor Guide](../../contributing/overview.md)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
|
||||
@@ -78,7 +78,7 @@ uv pip install -e .
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
[Contributor Guide](../../contributing/overview.md)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -45,6 +45,7 @@ FastVideo uses the Hugging Face Diffusers format for model organization:
|
||||
### Implementing Modules
|
||||
|
||||
Place new modules in the appropriate directories:
|
||||
|
||||
- Encoders: `fastvideo/models/encoders/`
|
||||
- VAEs: `fastvideo/models/vaes/`
|
||||
- Transformer models: `fastvideo/models/dits/`
|
||||
@@ -53,12 +54,15 @@ Place new modules in the appropriate directories:
|
||||
### Adapting Model Layers
|
||||
|
||||
#### Layer Replacements
|
||||
|
||||
Replace standard PyTorch layers with FastVideo optimized versions:
|
||||
|
||||
- nn.LayerNorm → fastvideo.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.layers.activation
|
||||
|
||||
#### Distributed Linear Layers
|
||||
|
||||
Use appropriate parallel layers for distribution:
|
||||
|
||||
```python
|
||||
@@ -91,6 +95,7 @@ self.out_proj = RowParallelLinear(
|
||||
```
|
||||
|
||||
### Attention Layers
|
||||
|
||||
Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
@@ -304,6 +309,7 @@ EntryClass = [MyCustomPipeline, MyOtherPipeline]
|
||||
```
|
||||
|
||||
The registry will automatically:
|
||||
|
||||
1. Scan all packages under `fastvideo/pipelines/`
|
||||
2. Look for `EntryClass` variables
|
||||
3. Register pipelines using their class names as identifiers
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# FastVideo CLI Inference
|
||||
|
||||
The FastVideo CLI provides a quick way to access the FastVideo inference pipeline for video generation. For more advanced usage,
|
||||
see the Python interface [here](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html).
|
||||
see the Python interface [here](examples/basic.md).
|
||||
|
||||
## Basic Usage
|
||||
|
||||
|
||||
@@ -74,4 +74,4 @@ if __name__ == '__main__':
|
||||
|
||||
## Performance Optimization
|
||||
|
||||
For configuring optimizations, please see our [optimizations guide](#inference-optimizations)
|
||||
For configuring optimizations, please see our [optimizations guide](optimizations.md)
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
This page contains step-by-step instructions to get you quickly started with video generation using FastVideo.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS**: Linux (Tested on Ubuntu 22.04+)
|
||||
- **Python**: 3.10-3.12
|
||||
- **CUDA**: 12.8
|
||||
@@ -21,9 +22,10 @@ conda activate fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
For advanced installation options, see the [Installation Guide](installation.md).
|
||||
For advanced installation options, see the [Installation Guide](../getting_started/installation.md).
|
||||
|
||||
## Generating Your First Video
|
||||
|
||||
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
@@ -60,9 +62,10 @@ python example.py
|
||||
The generated video will be saved in the current directory under `my_videos/`
|
||||
|
||||
More inference example scripts can be found in `scripts/inference/`
|
||||
|
||||
## Available Models
|
||||
|
||||
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
|
||||
Please see the [support matrix](support_matrix.md) for the list of supported models and their available optimizations.
|
||||
|
||||
## Image-to-Video Generation
|
||||
|
||||
@@ -96,20 +99,28 @@ if __name__ == '__main__':
|
||||
Common issues and their solutions:
|
||||
|
||||
### Out of Memory Errors
|
||||
|
||||
If you encounter CUDA out of memory errors:
|
||||
|
||||
- Reduce `num_frames` or video resolution
|
||||
- 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
|
||||
|
||||
To speed up generation:
|
||||
|
||||
- Reduce `num_inference_steps` (20-30 is usually sufficient)
|
||||
- Use half precision (`fp16`) for the VAE
|
||||
- Use multiple GPUs if available
|
||||
|
||||
### Unexpected Results
|
||||
|
||||
If the generated video doesn't match your prompt:
|
||||
|
||||
- Try increasing `guidance_scale` (7.0-9.0 works well)
|
||||
- Make your prompt more detailed and specific
|
||||
- Experiment with different random seeds
|
||||
@@ -117,8 +128,8 @@ If the generated video doesn't match your prompt:
|
||||
|
||||
## Next Steps
|
||||
|
||||
- Learn about [Advanced Inference Configurations](#inference-configuration)
|
||||
- Learn about using [Optimizations](#inference-optimizations)
|
||||
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
|
||||
- Learn about [Advanced Inference Configurations](configuration.md)
|
||||
- Learn about using [Optimizations](optimizations.md)
|
||||
- See [Examples](examples/examples_inference_index.md) for more usage scenarios
|
||||
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
@@ -7,13 +7,13 @@ This page describes the various options for speeding up generation times in Fast
|
||||
|
||||
- Optimized Attention Backends
|
||||
|
||||
- [Flash Attention](#optimizations-flash)
|
||||
- [Sliding Tile Attention](#optimizations-sta)
|
||||
- [Sage Attention](#optimizations-sage)
|
||||
- [Sage Attention 3](#optimizations-sage3)
|
||||
- [Flash Attention](#flash-attention)
|
||||
- [Sliding Tile Attention](#sliding-tile-attention)
|
||||
- [Sage Attention](#sage-attention)
|
||||
- [Sage Attention 3](#sage-attention-3)
|
||||
|
||||
- Caching Techniques
|
||||
- [TeaCache](#optimizations-teacache)
|
||||
- [TeaCache](#teacache)
|
||||
|
||||
## Attention Backends
|
||||
|
||||
@@ -74,7 +74,7 @@ python setup.py install
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
Please see [this page](../attention/sta/index.md) for more installation instructions.
|
||||
|
||||
### Video Sparse Attention
|
||||
|
||||
@@ -85,7 +85,7 @@ git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
Please see [this page](#vsa-installation) for more installation instructions.
|
||||
Please see [this page](../attention/vsa/index.md) for more installation instructions.
|
||||
|
||||
### Sage Attention
|
||||
|
||||
|
||||
@@ -40,20 +40,26 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
|
||||
}
|
||||
</style>
|
||||
|
||||
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA |
|
||||
|------------|---------------------|-------------|----------|-------------------|-----------|-----|
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
|
||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
|
||||
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
|
||||
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ |
|
||||
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ |
|
||||
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA | BSA |
|
||||
|------------|---------------------|-------------|----------|-------------------|-----------|-----|-----|
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| TurboWan2.1 T2V 1.3B | `loayrashid/TurboWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| TurboWan2.1 T2V 14B | `loayrashid/TurboWan2.1-T2V-14B-Diffusers` | 480P, 720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| LongCat T2V 13.6B | See note** | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
|
||||
| Matrix Game 2.0 Base | `FastVideo/Matrix-Game-2.0-Base-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 GTA | `FastVideo/Matrix-Game-2.0-GTA-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 TempleRun | `FastVideo/Matrix-Game-2.0-TempleRun-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
@@ -64,3 +70,13 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
|
||||
|
||||
### Sliding Tile Attention
|
||||
- Currently only Hopper GPUs (H100s) are supported.
|
||||
|
||||
### TurboWan2.1 (TurboDiffusion)
|
||||
- Uses TurboDiffusionPipeline with RCM scheduler for 1-4 step generation
|
||||
- Requires SLA attention backend: `export FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
|
||||
- Uses `guidance_scale=1.0` (no classifier-free guidance)
|
||||
|
||||
### Matrix Game 2.0
|
||||
- Image-to-video game world models with keyboard/mouse control input
|
||||
- Three variants available: Base (universal), GTA, and TempleRun
|
||||
- Each variant has different keyboard dimensions for control inputs
|
||||
|
||||
@@ -1,45 +1,130 @@
|
||||
# 🧱 Data Preprocessing
|
||||
|
||||
# 🧱 Data Preprocess
|
||||
To save GPU memory during training, FastVideo precomputes text embeddings and VAE latents. This eliminates the need to load the text encoder and VAE during training.
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
## Quick Start
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
Download the sample dataset and run preprocessing:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=data/mini_i2v_dataset --repo_type=dataset
|
||||
# Download the crush-smol dataset
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id "wlsaidhi/crush-smol-merged" \
|
||||
--local_dir "data/crush-smol" \
|
||||
--repo_type "dataset"
|
||||
|
||||
# Run preprocessing
|
||||
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v_new.sh
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
## Preprocessing Pipeline
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
The new preprocessing pipeline supports multiple dataset formats and video loaders:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
```bash
|
||||
GPU_NUM=2
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATASET_PATH="data/crush-smol/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path $MODEL_PATH \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 77 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.samples_per_file 8 \
|
||||
--preprocess.flush_frequency 8 \
|
||||
--preprocess.video_length_tolerance_range 5
|
||||
```
|
||||
|
||||
## Process your own dataset
|
||||
### Key Parameters
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
| Parameter | Description |
|
||||
|-----------|-------------|
|
||||
| `--workload_type` | Task type: `t2v` (text-to-video) or `i2v` (image-to-video) |
|
||||
| `--preprocess.dataset_type` | Input format: `hf` (HuggingFace) or `merged` (local folder) |
|
||||
| `--preprocess.dataset_path` | Path to dataset (HF repo ID or local folder) |
|
||||
| `--preprocess.dataset_output_dir` | Output directory for Parquet files |
|
||||
| `--preprocess.video_loader_type` | Video decoder: `torchcodec` or `torchvision` |
|
||||
| `--preprocess.max_height` / `max_width` | Target resolution for videos |
|
||||
| `--preprocess.num_frames` | Number of frames to extract per video |
|
||||
| `--preprocess.train_fps` | Target FPS for frame extraction |
|
||||
|
||||
## Dataset Formats
|
||||
|
||||
### Merged Dataset (Local Folder)
|
||||
|
||||
Structure your dataset as follows:
|
||||
|
||||
```
|
||||
path_to_your_dataset_folder/
|
||||
your_dataset/
|
||||
├── videos/
|
||||
│ ├── video_001.mp4
|
||||
│ ├── video_002.mp4
|
||||
│ └── ...
|
||||
└── videos2caption.json
|
||||
```
|
||||
|
||||
The `videos2caption.json` maps video filenames to captions:
|
||||
|
||||
```json
|
||||
[
|
||||
{"path": "video_001.mp4", "cap": "A cat playing with yarn..."},
|
||||
{"path": "video_002.mp4", "cap": "Ocean waves at sunset..."}
|
||||
]
|
||||
```
|
||||
|
||||
### HuggingFace Dataset
|
||||
|
||||
Use `--preprocess.dataset_type hf` and point `--preprocess.dataset_path` to a HuggingFace dataset with `video` and `caption` columns.
|
||||
|
||||
## Creating Your Own Dataset
|
||||
|
||||
If you have raw videos and captions in separate files, generate the `videos2caption.json`:
|
||||
|
||||
```bash
|
||||
python scripts/dataset_preparation/prepare_json_file.py \
|
||||
--data_folder path/to/your_raw_data/ \
|
||||
--output path/to/output_folder
|
||||
```
|
||||
|
||||
Your raw data folder should contain:
|
||||
|
||||
```
|
||||
your_raw_data/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
│ ├── 1.mp4
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
│ └── ...
|
||||
├── videos.txt # list of video filenames
|
||||
└── prompt.txt # corresponding captions (one per line)
|
||||
```
|
||||
|
||||
To generate the `videos2caption.json` and `merge.txt`, run
|
||||
## Output Format
|
||||
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
Preprocessing outputs Parquet files in the `combined_parquet_dataset/` subdirectory containing:
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
- `vae_latent_bytes` — VAE-encoded video latent
|
||||
- `text_embedding_bytes` — text encoder output
|
||||
- `clip_feature_bytes` — CLIP image features (I2V only)
|
||||
- `first_frame_latent_bytes` — first frame latent (I2V only)
|
||||
- Metadata: shapes, dtypes, and sample identifiers
|
||||
|
||||
```
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
```
|
||||
## Examples
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
See ready-to-run preprocessing scripts in the training examples:
|
||||
|
||||
- **T2V**: `examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v_new.sh`
|
||||
- **I2V**: `examples/training/finetune/wan_i2v_14B_480p/crush_smol/preprocess_wan_data_i2v_new.sh`
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
+153
-55
@@ -1,78 +1,176 @@
|
||||
# 🧠 Finetuning
|
||||
|
||||
# 🧠 Finetune
|
||||
## ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
This guide covers finetuning video diffusion models with FastVideo, including full finetuning and LoRA.
|
||||
|
||||
## Training Arguments
|
||||
|
||||
FastVideo training scripts use several argument groups:
|
||||
|
||||
### Training Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--max_train_steps` | Total training steps |
|
||||
| `--train_batch_size` | Batch size per GPU |
|
||||
| `--gradient_accumulation_steps` | Steps to accumulate before optimizer update |
|
||||
| `--num_latent_t` | Temporal latent dimension (reduce to save memory) |
|
||||
| `--num_height` / `--num_width` | Video resolution |
|
||||
| `--num_frames` | Number of frames per video |
|
||||
| `--output_dir` | Directory for checkpoints |
|
||||
|
||||
### Parallelism Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--num_gpus` | Total number of GPUs |
|
||||
| `--sp_size` | Sequence parallel size (increase to reduce memory per GPU) |
|
||||
| `--tp_size` | Tensor parallel size |
|
||||
| `--hsdp_replicate_dim` | HSDP replication dimension |
|
||||
| `--hsdp_shard_dim` | HSDP sharding dimension |
|
||||
|
||||
### Optimizer Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--learning_rate` | Base learning rate |
|
||||
| `--mixed_precision` | Precision mode (`bf16` recommended) |
|
||||
| `--weight_decay` | Weight decay for regularization |
|
||||
| `--max_grad_norm` | Gradient clipping threshold |
|
||||
|
||||
### Validation Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--log_validation` | Enable validation logging |
|
||||
| `--validation_dataset_file` | JSON file with validation prompts |
|
||||
| `--validation_steps` | Run validation every N steps |
|
||||
| `--validation_sampling_steps` | Inference steps for validation |
|
||||
| `--validation_guidance_scale` | CFG scale for validation |
|
||||
|
||||
## Full Finetuning
|
||||
|
||||
Full finetuning updates all model weights. This provides the best quality but requires more GPU memory.
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
# Example: Wan2.1 T2V 1.3B full finetune (4 GPUs)
|
||||
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
|
||||
```
|
||||
|
||||
Download the original model weights as specified in the [Distillation Section](../distillation/dmd.md):
|
||||
**Typical settings:**
|
||||
|
||||
Then you can run the finetune with:
|
||||
- Learning rate: `1e-5` to `5e-5`
|
||||
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
|
||||
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
## LoRA Finetuning
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
|
||||
|
||||
### LoRA-Specific Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--lora_training True` | Enable LoRA mode |
|
||||
| `--lora_rank` | Rank of LoRA adapters (16, 32, 64, 128) |
|
||||
|
||||
### Learning Rate for LoRA
|
||||
|
||||
**Important:** LoRA typically requires a **10–20× higher learning rate** than full finetuning because only the low-rank adapters are being trained while the base model is frozen.
|
||||
|
||||
| Training Mode | Recommended Learning Rate |
|
||||
|---------------|---------------------------|
|
||||
| Full finetune | `1e-5` to `5e-5` |
|
||||
| LoRA | `1e-4` to `2e-4` |
|
||||
|
||||
### Example LoRA Training
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
# Example: Wan2.1 T2V 1.3B LoRA finetune (1 GPU)
|
||||
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v_lora.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
Key differences from full finetune:
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
- Add `--lora_training True --lora_rank 32`
|
||||
- Use higher learning rate (10–20× full finetune)
|
||||
- Can run on fewer GPUs (even single GPU)
|
||||
- Outputs adapter weights instead of full model
|
||||
|
||||
## LoRA Extraction and Merging
|
||||
|
||||
FastVideo provides tools to extract LoRA adapters from finetuned models and merge them back.
|
||||
|
||||
### Extract LoRA Adapter
|
||||
|
||||
Extract a LoRA adapter by comparing a finetuned model to its base:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
python scripts/lora_extraction/extract_lora.py \
|
||||
--base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--ft path/to/your/finetuned_model \
|
||||
--out adapter_r32.safetensors \
|
||||
--rank 32
|
||||
```
|
||||
|
||||
### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--base` | Base model (HuggingFace ID or local path) |
|
||||
| `--ft` | Finetuned model path |
|
||||
| `--out` | Output adapter file (.safetensors) |
|
||||
| `--rank` | LoRA rank (16, 32, 64, 128) |
|
||||
| `--full-rank` | Extract full-rank adapter (optional) |
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
### Merge LoRA Adapter
|
||||
|
||||
### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
Merge an adapter back into a base model:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
python scripts/lora_extraction/merge_lora.py \
|
||||
--base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--adapter adapter_r32.safetensors \
|
||||
--ft path/to/your/finetuned_model \
|
||||
--output merged_model
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--base` | Base model path |
|
||||
| `--adapter` | LoRA adapter file |
|
||||
| `--ft` | Finetuned model (for config reference) |
|
||||
| `--output` | Output directory for merged model |
|
||||
|
||||
### Validate Merged Model
|
||||
|
||||
Compare the merged model against the original finetuned model:
|
||||
|
||||
```bash
|
||||
python scripts/lora_extraction/lora_inference_comparison.py \
|
||||
--base merged_model \
|
||||
--ft path/to/your/finetuned_model \
|
||||
--adapter NONE \
|
||||
--output-dir results \
|
||||
--prompt "A cat sitting on a windowsill" \
|
||||
--compute-ssim \
|
||||
--compute-lpips
|
||||
```
|
||||
|
||||
## Training Examples
|
||||
|
||||
Ready-to-run training scripts are available for multiple models:
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
| Model | Type | Example |
|
||||
|-------|------|---------|
|
||||
| Wan2.1 T2V 1.3B | T2V | `examples/training/finetune/wan_t2v_1.3B/crush_smol/` |
|
||||
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
|
||||
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
|
||||
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
- `preprocess_*.sh` — run preprocessing
|
||||
- `finetune_*.sh` — full finetune launcher
|
||||
- `finetune_*_lora.sh` — LoRA finetune launcher
|
||||
- `validation.json` — validation prompts
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# Training Overview
|
||||
|
||||
FastVideo supports finetuning video diffusion models on custom datasets. This page explains what data you need and how to get started.
|
||||
|
||||
## Data Requirements
|
||||
|
||||
To save GPU memory during training, FastVideo precomputes embeddings and latents ahead of time. This eliminates the need to load the text encoder and VAE during training, significantly reducing memory usage.
|
||||
|
||||
### Text-to-Video (T2V) Finetuning
|
||||
|
||||
For T2V models, you need:
|
||||
|
||||
| Component | Description |
|
||||
|-----------|-------------|
|
||||
| **Text embeddings** | Precomputed embeddings from the model's text encoder (e.g., T5 or LLaMA). Stored as numpy arrays in Parquet files. |
|
||||
| **Video latents** | VAE-encoded representations of your training videos. Each video is encoded into a compressed latent tensor. |
|
||||
|
||||
### Image-to-Video (I2V) Finetuning
|
||||
|
||||
For I2V models, you need everything from T2V plus additional image conditioning. Note that not all I2V architectures require encoded images—this depends on how the model conditions on the input frame. Wan2.1 and Wan2.2 A14B I2V models do require these additional components:
|
||||
|
||||
| Component | Description |
|
||||
|-----------|-------------|
|
||||
| **Text embeddings** | Same as T2V—precomputed from the text encoder. |
|
||||
| **Video latents** | Same as T2V—VAE-encoded video representations. |
|
||||
| **First frame latent** | VAE-encoded representation of the first frame, used as the conditioning image. |
|
||||
| **CLIP features** | Image embeddings from a CLIP vision encoder for the conditioning frame. |
|
||||
|
||||
## Preprocessing
|
||||
|
||||
Before training, you need to preprocess your raw videos and captions into Parquet files containing precomputed latents and embeddings.
|
||||
|
||||
FastVideo supports two input formats:
|
||||
|
||||
- **HuggingFace datasets** — load directly from HF Hub or local HF datasets
|
||||
- **Merged datasets** — local folder with videos and a `videos2caption.json` metadata file
|
||||
|
||||
**→ See [Data Preprocessing](data_preprocess.md) for full details and examples.**
|
||||
|
||||
## Training Examples
|
||||
|
||||
Ready-to-run examples with preprocessing scripts, training launchers, and validation configs are available for multiple models and datasets:
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
- `preprocess_*.sh` — run preprocessing
|
||||
- `finetune_*.sh` — launch training (full finetune or LoRA)
|
||||
- `validation.json` — validation prompts for checkpoints
|
||||
|
||||
## Training Methods
|
||||
|
||||
FastVideo supports several training approaches:
|
||||
|
||||
| Method | Use Case |
|
||||
|--------|----------|
|
||||
| **Full finetune** | Adapt entire model to a new domain or style |
|
||||
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
|
||||
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
|
||||
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
|
||||
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
@@ -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
|
||||
@@ -0,0 +1,19 @@
|
||||
# Self-Forcing Distillation for SFWan2.1 T2V 1.3B
|
||||
|
||||
These scripts demonstrate self-forcing distillation (SFwan) for the causal Wan2.1 T2V 1.3B model. The workflow mirrors DMD2 while injecting self-forcing blocks so the student can autoregressively refine later frames.
|
||||
|
||||
## Run the recipe
|
||||
1. Download the preprocessed text-video dataset:
|
||||
```bash
|
||||
bash examples/distill/SFWan2.1-T2V/download_dataset.sh
|
||||
```
|
||||
2. (Optional) Regenerate parquet shards locally:
|
||||
```bash
|
||||
bash examples/distill/SFWan2.1-T2V/preprocess_data.sh
|
||||
```
|
||||
3. Launch self-forcing distillation with your cluster settings:
|
||||
```bash
|
||||
sbatch examples/distill/SFWan2.1-T2V/distill_dmd_t2v_1.3B.sh
|
||||
```
|
||||
|
||||
Update the dataset paths and wandb credentials inside the script before running on your environment.
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def main():
|
||||
# Point this to your local diffusers model dir (or replace with a HF model ID).
|
||||
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A high-definition video captures the precision of robotic welding in an industrial setting. The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. The welding process is in full swing, with bright sparks and intense light illuminating the scene, creating a vivid display of blue and white hues. A significant amount of smoke billows around the welding area, partially obscuring the view but emphasizing the heat and activity. The background reveals parts of the workshop environment, including a ventilation system and various pieces of machinery, indicating a busy and functional industrial workspace. As the video progresses, the robotic arm maintains its steady position, continuing the welding process and moving to its left. The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. The metal surface beneath the torch shows ongoing signs of heating and melting. The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, underscoring the ongoing nature of the welding operation."
|
||||
)
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
negative_prompt="",
|
||||
height=704,
|
||||
width=1280,
|
||||
num_frames=77,
|
||||
num_inference_steps=35,
|
||||
guidance_scale=7.0,
|
||||
fps=24,
|
||||
output_path="outputs_video/cosmos2_5_t2w.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
|
||||
from fastvideo.models.dits.matrix_game.utils import get_current_action_async, expand_action_to_frames
|
||||
|
||||
import torch
|
||||
import asyncio
|
||||
|
||||
# Available variants: "base_distilled_model", "gta_distilled_model", "templerun_distilled_model"
|
||||
# Each variant has different keyboard_dim:
|
||||
# - base_distilled_model: keyboard_dim=4
|
||||
# - gta_distilled_model: keyboard_dim=2
|
||||
# - templerun_distilled_model: keyboard_dim=7 (keyboard only, no mouse)
|
||||
MODEL_VARIANT = "base_distilled_model"
|
||||
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"mode": "universal",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"mode": "gta_drive",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"mode": "templerun",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
async def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
config = VARIANT_CONFIG[MODEL_VARIANT]
|
||||
|
||||
generator = StreamingVideoGenerator.from_pretrained(
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
max_blocks = 50
|
||||
num_frames = 597
|
||||
actions = {
|
||||
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
|
||||
"mouse": torch.zeros((num_frames, 2))
|
||||
}
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
mode = config["mode"]
|
||||
|
||||
generator.reset(
|
||||
prompt="",
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
)
|
||||
print("Initialization complete.")
|
||||
|
||||
for block_id in range(max_blocks):
|
||||
print(f"\n=== Block {block_id + 1}/{max_blocks} ===")
|
||||
|
||||
action = await get_current_action_async(mode)
|
||||
keyboard_cond, mouse_cond = expand_action_to_frames(action, 12)
|
||||
await generator.step_async(keyboard_cond, mouse_cond)
|
||||
|
||||
if (await asyncio.to_thread(input, "\nContinue? (y/n): ")).lower() == 'n':
|
||||
break
|
||||
|
||||
# Save final video
|
||||
generator.finalize()
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
import os
|
||||
|
||||
# Set SLA attention backend BEFORE fastvideo imports
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_turbodiffusion"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# TurboDiffusion: 1-4 step video generation using RCM scheduler + SLA attention
|
||||
# FastVideo will automatically use TurboDiffusionPipeline when specified
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
|
||||
# set to false if using RTX 4090
|
||||
# pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
# 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 "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic."
|
||||
)
|
||||
video2 = generator.generate_video(
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,49 @@
|
||||
import os
|
||||
|
||||
# Set SLA attention backend BEFORE fastvideo imports
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_turbodiffusion_14B"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# TurboDiffusion 14B: 1-4 step video generation using RCM scheduler + SLA attention
|
||||
# FastVideo will automatically use TurboDiffusionPipeline when specified
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers",
|
||||
# 14B model needs more GPUs
|
||||
num_gpus=2,
|
||||
)
|
||||
|
||||
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 "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic."
|
||||
)
|
||||
video2 = generator.generate_video(
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1,684 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Request, HTTPException
|
||||
from fastapi.responses import HTMLResponse, FileResponse
|
||||
|
||||
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
|
||||
from fastvideo.models.dits.matrix_game.utils import expand_action_to_frames
|
||||
|
||||
|
||||
VARIANT_CONFIG = {
|
||||
"Matrix-Game-2.0-Base": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"mode": "universal",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"Matrix-Game-2.0-GTA": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"mode": "gta_drive",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"Matrix-Game-2.0-TempleRun": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"mode": "templerun",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
MODEL_PATH_MAPPING = {
|
||||
name: config["model_path"] for name, config in VARIANT_CONFIG.items()
|
||||
}
|
||||
|
||||
|
||||
CAM_VALUE = 0.1
|
||||
KEYBOARD_MAP_UNIVERSAL = {
|
||||
"W (Forward)": [1, 0, 0, 0],
|
||||
"S (Back)": [0, 1, 0, 0],
|
||||
"A (Left)": [0, 0, 1, 0],
|
||||
"D (Right)": [0, 0, 0, 1],
|
||||
"Q (Stop)": [0, 0, 0, 0],
|
||||
}
|
||||
KEYBOARD_MAP_GTA = {
|
||||
"W (Forward)": [1, 0],
|
||||
"S (Back)": [0, 1],
|
||||
"Q (Stop)": [0, 0],
|
||||
}
|
||||
KEYBOARD_MAP_TEMPLERUN = {
|
||||
"Q (Run)": [1, 0, 0, 0, 0, 0, 0],
|
||||
"W (Jump)": [0, 1, 0, 0, 0, 0, 0],
|
||||
"S (Slide)": [0, 0, 1, 0, 0, 0, 0],
|
||||
"Z (Turn Left)": [0, 0, 0, 1, 0, 0, 0],
|
||||
"C (Turn Right)": [0, 0, 0, 0, 1, 0, 0],
|
||||
"A (Left)": [0, 0, 0, 0, 0, 1, 0],
|
||||
"D (Right)": [0, 0, 0, 0, 0, 0, 1],
|
||||
}
|
||||
|
||||
|
||||
CAMERA_MAP_UNIVERSAL = {
|
||||
"U (Center)": [0, 0],
|
||||
"I (Up)": [CAM_VALUE, 0],
|
||||
"K (Down)": [-CAM_VALUE, 0],
|
||||
"J (Left)": [0, -CAM_VALUE],
|
||||
"L (Right)": [0, CAM_VALUE],
|
||||
}
|
||||
CAMERA_MAP_GTA = {
|
||||
"Q (Straight)": [0, 0],
|
||||
"A (Steer Left)": [0, -CAM_VALUE],
|
||||
"D (Steer Right)": [0, CAM_VALUE],
|
||||
}
|
||||
|
||||
def setup_model_environment(model_path: str) -> None:
|
||||
# if "fullattn" in model_path.lower():
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
# else:
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
|
||||
|
||||
def create_timing_display(inference_time, total_time, stage_execution_times, num_frames):
|
||||
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
|
||||
|
||||
timing_html = f"""
|
||||
<div style="margin: 10px 0;">
|
||||
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
|
||||
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
|
||||
<div class="timing-card timing-card-highlight">
|
||||
<div style="font-size: 20px;">🚀</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
|
||||
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🧠</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
|
||||
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🎬</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
|
||||
<div style="font-size: 16px; color: #dc2626;">N/A</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🌐</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
|
||||
<div style="font-size: 16px; color: #059669;">N/A</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">📊</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
|
||||
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
|
||||
</div>
|
||||
</div>"""
|
||||
|
||||
if inference_time > 0:
|
||||
fps = num_frames / inference_time
|
||||
timing_html += f"""
|
||||
<div class="performance-card" style="margin-top: 15px;">
|
||||
<span style="font-weight: bold;">Generation Speed: </span>
|
||||
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
|
||||
</div>"""
|
||||
|
||||
return timing_html + "</div>"
|
||||
|
||||
def get_action_tensors(mode: str, keyboard_key: str, mouse_key: str | None):
|
||||
if mode == "universal":
|
||||
keyboard = torch.tensor(KEYBOARD_MAP_UNIVERSAL.get(keyboard_key, [0, 0, 0, 0])).cuda()
|
||||
mouse = torch.tensor(CAMERA_MAP_UNIVERSAL.get(mouse_key, [0, 0])).cuda()
|
||||
elif mode == "gta_drive":
|
||||
keyboard = torch.tensor(KEYBOARD_MAP_GTA.get(keyboard_key, [0, 0])).cuda()
|
||||
mouse = torch.tensor(CAMERA_MAP_GTA.get(mouse_key, [0, 0])).cuda()
|
||||
elif mode == "templerun":
|
||||
keyboard = torch.tensor(KEYBOARD_MAP_TEMPLERUN.get(keyboard_key, [1, 0, 0, 0, 0, 0, 0])).cuda()
|
||||
mouse = None
|
||||
else:
|
||||
raise ValueError(f"Unknown mode: {mode}")
|
||||
|
||||
return {"keyboard": keyboard, "mouse": mouse}
|
||||
|
||||
def create_gradio_interface(generators: dict[str, StreamingVideoGenerator], loaded_model_name: str):
|
||||
initial_config = VARIANT_CONFIG.get(loaded_model_name, VARIANT_CONFIG["Matrix-Game-2.0-Base"])
|
||||
initial_mode = initial_config["mode"]
|
||||
|
||||
if initial_mode == "universal":
|
||||
initial_kb_choices = list(KEYBOARD_MAP_UNIVERSAL.keys())
|
||||
initial_mouse_choices = list(CAMERA_MAP_UNIVERSAL.keys())
|
||||
initial_mouse_visible = True
|
||||
elif initial_mode == "gta_drive":
|
||||
initial_kb_choices = list(KEYBOARD_MAP_GTA.keys())
|
||||
initial_mouse_choices = list(CAMERA_MAP_GTA.keys())
|
||||
initial_mouse_visible = True
|
||||
else: # templerun
|
||||
initial_kb_choices = list(KEYBOARD_MAP_TEMPLERUN.keys())
|
||||
initial_mouse_choices = []
|
||||
initial_mouse_visible = False
|
||||
|
||||
theme = gr.themes.Base().set(
|
||||
button_primary_background_fill="#2563eb",
|
||||
button_primary_background_fill_hover="#1d4ed8",
|
||||
button_primary_text_color="white",
|
||||
slider_color="#2563eb",
|
||||
checkbox_background_color_selected="#2563eb",
|
||||
)
|
||||
|
||||
with gr.Blocks(title="FastVideo - Matrix Game 2.0", theme=theme) as demo:
|
||||
game_state = gr.State({
|
||||
"initialized": False,
|
||||
"current_model": None,
|
||||
"block_idx": 0,
|
||||
"max_blocks": 50,
|
||||
})
|
||||
|
||||
# Header
|
||||
gr.Image("assets/full.svg", show_label=False, container=False, height=80)
|
||||
|
||||
gr.HTML("""
|
||||
<div style="text-align: center; margin-bottom: 10px;">
|
||||
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
|
||||
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
with gr.Accordion("🎥 What Is FastVideo?", open=False):
|
||||
gr.HTML("""
|
||||
<div style="padding: 20px; line-height: 1.6;">
|
||||
<p style="font-size: 16px; margin-bottom: 15px;">
|
||||
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
</p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
# Model Selection
|
||||
with gr.Row():
|
||||
model_selection = gr.Dropdown(
|
||||
choices=[loaded_model_name],
|
||||
value=loaded_model_name,
|
||||
label="Select Model",
|
||||
interactive=False
|
||||
)
|
||||
|
||||
|
||||
# Main Layout
|
||||
with gr.Row(equal_height=True, elem_classes="main-content-row"):
|
||||
with gr.Column(scale=1, elem_classes="advanced-options-column"):
|
||||
with gr.Group():
|
||||
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Game Controls</div>")
|
||||
|
||||
with gr.Group():
|
||||
gr.HTML("<div style='font-size: 14px; margin-bottom: 5px; font-weight: bold;'>🎮 Keyboard Control</div>")
|
||||
keyboard_action = gr.Radio(
|
||||
choices=initial_kb_choices,
|
||||
value=initial_kb_choices[0] if initial_kb_choices else None,
|
||||
label="Movement",
|
||||
show_label=False,
|
||||
interactive=True
|
||||
)
|
||||
|
||||
with gr.Group(visible=initial_mouse_visible) as mouse_group:
|
||||
gr.HTML("<div style='font-size: 14px; margin-bottom: 5px; font-weight: bold;'>🖱️ Mouse/Camera Control</div>")
|
||||
mouse_action = gr.Radio(
|
||||
choices=initial_mouse_choices if initial_mouse_visible else [],
|
||||
value=initial_mouse_choices[0] if initial_mouse_choices else None,
|
||||
label="Camera",
|
||||
show_label=False,
|
||||
interactive=True
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
action_btn = gr.Button("Start", variant="primary")
|
||||
stop_btn = gr.Button("Stop", variant="stop")
|
||||
|
||||
gr.HTML("<div style='margin-top: 15px;'></div>")
|
||||
|
||||
seed = gr.Slider(
|
||||
label="Seed",
|
||||
minimum=0,
|
||||
maximum=1000000,
|
||||
step=1,
|
||||
value=1024,
|
||||
)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
block_counter = gr.Textbox(label="Progress", value="Block: 0 / 50", interactive=False, lines=1)
|
||||
|
||||
|
||||
# Right Column: Video Output
|
||||
with gr.Column(scale=1, elem_classes="video-column"):
|
||||
video_output = gr.Video(
|
||||
label="Generated Video",
|
||||
show_label=True,
|
||||
height=466,
|
||||
width=600,
|
||||
container=True,
|
||||
elem_classes="video-component",
|
||||
autoplay=True
|
||||
)
|
||||
|
||||
# Styles
|
||||
gr.HTML("""
|
||||
<style>
|
||||
.center-button {
|
||||
display: flex !important;
|
||||
justify-content: center !important;
|
||||
height: 100% !important;
|
||||
padding-top: 1.4em !important;
|
||||
}
|
||||
|
||||
.gradio-container {
|
||||
max-width: 1200px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.main {
|
||||
max-width: 1200px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.gr-form, .gr-box, .gr-group {
|
||||
max-width: 1200px !important;
|
||||
}
|
||||
|
||||
.gr-video {
|
||||
max-width: 500px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.main-content-row {
|
||||
display: flex !important;
|
||||
align-items: flex-start !important;
|
||||
min-height: 500px !important;
|
||||
gap: 20px !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
display: flex !important;
|
||||
flex-direction: column !important;
|
||||
flex: 1 !important;
|
||||
min-height: 400px !important;
|
||||
align-items: stretch !important;
|
||||
}
|
||||
|
||||
.video-column > * {
|
||||
margin-top: 0 !important;
|
||||
}
|
||||
|
||||
.video-column .gr-video,
|
||||
.video-component {
|
||||
margin-top: 0 !important;
|
||||
padding-top: 0 !important;
|
||||
}
|
||||
|
||||
.video-column .gr-video .gr-form {
|
||||
margin-top: 0 !important;
|
||||
}
|
||||
|
||||
.advanced-options-column .gr-group,
|
||||
.video-column .gr-video {
|
||||
margin-top: 0 !important;
|
||||
vertical-align: top !important;
|
||||
}
|
||||
|
||||
.advanced-options-column > *:last-child,
|
||||
.video-column > *:last-child {
|
||||
flex-grow: 0 !important;
|
||||
}
|
||||
|
||||
@media (max-width: 1400px) {
|
||||
.main-content-row {
|
||||
min-height: 600px !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
min-height: 600px !important;
|
||||
}
|
||||
}
|
||||
|
||||
@media (max-width: 1200px) {
|
||||
.main-content-row {
|
||||
flex-direction: column !important;
|
||||
align-items: stretch !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
min-height: auto !important;
|
||||
width: 100% !important;
|
||||
}
|
||||
}
|
||||
|
||||
.timing-card {
|
||||
background: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color) !important;
|
||||
padding: 10px;
|
||||
border-radius: 8px;
|
||||
text-align: center;
|
||||
min-height: 80px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.timing-card-highlight {
|
||||
background: var(--background-fill-primary) !important;
|
||||
border: 2px solid var(--color-accent) !important;
|
||||
}
|
||||
|
||||
.performance-card {
|
||||
background: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color) !important;
|
||||
padding: 10px;
|
||||
border-radius: 6px;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
.gr-number input[readonly] {
|
||||
background-color: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color-subdued) !important;
|
||||
cursor: default !important;
|
||||
text-align: center !important;
|
||||
font-weight: 500 !important;
|
||||
}
|
||||
</style>
|
||||
""")
|
||||
|
||||
# UI update based on model selection
|
||||
def on_model_change(model_name):
|
||||
config = VARIANT_CONFIG.get(model_name, VARIANT_CONFIG["Matrix-Game-2.0-Base"])
|
||||
mode = config["mode"]
|
||||
|
||||
if mode == "universal":
|
||||
kb_choices = list(KEYBOARD_MAP_UNIVERSAL.keys())
|
||||
mouse_choices = list(CAMERA_MAP_UNIVERSAL.keys())
|
||||
mouse_visible = True
|
||||
elif mode == "gta_drive":
|
||||
kb_choices = list(KEYBOARD_MAP_GTA.keys())
|
||||
mouse_choices = list(CAMERA_MAP_GTA.keys())
|
||||
mouse_visible = True
|
||||
else: # templerun
|
||||
kb_choices = list(KEYBOARD_MAP_TEMPLERUN.keys())
|
||||
mouse_choices = []
|
||||
mouse_visible = False
|
||||
|
||||
return (
|
||||
gr.update(choices=kb_choices, value=kb_choices[0] if kb_choices else None),
|
||||
gr.update(choices=mouse_choices, value=mouse_choices[0] if mouse_choices else None, visible=mouse_visible),
|
||||
gr.update(visible=mouse_visible),
|
||||
)
|
||||
|
||||
model_selection.change(
|
||||
fn=on_model_change,
|
||||
inputs=model_selection,
|
||||
outputs=[keyboard_action, mouse_action, mouse_group]
|
||||
)
|
||||
|
||||
def start_game(model_name, seed_val, randomize, state):
|
||||
if randomize:
|
||||
seed_val = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
config = VARIANT_CONFIG.get(model_name)
|
||||
if not config:
|
||||
return state, seed_val, "Block: 0 / 50", None, "", gr.update(), gr.update()
|
||||
|
||||
generator = generators.get(config["model_path"])
|
||||
if not generator:
|
||||
return state, seed_val, "Block: 0 / 50", None, "", gr.update(), gr.update()
|
||||
|
||||
# If already initialized, clean up first
|
||||
if state.get("initialized"):
|
||||
try:
|
||||
# Clear accumulated frames without saving
|
||||
generator.accumulated_frames = []
|
||||
generator.executor.execute_streaming_clear()
|
||||
except Exception as e:
|
||||
print(f"Warning: cleanup error: {e}")
|
||||
|
||||
# Streaming parameters
|
||||
num_latent_frames_per_block = 3
|
||||
max_blocks = 50
|
||||
total_latent_frames = num_latent_frames_per_block * max_blocks
|
||||
num_frames = (total_latent_frames - 1) * 4 + 1
|
||||
|
||||
actions = {
|
||||
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
|
||||
"mouse": torch.zeros((num_frames, 2))
|
||||
}
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
|
||||
output_dir = os.path.abspath("outputs/matrixgame")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
video_path = os.path.join(output_dir, f"video_{int(time.time())}.mp4")
|
||||
|
||||
generator.reset(
|
||||
prompt="",
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
output_path=video_path,
|
||||
)
|
||||
|
||||
new_state = {
|
||||
"initialized": True,
|
||||
"current_model": model_name,
|
||||
"block_idx": 0,
|
||||
"max_blocks": max_blocks,
|
||||
"video_path": video_path,
|
||||
"frames_per_block": num_latent_frames_per_block * 4,
|
||||
"mode": config["mode"],
|
||||
"seed": seed_val,
|
||||
}
|
||||
|
||||
return new_state, seed_val, "Block: 0 / 50", None, gr.update(value="Step"), gr.update(interactive=True)
|
||||
|
||||
async def step_game(keyboard_key, mouse_key, model_name, state):
|
||||
if not state.get("initialized"):
|
||||
return state, state.get("seed", 0), "Block: 0 / 50", None, gr.update(), gr.update()
|
||||
|
||||
# total_start_time = time.time()
|
||||
config = VARIANT_CONFIG.get(model_name)
|
||||
generator = generators.get(config["model_path"])
|
||||
mode = state["mode"]
|
||||
frames_per_block = state["frames_per_block"]
|
||||
|
||||
# Parse inputs to tensors
|
||||
action = get_action_tensors(mode, keyboard_key, mouse_key)
|
||||
keyboard_cond, mouse_cond = expand_action_to_frames(action, frames_per_block)
|
||||
|
||||
# run step async
|
||||
# inference_start_time = time.time()
|
||||
frames, block_future = await generator.step_async(keyboard_cond, mouse_cond)
|
||||
# inference_time = time.time() - inference_start_time
|
||||
|
||||
# wait for block file to be written
|
||||
block_path = await asyncio.to_thread(block_future.result) if block_future else None
|
||||
state["block_idx"] = generator.block_idx
|
||||
block_str = f"Block: {state['block_idx']} / {state['max_blocks']}"
|
||||
|
||||
# total_time = time.time() - total_start_time
|
||||
|
||||
# Timing breakdown
|
||||
# timing_html = create_timing_display(inference_time, total_time, [], frames_per_block)
|
||||
|
||||
return state, state.get("seed", 0), block_str, block_path, gr.update(), gr.update()
|
||||
|
||||
def stop_game(model_name, state):
|
||||
if not state.get("initialized"):
|
||||
return {"initialized": False}, 0, "Block: 0 / 50", None, gr.update(value="Start"), gr.update(interactive=False)
|
||||
|
||||
config = VARIANT_CONFIG.get(model_name)
|
||||
generator = generators.get(config["model_path"])
|
||||
|
||||
final_path = state.get("video_path")
|
||||
generator.finalize(final_path)
|
||||
|
||||
return {"initialized": False}, state.get("seed", 0), "Block: 0 / 50", final_path, gr.update(value="Start"), gr.update(interactive=False)
|
||||
|
||||
async def handle_action(keyboard_key, mouse_key, model_name, seed_val, randomize, state):
|
||||
if not state.get("initialized"):
|
||||
return start_game(model_name, seed_val, randomize, state)
|
||||
else:
|
||||
return await step_game(keyboard_key, mouse_key, model_name, state)
|
||||
|
||||
action_btn.click(
|
||||
fn=handle_action,
|
||||
inputs=[keyboard_action, mouse_action, model_selection, seed, randomize_seed, game_state],
|
||||
outputs=[game_state, seed_output, block_counter, video_output, action_btn, stop_btn]
|
||||
)
|
||||
|
||||
stop_btn.click(
|
||||
fn=stop_game,
|
||||
inputs=[model_selection, game_state],
|
||||
outputs=[game_state, seed_output, block_counter, video_output, action_btn, stop_btn]
|
||||
)
|
||||
|
||||
gr.HTML("""
|
||||
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
|
||||
<p style="font-size: 16px; margin: 0;">Note that this demo is meant to showcase Matrix Game's quality and that under a large number of requests, generation speed may be affected.</p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Matrix Game Gradio Demo")
|
||||
parser.add_argument("--model", type=str, default="Matrix-Game-2.0-Base",
|
||||
choices=list(VARIANT_CONFIG.keys()),
|
||||
help="Model variant to load")
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0")
|
||||
parser.add_argument("--port", type=int, default=7860)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Load the selected model
|
||||
config = VARIANT_CONFIG[args.model]
|
||||
model_path = config["model_path"]
|
||||
|
||||
print(f"Loading model: {model_path}")
|
||||
setup_model_environment(model_path)
|
||||
generator = StreamingVideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
generators = {model_path: generator}
|
||||
|
||||
demo = create_gradio_interface(generators, args.model)
|
||||
|
||||
print(f"Starting Gradio at http://{args.host}:{args.port}")
|
||||
|
||||
# FastAPI Wrapper
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/logo.png")
|
||||
def get_logo():
|
||||
return FileResponse(
|
||||
"assets/full.svg",
|
||||
media_type="image/svg+xml",
|
||||
headers={
|
||||
"Cache-Control": "public, max-age=3600",
|
||||
"Access-Control-Allow-Origin": "*"
|
||||
}
|
||||
)
|
||||
|
||||
@app.get("/favicon.ico")
|
||||
def get_favicon():
|
||||
favicon_path = "assets/icon-simple.svg"
|
||||
|
||||
if os.path.exists(favicon_path):
|
||||
return FileResponse(
|
||||
favicon_path,
|
||||
media_type="image/svg+xml",
|
||||
headers={
|
||||
"Cache-Control": "public, max-age=3600",
|
||||
"Access-Control-Allow-Origin": "*"
|
||||
}
|
||||
)
|
||||
else:
|
||||
raise HTTPException(status_code=404, detail="Favicon not found")
|
||||
|
||||
@app.get("/", response_class=HTMLResponse)
|
||||
def index(request: Request):
|
||||
base_url = str(request.base_url).rstrip('/')
|
||||
return f"""
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
|
||||
<title>FastVideo - Matrix Game 2.0</title>
|
||||
<meta name="title" content="MatrixGame2.0">
|
||||
<meta name="description" content="Make video generation go blurrrrrrr">
|
||||
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, Matrix Game 2.0">
|
||||
|
||||
<meta property="og:type" content="website">
|
||||
<meta property="og:url" content="{base_url}/">
|
||||
<meta property="og:title" content="FastVideo - Matrix Game 2.0">
|
||||
<meta property="og:description" content="Make video generation go blurrrrrrr">
|
||||
<meta property="og:image" content="{base_url}/logo.png">
|
||||
<meta property="og:image:width" content="1200">
|
||||
<meta property="og:image:height" content="630">
|
||||
<meta property="og:site_name" content="MatrixGame2.0">
|
||||
|
||||
<meta property="twitter:card" content="summary_large_image">
|
||||
<meta property="twitter:url" content="{base_url}/">
|
||||
<meta property="twitter:title" content="MatrixGame2.0">
|
||||
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
|
||||
<meta property="twitter:image" content="{base_url}/logo.png">
|
||||
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
|
||||
<link rel="icon" type="image/png" sizes="16x16" href="/favicon.ico">
|
||||
<link rel="apple-touch-icon" href="/favicon.ico">
|
||||
<style>
|
||||
body, html {{
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
height: 100%;
|
||||
overflow: hidden;
|
||||
}}
|
||||
iframe {{
|
||||
width: 100%;
|
||||
height: 100vh;
|
||||
border: none;
|
||||
}}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
|
||||
</body>
|
||||
</html>
|
||||
"""
|
||||
|
||||
app = gr.mount_gradio_app(
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
|
||||
)
|
||||
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,46 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import argparse
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
|
||||
|
||||
|
||||
def main(text_encoder_path: str):
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
# AbsMaxFP8 is the quantization method used by ComfyUI;
|
||||
# check fastvideo/layers/quantization/* for more quantization methods
|
||||
override_text_encoder_quant="AbsMaxFP8",
|
||||
# for Wan 2.2, this is the path to "umt5_xxl_fp8_e4m3fn_scaled.safetensors"
|
||||
override_text_encoder_safetensors=text_encoder_path,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
)
|
||||
|
||||
# I2V is triggered just by passing in an image_path argument
|
||||
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."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
video = generator.generate_video(
|
||||
prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--text_encoder_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the quantized text encoder safetensors file.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args.text_encoder_path)
|
||||
@@ -10,6 +10,8 @@ if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
else()
|
||||
enable_language(CUDA)
|
||||
# Ensure CUDA toolkit targets (CUDA::cudart, CUDA::cuda_driver, etc.) are available.
|
||||
find_package(CUDAToolkit REQUIRED)
|
||||
endif()
|
||||
|
||||
# Import common utils if needed, but we keep it simple for now
|
||||
@@ -153,6 +155,30 @@ if(BUILD_CXX_KERNELS)
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${CUDA_FLAGS}>
|
||||
)
|
||||
|
||||
# Link against Torch libraries to avoid undefined symbols at import time
|
||||
# (e.g., torch::autograd vtables) when loading the extension module.
|
||||
target_link_libraries(fastvideo_kernel_ops PRIVATE ${TORCH_LIBRARIES})
|
||||
|
||||
# Also link against libtorch_python to satisfy Python-binding symbols
|
||||
# (e.g., torch::PyWarningHandler) required by torch/extension.h.
|
||||
execute_process(
|
||||
COMMAND "${Python_EXECUTABLE}" -c "import torch; from pathlib import Path; p=Path(torch.__file__).parent/'lib'; m=sorted(p.glob('libtorch_python*')); print(str(m[0]) if m else '')"
|
||||
OUTPUT_VARIABLE TORCH_PYTHON_LIBRARY_PATH
|
||||
OUTPUT_STRIP_TRAILING_WHITESPACE
|
||||
ERROR_QUIET
|
||||
)
|
||||
if(TORCH_PYTHON_LIBRARY_PATH)
|
||||
message(STATUS "TORCH_PYTHON_LIBRARY_PATH: ${TORCH_PYTHON_LIBRARY_PATH}")
|
||||
target_link_libraries(fastvideo_kernel_ops PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
|
||||
else()
|
||||
message(WARNING "Could not locate libtorch_python; fastvideo_kernel_ops may fail to import.")
|
||||
endif()
|
||||
|
||||
# Link CUDA runtime + driver explicitly (fixes missing symbols like cuGetErrorString at import time)
|
||||
if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
target_link_libraries(fastvideo_kernel_ops PRIVATE CUDA::cudart CUDA::cuda_driver)
|
||||
endif()
|
||||
|
||||
# We install it to fastvideo_kernel/_C so we can load it to register the ops
|
||||
install(TARGETS fastvideo_kernel_ops LIBRARY DESTINATION fastvideo_kernel/_C)
|
||||
endif()
|
||||
|
||||
@@ -34,7 +34,7 @@ from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_att
|
||||
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
|
||||
|
||||
# Example: Video Sparse Attention (with Triton fallback)
|
||||
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
|
||||
out = video_sparse_attn(q, k, v, block_sizes, block_sizes, topk=5)
|
||||
|
||||
# Example: VMoBA
|
||||
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
|
||||
|
||||
@@ -639,7 +639,8 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
|
||||
// store kq and vq
|
||||
|
||||
// ensuring all writes are finished
|
||||
// ! the following two line seems unnecessary.
|
||||
// tma::store_async_wait(); // ensure qg is finished
|
||||
__syncthreads();
|
||||
|
||||
warpgroup::store(kg_smem[0], kg_reg);
|
||||
@@ -660,145 +661,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
tma::store_async_wait();
|
||||
}
|
||||
|
||||
|
||||
template<int D>
|
||||
void block_sparse_attention_forward_impl(
|
||||
bf16* d_q, bf16* d_k, bf16* d_v, float* d_l, bf16* d_o,
|
||||
int batch, int qo_heads, int kv_heads, int seq_len, int hr,
|
||||
int max_kv_blocks_per_q,
|
||||
int32_t* q2k_block_sparse_index_ptr,
|
||||
int32_t* q2k_block_sparse_num_ptr,
|
||||
int32_t* block_size_ptr,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
using K = fwd_attend_ker_tile_dims<D>;
|
||||
using q_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
using k_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using v_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
|
||||
using o_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<D>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
|
||||
globals g{
|
||||
qg_arg, kg_arg, vg_arg, lg_arg, og_arg,
|
||||
static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q),
|
||||
q2k_block_sparse_index_ptr, q2k_block_sparse_num_ptr, block_size_ptr
|
||||
};
|
||||
|
||||
// Shared memory size for the kernel
|
||||
// 54000 bytes is calibrated for H100 shared memory constraints for these tile sizes
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<D>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
|
||||
fwd_attend_ker<D><<<grid, (128), mem_size, stream>>>(g);
|
||||
}
|
||||
|
||||
template<int D>
|
||||
void block_sparse_attention_backward_impl(
|
||||
bf16* d_q, bf16* d_k, bf16* d_v, bf16* d_o, bf16* d_og, float* d_l, float* d_d, float* d_qg, float* d_kg, float* d_vg,
|
||||
int batch, int qo_heads, int kv_heads, int seq_len, int hr, int max_q_blocks_per_kv,
|
||||
int32_t* k2q_block_sparse_index_ptr,
|
||||
int32_t* k2q_block_sparse_num_ptr,
|
||||
int32_t* block_size_ptr,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
using G = bwd_attend_ker_tile_dims<D>;
|
||||
using og_tile = st_bf<4*16, D>;
|
||||
using o_tile = st_bf<4*16, D>;
|
||||
using d_tile = col_vec<st_fl<4*16, D>>;
|
||||
|
||||
using og_global = gl<bf16, -1, -1, -1, -1, og_tile>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
using d_global = gl<float, -1, -1, -1, -1, d_tile>;
|
||||
|
||||
using prep_globals = bwd_prep_globals<D>;
|
||||
|
||||
constexpr int mem_size_prep = kittens::MAX_SHARED_MEMORY;
|
||||
int threads_prep = PREP_NUM_WARPS * kittens::WARP_THREADS;
|
||||
dim3 grid_bwd_prep(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
bwd_attend_prep_ker<D>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size_prep
|
||||
);
|
||||
bwd_attend_prep_ker<D><<<grid_bwd_prep, threads_prep, mem_size_prep, stream>>>(bwd_g);
|
||||
|
||||
using bwd_q_tile = st_bf<G::tile_h_qo, G::tile_width>;
|
||||
using bwd_k_tile = st_bf<G::tile_h, G::tile_width>;
|
||||
using bwd_v_tile = st_bf<G::tile_h, G::tile_width>;
|
||||
using bwd_og_tile = st_bf<G::tile_h_qo, G::tile_width>;
|
||||
using bwd_qg_tile = st_fl<G::tile_h_qo, G::tile_width>;
|
||||
using bwd_kg_tile = st_fl<G::tile_h, G::tile_width>;
|
||||
using bwd_vg_tile = st_fl<G::tile_h, G::tile_width>;
|
||||
using bwd_l_tile = row_vec<st_fl<G::tile_h_qo, G::tile_h>>;
|
||||
using bwd_d_tile = row_vec<st_fl<G::tile_h_qo, G::tile_h>>;
|
||||
|
||||
using bwd_q_global = gl<bf16, -1, -1, -1, -1, bwd_q_tile>;
|
||||
using bwd_k_global = gl<bf16, -1, -1, -1, -1, bwd_k_tile>;
|
||||
using bwd_v_global = gl<bf16, -1, -1, -1, -1, bwd_v_tile>;
|
||||
using bwd_og_global = gl<bf16, -1, -1, -1, -1, bwd_og_tile>;
|
||||
using bwd_qg_global = gl<float, -1, -1, -1, -1, bwd_qg_tile>;
|
||||
using bwd_kg_global = gl<float, -1, -1, -1, -1, bwd_kg_tile>;
|
||||
using bwd_vg_global = gl<float, -1, -1, -1, -1, bwd_vg_tile>;
|
||||
using bwd_l_global = gl<float, -1, -1, -1, -1, bwd_l_tile>;
|
||||
using bwd_d_global = gl<float, -1, -1, -1, -1, bwd_d_tile>;
|
||||
|
||||
using bwd_global_args = bwd_globals<D>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
bwd_global_args bwd_global{bwd_q_arg, bwd_k_arg, bwd_v_arg, bwd_og_arg, bwd_qg_arg, bwd_kg_arg, bwd_vg_arg, bwd_l_arg, bwd_d_arg,
|
||||
static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_q_blocks_per_kv),
|
||||
k2q_block_sparse_index_ptr, k2q_block_sparse_num_ptr, block_size_ptr};
|
||||
|
||||
dim3 grid_bwd_main(seq_len/64, qo_heads, batch);
|
||||
int threads_main = 128;
|
||||
// Calibrated shared memory sizes for different head dimensions
|
||||
int bwd_mem_size = (D == 64) ? 72000 : 113000;
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
bwd_attend_ker<D>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
bwd_mem_size
|
||||
);
|
||||
bwd_attend_ker<D><<<grid_bwd_main, threads_main, bwd_mem_size, stream>>>(bwd_global);
|
||||
}
|
||||
|
||||
#include "pyutils/torch_helpers.cuh"
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <iostream>
|
||||
@@ -810,23 +672,32 @@ block_sparse_attention_forward(
|
||||
torch::Tensor v,
|
||||
torch::Tensor q2k_block_sparse_index,
|
||||
torch::Tensor q2k_block_sparse_num,
|
||||
torch::Tensor block_size
|
||||
torch::Tensor kv_block_size
|
||||
)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
// q shape: (batch, qo_heads, q_seq_len, head_dim)
|
||||
// k shape: (batch, kv_heads, kv_seq_len, head_dim)
|
||||
// v shape: (batch, kv_heads, kv_seq_len, head_dim)
|
||||
// q2k_block_sparse_index shape: (batch, qo_heads, num_q_blocks, max_kv_blocks_per_q)
|
||||
// q2k_block_sparse_num shape: (batch, qo_heads, num_q_blocks)
|
||||
// kv_block_size shape: (num_kv_blocks) This does not need other dimensions because across all batch/heads the padding is the same.
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto seq_len = q.size(2);
|
||||
auto q_seq_len = q.size(2);
|
||||
auto kv_seq_len = k.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
|
||||
auto num_q_blocks = block_size.size(0);
|
||||
auto num_q_blocks = q2k_block_sparse_index.size(2);
|
||||
auto num_kv_blocks = kv_block_size.size(0);
|
||||
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
|
||||
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
|
||||
TORCH_CHECK(num_q_blocks * BLOCK_M == q_seq_len, "This kernel supports variable q block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_kv_blocks * BLOCK_M == kv_seq_len, "This kernel supports variable kv block size, but it assumes the input sequence is properly padded.");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -834,11 +705,9 @@ block_sparse_attention_forward(
|
||||
TORCH_CHECK(q2k_block_sparse_index.size(0) == batch, "q2k_block_sparse_index batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(0) == batch, "q2k_block_sparse_num batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_index.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_index idx 2 - must match seq_len / BLOCK_M");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_num idx 2 - must match seq_len / BLOCK_M");
|
||||
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(2) == num_q_blocks, "q2k_block_sparse_num idx 2 - must match num_q_blocks");
|
||||
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
@@ -864,12 +733,12 @@ block_sparse_attention_forward(
|
||||
// for the returned outputs
|
||||
torch::Tensor o = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(head_dim)}, v.options());
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
@@ -880,32 +749,110 @@ block_sparse_attention_forward(
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
// Temporated implementation to avoid code duplication between head_dim=64 and 128
|
||||
if (head_dim == 64) {
|
||||
block_sparse_attention_forward_impl<64>(
|
||||
d_q, d_k, d_v, d_l, d_o,
|
||||
batch, qo_heads, kv_heads, seq_len, hr,
|
||||
max_kv_blocks_per_q,
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<64>::kv_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<64>::kv_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<64>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
|
||||
globals g{
|
||||
qg_arg,
|
||||
kg_arg,
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(q_seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr()),
|
||||
stream
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<64>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
} else if (head_dim == 128) {
|
||||
block_sparse_attention_forward_impl<128>(
|
||||
d_q, d_k, d_v, d_l, d_o,
|
||||
batch, qo_heads, kv_heads, seq_len, hr,
|
||||
max_kv_blocks_per_q,
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
|
||||
fwd_attend_ker<64><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
if (head_dim == 128) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<128>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
|
||||
globals g{
|
||||
qg_arg,
|
||||
kg_arg,
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(q_seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr()),
|
||||
stream
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
} else {
|
||||
TORCH_CHECK(false, "Unsupported head_dim: ", head_dim, ". Only 64 and 128 are supported.");
|
||||
|
||||
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return {o, l_vec};
|
||||
@@ -921,7 +868,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
torch::Tensor og,
|
||||
torch::Tensor k2q_block_sparse_index,
|
||||
torch::Tensor k2q_block_sparse_num,
|
||||
torch::Tensor block_size)
|
||||
torch::Tensor kv_block_size)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
@@ -930,11 +877,23 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
CHECK_INPUT(o);
|
||||
CHECK_INPUT(og);
|
||||
|
||||
// q: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// k: [batch, kv_heads, kv_seq_len, head_dim]
|
||||
// v: [batch, kv_heads, kv_seq_len, head_dim]
|
||||
// o: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// l_vec: [batch, qo_heads, q_seq_len, 1]
|
||||
// og: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// k2q_block_sparse_index: [batch, kv_heads, num_kv_blocks, max_num_q_blocks]
|
||||
// k2q_block_sparse_num: [batch, kv_heads, num_kv_blocks]
|
||||
// kv_block_size: [num_kv_blocks]
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto seq_len = q.size(2);
|
||||
auto q_seq_len = q.size(2);
|
||||
auto kv_seq_len = k.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
|
||||
auto num_kv_blocks = kv_block_size.size(0);
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index.size(2) must match num_kv_blocks (kv_block_size.size(0))");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -945,23 +904,18 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(0) == batch, "k2q_block_sparse_index batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(0) == batch, "k2q_block_sparse_num batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(l_vec.size(2) == seq_len, "L sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(o.size(2) == seq_len, "O sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(og.size(2) == seq_len, "OG sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_index idx 2 - must match seq_len / BLOCK_N");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_num idx 2 - must match seq_len / BLOCK_N");
|
||||
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K sequence length");
|
||||
TORCH_CHECK(l_vec.size(2) == q_seq_len, "L sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(o.size(2) == q_seq_len, "O sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(og.size(2) == q_seq_len, "OG sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(2) == num_kv_blocks, "k2q_block_sparse_num idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(o.size(3) == head_dim, "O head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(og.size(3) == head_dim, "OG head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
|
||||
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
@@ -988,20 +942,20 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
torch::Tensor qg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
torch::Tensor kg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(kv_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(kv_seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
torch::Tensor vg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(kv_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(kv_seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
|
||||
torch::Tensor d_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(1)}, l_vec.options());
|
||||
|
||||
float* qg_ptr = qg.data_ptr<float>();
|
||||
@@ -1030,7 +984,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
// cudaStreamSynchronize(stream);
|
||||
|
||||
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
|
||||
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
dim3 grid_bwd(q_seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
if (head_dim == 64) {
|
||||
using og_tile = st_bf<4*16, 64>;
|
||||
@@ -1043,9 +997,9 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_prep_globals = bwd_prep_globals<64>;
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
|
||||
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
@@ -1082,15 +1036,15 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_global_args = bwd_globals<64>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
|
||||
bwd_global_args bwd_global{bwd_q_arg,
|
||||
bwd_k_arg,
|
||||
@@ -1101,14 +1055,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_vg_arg,
|
||||
bwd_l_arg,
|
||||
bwd_d_arg,
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(kv_seq_len), // N is not used in the kernel
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
@@ -1147,9 +1101,9 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_prep_globals = bwd_prep_globals<128>;
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
|
||||
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
@@ -1186,15 +1140,15 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_global_args = bwd_globals<128>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
|
||||
bwd_global_args bwd_global{bwd_q_arg,
|
||||
bwd_k_arg,
|
||||
@@ -1205,14 +1159,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_vg_arg,
|
||||
bwd_l_arg,
|
||||
bwd_d_arg,
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(kv_seq_len), // N is not used in the kernel
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
@@ -1233,4 +1187,4 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
return {qg, kg, vg};
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include "common/common.hpp"
|
||||
#include "norm/layernorm.hpp"
|
||||
|
||||
@@ -14,10 +15,6 @@ auto layer_norm(
|
||||
std::optional<at::Tensor const> const B,
|
||||
std::optional<at::Tensor> Output
|
||||
) {
|
||||
using ElementIn = float;
|
||||
using ElementOut = float;
|
||||
using ElementWeight = float;
|
||||
|
||||
int64_t const m = Input.size(0);
|
||||
int64_t const n = Input.size(1);
|
||||
torch::Device const input_device = Input.device();
|
||||
@@ -26,31 +23,70 @@ auto layer_norm(
|
||||
Output.emplace(
|
||||
torch::empty(
|
||||
{m, n},
|
||||
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
|
||||
torch::TensorOptions().device(input_device).dtype(Input.scalar_type())
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
TORCH_CHECK(Output.value().scalar_type() == Input.scalar_type(),
|
||||
"Output dtype must match Input dtype. Got Output=",
|
||||
Output.value().scalar_type(), ", Input=", Input.scalar_type());
|
||||
if (W.has_value()) {
|
||||
TORCH_CHECK(W.value().scalar_type() == Input.scalar_type(),
|
||||
"W dtype must match Input dtype. Got W=",
|
||||
W.value().scalar_type(), ", Input=", Input.scalar_type());
|
||||
}
|
||||
if (B.has_value()) {
|
||||
TORCH_CHECK(B.value().scalar_type() == Input.scalar_type(),
|
||||
"B dtype must match Input dtype. Got B=",
|
||||
B.value().scalar_type(), ", Input=", Input.scalar_type());
|
||||
}
|
||||
|
||||
void *Iptr = Input.data_ptr();
|
||||
void *Wptr = W.has_value() ? W.value().data_ptr() : nullptr;
|
||||
void *Bptr = B.has_value() ? B.value().data_ptr() : nullptr;
|
||||
void *Optr = Output.value().data_ptr();
|
||||
|
||||
BOOL_SWITCH(B.has_value(), BIAS, [&]{
|
||||
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
layernorm<
|
||||
ElementIn, ElementOut, ElementWeight,
|
||||
AFFINE, BIAS,
|
||||
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA> (
|
||||
Iptr, Wptr, Bptr,
|
||||
Optr, eps, m, n,
|
||||
at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
if (Input.scalar_type() == at::kHalf) {
|
||||
using ElementIn = cutlass::half_t;
|
||||
using ElementOut = cutlass::half_t;
|
||||
using ElementWeight = cutlass::half_t;
|
||||
BOOL_SWITCH(B.has_value(), BIAS, [&]{
|
||||
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
layernorm<ElementIn, ElementOut, ElementWeight, AFFINE, BIAS, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
|
||||
Iptr, Wptr, Bptr, Optr, eps, m, n, stream);
|
||||
});
|
||||
});
|
||||
});
|
||||
});
|
||||
} else if (Input.scalar_type() == at::kBFloat16) {
|
||||
using ElementIn = cutlass::bfloat16_t;
|
||||
using ElementOut = cutlass::bfloat16_t;
|
||||
using ElementWeight = cutlass::bfloat16_t;
|
||||
BOOL_SWITCH(B.has_value(), BIAS, [&]{
|
||||
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
layernorm<ElementIn, ElementOut, ElementWeight, AFFINE, BIAS, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
|
||||
Iptr, Wptr, Bptr, Optr, eps, m, n, stream);
|
||||
});
|
||||
});
|
||||
});
|
||||
} else if (Input.scalar_type() == at::kFloat) {
|
||||
using ElementIn = float;
|
||||
using ElementOut = float;
|
||||
using ElementWeight = float;
|
||||
BOOL_SWITCH(B.has_value(), BIAS, [&]{
|
||||
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
layernorm<ElementIn, ElementOut, ElementWeight, AFFINE, BIAS, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
|
||||
Iptr, Wptr, Bptr, Optr, eps, m, n, stream);
|
||||
});
|
||||
});
|
||||
});
|
||||
} else {
|
||||
TORCH_CHECK(false, "Unsupported dtype for layer_norm_cuda: ", Input.scalar_type());
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -68,9 +68,22 @@ public:
|
||||
// mean reduction
|
||||
float u = _reduce_sum(x, shared_data) / params.n;
|
||||
|
||||
// IMPORTANT:
|
||||
// Loader pads out-of-range lanes with 0. That is OK for the sum, but after
|
||||
// subtracting mean, those padded lanes become -u and would incorrectly
|
||||
// contribute to the variance. Mask them back to 0 before variance reduction.
|
||||
// We launch exactly NumThrPerCta threads for a 1xMaxHiddenSize tile,
|
||||
// so each thread is responsible for a contiguous chunk in N.
|
||||
int thr_n_offset = tidx * NumElementPerThread;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
x[i] -= u;
|
||||
for (int i = 0; i < NumElementPerThread; ++i) {
|
||||
int idx = thr_n_offset + i;
|
||||
if (idx < params.n) {
|
||||
x[i] -= u;
|
||||
} else {
|
||||
x[i] = 0.f;
|
||||
}
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
// var reduction
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include <pybind11/pybind11.h>
|
||||
|
||||
#include "common/common.hpp"
|
||||
@@ -16,10 +17,6 @@ auto rms_norm(
|
||||
std::optional<at::Tensor>& Output
|
||||
) {
|
||||
|
||||
using ElementIn = float;
|
||||
using ElementOut = float;
|
||||
using ElementWeight = float;
|
||||
|
||||
int64_t const m = Input.size(0);
|
||||
int64_t const n = Input.size(1);
|
||||
torch::Device const input_device = Input.device();
|
||||
@@ -28,27 +25,51 @@ auto rms_norm(
|
||||
Output.emplace(
|
||||
torch::empty(
|
||||
{m, n},
|
||||
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
|
||||
torch::TensorOptions().device(input_device).dtype(Input.scalar_type())
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
TORCH_CHECK(Output.value().scalar_type() == Input.scalar_type(),
|
||||
"Output dtype must match Input dtype. Got Output=",
|
||||
Output.value().scalar_type(), ", Input=", Input.scalar_type());
|
||||
if (Weight.has_value()) {
|
||||
TORCH_CHECK(Weight.value().scalar_type() == Input.scalar_type(),
|
||||
"Weight dtype must match Input dtype. Got Weight=",
|
||||
Weight.value().scalar_type(), ", Input=", Input.scalar_type());
|
||||
}
|
||||
|
||||
void *Iptr = Input.data_ptr();
|
||||
void *Wptr = Weight.has_value() ? Weight.value().data_ptr() : nullptr;
|
||||
void *Optr = Output.value().data_ptr();
|
||||
|
||||
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
rmsnorm<
|
||||
ElementIn, ElementOut, ElementWeight,
|
||||
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA
|
||||
> (
|
||||
Iptr, Wptr,
|
||||
Optr,
|
||||
eps, m, n,
|
||||
at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
});
|
||||
if (Input.scalar_type() == at::kHalf) {
|
||||
using ElementIn = cutlass::half_t;
|
||||
using ElementOut = cutlass::half_t;
|
||||
using ElementWeight = cutlass::half_t;
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
rmsnorm<ElementIn, ElementOut, ElementWeight, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
|
||||
Iptr, Wptr, Optr, eps, m, n, at::cuda::getCurrentCUDAStream().stream());
|
||||
});
|
||||
} else if (Input.scalar_type() == at::kBFloat16) {
|
||||
using ElementIn = cutlass::bfloat16_t;
|
||||
using ElementOut = cutlass::bfloat16_t;
|
||||
using ElementWeight = cutlass::bfloat16_t;
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
rmsnorm<ElementIn, ElementOut, ElementWeight, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
|
||||
Iptr, Wptr, Optr, eps, m, n, at::cuda::getCurrentCUDAStream().stream());
|
||||
});
|
||||
} else if (Input.scalar_type() == at::kFloat) {
|
||||
using ElementIn = float;
|
||||
using ElementOut = float;
|
||||
using ElementWeight = float;
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
rmsnorm<ElementIn, ElementOut, ElementWeight, MAX_HIDDEN_SIZE, NUM_THR_PER_CTA>(
|
||||
Iptr, Wptr, Optr, eps, m, n, at::cuda::getCurrentCUDAStream().stream());
|
||||
});
|
||||
} else {
|
||||
TORCH_CHECK(false, "Unsupported dtype for rms_norm_cuda: ", Input.scalar_type());
|
||||
}
|
||||
|
||||
|
||||
return Output;
|
||||
|
||||
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.2.1"
|
||||
version = "0.2.4"
|
||||
description = "Unified CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -0,0 +1,298 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _get_sm90_ops():
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
|
||||
except Exception:
|
||||
return None, None
|
||||
return (
|
||||
getattr(fastvideo_kernel_ops, "block_sparse_fwd", None),
|
||||
getattr(fastvideo_kernel_ops, "block_sparse_bwd", None),
|
||||
)
|
||||
|
||||
|
||||
def _is_sm90() -> bool:
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
return major == 9 and minor == 0
|
||||
|
||||
|
||||
def _force_triton() -> bool:
|
||||
# Force Triton even on SM90 and even if the compiled extension is available.
|
||||
# Useful for CI / debugging / parity testing.
|
||||
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
|
||||
|
||||
|
||||
def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Pure-torch (no triton) conversion:
|
||||
block_map: [B, H, Q, KV] bool (or [H, Q, KV] which will be treated as B=1)
|
||||
returns:
|
||||
index: [B, H, Q, KV] int32 (packed KV indices, -1 padding)
|
||||
num: [B, H, Q] int32 (#kv blocks per q block)
|
||||
"""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
if block_map.dim() != 4:
|
||||
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
|
||||
if block_map.dtype != torch.bool:
|
||||
block_map = block_map.to(torch.bool)
|
||||
|
||||
B, H, Q, KV = block_map.shape
|
||||
index = torch.full((B, H, Q, KV), -1, dtype=torch.int32, device=block_map.device)
|
||||
num = torch.zeros((B, H, Q), dtype=torch.int32, device=block_map.device)
|
||||
|
||||
# Small sizes in practice (B=1, H<=16, Q/KV<=64), so a Python loop is fine.
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for q in range(Q):
|
||||
kv_idx = torch.nonzero(block_map[b, h, q], as_tuple=False).flatten().to(torch.int32)
|
||||
n = int(kv_idx.numel())
|
||||
if n:
|
||||
index[b, h, q, :n] = kv_idx
|
||||
num[b, h, q] = n
|
||||
return index, num
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_triton",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index_torch(block_map)
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
triton_block_sparse_attn_forward,
|
||||
)
|
||||
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_backward_triton",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def block_sparse_attn_backward_triton(
|
||||
grad_output: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output = grad_output.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index_torch(block_map)
|
||||
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
triton_block_sparse_attn_backward,
|
||||
)
|
||||
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(
|
||||
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_triton")
|
||||
def _block_sparse_attn_backward_triton_fake(
|
||||
grad_output: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _backward_triton(ctx, grad_o, grad_M):
|
||||
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _setup_context_triton(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
o, M = output
|
||||
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_sm90",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def block_sparse_attn_sm90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
block_sparse_fwd, _ = _get_sm90_ops()
|
||||
if block_sparse_fwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
|
||||
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index_torch(block_map)
|
||||
|
||||
o_padded, lse_padded = block_sparse_fwd(
|
||||
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
|
||||
)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_sm90")
|
||||
def _block_sparse_attn_sm90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q_padded)
|
||||
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
|
||||
return o, lse
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_backward_sm90",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def block_sparse_attn_backward_sm90(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
_, block_sparse_bwd = _get_sm90_ops()
|
||||
if block_sparse_bwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
|
||||
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
dq, dk, dv = block_sparse_bwd(
|
||||
q_padded,
|
||||
k_padded,
|
||||
v_padded,
|
||||
o_padded,
|
||||
lse_padded,
|
||||
grad_output_padded,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes.int(),
|
||||
)
|
||||
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
|
||||
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
|
||||
def _block_sparse_attn_backward_sm90_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q_padded)
|
||||
dk = torch.empty_like(k_padded)
|
||||
dv = torch.empty_like(v_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _backward_sm90(ctx, grad_o, grad_lse):
|
||||
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_sm90(
|
||||
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _setup_context_sm90(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
o, lse = output
|
||||
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
|
||||
|
||||
|
||||
def block_sparse_attn(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Unified block-sparse attention op with autograd support.
|
||||
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
|
||||
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
|
||||
"""
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
|
||||
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
|
||||
# Triton path: generally assumes q/k/v share the same padded length
|
||||
if q.shape[2] != k.shape[2] or q.shape[2] != v.shape[2]:
|
||||
raise RuntimeError("Triton fallback requires q/k/v to have the same padded length.")
|
||||
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import math
|
||||
import torch
|
||||
from .block_sparse_attn import block_sparse_attn
|
||||
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
|
||||
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
|
||||
from .triton_kernels.index import map_to_index
|
||||
@@ -45,14 +46,22 @@ def sliding_tile_attention(
|
||||
flag = shape_map[seq_shape]
|
||||
|
||||
for head_idx, (t, h, w) in enumerate(window_size):
|
||||
# Per-head slices are not contiguous in the batch dimension when batch>1
|
||||
# (they keep the original head-stride). The TK kernel assumes contiguous
|
||||
# [B, H, S, D] layout, so we materialize a contiguous [B,1,S,D] view.
|
||||
q_h = q[:, head_idx:head_idx + 1].contiguous()
|
||||
k_h = k[:, head_idx:head_idx + 1].contiguous()
|
||||
v_h = v[:, head_idx:head_idx + 1].contiguous()
|
||||
o_h = torch.empty_like(q_h)
|
||||
sta_fwd(
|
||||
q[:, head_idx:head_idx + 1], k[:, head_idx:head_idx + 1],
|
||||
v[:, head_idx:head_idx + 1], output[:, head_idx:head_idx + 1],
|
||||
q_h, k_h,
|
||||
v_h, o_h,
|
||||
t, h, w, text_length, False, has_text, flag
|
||||
)
|
||||
output[:, head_idx:head_idx + 1] = o_h
|
||||
|
||||
if has_text:
|
||||
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
|
||||
sta_fwd(q.contiguous(), k.contiguous(), v.contiguous(), output, 3, 3, 3, text_length, True, True, flag)
|
||||
|
||||
return output[:, :, :seq_length]
|
||||
|
||||
@@ -62,6 +71,7 @@ def video_sparse_attn(
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
q_variable_block_sizes: torch.Tensor,
|
||||
topk: int,
|
||||
block_size: int | tuple = 64,
|
||||
compress_attn_weight: torch.Tensor = None,
|
||||
@@ -70,14 +80,42 @@ def video_sparse_attn(
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
batch, heads, seq_len, dim = q.shape
|
||||
batch, heads, q_seq_len, dim = q.shape
|
||||
kv_seq_len = k.shape[2]
|
||||
if v.shape[2] != kv_seq_len:
|
||||
raise ValueError(
|
||||
f"Expected k and v to have the same sequence length, got "
|
||||
f"k.shape[2]={kv_seq_len}, v.shape[2]={v.shape[2]}"
|
||||
)
|
||||
if k.shape[0] != batch or v.shape[0] != batch or k.shape[1] != heads or v.shape[1] != heads:
|
||||
raise ValueError("Expected q/k/v to have the same batch and head dimensions.")
|
||||
|
||||
if q_seq_len % block_elements != 0 or kv_seq_len % block_elements != 0:
|
||||
raise ValueError(
|
||||
f"q_seq_len and kv_seq_len must be divisible by block_elements={block_elements}, "
|
||||
f"got q_seq_len={q_seq_len}, kv_seq_len={kv_seq_len}"
|
||||
)
|
||||
q_num_blocks = q_seq_len // block_elements
|
||||
kv_num_blocks = kv_seq_len // block_elements
|
||||
|
||||
if variable_block_sizes.numel() != kv_num_blocks:
|
||||
raise ValueError(
|
||||
f"variable_block_sizes must have length kv_num_blocks={kv_num_blocks}, "
|
||||
f"got {variable_block_sizes.numel()}"
|
||||
)
|
||||
|
||||
if q_variable_block_sizes.numel() != q_num_blocks:
|
||||
raise ValueError(
|
||||
f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
|
||||
f"got {q_variable_block_sizes.numel()}"
|
||||
)
|
||||
|
||||
# Compression branch
|
||||
q_c = q.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
k_c = k.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
v_c = v.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
q_c = q.view(batch, heads, q_num_blocks, block_elements, dim)
|
||||
k_c = k.view(batch, heads, kv_num_blocks, block_elements, dim)
|
||||
v_c = v.view(batch, heads, kv_num_blocks, block_elements, dim)
|
||||
|
||||
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
|
||||
q_c = (q_c.float().sum(dim=3) / q_variable_block_sizes.view(1, 1, -1, 1)).to(
|
||||
q.dtype)
|
||||
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
|
||||
k.dtype)
|
||||
@@ -88,9 +126,9 @@ def video_sparse_attn(
|
||||
attn = torch.softmax(scores, dim=-1)
|
||||
out_c = torch.matmul(attn, v_c)
|
||||
|
||||
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
|
||||
out_c = out_c.view(batch, heads, q_num_blocks, 1, dim)
|
||||
out_c = out_c.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch, heads, seq_len, dim)
|
||||
1).view(batch, heads, q_seq_len, dim)
|
||||
|
||||
# Sparse branch
|
||||
topk_idx = torch.topk(scores, topk, dim=-1).indices
|
||||
@@ -100,12 +138,17 @@ def video_sparse_attn(
|
||||
idx, num = map_to_index(mask)
|
||||
|
||||
if block_sparse_fwd is not None:
|
||||
out_s = block_sparse_fwd(
|
||||
q, k, v, idx, num, variable_block_sizes.int()
|
||||
)[0] # block_sparse_fwd returns vector<Tensor>
|
||||
# Use autograd-enabled wrapper so backward works (and still uses SM90 kernel when available)
|
||||
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
else:
|
||||
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num,
|
||||
variable_block_sizes)
|
||||
if q_seq_len != kv_seq_len:
|
||||
raise RuntimeError(
|
||||
"q/k have different lengths, but the compiled CUDA kernel (block_sparse_fwd) "
|
||||
"is not available. The Triton fallback currently requires q and k/v to have "
|
||||
"the same padded length."
|
||||
)
|
||||
# Triton-only forward (kept for environments without the wrapper deps)
|
||||
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
|
||||
@@ -0,0 +1,311 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from TurboDiffusion SLA implementation
|
||||
# Copyright (c) 2025 by SLA team.
|
||||
#
|
||||
# Citation:
|
||||
# @article{zhang2025sla,
|
||||
# title={SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse-Linear Attention},
|
||||
# author={Jintao Zhang and Haoxu Wang and Kai Jiang and Shuo Yang and Kaiwen Zheng and
|
||||
# Haocheng Xi and Ziteng Wang and Hongzhou Zhu and Min Zhao and Ion Stoica and
|
||||
# Joseph E. Gonzalez and Jun Zhu and Jianfei Chen},
|
||||
# journal={arXiv preprint arXiv:2509.24006},
|
||||
# year={2025}
|
||||
# }
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd(
|
||||
Q, K, V,
|
||||
qk_scale: tl.constexpr,
|
||||
topk: tl.constexpr,
|
||||
LUT, LSE, OS,
|
||||
L: tl.constexpr,
|
||||
M_BLOCKS: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
idx_m = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
qkv_offset = idx_bh * L * D
|
||||
lut_offset = (idx_bh * M_BLOCKS + idx_m) * topk
|
||||
lse_offset = idx_bh * L
|
||||
offs_m = idx_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
Q_ptrs = Q + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
K_ptrs = K + qkv_offset + offs_n[None, :] * D + offs_d[:, None]
|
||||
V_ptrs = V + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
OS_ptrs = OS + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
LUT_ptr = LUT + lut_offset
|
||||
LSE_ptrs = LSE + lse_offset + offs_m
|
||||
|
||||
m_i = tl.full([BLOCK_M], -float('inf'), dtype=tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
|
||||
o_s = tl.zeros([BLOCK_M, D], dtype=tl.float32)
|
||||
|
||||
q = tl.load(Q_ptrs, mask=offs_m[:, None] < L)
|
||||
for block_idx in tl.range(topk):
|
||||
idx_n = tl.load(LUT_ptr + block_idx)
|
||||
n_mask = offs_n < L - idx_n * BLOCK_N
|
||||
|
||||
k = tl.load(K_ptrs + idx_n * BLOCK_N * D, mask=n_mask[None, :])
|
||||
qk = tl.dot(q, k) * (qk_scale * 1.4426950408889634) # = 1 / ln(2)
|
||||
if L - idx_n * BLOCK_N < BLOCK_N:
|
||||
qk = tl.where(n_mask[None, :], qk, float("-inf"))
|
||||
|
||||
v = tl.load(V_ptrs + idx_n * BLOCK_N * D, mask=n_mask[:, None])
|
||||
local_m = tl.max(qk, 1)
|
||||
new_m = tl.maximum(m_i, local_m)
|
||||
qk = qk - new_m[:, None]
|
||||
|
||||
p = tl.math.exp2(qk)
|
||||
l_ij = tl.sum(p, 1)
|
||||
alpha = tl.math.exp2(m_i - new_m)
|
||||
o_s = o_s * alpha[:, None]
|
||||
o_s += tl.dot(p.to(v.dtype), v)
|
||||
|
||||
l_i = l_i * alpha + l_ij
|
||||
m_i = new_m
|
||||
|
||||
o_s = o_s / l_i[:, None]
|
||||
tl.store(OS_ptrs, o_s.to(OS.type.element_ty), mask=offs_m[:, None] < L)
|
||||
|
||||
m_i += tl.math.log2(l_i)
|
||||
tl.store(LSE_ptrs, m_i, mask=offs_m < L)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(
|
||||
OS, DOS, DELTAS,
|
||||
L,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
):
|
||||
idx_m = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
OS += idx_bh * L * D
|
||||
DOS += idx_bh * L * D
|
||||
DELTAS += idx_bh * L
|
||||
|
||||
offs_m = idx_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
o_s = tl.load(OS + offs_m[:, None] * D + offs_d[None, :], mask=offs_m[:, None] < L)
|
||||
do_s = tl.load(DOS + offs_m[:, None] * D + offs_d[None, :], mask=offs_m[:, None] < L)
|
||||
|
||||
delta_s = tl.sum(o_s * do_s, axis=1).to(DELTAS.type.element_ty)
|
||||
tl.store(DELTAS + offs_m, delta_s, mask=offs_m < L)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(
|
||||
Q, K, V, LSE, DELTAS,
|
||||
DOS, DQ, LUT,
|
||||
qk_scale: tl.constexpr,
|
||||
topk: tl.constexpr,
|
||||
L: tl.constexpr,
|
||||
M_BLOCKS: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
idx_m = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
offs_m = idx_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
qkv_offset = idx_bh * L * D
|
||||
lse_offset = idx_bh * L
|
||||
lut_offset = (idx_bh * M_BLOCKS + idx_m) * topk
|
||||
|
||||
Q_ptrs = Q + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
K_ptrs = K + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
V_ptrs = V + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
DQ_ptrs = DQ + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
DOS_ptrs = DOS + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
LSE_ptrs = LSE + lse_offset + offs_m
|
||||
DELTAS_ptrs = DELTAS + lse_offset + offs_m
|
||||
LUT_ptr = LUT + lut_offset
|
||||
|
||||
q = tl.load(Q_ptrs, mask=offs_m[:, None] < L)
|
||||
do_s = tl.load(DOS_ptrs, mask=offs_m[:, None] < L)
|
||||
delta_s = tl.load(DELTAS_ptrs, mask=offs_m < L)
|
||||
lse = tl.load(LSE_ptrs, mask=offs_m < L, other=float("inf"))
|
||||
|
||||
dq = tl.zeros([BLOCK_M, D], dtype=tl.float32)
|
||||
for block_idx in tl.range(topk, num_stages=2):
|
||||
idx_n = tl.load(LUT_ptr + block_idx)
|
||||
n_mask = offs_n < L - idx_n * BLOCK_N
|
||||
|
||||
k = tl.load(K_ptrs + idx_n * BLOCK_N * D, mask=n_mask[:, None])
|
||||
v = tl.load(V_ptrs + idx_n * BLOCK_N * D, mask=n_mask[:, None])
|
||||
qk = tl.dot(q, k.T) * (qk_scale * 1.4426950408889634)
|
||||
p = tl.math.exp2(qk - lse[:, None])
|
||||
p = tl.where(n_mask[None, :], p, 0.0)
|
||||
|
||||
dp = tl.dot(do_s, v.T).to(tl.float32)
|
||||
ds = p * (dp - delta_s[:, None])
|
||||
dq += tl.dot(ds.to(k.dtype), k)
|
||||
tl.store(DQ_ptrs, dq * qk_scale, mask=offs_m[:, None] < L)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(
|
||||
Q, K, V, DOS, DK, DV,
|
||||
qk_scale, KBID, LSE, DELTAS,
|
||||
L: tl.constexpr,
|
||||
M_BLOCKS: tl.constexpr,
|
||||
N_BLOCKS: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
BLOCK_SLICE_FACTOR: tl.constexpr,
|
||||
):
|
||||
BLOCK_M2: tl.constexpr = BLOCK_M // BLOCK_SLICE_FACTOR
|
||||
|
||||
idx_n = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
offs_n = idx_n * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
offs_m = tl.arange(0, BLOCK_M2)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
qkv_offset = idx_bh * L * D
|
||||
kbid_offset = idx_bh * M_BLOCKS * N_BLOCKS
|
||||
lse_offset = idx_bh * L
|
||||
|
||||
Q_ptrs = Q + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
K_ptrs = K + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
V_ptrs = V + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
DOS_ptrs = DOS + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
DK_ptrs = DK + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
DV_ptrs = DV + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
LSE_ptrs = LSE + lse_offset + offs_m
|
||||
DELTAS_ptrs = DELTAS + lse_offset + offs_m
|
||||
KBID_ptr = KBID + kbid_offset + idx_n
|
||||
|
||||
k = tl.load(K_ptrs, mask=offs_n[:, None] < L)
|
||||
v = tl.load(V_ptrs, mask=offs_n[:, None] < L)
|
||||
|
||||
dk = tl.zeros([BLOCK_N, D], dtype=tl.float32)
|
||||
dv = tl.zeros([BLOCK_N, D], dtype=tl.float32)
|
||||
for idx_m in tl.range(0, L, BLOCK_M2):
|
||||
kbid = tl.load(KBID_ptr)
|
||||
if kbid == 1:
|
||||
m_mask = offs_m < L - idx_m
|
||||
q = tl.load(Q_ptrs, mask=m_mask[:, None])
|
||||
lse = tl.load(LSE_ptrs, mask=m_mask, other=float("inf"))
|
||||
qkT = tl.dot(k, q.T) * (qk_scale * 1.4426950408889634)
|
||||
pT = tl.math.exp2(qkT - lse[None, :])
|
||||
pT = tl.where(offs_n[:, None] < L, pT, 0.0)
|
||||
|
||||
do = tl.load(DOS_ptrs, mask=m_mask[:, None])
|
||||
dv += tl.dot(pT.to(do.dtype), do)
|
||||
delta = tl.load(DELTAS_ptrs, mask=m_mask)
|
||||
dpT = tl.dot(v, tl.trans(do))
|
||||
dsT = pT * (dpT - delta[None, :])
|
||||
dk += tl.dot(dsT.to(q.dtype), q)
|
||||
|
||||
Q_ptrs += BLOCK_M2 * D
|
||||
DOS_ptrs += BLOCK_M2 * D
|
||||
LSE_ptrs += BLOCK_M2
|
||||
DELTAS_ptrs += BLOCK_M2
|
||||
if (idx_m + BLOCK_M2) % BLOCK_M == 0:
|
||||
KBID_ptr += N_BLOCKS
|
||||
|
||||
tl.store(DK_ptrs, dk * qk_scale, mask=offs_n[:, None] < L)
|
||||
tl.store(DV_ptrs, dv, mask=offs_n[:, None] < L)
|
||||
|
||||
|
||||
class _attention(torch.autograd.Function):
|
||||
"""Sparse attention forward/backward with autograd support."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, k_block_id, lut, topk, BLOCK_M, BLOCK_N, qk_scale=None):
|
||||
assert q.is_contiguous() and k.is_contiguous() and v.is_contiguous()
|
||||
assert k_block_id.is_contiguous() and lut.is_contiguous()
|
||||
|
||||
assert BLOCK_M == 64 or BLOCK_M == 128
|
||||
assert BLOCK_N == 64
|
||||
|
||||
B, H, L, D = q.shape
|
||||
if qk_scale is None:
|
||||
qk_scale = D**-0.5
|
||||
|
||||
M_BLOCKS = triton.cdiv(L, BLOCK_M)
|
||||
|
||||
o_s = torch.empty_like(v)
|
||||
lse = torch.empty(q.shape[:-1], device=q.device, dtype=torch.float32)
|
||||
|
||||
grid = (M_BLOCKS, B * H)
|
||||
_attn_fwd[grid](
|
||||
q, k, v, qk_scale, topk,
|
||||
lut, lse, o_s,
|
||||
L, M_BLOCKS,
|
||||
D, BLOCK_M, BLOCK_N,
|
||||
num_warps=4 if q.shape[-1] == 64 else 8,
|
||||
num_stages=3
|
||||
)
|
||||
|
||||
ctx.save_for_backward(q, k, v, k_block_id, lut, lse, o_s)
|
||||
ctx.qk_scale = qk_scale
|
||||
ctx.topk = topk
|
||||
ctx.BLOCK_M = BLOCK_M
|
||||
ctx.BLOCK_N = BLOCK_N
|
||||
return o_s
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, do_s):
|
||||
q, k, v, k_block_id, lut, lse, o_s = ctx.saved_tensors
|
||||
do_s = do_s.contiguous()
|
||||
|
||||
BLOCK_M, BLOCK_N = ctx.BLOCK_M, ctx.BLOCK_N
|
||||
B, H, L, D = q.shape
|
||||
|
||||
M_BLOCKS = triton.cdiv(L, BLOCK_M)
|
||||
N_BLOCKS = triton.cdiv(L, BLOCK_N)
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
delta_s = torch.empty_like(lse)
|
||||
|
||||
grid = (M_BLOCKS, B * H)
|
||||
_attn_bwd_preprocess[grid](
|
||||
o_s, do_s, delta_s,
|
||||
L, D, BLOCK_M,
|
||||
)
|
||||
|
||||
grid = (M_BLOCKS, B * H)
|
||||
_attn_bwd_dq[grid](
|
||||
q, k, v, lse, delta_s,
|
||||
do_s, dq, lut,
|
||||
ctx.qk_scale, ctx.topk,
|
||||
L, M_BLOCKS,
|
||||
D, BLOCK_M, BLOCK_N,
|
||||
num_warps=4 if q.shape[-1] == 64 else 8,
|
||||
num_stages=4 if q.shape[-1] == 64 else 5
|
||||
)
|
||||
|
||||
grid = (N_BLOCKS, B * H)
|
||||
_attn_bwd_dkdv[grid](
|
||||
q, k, v, do_s, dk, dv,
|
||||
ctx.qk_scale, k_block_id, lse, delta_s,
|
||||
L, M_BLOCKS, N_BLOCKS,
|
||||
D, BLOCK_M, BLOCK_N,
|
||||
BLOCK_SLICE_FACTOR=BLOCK_M // 64,
|
||||
num_warps=4 if q.shape[-1] == 64 else 8,
|
||||
num_stages=4 if q.shape[-1] == 64 else 5
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, None, None, None, None, None
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.2.1"
|
||||
__version__ = "0.2.4"
|
||||
|
||||
@@ -42,37 +42,13 @@ def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, q
|
||||
q_padded = vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
|
||||
k_padded = vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
v_padded = vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
# Use raw kernel or triton
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops
|
||||
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
|
||||
except ImportError:
|
||||
raw_kernel = None
|
||||
# Use autograd-enabled wrapper (internally dispatches to SM90 kernel or Triton)
|
||||
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
|
||||
output_padded, _aux = block_sparse_attn(
|
||||
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
|
||||
)
|
||||
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index
|
||||
|
||||
|
||||
# Convert mask to indices
|
||||
# block_sparse_mask is [H, M, N] bool
|
||||
# We need to map it to index.
|
||||
# block_sparse_mask needs to be expanded/reshaped?
|
||||
# generate_block_sparse_mask_for_function returns [H, NumBlocksQ, NumBlocksKV]
|
||||
|
||||
# Ops.py logic:
|
||||
# mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
|
||||
# idx, num = map_to_index(mask)
|
||||
|
||||
idx, num = map_to_index(block_sparse_mask.unsqueeze(0)) # Add batch dim [1, H, M, N]
|
||||
|
||||
if raw_kernel:
|
||||
out_s = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())
|
||||
output = out_s[0]
|
||||
else:
|
||||
# Fallback to triton testing if C++ not available
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
|
||||
output, _ = triton_block_sparse_attn_forward(q_padded, k_padded, v_padded, idx, num, variable_block_sizes)
|
||||
|
||||
output = output[:, :, q_non_pad_index, :]
|
||||
output = output_padded[:, :, q_non_pad_index, :]
|
||||
output.backward(dO)
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
@@ -264,7 +240,6 @@ def generate_error_graphs_qkdiff(h, d, error_mode='all'):
|
||||
|
||||
print("-" * 150)
|
||||
|
||||
@pytest.mark.skip()
|
||||
def test_video_sparse_attention_backward():
|
||||
if not torch.cuda.is_available():
|
||||
return
|
||||
|
||||
@@ -3,6 +3,7 @@ import sys
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
|
||||
from .utils import (
|
||||
generate_block_sparse_mask_for_function,
|
||||
@@ -57,23 +58,14 @@ def block_sparse_forward_test(
|
||||
k_padded = ref.vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
v_padded = ref.vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
|
||||
# Use raw kernel or triton
|
||||
# Use autograd-enabled wrapper (internally dispatches SM90 C++ vs Triton)
|
||||
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops
|
||||
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
|
||||
except ImportError:
|
||||
raw_kernel = None
|
||||
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index
|
||||
idx, num = map_to_index(block_sparse_mask)
|
||||
|
||||
if raw_kernel:
|
||||
out_padded = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())[0]
|
||||
else:
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
|
||||
out_padded, _ = triton_block_sparse_attn_forward(
|
||||
q_padded, k_padded, v_padded, idx, num, variable_block_sizes
|
||||
out_padded, _aux = block_sparse_attn(
|
||||
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
|
||||
)
|
||||
except RuntimeError as e:
|
||||
pytest.skip(str(e))
|
||||
|
||||
# Remove padding on the query side
|
||||
out = out_padded[:, :, q_non_pad_index, :]
|
||||
@@ -156,11 +148,6 @@ def run_forward_qk_diff(
|
||||
) -> Tuple[float, float]:
|
||||
"""
|
||||
Forward-only correctness test for the case S_q != S_kv.
|
||||
|
||||
NOTE:
|
||||
- The Triton backend supports different Q/KV logical lengths via padding.
|
||||
- The SM90 (H100) CUDA backend currently assumes the same number of blocks
|
||||
for Q and KV, so we skip this test there.
|
||||
"""
|
||||
assert torch.cuda.is_available(), "VSA kernels require CUDA"
|
||||
|
||||
|
||||
@@ -0,0 +1,587 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SLA (Sparse-Linear Attention) backend for FastVideo
|
||||
# Adapted from TurboDiffusion SLA implementation
|
||||
#
|
||||
# Copyright (c) 2025 by SLA team.
|
||||
# Citation:
|
||||
# @article{zhang2025sla,
|
||||
# title={SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse-Linear Attention},
|
||||
# author={Jintao Zhang and Haoxu Wang and Kai Jiang and Shuo Yang and Kaiwen Zheng and
|
||||
# Haocheng Xi and Ziteng Wang and Hongzhou Zhu and Min Zhao and Ion Stoica and
|
||||
# Joseph E. Gonzalez and Jun Zhu and Jianfei Chen},
|
||||
# journal={arXiv preprint arXiv:2509.24006},
|
||||
# year={2025}
|
||||
# }
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo_kernel.triton_kernels.sla_triton import _attention
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# ============================================================================
|
||||
# SLA Utility functions (moved from sla_kernels/utils.py)
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@triton.jit
|
||||
def compress_kernel(
|
||||
X,
|
||||
XM,
|
||||
L: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_L: tl.constexpr,
|
||||
):
|
||||
idx_l = tl.program_id(0)
|
||||
idx_bh = tl.program_id(1)
|
||||
|
||||
offs_l = idx_l * BLOCK_L + tl.arange(0, BLOCK_L)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
x_offset = idx_bh * L * D
|
||||
xm_offset = idx_bh * ((L + BLOCK_L - 1) // BLOCK_L) * D
|
||||
x = tl.load(X + x_offset + offs_l[:, None] * D + offs_d[None, :],
|
||||
mask=offs_l[:, None] < L)
|
||||
|
||||
nx = min(BLOCK_L, L - idx_l * BLOCK_L)
|
||||
x_mean = tl.sum(x, axis=0, dtype=tl.float32) / nx
|
||||
tl.store(XM + xm_offset + idx_l * D + offs_d,
|
||||
x_mean.to(XM.dtype.element_ty))
|
||||
|
||||
|
||||
def mean_pool(x: torch.Tensor, BLK: int) -> torch.Tensor:
|
||||
"""Mean pool tensor along sequence dimension with block size BLK."""
|
||||
assert x.is_contiguous()
|
||||
|
||||
B, H, L, D = x.shape
|
||||
L_BLOCKS = (L + BLK - 1) // BLK
|
||||
x_mean = torch.empty((B, H, L_BLOCKS, D), device=x.device, dtype=x.dtype)
|
||||
|
||||
grid = (L_BLOCKS, B * H)
|
||||
compress_kernel[grid](x, x_mean, L, D, BLK)
|
||||
return x_mean
|
||||
|
||||
|
||||
def get_block_map(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
topk_ratio: float,
|
||||
BLKQ: int = 64,
|
||||
BLKK: int = 64,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""Compute sparse block map for attention based on QK similarity.
|
||||
|
||||
Args:
|
||||
q: Query tensor of shape (B, H, L, D)
|
||||
k: Key tensor of shape (B, H, L, D)
|
||||
topk_ratio: Ratio of key blocks to attend to (0-1)
|
||||
BLKQ: Query block size
|
||||
BLKK: Key block size
|
||||
|
||||
Returns:
|
||||
sparse_map: Binary mask of shape (B, H, num_q_blocks, num_k_blocks)
|
||||
lut: Top-k indices of shape (B, H, num_q_blocks, topk)
|
||||
topk: Number of key blocks selected
|
||||
"""
|
||||
arg_k = k - torch.mean(
|
||||
k, dim=-2, keepdim=True) # smooth-k technique from SageAttention
|
||||
pooled_qblocks = mean_pool(q, BLKQ)
|
||||
pooled_kblocks = mean_pool(arg_k, BLKK)
|
||||
pooled_score = pooled_qblocks @ pooled_kblocks.transpose(-1, -2)
|
||||
|
||||
K = pooled_score.shape[-1]
|
||||
topk = min(K, int(topk_ratio * K))
|
||||
lut = torch.topk(pooled_score, topk, dim=-1, sorted=False).indices
|
||||
|
||||
sparse_map = torch.zeros_like(pooled_score, dtype=torch.int8)
|
||||
sparse_map.scatter_(-1, lut, 1)
|
||||
return sparse_map, lut, topk
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# SLA Backend classes
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class SLAAttentionBackend(AttentionBackend):
|
||||
"""Sparse-Linear Attention backend."""
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SLAAttentionImpl"]:
|
||||
return SLAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["SLAAttentionMetadata"]:
|
||||
return SLAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["SLAAttentionMetadataBuilder"]:
|
||||
return SLAAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class SLAAttentionMetadata(AttentionMetadata):
|
||||
"""Metadata for SLA attention."""
|
||||
current_timestep: int
|
||||
topk_ratio: float = 0.5 # Ratio of key blocks to attend to
|
||||
|
||||
|
||||
class SLAAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
"""Builder for SLA attention metadata."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
topk_ratio: float = 0.5,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> SLAAttentionMetadata:
|
||||
return SLAAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
topk_ratio=topk_ratio,
|
||||
)
|
||||
|
||||
|
||||
class SLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
"""SLA attention implementation with learnable linear projection.
|
||||
|
||||
This implementation combines sparse attention with linear attention,
|
||||
using a learnable projection to blend the outputs. The sparse attention
|
||||
uses a block-sparse pattern determined by QK similarity.
|
||||
|
||||
Args:
|
||||
num_heads: Number of attention heads
|
||||
head_size: Dimension of each head
|
||||
topk_ratio: Ratio of key blocks to attend to (0-1), default 0.5
|
||||
feature_map: Feature map for linear attention ('softmax', 'elu', 'relu')
|
||||
BLKQ: Query block size for sparse attention
|
||||
BLKK: Key block size for sparse attention
|
||||
use_bf16: Whether to use bfloat16 for computation
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool = False,
|
||||
softmax_scale: float | None = None,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
# SLA-specific parameters - matched to TurboDiffusion defaults
|
||||
topk_ratio: float = 0.1, # TurboDiffusion uses topk=0.1
|
||||
feature_map: str = "softmax",
|
||||
BLKQ: int = 128, # TurboDiffusion uses BLKQ=128
|
||||
BLKK: int = 64, # TurboDiffusion uses BLKK=64
|
||||
use_bf16: bool = True,
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
nn.Module.__init__(self)
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.softmax_scale = softmax_scale if softmax_scale else head_size**-0.5
|
||||
self.causal = causal
|
||||
self.prefix = prefix
|
||||
|
||||
# SLA-specific config
|
||||
self.topk_ratio = topk_ratio
|
||||
self.BLKQ = BLKQ
|
||||
self.BLKK = BLKK
|
||||
self.dtype = torch.bfloat16 if use_bf16 else torch.float16
|
||||
|
||||
# Learnable linear projection for combining sparse + linear attention
|
||||
self.proj_l = nn.Linear(head_size, head_size, dtype=torch.float32)
|
||||
|
||||
# Feature map for linear attention
|
||||
# Type annotation for callables
|
||||
self.feature_map_q: Callable[[torch.Tensor], torch.Tensor]
|
||||
self.feature_map_k: Callable[[torch.Tensor], torch.Tensor]
|
||||
if feature_map == "elu":
|
||||
self.feature_map_q = lambda x: F.elu(x) + 1
|
||||
self.feature_map_k = lambda x: F.elu(x) + 1
|
||||
elif feature_map == "relu":
|
||||
self.feature_map_q = F.relu
|
||||
self.feature_map_k = F.relu
|
||||
elif feature_map == "softmax":
|
||||
self.feature_map_q = lambda x: F.softmax(x, dim=-1)
|
||||
self.feature_map_k = lambda x: F.softmax(x, dim=-1)
|
||||
else:
|
||||
raise ValueError(f"Unknown feature map: {feature_map}")
|
||||
|
||||
self._init_weights()
|
||||
|
||||
def _init_weights(self) -> None:
|
||||
"""Initialize projection weights to zero for residual-like behavior."""
|
||||
with torch.no_grad():
|
||||
nn.init.zeros_(self.proj_l.weight)
|
||||
nn.init.zeros_(self.proj_l.bias) # type: ignore[arg-type]
|
||||
|
||||
def _calc_linear_attention(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Compute linear attention: (Q @ K^T @ V) / normalizer.
|
||||
|
||||
Args:
|
||||
q: Query tensor (B, H, L, D) after feature map
|
||||
k: Key tensor (B, H, L, D) after feature map
|
||||
v: Value tensor (B, H, L, D)
|
||||
|
||||
Returns:
|
||||
Linear attention output (B, H, L, D)
|
||||
"""
|
||||
kvsum = k.transpose(-1, -2) @ v # (B, H, D, D)
|
||||
ksum = torch.sum(k, dim=-2, keepdim=True) # (B, H, 1, D)
|
||||
return (q @ kvsum) / (1e-5 + (q * ksum).sum(dim=-1, keepdim=True))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass for SLA attention.
|
||||
|
||||
Input tensors are in FastVideo format: (B, L, H, D)
|
||||
Internally converted to SLA format: (B, H, L, D)
|
||||
|
||||
Args:
|
||||
query: Query tensor (B, L, H, D)
|
||||
key: Key tensor (B, L, H, D)
|
||||
value: Value tensor (B, L, H, D)
|
||||
attn_metadata: Attention metadata
|
||||
|
||||
Returns:
|
||||
Output tensor (B, L, H, D)
|
||||
"""
|
||||
original_dtype = query.dtype
|
||||
|
||||
# Convert from FastVideo format (B, L, H, D) to SLA format (B, H, L, D)
|
||||
q = query.transpose(1, 2).contiguous()
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Get topk ratio from metadata if available
|
||||
topk_ratio = self.topk_ratio
|
||||
if hasattr(attn_metadata, 'topk_ratio'):
|
||||
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
|
||||
|
||||
# Compute block-sparse attention pattern
|
||||
sparse_map, lut, real_topk = get_block_map(q,
|
||||
k,
|
||||
topk_ratio=topk_ratio,
|
||||
BLKQ=self.BLKQ,
|
||||
BLKK=self.BLKK)
|
||||
|
||||
# Convert to compute dtype
|
||||
q = q.to(self.dtype)
|
||||
k = k.to(self.dtype)
|
||||
v = v.to(self.dtype)
|
||||
|
||||
# Sparse attention
|
||||
o_s = _attention.apply(q, k, v, sparse_map, lut, real_topk, self.BLKQ,
|
||||
self.BLKK)
|
||||
|
||||
# Linear attention with feature maps
|
||||
q_linear = self.feature_map_q(q).contiguous().to(self.dtype)
|
||||
k_linear = self.feature_map_k(k).contiguous().to(self.dtype)
|
||||
o_l = self._calc_linear_attention(q_linear, k_linear, v)
|
||||
|
||||
# Project linear attention output and combine
|
||||
with torch.amp.autocast('cuda', dtype=self.dtype):
|
||||
o_l = self.proj_l(o_l)
|
||||
|
||||
# Combine sparse and linear outputs
|
||||
output = (o_s + o_l).to(original_dtype)
|
||||
|
||||
# Convert back to FastVideo format (B, L, H, D)
|
||||
output = output.transpose(1, 2)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
# Check if spas_sage_attn is available for SageSLA
|
||||
SAGESLA_ENABLED = True
|
||||
try:
|
||||
import spas_sage_attn._qattn as qattn
|
||||
import spas_sage_attn._fused as fused
|
||||
from spas_sage_attn.utils import get_vanilla_qk_quant, block_map_lut_triton
|
||||
except ImportError:
|
||||
SAGESLA_ENABLED = False
|
||||
|
||||
SAGE2PP_ENABLED = True
|
||||
try:
|
||||
from spas_sage_attn._qattn import qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold
|
||||
except ImportError:
|
||||
SAGE2PP_ENABLED = False
|
||||
|
||||
|
||||
class SageSLAAttentionBackend(AttentionBackend):
|
||||
"""Quantized Sparse-Linear Attention backend using SageAttention kernels."""
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_SLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SageSLAAttentionImpl"]:
|
||||
return SageSLAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["SLAAttentionMetadata"]:
|
||||
return SLAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["SLAAttentionMetadataBuilder"]:
|
||||
return SLAAttentionMetadataBuilder
|
||||
|
||||
|
||||
def _get_cuda_arch(device_index: int) -> str:
|
||||
"""Get CUDA architecture string for the given device."""
|
||||
major, minor = torch.cuda.get_device_capability(device_index)
|
||||
return f"sm{major}{minor}"
|
||||
|
||||
|
||||
class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
"""SageSLA attention implementation using quantized SageAttention kernels.
|
||||
|
||||
This uses INT8 quantization for Q/K and FP8 for V to achieve better performance
|
||||
while maintaining accuracy. Requires spas_sage_attn package.
|
||||
|
||||
Args:
|
||||
num_heads: Number of attention heads
|
||||
head_size: Dimension of each head (must be 64 or 128)
|
||||
topk_ratio: Ratio of key blocks to attend to (0-1), default 0.5
|
||||
feature_map: Feature map for linear attention ('softmax', 'elu', 'relu')
|
||||
use_bf16: Whether to use bfloat16 for computation
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool = False,
|
||||
softmax_scale: float | None = None,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
# SageSLA-specific parameters
|
||||
topk_ratio: float = 0.5,
|
||||
feature_map: str = "softmax",
|
||||
use_bf16: bool = True,
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
nn.Module.__init__(self)
|
||||
|
||||
if not SAGESLA_ENABLED:
|
||||
raise ImportError(
|
||||
"SageSLA requires spas_sage_attn. "
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git"
|
||||
)
|
||||
|
||||
assert head_size in [
|
||||
64, 128
|
||||
], f"SageSLA requires head_size in [64, 128], got {head_size}"
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.softmax_scale = softmax_scale if softmax_scale else head_size**-0.5
|
||||
self.causal = causal
|
||||
self.prefix = prefix
|
||||
|
||||
# SageSLA-specific config
|
||||
self.topk_ratio = topk_ratio
|
||||
self.dtype = torch.bfloat16 if use_bf16 else torch.float16
|
||||
|
||||
# Learnable linear projection for combining sparse + linear attention
|
||||
self.proj_l = nn.Linear(head_size, head_size, dtype=torch.float32)
|
||||
|
||||
# Feature map for linear attention
|
||||
# Type annotation for callables
|
||||
self.feature_map_q: Callable[[torch.Tensor], torch.Tensor]
|
||||
self.feature_map_k: Callable[[torch.Tensor], torch.Tensor]
|
||||
if feature_map == "elu":
|
||||
self.feature_map_q = lambda x: F.elu(x) + 1
|
||||
self.feature_map_k = lambda x: F.elu(x) + 1
|
||||
elif feature_map == "relu":
|
||||
self.feature_map_q = F.relu
|
||||
self.feature_map_k = F.relu
|
||||
elif feature_map == "softmax":
|
||||
self.feature_map_q = lambda x: F.softmax(x, dim=-1)
|
||||
self.feature_map_k = lambda x: F.softmax(x, dim=-1)
|
||||
else:
|
||||
raise ValueError(f"Unknown feature map: {feature_map}")
|
||||
|
||||
self._init_weights()
|
||||
|
||||
def _init_weights(self) -> None:
|
||||
"""Initialize projection weights to zero for residual-like behavior."""
|
||||
with torch.no_grad():
|
||||
nn.init.zeros_(self.proj_l.weight)
|
||||
nn.init.zeros_(self.proj_l.bias) # type: ignore[arg-type]
|
||||
|
||||
def _calc_linear_attention(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Compute linear attention: (Q @ K^T @ V) / normalizer."""
|
||||
kvsum = k.transpose(-1, -2) @ v
|
||||
ksum = torch.sum(k, dim=-2, keepdim=True)
|
||||
return (q @ kvsum) / (1e-5 + (q * ksum).sum(dim=-1, keepdim=True))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass for SageSLA attention with quantized kernels.
|
||||
|
||||
Input tensors are in FastVideo format: (B, L, H, D)
|
||||
|
||||
Args:
|
||||
query: Query tensor (B, L, H, D)
|
||||
key: Key tensor (B, L, H, D)
|
||||
value: Value tensor (B, L, H, D)
|
||||
attn_metadata: Attention metadata
|
||||
|
||||
Returns:
|
||||
Output tensor (B, L, H, D)
|
||||
"""
|
||||
original_dtype = query.dtype
|
||||
|
||||
# Convert from FastVideo format (B, L, H, D) to SLA format (B, H, L, D)
|
||||
q = query.transpose(1, 2).contiguous()
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Get topk ratio from metadata if available
|
||||
topk_ratio = self.topk_ratio
|
||||
if hasattr(attn_metadata, 'topk_ratio'):
|
||||
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
|
||||
|
||||
# Determine block sizes based on GPU architecture
|
||||
arch = _get_cuda_arch(q.device.index)
|
||||
if arch == "sm90":
|
||||
BLKQ, BLKK = 64, 128
|
||||
else:
|
||||
BLKQ, BLKK = 128, 64
|
||||
|
||||
# Compute block-sparse attention pattern
|
||||
sparse_map, lut, real_topk = get_block_map(q,
|
||||
k,
|
||||
topk_ratio=topk_ratio,
|
||||
BLKQ=BLKQ,
|
||||
BLKK=BLKK)
|
||||
|
||||
# Convert to compute dtype
|
||||
q = q.to(self.dtype)
|
||||
k = k.to(self.dtype)
|
||||
v = v.to(self.dtype)
|
||||
|
||||
# ========== SPARGE QUANTIZED ATTENTION ==========
|
||||
km = k.mean(dim=-2, keepdim=True)
|
||||
headdim = q.size(-1)
|
||||
scale = 1.0 / (headdim**0.5)
|
||||
|
||||
# Quantize Q, K to INT8
|
||||
q_int8, q_scale, k_int8, k_scale = get_vanilla_qk_quant(
|
||||
q, k, km, BLKQ, BLKK)
|
||||
lut_triton, valid_block_num = block_map_lut_triton(sparse_map)
|
||||
|
||||
# Quantize V to FP8
|
||||
b, h_kv, kv_len, head_dim = v.shape
|
||||
padded_len = (kv_len + 127) // 128 * 128
|
||||
v_transposed_permutted = torch.empty((b, h_kv, head_dim, padded_len),
|
||||
dtype=v.dtype,
|
||||
device=v.device)
|
||||
fused.transpose_pad_permute_cuda(v, v_transposed_permutted, 1)
|
||||
v_fp8 = torch.empty(v_transposed_permutted.shape,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=v.device)
|
||||
v_scale = torch.empty((b, h_kv, head_dim),
|
||||
dtype=torch.float32,
|
||||
device=v.device)
|
||||
fused.scale_fuse_quant_cuda(v_transposed_permutted, v_fp8, v_scale,
|
||||
kv_len, 2.25, 1)
|
||||
|
||||
# Sparse attention with quantized kernels
|
||||
o_s = torch.empty_like(q)
|
||||
if arch == "sm90":
|
||||
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_sm90(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
q_scale, k_scale, v_scale, 1, False, 1, scale)
|
||||
else:
|
||||
pvthreshold = torch.full((q.shape[-3], ),
|
||||
1e6,
|
||||
dtype=torch.float32,
|
||||
device=q.device)
|
||||
if SAGE2PP_ENABLED:
|
||||
qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
pvthreshold, q_scale, k_scale, v_scale, 1, False, 1, scale,
|
||||
0)
|
||||
else:
|
||||
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
pvthreshold, q_scale, k_scale, v_scale, 1, False, 1, scale,
|
||||
0)
|
||||
# ========== END SPARGE ==========
|
||||
|
||||
# Linear attention with feature maps
|
||||
q_linear = self.feature_map_q(q).contiguous().to(self.dtype)
|
||||
k_linear = self.feature_map_k(k).contiguous().to(self.dtype)
|
||||
o_l = self._calc_linear_attention(q_linear, k_linear, v)
|
||||
|
||||
# Project linear attention output and combine
|
||||
with torch.amp.autocast('cuda', dtype=self.dtype):
|
||||
o_l = self.proj_l(o_l)
|
||||
|
||||
# Combine sparse and linear outputs
|
||||
output = (o_s + o_l).to(original_dtype)
|
||||
|
||||
# Convert back to FastVideo format (B, L, H, D)
|
||||
output = output.transpose(1, 2)
|
||||
|
||||
return output
|
||||
@@ -276,8 +276,9 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
variable_block_sizes=attn_metadata.variable_block_sizes,
|
||||
topk=cur_topk,
|
||||
attn_metadata.variable_block_sizes,
|
||||
attn_metadata.variable_block_sizes,
|
||||
cur_topk,
|
||||
block_size=VSA_TILE_SIZE,
|
||||
compress_attn_weight=gate_compress).transpose(1, 2)
|
||||
|
||||
|
||||
@@ -50,6 +50,10 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
# Register attn_impl as submodule if it has learnable parameters (e.g., SLA's proj_l)
|
||||
# This ensures its parameters are included in state_dict() for saving/loading
|
||||
if isinstance(self.attn_impl, nn.Module):
|
||||
self.add_module('attn_impl', self.attn_impl)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
|
||||
@@ -2,5 +2,16 @@ from fastvideo.configs.models.base import ModelConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig)
|
||||
|
||||
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
__all__ = [
|
||||
"ModelConfig",
|
||||
"VAEConfig",
|
||||
"DiTConfig",
|
||||
"EncoderConfig",
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.configs.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 audio VAE and vocoder configuration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioArchConfig(ArchConfig):
|
||||
architectures: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioEncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioEncoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioDecoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioDecoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VocoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2Vocoder"]))
|
||||
@@ -3,11 +3,12 @@ from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
|
||||
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
|
||||
"LongCatVideoConfig"
|
||||
"LongCatVideoConfig", "LTX2VideoConfig"
|
||||
]
|
||||
|
||||
@@ -18,7 +18,8 @@ class DiTArchConfig(ArchConfig):
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE)
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE, AttentionBackendEnum.SLA_ATTN,
|
||||
AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 Transformer configuration for native FastVideo integration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_ltx2_blocks(name: str, _module) -> bool:
|
||||
"""FSDP shard condition for LTX-2 transformer blocks."""
|
||||
return "transformer_blocks" in name
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VideoArchConfig(DiTArchConfig):
|
||||
"""Architecture configuration for LTX-2 video transformer."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_ltx2_blocks])
|
||||
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
|
||||
|
||||
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo)
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^model\.diffusion_model\.(.*)$": r"model.\1",
|
||||
r"^diffusion_model\.(.*)$": r"model.\1",
|
||||
r"^model\.(.*)$": r"model.\1",
|
||||
r"^(.*)$": r"model.\1",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
lora_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Core transformer settings (defaults from LTX-2 metadata)
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
num_layers: int = 48
|
||||
cross_attention_dim: int = 4096
|
||||
caption_channels: int = 3840
|
||||
norm_eps: float = 1e-6
|
||||
attention_type: str = "default"
|
||||
rope_type: str = "split"
|
||||
double_precision_rope: bool = True
|
||||
|
||||
positional_embedding_theta: float = 10000.0
|
||||
positional_embedding_max_pos: list[int] = field(
|
||||
default_factory=lambda: [20, 2048, 2048])
|
||||
timestep_scale_multiplier: int = 1000
|
||||
use_middle_indices_grid: bool = True
|
||||
|
||||
# Patchification (video-only path)
|
||||
patch_size: tuple[int, int, int] = (1, 1, 1)
|
||||
num_channels_latents: int = 128
|
||||
in_channels: int | None = None
|
||||
out_channels: int | None = None
|
||||
|
||||
# Audio defaults (reserved for joint AV ports)
|
||||
audio_num_attention_heads: int = 32
|
||||
audio_attention_head_dim: int = 64
|
||||
audio_in_channels: int = 128
|
||||
audio_out_channels: int = 128
|
||||
audio_cross_attention_dim: int = 2048
|
||||
audio_positional_embedding_max_pos: list[int] = field(
|
||||
default_factory=lambda: [20])
|
||||
av_ca_timestep_scale_multiplier: int = 1
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
patch_volume = self.patch_size[0] * self.patch_size[
|
||||
1] * self.patch_size[2]
|
||||
if self.in_channels is None:
|
||||
self.in_channels = self.num_channels_latents * patch_volume
|
||||
if self.out_channels is None:
|
||||
self.out_channels = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VideoConfig(DiTConfig):
|
||||
"""Main configuration for LTX-2 transformer."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
|
||||
prefix: str = "ltx2"
|
||||
@@ -7,10 +7,12 @@ from fastvideo.configs.models.encoders.clip import (
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
|
||||
"Qwen2_5_VLConfig"
|
||||
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2GemmaArchConfig(TextEncoderArchConfig):
|
||||
architectures: list[str] = field(
|
||||
default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
|
||||
hidden_size: int = 3840
|
||||
num_hidden_layers: int = 48
|
||||
num_attention_heads: int = 30
|
||||
text_len: int = 1024
|
||||
pad_token_id: int = 0
|
||||
eos_token_id: int = 2
|
||||
|
||||
gemma_model_path: str = ""
|
||||
gemma_dtype: str = "bfloat16"
|
||||
padding_side: str = "left"
|
||||
|
||||
feature_extractor_in_features: int = 3840 * 49
|
||||
feature_extractor_out_features: int = 3840
|
||||
|
||||
connector_num_attention_heads: int = 30
|
||||
connector_attention_head_dim: int = 128
|
||||
connector_num_layers: int = 2
|
||||
connector_positional_embedding_theta: float = 10000.0
|
||||
connector_positional_embedding_max_pos: list[int] = field(
|
||||
default_factory=lambda: [4096])
|
||||
connector_rope_type: str = "split"
|
||||
connector_double_precision_rope: bool = False
|
||||
connector_num_learnable_registers: int | None = 128
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2GemmaConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=LTX2GemmaArchConfig)
|
||||
|
||||
prefix: str = "ltx2_gemma"
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for Reason1 (Qwen2.5-VL) text encoder."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Reason1ArchConfig(TextEncoderArchConfig):
|
||||
"""Architecture settings (defaults match Qwen2.5-VL-7B-Instruct)."""
|
||||
|
||||
architectures: list[str] = field(
|
||||
default_factory=lambda: ["Qwen2_5_VLForConditionalGeneration"])
|
||||
model_type: str = "qwen2_5_vl"
|
||||
|
||||
vocab_size: int = 152064
|
||||
hidden_size: int = 3584
|
||||
num_hidden_layers: int = 28
|
||||
num_attention_heads: int = 28
|
||||
num_key_value_heads: int = 4
|
||||
intermediate_size: int = 18944
|
||||
|
||||
text_len: int = 512
|
||||
hidden_state_skip_layer: int = 0
|
||||
bos_token_id: int = 151643
|
||||
pad_token_id: int = 151643
|
||||
eos_token_id: int = 151645
|
||||
|
||||
image_token_id: int = 151655
|
||||
video_token_id: int = 151656
|
||||
vision_token_id: int = 151654
|
||||
vision_start_token_id: int = 151652
|
||||
vision_end_token_id: int = 151653
|
||||
|
||||
vision_config: dict[str, Any] | None = None
|
||||
|
||||
rope_theta: float = 1000000.0
|
||||
rope_scaling: dict[str, Any] | None = field(default_factory=lambda: {
|
||||
"type": "mrope",
|
||||
"mrope_section": [16, 24, 24]
|
||||
})
|
||||
max_position_embeddings: int = 128000
|
||||
max_window_layers: int = 28
|
||||
|
||||
embedding_concat_strategy: str = "mean_pooling"
|
||||
n_layers_per_group: int = 5
|
||||
num_embedding_padding_tokens: int = 512
|
||||
|
||||
attention_dropout: float = 0.0
|
||||
hidden_act: str = "silu"
|
||||
initializer_range: float = 0.02
|
||||
rms_norm_eps: float = 1e-6
|
||||
|
||||
use_sliding_window: bool = False
|
||||
sliding_window: int = 32768
|
||||
|
||||
tie_word_embeddings: bool = False
|
||||
use_cache: bool = False
|
||||
output_hidden_states: bool = True
|
||||
|
||||
torch_dtype: str = "bfloat16"
|
||||
_attn_implementation: str = "flash_attention_2"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Reason1Config(TextEncoderConfig):
|
||||
"""Reason1 text encoder config."""
|
||||
|
||||
arch_config: Reason1ArchConfig = field(default_factory=Reason1ArchConfig)
|
||||
tokenizer_type: str = "Qwen/Qwen2.5-VL-7B-Instruct"
|
||||
@@ -1,6 +1,8 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
@@ -9,5 +11,7 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"StepVideoVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Cosmos 2.5 (Wan2.1-style) VAE config and checkpoint-key mapping."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25VAEArchConfig(VAEArchConfig):
|
||||
_name_or_path: str = ""
|
||||
base_dim: int = 96
|
||||
decoder_base_dim: int | None = None
|
||||
z_dim: int = 16
|
||||
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: tuple[float, ...] = ()
|
||||
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||
dropout: float = 0.0
|
||||
is_residual: bool = False
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
patch_size: int | None = None
|
||||
scale_factor_temporal: int = 4
|
||||
scale_factor_spatial: int = 8
|
||||
clip_output: bool = True
|
||||
|
||||
latents_mean: tuple[float, ...] = (
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
0.1075,
|
||||
-0.1745,
|
||||
0.9653,
|
||||
-0.1517,
|
||||
1.5508,
|
||||
0.4134,
|
||||
-0.0715,
|
||||
0.5517,
|
||||
-0.3632,
|
||||
-0.1922,
|
||||
-0.9497,
|
||||
0.2503,
|
||||
-0.2921,
|
||||
)
|
||||
latents_std: tuple[float, ...] = (
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
2.6558,
|
||||
1.2196,
|
||||
1.7708,
|
||||
2.6052,
|
||||
2.0743,
|
||||
3.2687,
|
||||
2.1526,
|
||||
2.8652,
|
||||
1.5579,
|
||||
1.6382,
|
||||
1.1253,
|
||||
2.8251,
|
||||
1.9160,
|
||||
)
|
||||
|
||||
# Simple 1:1 renames. More complex decoder remapping is handled by
|
||||
# `map_official_key()`.
|
||||
param_names_mapping: dict[str, str] = field(
|
||||
default_factory=lambda: {
|
||||
r"^conv1\.(.*)$": r"quant_conv.\1",
|
||||
r"^conv2\.(.*)$": r"post_quant_conv.\1",
|
||||
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
|
||||
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
|
||||
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
|
||||
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
|
||||
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
|
||||
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def map_official_key(key: str) -> str | None:
|
||||
"""Map a single official checkpoint key into FastVideo key space."""
|
||||
|
||||
def map_residual_subkey(prefix: str, sub: str) -> str | None:
|
||||
if re.match(r"^residual\.0\.gamma$", sub):
|
||||
return f"{prefix}.norm1.gamma"
|
||||
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv1.{m.group(1)}"
|
||||
if re.match(r"^residual\.3\.gamma$", sub):
|
||||
return f"{prefix}.norm2.gamma"
|
||||
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv2.{m.group(1)}"
|
||||
m = re.match(r"^shortcut\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv_shortcut.{m.group(1)}"
|
||||
return None
|
||||
|
||||
def map_attn_subkey(prefix: str, sub: str) -> str | None:
|
||||
if re.match(r"^norm\.gamma$", sub):
|
||||
return f"{prefix}.norm.gamma"
|
||||
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.to_qkv.{m.group(1)}"
|
||||
m = re.match(r"^proj\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.proj.{m.group(1)}"
|
||||
return None
|
||||
|
||||
def map_resample_subkey(prefix: str, sub: str) -> str | None:
|
||||
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.resample.1.{m.group(1)}"
|
||||
m = re.match(r"^time_conv\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.time_conv.{m.group(1)}"
|
||||
return None
|
||||
|
||||
m = re.match(r"^conv1\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"quant_conv.{m.group(1)}"
|
||||
m = re.match(r"^conv2\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"post_quant_conv.{m.group(1)}"
|
||||
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.conv_in.{m.group(2)}"
|
||||
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.norm_out.gamma"
|
||||
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.conv_out.{m.group(2)}"
|
||||
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
|
||||
if m:
|
||||
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0",
|
||||
m.group(2))
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
|
||||
if m:
|
||||
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0",
|
||||
m.group(2))
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
|
||||
if m:
|
||||
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1",
|
||||
m.group(2))
|
||||
|
||||
m = re.match(r"^encoder\.downsamples\.(\d+)\.(.*)$", key)
|
||||
if m:
|
||||
idx = int(m.group(1))
|
||||
sub = m.group(2)
|
||||
if sub.startswith("residual.") or sub.startswith("shortcut."):
|
||||
return map_residual_subkey(f"encoder.down_blocks.{idx}", sub)
|
||||
if sub.startswith("resample.") or sub.startswith("time_conv."):
|
||||
return map_resample_subkey(f"encoder.down_blocks.{idx}", sub)
|
||||
return None
|
||||
|
||||
m = re.match(r"^decoder\.upsamples\.(\d+)\.(.*)$", key)
|
||||
if m:
|
||||
uidx = int(m.group(1))
|
||||
sub = m.group(2)
|
||||
|
||||
if uidx in (0, 1, 2):
|
||||
block_i, res_i = 0, uidx
|
||||
elif uidx == 3:
|
||||
block_i, res_i = 0, None
|
||||
elif uidx in (4, 5, 6):
|
||||
block_i, res_i = 1, uidx - 4
|
||||
elif uidx == 7:
|
||||
block_i, res_i = 1, None
|
||||
elif uidx in (8, 9, 10):
|
||||
block_i, res_i = 2, uidx - 8
|
||||
elif uidx == 11:
|
||||
block_i, res_i = 2, None
|
||||
elif uidx in (12, 13, 14):
|
||||
block_i, res_i = 3, uidx - 12
|
||||
else:
|
||||
return None
|
||||
|
||||
if res_i is None:
|
||||
return map_resample_subkey(
|
||||
f"decoder.up_blocks.{block_i}.upsamplers.0",
|
||||
sub,
|
||||
)
|
||||
|
||||
return map_residual_subkey(
|
||||
f"decoder.up_blocks.{block_i}.resnets.{res_i}",
|
||||
sub,
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
temporal_compression_ratio: int = 4
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
def __post_init__(self):
|
||||
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(
|
||||
self.latents_std).view(1, self.z_dim, 1, 1, 1)
|
||||
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
self.temporal_compression_ratio = self.scale_factor_temporal
|
||||
self.spatial_compression_ratio = self.scale_factor_spatial
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25VAEConfig(VAEConfig):
|
||||
"""Cosmos2.5 VAE config."""
|
||||
|
||||
arch_config: Cosmos25VAEArchConfig = field(
|
||||
default_factory=Cosmos25VAEArchConfig)
|
||||
|
||||
use_feature_cache: bool = True
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.blend_num_frames = (self.tile_sample_min_num_frames -
|
||||
self.tile_sample_stride_num_frames) * 2
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 VAE configuration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VAEArchConfig(VAEArchConfig):
|
||||
# Mirrors LTX-2 safetensors metadata config under "vae"
|
||||
_class_name: str = "CausalVideoAutoencoder"
|
||||
dims: int = 3
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 128
|
||||
encoder_blocks: list = field(default_factory=list)
|
||||
decoder_blocks: list = field(default_factory=list)
|
||||
patch_size: int = 4
|
||||
norm_layer: str = "pixel_norm"
|
||||
latent_log_var: str = "uniform"
|
||||
encoder_spatial_padding_mode: str = "zeros"
|
||||
decoder_spatial_padding_mode: str = "reflect"
|
||||
causal_decoder: bool = False
|
||||
timestep_conditioning: bool = True
|
||||
use_quant_conv: bool = False
|
||||
scaling_factor: float = 1.0
|
||||
normalize_latent_channels: bool = False
|
||||
|
||||
# Match FastVideo naming for compression ratios (LTX-2 default)
|
||||
temporal_compression_ratio: int = 8
|
||||
spatial_compression_ratio: int = 32
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
|
||||
|
||||
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
|
||||
ltx2_spatial_tile_size_in_pixels: int = 512
|
||||
ltx2_spatial_tile_overlap_in_pixels: int = 64
|
||||
ltx2_temporal_tile_size_in_frames: int = 64
|
||||
ltx2_temporal_tile_overlap_in_frames: int = 24
|
||||
@@ -1,8 +1,10 @@
|
||||
from fastvideo.configs.pipelines.base import (PipelineConfig,
|
||||
SlidingTileAttnConfig)
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
@@ -15,5 +17,6 @@ __all__ = [
|
||||
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "get_pipeline_config_cls_from_name"
|
||||
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25ArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1Config, Reason1ArchConfig
|
||||
from fastvideo.configs.models.vaes import Cosmos25VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
|
||||
|
||||
def _identity_preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def reason1_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
hidden_states = getattr(outputs, "hidden_states", None)
|
||||
if hidden_states is None:
|
||||
raise ValueError("Reason1 postprocess requires outputs.hidden_states")
|
||||
|
||||
hs = list(hidden_states)[1:]
|
||||
normed = []
|
||||
for h in hs:
|
||||
h = h.float()
|
||||
h = (h - h.mean(dim=-1, keepdim=True)) / (h.std(dim=-1, keepdim=True) +
|
||||
1e-8)
|
||||
normed.append(h)
|
||||
return torch.cat(normed, dim=-1).to(hidden_states[0].dtype)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25Config(PipelineConfig):
|
||||
"""Configuration for Cosmos 2.5 (Predict2.5) video generation pipeline."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=lambda: Cosmos25VideoConfig(
|
||||
arch_config=Cosmos25ArchConfig(
|
||||
num_attention_heads=16,
|
||||
attention_head_dim=128,
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
num_layers=28,
|
||||
patch_size=[1, 2, 2],
|
||||
max_size=[128, 240, 240],
|
||||
rope_scale=[1.0, 3.0, 3.0],
|
||||
text_embed_dim=1024,
|
||||
mlp_ratio=4.0,
|
||||
adaln_lora_dim=256,
|
||||
use_adaln_lora=True,
|
||||
concat_padding_mask=True,
|
||||
extra_pos_embed_type=None,
|
||||
use_crossattn_projection=True,
|
||||
rope_enable_fps_modulation=False,
|
||||
qk_norm="rms_norm",
|
||||
)))
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=Cosmos25VAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (Reason1Config(arch_config=Reason1ArchConfig(
|
||||
embedding_concat_strategy="full_concat")), ))
|
||||
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (_identity_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(reason1_postprocess_text, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", ))
|
||||
|
||||
embedded_cfg_scale: float = 0.0
|
||||
flow_shift: float = 5.0
|
||||
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
STA_mode: STA_Mode = STA_Mode.NONE
|
||||
skip_time_steps: int = 0
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self._vae_latent_dim = 16
|
||||
@@ -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()))
|
||||
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
LTX2AudioDecoderConfig, LTX2VocoderConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.configs.models.dits import LTX2VideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, LTX2GemmaConfig
|
||||
from fastvideo.configs.models.vaes import LTX2VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def ltx2_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2T2VConfig(PipelineConfig):
|
||||
"""Configuration for LTX-2 T2V pipeline."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=LTX2VideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=LTX2VAEConfig)
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LTX2GemmaConfig(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(ltx2_postprocess_text, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", ))
|
||||
|
||||
audio_decoder_config: ModelConfig = field(
|
||||
default_factory=LTX2AudioDecoderConfig)
|
||||
vocoder_config: ModelConfig = field(default_factory=LTX2VocoderConfig)
|
||||
audio_decoder_precision: str = "bf16"
|
||||
vocoder_precision: str = "bf16"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -6,10 +6,15 @@ from collections.abc import Callable
|
||||
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.configs.pipelines.turbodiffusion import (
|
||||
TurboDiffusionT2V_1_3B_Config, TurboDiffusionT2V_14B_Config,
|
||||
TurboDiffusionI2V_A14B_Config)
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
@@ -52,14 +57,32 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World": CosmosConfig,
|
||||
"KyleShao/Cosmos-Predict2.5-2B-Diffusers": Cosmos25Config,
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGameI2V480PConfig,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGameI2V480PConfig,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGameI2V480PConfig,
|
||||
# LongCat Video models
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
|
||||
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
|
||||
# LTX-2 models
|
||||
"Lightricks/LTX-2": LTX2T2VConfig,
|
||||
"converted/ltx2_diffusers": LTX2T2VConfig,
|
||||
# TurboDiffusion models
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers": TurboDiffusionT2V_1_3B_Config,
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers": TurboDiffusionT2V_14B_Config,
|
||||
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers": TurboDiffusionI2V_A14B_Config,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"longcatimagetovideo":
|
||||
lambda id: "longcatimagetovideo" in id.lower(),
|
||||
"longcatvideocontinuation":
|
||||
lambda id: "longcatvideocontinuation" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
"hunyuan":
|
||||
lambda id: "hunyuan" in id.lower(),
|
||||
"hunyuan15":
|
||||
@@ -77,15 +100,23 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"stepvideo":
|
||||
lambda id: "stepvideo" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
lambda id: "cosmos" in id.lower() and ("2.5" not in id.lower(
|
||||
) and "2_5" not in id.lower() and "25" not in id.lower()),
|
||||
"cosmos25":
|
||||
lambda id: "cosmos25" in id.lower(),
|
||||
"turbodiffusion":
|
||||
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
|
||||
"ltx2":
|
||||
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"longcatimagetovideo": LongCatT2V480PConfig,
|
||||
"longcatvideocontinuation": LongCatT2V480PConfig,
|
||||
"longcat": LongCatT2V480PConfig,
|
||||
"cosmos25": Cosmos25Config,
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"matrixgame": MatrixGameI2V480PConfig,
|
||||
@@ -96,7 +127,9 @@ 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,
|
||||
"ltx2": LTX2T2VConfig,
|
||||
# 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",
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos_Predict2_5_2B_Diffusers_SamplingParam(SamplingParam):
|
||||
"""Defaults for Cosmos 2.5 (Predict2.5) text-to-video diffusers-format model."""
|
||||
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 121
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 7.0
|
||||
# Official Cosmos2.5 sampling uses empty string as unconditional.
|
||||
negative_prompt: str = ""
|
||||
num_inference_steps: int = 35
|
||||
@@ -0,0 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2SamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled T2V.
|
||||
"""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 1024
|
||||
width: int = 1536
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 8
|
||||
guidance_scale: float = 1.0
|
||||
# No default negative_prompt for distilled models
|
||||
negative_prompt: str = ""
|
||||
@@ -9,6 +9,8 @@ from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hun
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
|
||||
from fastvideo.configs.sample.cosmos2_5 import Cosmos_Predict2_5_2B_Diffusers_SamplingParam
|
||||
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
@@ -26,6 +28,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,
|
||||
@@ -79,11 +86,27 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World":
|
||||
Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
|
||||
# Cosmos2.5
|
||||
"KyleShao/Cosmos-Predict2.5-2B-Diffusers":
|
||||
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
|
||||
|
||||
# MatrixGame2.0 models
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
|
||||
|
||||
# TurboDiffusion models
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers":
|
||||
TurboDiffusionT2V_1_3B_SamplingParam,
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers":
|
||||
TurboDiffusionT2V_14B_SamplingParam,
|
||||
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers":
|
||||
TurboDiffusionI2V_A14B_SamplingParam,
|
||||
|
||||
# LTX-2 models
|
||||
"Lightricks/LTX-2": LTX2SamplingParam,
|
||||
"FastVideo/LTX2-Distilled-Diffusers": LTX2SamplingParam,
|
||||
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -105,6 +128,14 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"matrixgame":
|
||||
lambda id: "matrixgame" in id.lower() or "matrix-game" in id.lower(),
|
||||
"turbodiffusion":
|
||||
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
|
||||
"cosmos25":
|
||||
lambda id: "cosmos2_5" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower() and "2_5" not in id.lower(),
|
||||
"ltx2":
|
||||
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -121,6 +152,11 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam,
|
||||
"matrixgame": MatrixGame2_SamplingParam,
|
||||
"turbodiffusion":
|
||||
TurboDiffusionT2V_1_3B_SamplingParam, # Default to T2V for fallback
|
||||
"cosmos25": Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
|
||||
"cosmos": Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
"ltx2": LTX2SamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
@@ -144,9 +180,6 @@ def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal sampling param for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,288 @@
|
||||
import asyncio
|
||||
import os
|
||||
from concurrent.futures import Future, ThreadPoolExecutor
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ForwardBatch
|
||||
from fastvideo.utils import align_to, shallow_asdict
|
||||
from fastvideo.worker.executor import Executor
|
||||
from fastvideo.worker.multiproc_executor import MultiprocExecutor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class IncrementalVideoWriter:
|
||||
|
||||
def __init__(self, path: str, fps: int = 24, block_dir: str | None = None):
|
||||
self._executor = ThreadPoolExecutor(max_workers=2,
|
||||
thread_name_prefix="video_write_")
|
||||
self._path = path
|
||||
self._writer = imageio.get_writer(path, fps=fps, format="mp4")
|
||||
self._pending_main: Future | None = None
|
||||
self._block_dir = block_dir
|
||||
self._block_idx = 0
|
||||
self._fps = fps
|
||||
|
||||
@property
|
||||
def path(self) -> str:
|
||||
return self._path
|
||||
|
||||
def add_frames(self, frames: list[np.ndarray]) -> Future | None:
|
||||
# Wait for previous main video write to complete
|
||||
if self._pending_main is not None:
|
||||
self._pending_main.result()
|
||||
|
||||
# Copy frames to avoid race conditions
|
||||
frames_copy = [f.copy() for f in frames]
|
||||
self._pending_main = self._executor.submit(self._write_frames,
|
||||
frames_copy)
|
||||
|
||||
# Write block file if block_dir is set
|
||||
block_future = None
|
||||
if self._block_dir:
|
||||
self._block_idx += 1
|
||||
block_path = os.path.join(self._block_dir,
|
||||
f"b{self._block_idx}.mp4")
|
||||
block_future = self._executor.submit(self._write_block, frames_copy,
|
||||
block_path)
|
||||
return block_future
|
||||
|
||||
def _write_frames(self, frames: list[np.ndarray]) -> None:
|
||||
for frame in frames:
|
||||
self._writer.append_data(frame)
|
||||
|
||||
def _write_block(self, frames: list[np.ndarray], path: str) -> str:
|
||||
imageio.mimsave(path, frames, fps=self._fps)
|
||||
return path
|
||||
|
||||
def close(self) -> None:
|
||||
if self._pending_main is not None:
|
||||
self._pending_main.result()
|
||||
self._pending_main = None
|
||||
if self._writer:
|
||||
self._writer.close()
|
||||
self._writer = None
|
||||
self._executor.shutdown(wait=True)
|
||||
|
||||
|
||||
class StreamingVideoGenerator(VideoGenerator):
|
||||
"""
|
||||
This class extends VideoGenerator with streaming capabilities,
|
||||
allowing incremental video generation with step-by-step control.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
executor_class: type[Executor],
|
||||
log_stats: bool,
|
||||
use_queue_mode: bool = True):
|
||||
super().__init__(fastvideo_args, executor_class, log_stats)
|
||||
self.accumulated_frames: list[np.ndarray] = []
|
||||
self.sampling_param: SamplingParam | None = None
|
||||
self.batch: ForwardBatch | None = None
|
||||
self._use_queue_mode = use_queue_mode and isinstance(
|
||||
self.executor, MultiprocExecutor)
|
||||
self.writer: IncrementalVideoWriter | None = None
|
||||
self.block_dir: str | None = None
|
||||
self.block_idx: int = 0
|
||||
|
||||
@classmethod
|
||||
def from_fastvideo_args(
|
||||
cls, fastvideo_args: FastVideoArgs) -> "StreamingVideoGenerator":
|
||||
executor_class = Executor.get_class(fastvideo_args)
|
||||
return cls(
|
||||
fastvideo_args=fastvideo_args,
|
||||
executor_class=executor_class,
|
||||
log_stats=False,
|
||||
)
|
||||
|
||||
def reset(
|
||||
self,
|
||||
prompt: str = "A gameplay video of a cyberpunk city",
|
||||
image_path: str | None = None,
|
||||
num_frames: int = 120, # Default max frames
|
||||
**kwargs):
|
||||
self.accumulated_frames = []
|
||||
self.block_idx = 0
|
||||
self.block_dir = None
|
||||
if self.writer:
|
||||
self.writer.close()
|
||||
self.writer = None
|
||||
self.executor.execute_streaming_clear()
|
||||
|
||||
# Handle batch processing from text file
|
||||
if self.sampling_param is None:
|
||||
self.sampling_param = SamplingParam.from_pretrained(
|
||||
self.fastvideo_args.model_path)
|
||||
|
||||
self.sampling_param.update(kwargs)
|
||||
self.sampling_param.prompt = prompt
|
||||
if image_path:
|
||||
self.sampling_param.image_path = image_path
|
||||
self.sampling_param.num_frames = num_frames
|
||||
|
||||
if "output_path" in kwargs:
|
||||
output_path = self._prepare_output_path(kwargs["output_path"],
|
||||
prompt)
|
||||
# Create block directory for individual block files
|
||||
block_dir = output_path.replace(".mp4", "")
|
||||
os.makedirs(block_dir, exist_ok=True)
|
||||
self.block_dir = block_dir
|
||||
self.writer = IncrementalVideoWriter(output_path,
|
||||
fps=24,
|
||||
block_dir=block_dir)
|
||||
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
self.sampling_param.height = align_to(self.sampling_param.height, 16)
|
||||
self.sampling_param.width = align_to(self.sampling_param.width, 16)
|
||||
|
||||
latents_size = [(self.sampling_param.num_frames - 1) // 4 + 1,
|
||||
self.sampling_param.height // 8,
|
||||
self.sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
self.sampling_param.return_frames = True
|
||||
self.sampling_param.save_video = False
|
||||
|
||||
self.batch = ForwardBatch(
|
||||
**shallow_asdict(self.sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
)
|
||||
|
||||
if self._use_queue_mode:
|
||||
self.executor.submit_reset(self.batch, fastvideo_args)
|
||||
result = self.executor.wait_result()
|
||||
if result.error:
|
||||
raise result.error
|
||||
else:
|
||||
self.executor.execute_streaming_reset(self.batch, fastvideo_args)
|
||||
|
||||
def step(
|
||||
self, keyboard_cond: torch.Tensor,
|
||||
mouse_cond: torch.Tensor) -> tuple[list[np.ndarray], Future | None]:
|
||||
if self.batch is None:
|
||||
raise RuntimeError("Call reset() before step()")
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.submit_step(keyboard_cond, mouse_cond)
|
||||
result = self.executor.wait_result()
|
||||
if result.error:
|
||||
raise result.error
|
||||
output_batch = result.output_batch
|
||||
else:
|
||||
# Fallback to RPC-based
|
||||
output_batch = self.executor.execute_streaming_step(
|
||||
keyboard_action=keyboard_cond, mouse_action=mouse_cond)
|
||||
|
||||
frames = self._process_output_batch(output_batch)
|
||||
block_future = None
|
||||
if len(frames) > 0:
|
||||
self.accumulated_frames.extend(frames)
|
||||
self.block_idx += 1
|
||||
if self.writer:
|
||||
# Returns Future for block file, or None if no block_dir
|
||||
block_future = self.writer.add_frames(frames)
|
||||
|
||||
return frames, block_future
|
||||
|
||||
async def step_async(
|
||||
self, keyboard_cond: torch.Tensor,
|
||||
mouse_cond: torch.Tensor) -> tuple[list[np.ndarray], Future | None]:
|
||||
if self.batch is None:
|
||||
raise RuntimeError("Call reset() before step_async()")
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.submit_step(keyboard_cond, mouse_cond)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
result = await loop.run_in_executor(None, self.executor.wait_result)
|
||||
|
||||
if result.error:
|
||||
raise result.error
|
||||
output_batch = result.output_batch
|
||||
else:
|
||||
# Fallback to RPC-based
|
||||
output_batch = await self.executor.execute_streaming_step_async(
|
||||
keyboard_action=keyboard_cond,
|
||||
mouse_action=mouse_cond,
|
||||
)
|
||||
|
||||
frames = self._process_output_batch(output_batch)
|
||||
block_future = None
|
||||
if len(frames) > 0:
|
||||
self.accumulated_frames.extend(frames)
|
||||
self.block_idx += 1
|
||||
if self.writer:
|
||||
block_future = self.writer.add_frames(frames)
|
||||
|
||||
return frames, block_future
|
||||
|
||||
def finalize(self,
|
||||
output_path: str = "streaming_output.mp4",
|
||||
fps: int = 24) -> str:
|
||||
if not self.accumulated_frames:
|
||||
logger.warning("No frames to save.")
|
||||
return ""
|
||||
|
||||
if self.writer:
|
||||
output_path = self.writer.path
|
||||
self.writer.close()
|
||||
self.writer = None
|
||||
logger.info("Saved video to %s", output_path)
|
||||
else:
|
||||
imageio.mimsave(output_path,
|
||||
self.accumulated_frames,
|
||||
fps=fps,
|
||||
format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.submit_clear()
|
||||
else:
|
||||
self.executor.execute_streaming_clear()
|
||||
self.accumulated_frames = []
|
||||
return output_path
|
||||
|
||||
def _process_output_batch(self,
|
||||
output_batch: ForwardBatch) -> list[np.ndarray]:
|
||||
if output_batch.output is None:
|
||||
return []
|
||||
|
||||
samples = output_batch.output
|
||||
# [B, C, T, H, W] or [1, C, T, H, W]
|
||||
if len(samples.shape) == 5:
|
||||
# Rearrange to [T, B, C, H, W] for processing loop
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
else:
|
||||
logger.warning("Unexpected output shape: %s", samples.shape)
|
||||
return []
|
||||
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=1)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).cpu().numpy().astype(np.uint8))
|
||||
|
||||
return frames
|
||||
|
||||
def shutdown(self):
|
||||
if self.writer:
|
||||
self.writer.close()
|
||||
self.writer = None
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.disable_streaming()
|
||||
|
||||
super().shutdown()
|
||||
@@ -18,6 +18,8 @@ import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -389,6 +391,11 @@ class VideoGenerator:
|
||||
if batch.save_video:
|
||||
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
audio = output_batch.extra.get("audio")
|
||||
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
|
||||
if (audio is not None and audio_sample_rate is not None and
|
||||
not self._mux_audio(output_path, audio, audio_sample_rate)):
|
||||
logger.warning("Audio mux failed; saved video without audio.")
|
||||
|
||||
if batch.return_frames:
|
||||
return frames
|
||||
@@ -396,6 +403,7 @@ class VideoGenerator:
|
||||
return {
|
||||
"samples": samples,
|
||||
"frames": frames,
|
||||
"audio": output_batch.extra.get("audio"),
|
||||
"prompts": prompt,
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
@@ -405,6 +413,98 @@ class VideoGenerator:
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mux_audio(
|
||||
video_path: str,
|
||||
audio: torch.Tensor | np.ndarray,
|
||||
sample_rate: int,
|
||||
) -> bool:
|
||||
"""Mux audio into video using PyAV."""
|
||||
try:
|
||||
import av
|
||||
except ImportError:
|
||||
logger.warning("PyAV not installed; cannot mux audio. "
|
||||
"Install with: pip install av")
|
||||
return False
|
||||
|
||||
if torch.is_tensor(audio):
|
||||
audio_np = audio.detach().cpu().float().numpy()
|
||||
else:
|
||||
audio_np = np.asarray(audio, dtype=np.float32)
|
||||
|
||||
if audio_np.ndim == 1:
|
||||
audio_np = audio_np[:, None]
|
||||
elif audio_np.ndim == 2:
|
||||
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
|
||||
audio_np = audio_np.T
|
||||
else:
|
||||
logger.warning("Unexpected audio shape %s; skipping mux.",
|
||||
audio_np.shape)
|
||||
return False
|
||||
|
||||
audio_np = np.clip(audio_np, -1.0, 1.0)
|
||||
audio_int16 = (audio_np * 32767.0).astype(np.int16)
|
||||
num_channels = audio_int16.shape[1]
|
||||
layout = "stereo" if num_channels == 2 else "mono"
|
||||
|
||||
try:
|
||||
import wave
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
out_path = os.path.join(tmpdir, "muxed.mp4")
|
||||
wav_path = os.path.join(tmpdir, "audio.wav")
|
||||
|
||||
# Write audio to WAV file
|
||||
with wave.open(wav_path, "wb") as wav_file:
|
||||
wav_file.setnchannels(num_channels)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(sample_rate)
|
||||
wav_file.writeframes(audio_int16.tobytes())
|
||||
|
||||
# Open input video and audio
|
||||
input_video = av.open(video_path)
|
||||
input_audio = av.open(wav_path)
|
||||
|
||||
# Create output with both streams
|
||||
output = av.open(out_path, mode="w")
|
||||
|
||||
# Add video stream (copy codec from input)
|
||||
in_video_stream = input_video.streams.video[0]
|
||||
out_video_stream = output.add_stream(
|
||||
codec_name=in_video_stream.codec_context.name,
|
||||
rate=in_video_stream.average_rate,
|
||||
)
|
||||
out_video_stream.width = in_video_stream.width
|
||||
out_video_stream.height = in_video_stream.height
|
||||
out_video_stream.pix_fmt = in_video_stream.pix_fmt
|
||||
|
||||
# Add audio stream (AAC)
|
||||
out_audio_stream = output.add_stream("aac", rate=sample_rate)
|
||||
out_audio_stream.layout = layout
|
||||
|
||||
# Remux video (decode and re-encode to be safe)
|
||||
for frame in input_video.decode(video=0):
|
||||
for packet in out_video_stream.encode(frame):
|
||||
output.mux(packet)
|
||||
for packet in out_video_stream.encode():
|
||||
output.mux(packet)
|
||||
|
||||
# Encode audio
|
||||
for frame in input_audio.decode(audio=0):
|
||||
frame.pts = None # Let encoder assign PTS
|
||||
for packet in out_audio_stream.encode(frame):
|
||||
output.mux(packet)
|
||||
for packet in out_audio_stream.encode():
|
||||
output.mux(packet)
|
||||
|
||||
input_video.close()
|
||||
input_audio.close()
|
||||
output.close()
|
||||
shutil.move(out_path, video_path)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning("Audio mux failed: %s", e)
|
||||
return False
|
||||
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None) -> None:
|
||||
|
||||
+123
-2
@@ -12,6 +12,7 @@ from typing import Any, TYPE_CHECKING
|
||||
from fastvideo.configs.configs import PreprocessConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.configs.utils import clean_cli_args
|
||||
from fastvideo.layers.quantization import QUANTIZATION_METHODS, QuantizationMethods
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
@@ -131,7 +132,8 @@ class FastVideoArgs:
|
||||
|
||||
# CPU offload parameters
|
||||
dit_cpu_offload: bool = True
|
||||
use_fsdp_inference: bool = True
|
||||
use_fsdp_inference: bool = False
|
||||
dit_layerwise_offload: bool = True
|
||||
text_encoder_cpu_offload: bool = True
|
||||
image_encoder_cpu_offload: bool = True
|
||||
vae_cpu_offload: bool = True
|
||||
@@ -164,12 +166,24 @@ class FastVideoArgs:
|
||||
# Prompt text file for batch processing
|
||||
prompt_txt: str | None = None
|
||||
|
||||
# LTX-2 VAE tiling overrides
|
||||
ltx2_vae_tiling: bool | None = None
|
||||
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
|
||||
ltx2_vae_temporal_tile_size_in_frames: int | None = None
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
|
||||
ltx2_initial_latent_path: str | None = None
|
||||
|
||||
# model paths for correct deallocation
|
||||
model_paths: dict[str, str] = field(default_factory=dict)
|
||||
model_loaded: dict[str, bool] = field(default_factory=lambda: {
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
|
||||
override_text_encoder_safetensors: str | None = None # path to safetensors file for text encoder override
|
||||
override_text_encoder_quant: QuantizationMethods = None
|
||||
|
||||
override_transformer_cls_name: str | None = None
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
|
||||
@@ -197,8 +211,44 @@ class FastVideoArgs:
|
||||
logger.error("Failed to load V-MoBA config from %s: %s",
|
||||
self.moba_config_path, e)
|
||||
raise
|
||||
self._apply_ltx2_vae_overrides()
|
||||
self.check_fastvideo_args()
|
||||
|
||||
def _apply_ltx2_vae_overrides(self) -> None:
|
||||
if self.pipeline_config is None:
|
||||
return
|
||||
vae_config = self.pipeline_config.vae_config
|
||||
has_any = any(value is not None for value in (
|
||||
self.ltx2_vae_spatial_tile_size_in_pixels,
|
||||
self.ltx2_vae_spatial_tile_overlap_in_pixels,
|
||||
self.ltx2_vae_temporal_tile_size_in_frames,
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames,
|
||||
))
|
||||
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
|
||||
"vae_tiling"):
|
||||
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
|
||||
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
|
||||
self.pipeline_config.vae_tiling = True
|
||||
|
||||
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
|
||||
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
|
||||
vae_config.ltx2_spatial_tile_size_in_pixels = (
|
||||
self.ltx2_vae_spatial_tile_size_in_pixels)
|
||||
if hasattr(
|
||||
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
|
||||
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
|
||||
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
|
||||
self.ltx2_vae_spatial_tile_overlap_in_pixels)
|
||||
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
|
||||
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
|
||||
vae_config.ltx2_temporal_tile_size_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_size_in_frames)
|
||||
if hasattr(
|
||||
vae_config, "ltx2_temporal_tile_overlap_in_frames"
|
||||
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
|
||||
vae_config.ltx2_temporal_tile_overlap_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames)
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
# Model and path configuration
|
||||
@@ -319,6 +369,44 @@ class FastVideoArgs:
|
||||
"Path to a text file containing prompts (one per line) for batch processing",
|
||||
)
|
||||
|
||||
# LTX-2 VAE tiling overrides
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-tiling",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.ltx2_vae_tiling,
|
||||
help="Enable LTX-2 VAE tiling overrides.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-spatial-tile-size-in-pixels",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
|
||||
help="LTX-2 VAE spatial tile size in pixels.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-spatial-tile-overlap-in-pixels",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
|
||||
help="LTX-2 VAE spatial tile overlap in pixels.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-temporal-tile-size-in-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
|
||||
help="LTX-2 VAE temporal tile size in frames.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-temporal-tile-overlap-in-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
|
||||
help="LTX-2 VAE temporal tile overlap in frames.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-initial-latent-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_initial_latent_path,
|
||||
help="Path to load/save a precomputed LTX-2 initial latent.",
|
||||
)
|
||||
|
||||
# LoRA parameters (inference-time adapter loading)
|
||||
parser.add_argument(
|
||||
"--lora-path",
|
||||
@@ -416,11 +504,18 @@ class FastVideoArgs:
|
||||
help=
|
||||
"Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-layerwise-offload",
|
||||
action=StoreBoolean,
|
||||
help="Enable layerwise CPU offload with async H2D prefetch overlap.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--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",
|
||||
@@ -476,6 +571,19 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-text-encoder-safetensors",
|
||||
type=str,
|
||||
default=FastVideoArgs.override_text_encoder_safetensors,
|
||||
help="Path to safetensors file for text encoder override",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-text-encoder-quant",
|
||||
type=str,
|
||||
choices=QUANTIZATION_METHODS,
|
||||
default=FastVideoArgs.override_text_encoder_quant,
|
||||
help="Quantization method for text encoder override",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-transformer-cls-name",
|
||||
type=str,
|
||||
@@ -585,6 +693,19 @@ class FastVideoArgs:
|
||||
|
||||
if current_platform.is_mps():
|
||||
self.use_fsdp_inference = False
|
||||
self.dit_layerwise_offload = False
|
||||
|
||||
if self.dit_layerwise_offload:
|
||||
if self.use_fsdp_inference:
|
||||
logger.warning(
|
||||
"dit_layerwise_offload is enabled, automatically disabling use_fsdp_inference."
|
||||
)
|
||||
self.use_fsdp_inference = False
|
||||
if self.dit_cpu_offload:
|
||||
logger.warning(
|
||||
"dit_layerwise_offload is enabled, automatically disabling dit_cpu_offload."
|
||||
)
|
||||
self.dit_cpu_offload = False
|
||||
|
||||
# Validate mode and inference_mode consistency
|
||||
assert isinstance(
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import functools
|
||||
from typing import Any
|
||||
from torch import nn
|
||||
|
||||
|
||||
class ForwardHook:
|
||||
"""
|
||||
Base class for forward hooks.
|
||||
Hooks are used in the way:
|
||||
modified_args, modified_kwargs = hook.pre_forward(module, *args, **kwargs)
|
||||
output = module.forward(*modified_args, **modified_kwargs)
|
||||
modified_output = hook.post_forward(module, output)
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def name(cls) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
def on_attach(self, module: nn.Module): # noqa: B027
|
||||
"""Called once when the hook is attached to the module."""
|
||||
pass
|
||||
|
||||
def on_detach(self, module: nn.Module): # noqa: B027
|
||||
"""
|
||||
Called once when the hook is detached from the module.
|
||||
Note: this function is not guaranteed to be called if the module is
|
||||
deleted before the hook is detached.
|
||||
"""
|
||||
pass
|
||||
|
||||
def pre_forward(self, module: nn.Module, *args,
|
||||
**kwargs) -> tuple[tuple[Any, ...], dict[str, Any]]:
|
||||
"""Called before the module's forward method is executed."""
|
||||
return args, kwargs
|
||||
|
||||
def post_forward(self, module: nn.Module, output: Any) -> Any:
|
||||
"""Called after the module's forward method is executed."""
|
||||
return output
|
||||
|
||||
|
||||
class ModuleHookManager:
|
||||
module_hook_attribute = "_hook_manager"
|
||||
|
||||
def __init__(self, module: nn.Module):
|
||||
self.module = module
|
||||
self.forward_hooks: dict[str, ForwardHook] = {}
|
||||
self.original_forward = module.forward
|
||||
|
||||
@classmethod
|
||||
def get_from(cls, module: nn.Module) -> "ModuleHookManager | None":
|
||||
if hasattr(module, cls.module_hook_attribute):
|
||||
return getattr(module, cls.module_hook_attribute)
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_from_or_default(cls, module: nn.Module) -> "ModuleHookManager":
|
||||
if not hasattr(module, cls.module_hook_attribute):
|
||||
setattr(module, cls.module_hook_attribute, cls(module))
|
||||
|
||||
def forward_hook_wrapper(mod: nn.Module, *args, **kwargs):
|
||||
manager: ModuleHookManager = getattr(mod,
|
||||
cls.module_hook_attribute)
|
||||
for hook in manager.forward_hooks.values():
|
||||
args, kwargs = hook.pre_forward(mod, *args, **kwargs)
|
||||
output = manager.original_forward(*args, **kwargs)
|
||||
for hook in reversed(manager.forward_hooks.values()):
|
||||
output = hook.post_forward(mod, output)
|
||||
return output
|
||||
|
||||
module.forward = functools.partial(forward_hook_wrapper, module)
|
||||
|
||||
return getattr(module, cls.module_hook_attribute)
|
||||
|
||||
@staticmethod
|
||||
def remove_from_manager(module: nn.Module) -> None:
|
||||
if hasattr(module, ModuleHookManager.module_hook_attribute):
|
||||
manager: ModuleHookManager = getattr(
|
||||
module, ModuleHookManager.module_hook_attribute)
|
||||
module.forward = manager.original_forward
|
||||
delattr(module, ModuleHookManager.module_hook_attribute)
|
||||
|
||||
def _check_manager_attached(self) -> None:
|
||||
if not hasattr(self.module, self.module_hook_attribute):
|
||||
raise ValueError("ModuleHookManager is not attached to the module.")
|
||||
if getattr(self.module, self.module_hook_attribute) is not self:
|
||||
raise ValueError(
|
||||
"ModuleHookManager attached to the module is different.")
|
||||
|
||||
def append_forward_hook(self, hook: ForwardHook):
|
||||
self._check_manager_attached()
|
||||
if hook.name() in self.forward_hooks:
|
||||
raise ValueError(
|
||||
f"Hook with name {hook.name()} is already registered.")
|
||||
# after python 3.7, dicts maintain insertion order
|
||||
self.forward_hooks[hook.name()] = hook
|
||||
hook.on_attach(self.module)
|
||||
|
||||
def replace_forward_hook(self,
|
||||
hook_name: str,
|
||||
new_hook: ForwardHook,
|
||||
run_on_attach: bool = True):
|
||||
self._check_manager_attached()
|
||||
if hook_name not in self.forward_hooks:
|
||||
raise ValueError(f"No hook with name {hook_name} found.")
|
||||
old_hook = self.forward_hooks[hook_name]
|
||||
if run_on_attach:
|
||||
old_hook.on_detach(self.module)
|
||||
self.forward_hooks[hook_name] = new_hook
|
||||
new_hook.on_attach(self.module)
|
||||
|
||||
def remove_forward_hook(self, hook_name: str, run_detach: bool = True):
|
||||
self._check_manager_attached()
|
||||
if hook_name not in self.forward_hooks:
|
||||
raise ValueError(f"No hook with name {hook_name} found.")
|
||||
if run_detach:
|
||||
self.forward_hooks[hook_name].on_detach(self.module)
|
||||
del self.forward_hooks[hook_name]
|
||||
|
||||
def get_forward_hook(self, hook_name: str) -> ForwardHook | None:
|
||||
return self.forward_hooks.get(hook_name, None)
|
||||
@@ -0,0 +1,164 @@
|
||||
from contextlib import contextmanager
|
||||
from typing import Any
|
||||
import torch
|
||||
from torch import nn
|
||||
from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _tensor_placeholder(tensor: torch.Tensor,
|
||||
device: torch.device) -> torch.Tensor:
|
||||
"""Create a rank-preserving empty placeholder on the specified device."""
|
||||
shape = (0, ) if tensor.ndim <= 0 else (0, ) * tensor.ndim
|
||||
return torch.empty(shape, device=device, dtype=tensor.dtype)
|
||||
|
||||
|
||||
class LayerwiseOffloadState:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
async_copy_stream: torch.cuda.Stream,
|
||||
device: torch.device,
|
||||
next_state: "LayerwiseOffloadState | None" = None,
|
||||
) -> None:
|
||||
self.async_copy_stream = async_copy_stream
|
||||
self.next_state = next_state
|
||||
self.gpu_named_parameters: dict[str, torch.Tensor] = {}
|
||||
self.cpu_named_parameters: dict[str, torch.Tensor] = {}
|
||||
self.module_ref: nn.Module = None # type: ignore
|
||||
self.device: torch.device = device
|
||||
|
||||
def _will_offload(self, name: str) -> bool:
|
||||
return True
|
||||
|
||||
@torch.compiler.disable
|
||||
def on_init(self, module: nn.Module):
|
||||
self.module_ref = module
|
||||
for name, param in self.module_ref.named_parameters():
|
||||
if self._will_offload(name):
|
||||
self.cpu_named_parameters[name] = (
|
||||
param.data.detach().to("cpu").pin_memory())
|
||||
param.data = _tensor_placeholder(param.data, self.device)
|
||||
|
||||
@torch.compiler.disable
|
||||
def wait_and_replace_params(self):
|
||||
torch.cuda.current_stream().wait_stream(self.async_copy_stream)
|
||||
# now gpu_named_parameters are ready
|
||||
for name, param in self.module_ref.named_parameters():
|
||||
if not self._will_offload(name):
|
||||
continue
|
||||
if name not in self.gpu_named_parameters:
|
||||
# first load with blocking load
|
||||
self.gpu_named_parameters[name] = self.cpu_named_parameters[
|
||||
name].to(self.device)
|
||||
param.data = self.gpu_named_parameters[name]
|
||||
|
||||
@torch.compiler.disable
|
||||
def prefetch_params(self):
|
||||
compute_stream = torch.cuda.current_stream()
|
||||
with torch.cuda.stream(self.async_copy_stream):
|
||||
for name, param in self.module_ref.named_parameters():
|
||||
if not self._will_offload(name):
|
||||
continue
|
||||
assert name not in self.gpu_named_parameters
|
||||
gpu_param = self.cpu_named_parameters[name].to(
|
||||
self.device, non_blocking=True)
|
||||
gpu_param.record_stream(
|
||||
compute_stream
|
||||
) # ensure tensor will not be freed until forward is completed
|
||||
self.gpu_named_parameters[name] = gpu_param
|
||||
|
||||
@torch.compiler.disable
|
||||
def release_gpu_params(self):
|
||||
for name, param in self.module_ref.named_parameters():
|
||||
if self._will_offload(name):
|
||||
param.data = _tensor_placeholder(param.data, self.device)
|
||||
del self.gpu_named_parameters[name]
|
||||
assert len(self.gpu_named_parameters) == 0
|
||||
|
||||
|
||||
class LayerwiseOffloadHook(ForwardHook):
|
||||
"""A hook that enables layerwise CPU offloading during forward pass."""
|
||||
|
||||
def __init__(self, state: LayerwiseOffloadState) -> None:
|
||||
self.state = state
|
||||
|
||||
def on_attach(self, module: nn.Module):
|
||||
self.state.on_init(module) # pyright: ignore
|
||||
|
||||
def on_detach(self, module: nn.Module):
|
||||
named_parameters = dict(module.named_parameters())
|
||||
for name, cpu_tensor in self.state.cpu_named_parameters.items():
|
||||
if name not in self.state.gpu_named_parameters:
|
||||
if name in named_parameters:
|
||||
named_parameters[name].data = cpu_tensor.to(
|
||||
device=self.state.device)
|
||||
else:
|
||||
logger.warning(
|
||||
"Parameter {} not found in module during detachment.",
|
||||
name,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def name(cls) -> str:
|
||||
return "LayerwiseOffloadHook"
|
||||
|
||||
def pre_forward(self, module: nn.Module, *args, **kwargs):
|
||||
self.state.wait_and_replace_params() # pyright: ignore
|
||||
if self.state.next_state is not None:
|
||||
self.state.next_state.prefetch_params() # pyright: ignore
|
||||
return args, kwargs
|
||||
|
||||
def post_forward(self, module: torch.nn.Module, output: Any):
|
||||
self.state.release_gpu_params() # pyright: ignore
|
||||
return output
|
||||
|
||||
@contextmanager
|
||||
def mutate_params_scope(self):
|
||||
try:
|
||||
# load params to GPU and keep them there
|
||||
self.state.wait_and_replace_params() # pyright: ignore
|
||||
yield
|
||||
finally:
|
||||
# instead of releasing, we should overwrite the original params since they have been modified
|
||||
self.state.cpu_named_parameters.clear()
|
||||
self.state.gpu_named_parameters.clear()
|
||||
self.state.on_init(self.state.module_ref) # pyright: ignore
|
||||
|
||||
|
||||
def enable_layerwise_offload(model: nn.Module, is_replace: bool = False):
|
||||
if torch.cuda.is_available():
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
else:
|
||||
logger.warning(
|
||||
"CUDA is not available. Layerwise offloading is disabled.")
|
||||
return
|
||||
state_list = []
|
||||
async_stream = torch.cuda.Stream()
|
||||
for name, submodule in model.named_children():
|
||||
if isinstance(submodule, nn.ModuleList):
|
||||
for idx, module_entry in enumerate(submodule):
|
||||
state = LayerwiseOffloadState(async_copy_stream=async_stream,
|
||||
device=device)
|
||||
state_list.append(state)
|
||||
hook_mgr = ModuleHookManager.get_from_or_default(module_entry)
|
||||
hook = LayerwiseOffloadHook(state)
|
||||
if is_replace:
|
||||
existing_hook = hook_mgr.forward_hooks.get(hook.name())
|
||||
if existing_hook is not None:
|
||||
hook_mgr.replace_forward_hook(hook.name(), hook)
|
||||
else:
|
||||
raise AssertionError(
|
||||
f"Expect hook exists in {name} for replacement.")
|
||||
else:
|
||||
hook_mgr.append_forward_hook(hook)
|
||||
break
|
||||
if len(state_list) == 0:
|
||||
raise ValueError(
|
||||
"No nn.ModuleList found in the model for layerwise offloading.")
|
||||
|
||||
# circular linking of states
|
||||
for i in range(len(state_list)):
|
||||
state_list[i].next_state = state_list[(i + 1) % len(state_list)]
|
||||
+286
-190
@@ -7,13 +7,20 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.distributed import (divide, get_tp_rank, get_tp_world_size,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from fastvideo.layers.quantization.base_config import (QuantizationConfig,
|
||||
QuantizeMethodBase)
|
||||
from fastvideo.distributed import (
|
||||
divide,
|
||||
get_tp_rank,
|
||||
get_tp_world_size,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
# yapf: disable
|
||||
from fastvideo.models.parameter import (BasevLLMParameter,
|
||||
BlockQuantScaleParameter,
|
||||
@@ -27,12 +34,22 @@ from fastvideo.models.utils import set_weight_attrs
|
||||
logger = init_logger(__name__)
|
||||
|
||||
WEIGHT_LOADER_V2_SUPPORTED = [
|
||||
"CompressedTensorsLinearMethod", "AWQMarlinLinearMethod", "AWQLinearMethod",
|
||||
"GPTQMarlinLinearMethod", "Fp8LinearMethod", "MarlinLinearMethod",
|
||||
"QQQLinearMethod", "GPTQMarlin24LinearMethod", "TPUInt8LinearMethod",
|
||||
"GPTQLinearMethod", "FBGEMMFp8LinearMethod", "ModelOptFp8LinearMethod",
|
||||
"IPEXAWQLinearMethod", "IPEXGPTQLinearMethod", "HQQMarlinMethod",
|
||||
"QuarkLinearMethod"
|
||||
"CompressedTensorsLinearMethod",
|
||||
"AWQMarlinLinearMethod",
|
||||
"AWQLinearMethod",
|
||||
"GPTQMarlinLinearMethod",
|
||||
"Fp8LinearMethod",
|
||||
"MarlinLinearMethod",
|
||||
"QQQLinearMethod",
|
||||
"GPTQMarlin24LinearMethod",
|
||||
"TPUInt8LinearMethod",
|
||||
"GPTQLinearMethod",
|
||||
"FBGEMMFp8LinearMethod",
|
||||
"ModelOptFp8LinearMethod",
|
||||
"IPEXAWQLinearMethod",
|
||||
"IPEXGPTQLinearMethod",
|
||||
"HQQMarlinMethod",
|
||||
"QuarkLinearMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -41,8 +58,8 @@ def adjust_scalar_to_fused_array(
|
||||
shard_id: str | int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""For fused modules (QKV and MLP) we have an array of length
|
||||
N that holds 1 scale for each "logical" matrix. So the param
|
||||
is an array of length N. The loaded_weight corresponds to
|
||||
one of the shards on disk. Here, we slice the param based on
|
||||
is an array of length N. The loaded_weight corresponds to
|
||||
one of the shards on disk. Here, we slice the param based on
|
||||
the shard_id for loading.
|
||||
"""
|
||||
qkv_idxs = {"q": 0, "k": 1, "v": 2}
|
||||
@@ -65,18 +82,23 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
"""Base class for different (maybe quantized) linear methods."""
|
||||
|
||||
@abstractmethod
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
"""Create weights for a linear layer.
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
"""Create weights for a linear layer.
|
||||
The weights will be set as attributes of the layer.
|
||||
|
||||
Args:
|
||||
layer: The layer that is using the LinearMethodBase factory.
|
||||
input_size_per_partition: Size of the weight input dim on rank X.
|
||||
output_partition_sizes: Sizes of the output dim of each logical
|
||||
output_partition_sizes: Sizes of the output dim of each logical
|
||||
weight on rank X. E.g., output_partition_sizes for QKVLinear
|
||||
is a list contains the width of Wq, Wk, Wv on rank X.
|
||||
input_size: Size of the input dim of the weight across all ranks.
|
||||
@@ -86,10 +108,12 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply the weights in layer to the input tensor.
|
||||
Expects create_weights to have been called before on the layer."""
|
||||
raise NotImplementedError
|
||||
@@ -98,28 +122,37 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
class UnquantizedLinearMethod(LinearMethodBase):
|
||||
"""Linear method without quantization."""
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False)
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
weight = Parameter(
|
||||
torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
output = F.linear(x, layer.weight, bias) if torch.cuda.is_available(
|
||||
) or bias is None else F.linear(
|
||||
x, layer.weight, bias.to(x.dtype)
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
output = (
|
||||
F.linear(x, layer.weight, bias) if torch.cuda.is_available()
|
||||
or bias is None else F.linear(x, layer.weight, bias.to(x.dtype))
|
||||
) # NOTE: this line assumes that we are using amp when using cuda and is needed to account for the fact that amp isn't supported in mps
|
||||
return output
|
||||
|
||||
@@ -157,8 +190,8 @@ class LinearBase(torch.nn.Module):
|
||||
self.quant_config = quant_config
|
||||
self.prefix = prefix
|
||||
if quant_config is None:
|
||||
self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
|
||||
)
|
||||
self.quant_method: QuantizeMethodBase | None = (
|
||||
UnquantizedLinearMethod())
|
||||
else:
|
||||
self.quant_method = quant_config.get_quant_method(self,
|
||||
prefix=prefix)
|
||||
@@ -181,29 +214,36 @@ class ReplicatedLinear(LinearBase):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
super().__init__(input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix=prefix)
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__(
|
||||
input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
# All the linear layer supports quant method.
|
||||
assert self.quant_method is not None
|
||||
self.quant_method.create_weights(self,
|
||||
self.input_size, [self.output_size],
|
||||
self.input_size,
|
||||
self.output_size,
|
||||
self.params_dtype,
|
||||
weight_loader=self.weight_loader)
|
||||
self.quant_method.create_weights(
|
||||
self,
|
||||
self.input_size,
|
||||
[self.output_size],
|
||||
self.input_size,
|
||||
self.output_size,
|
||||
self.params_dtype,
|
||||
weight_loader=self.weight_loader,
|
||||
)
|
||||
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
@@ -211,10 +251,13 @@ class ReplicatedLinear(LinearBase):
|
||||
self.output_size,
|
||||
dtype=self.params_dtype,
|
||||
))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
set_weight_attrs(
|
||||
self.bias,
|
||||
{
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
},
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
@@ -265,19 +308,21 @@ class ColumnParallelLinear(LinearBase):
|
||||
output_sizes: list of output sizes packed into one output, like for QKV
|
||||
the list would be size 3.
|
||||
prefix: The name of the layer in the state dict, including all parents
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
output_sizes: list[int] | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
output_sizes: list[int] | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
# Divide the weight matrix along the last dimension.
|
||||
self.tp_size = get_tp_world_size()
|
||||
self.input_size_per_partition = input_size
|
||||
@@ -290,8 +335,14 @@ class ColumnParallelLinear(LinearBase):
|
||||
for output_size in self.output_sizes
|
||||
]
|
||||
|
||||
super().__init__(input_size, output_size, skip_bias_add, params_dtype,
|
||||
quant_config, prefix)
|
||||
super().__init__(
|
||||
input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix,
|
||||
)
|
||||
|
||||
self.gather_output = gather_output
|
||||
|
||||
@@ -308,17 +359,21 @@ class ColumnParallelLinear(LinearBase):
|
||||
params_dtype=self.params_dtype,
|
||||
weight_loader=(
|
||||
self.weight_loader_v2 if self.quant_method.__class__.__name__
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader))
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader),
|
||||
)
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(
|
||||
self.output_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
set_weight_attrs(
|
||||
self.bias,
|
||||
{
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
},
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
@@ -401,32 +456,37 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_sizes: list[int],
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_sizes: list[int],
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
self.output_sizes = output_sizes
|
||||
tp_size = get_tp_world_size()
|
||||
assert all(output_size % tp_size == 0 for output_size in output_sizes)
|
||||
super().__init__(input_size=input_size,
|
||||
output_size=sum(output_sizes),
|
||||
bias=bias,
|
||||
gather_output=gather_output,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None) -> None:
|
||||
super().__init__(
|
||||
input_size=input_size,
|
||||
output_size=sum(output_sizes),
|
||||
bias=bias,
|
||||
gather_output=gather_output,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None,
|
||||
) -> None:
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
# Special case for AQLM codebooks.
|
||||
@@ -518,20 +578,22 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
if (isinstance(param, PackedColumnParameter | PackedvLLMParameter)
|
||||
and param.packed_dim == param.output_dim):
|
||||
shard_size, shard_offset = (
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
shard_size=shard_size, shard_offset=shard_offset))
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(param.output_dim,
|
||||
shard_offset, shard_size)
|
||||
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
|
||||
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None) -> None:
|
||||
def weight_loader_v2(
|
||||
self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None,
|
||||
) -> None:
|
||||
if loaded_shard_id is None:
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight,
|
||||
@@ -568,10 +630,12 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
|
||||
shard_size = self.output_sizes[loaded_shard_id] // tp_size
|
||||
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size)
|
||||
param.load_merged_column_weight(
|
||||
loaded_weight=loaded_weight,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size,
|
||||
)
|
||||
|
||||
|
||||
class QKVParallelLinear(ColumnParallelLinear):
|
||||
@@ -600,16 +664,18 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
hidden_size: int,
|
||||
head_size: int,
|
||||
total_num_heads: int,
|
||||
total_num_kv_heads: int | None = None,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
head_size: int,
|
||||
total_num_heads: int,
|
||||
total_num_kv_heads: int | None = None,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
self.hidden_size = hidden_size
|
||||
self.head_size = head_size
|
||||
self.total_num_heads = total_num_heads
|
||||
@@ -626,29 +692,31 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
self.num_kv_heads = divide(self.total_num_kv_heads, tp_size)
|
||||
self.num_kv_head_replicas = 1
|
||||
input_size = self.hidden_size
|
||||
output_size = (self.num_heads +
|
||||
2 * self.num_kv_heads) * tp_size * self.head_size
|
||||
output_size = ((self.num_heads + 2 * self.num_kv_heads) * tp_size *
|
||||
self.head_size)
|
||||
self.output_sizes = [
|
||||
self.num_heads * self.head_size * tp_size, # q_proj
|
||||
self.num_kv_heads * self.head_size * tp_size, # k_proj
|
||||
self.num_kv_heads * self.head_size * tp_size, # v_proj
|
||||
self.num_kv_heads * self.head_size * tp_size, # v_proj
|
||||
]
|
||||
|
||||
super().__init__(input_size=input_size,
|
||||
output_size=output_size,
|
||||
bias=bias,
|
||||
gather_output=False,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
super().__init__(
|
||||
input_size=input_size,
|
||||
output_size=output_size,
|
||||
bias=bias,
|
||||
gather_output=False,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
|
||||
shard_offset_mapping = {
|
||||
"q": 0,
|
||||
"k": self.num_heads * self.head_size,
|
||||
"v": (self.num_heads + self.num_kv_heads) * self.head_size,
|
||||
"total": (self.num_heads + 2 * self.num_kv_heads) * self.head_size
|
||||
"total": (self.num_heads + 2 * self.num_kv_heads) * self.head_size,
|
||||
}
|
||||
return shard_offset_mapping.get(loaded_shard_id)
|
||||
|
||||
@@ -663,7 +731,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
def _load_fused_module_from_checkpoint(self, param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor):
|
||||
"""
|
||||
Handle special case for models where QKV layers are already
|
||||
Handle special case for models where QKV layers are already
|
||||
fused on disk. In this case, we have no shard id. This function
|
||||
determmines the shard id by splitting these layers and then calls
|
||||
the weight loader using the shard id.
|
||||
@@ -674,31 +742,39 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offsets = [
|
||||
# (shard_id, shard_offset, shard_size)
|
||||
("q", 0, self.total_num_heads * self.head_size),
|
||||
("k", self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
("v",
|
||||
(self.total_num_heads + self.total_num_kv_heads) * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
(
|
||||
"k",
|
||||
self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
(
|
||||
"v",
|
||||
(self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
]
|
||||
|
||||
for shard_id, shard_offset, shard_size in shard_offsets:
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
if (isinstance(param, PackedColumnParameter | PackedvLLMParameter)
|
||||
and param.packed_dim == param.output_dim):
|
||||
shard_size, shard_offset = (
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
shard_size=shard_size, shard_offset=shard_offset))
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(param.output_dim,
|
||||
shard_offset, shard_size)
|
||||
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
|
||||
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None):
|
||||
def weight_loader_v2(
|
||||
self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None,
|
||||
):
|
||||
if loaded_shard_id is None: # special case for certain models
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0)
|
||||
@@ -715,17 +791,20 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offset = self._get_shard_offset_mapping(loaded_shard_id)
|
||||
shard_size = self._get_shard_size_mapping(loaded_shard_id)
|
||||
|
||||
param.load_qkv_weight(loaded_weight=loaded_weight,
|
||||
num_heads=self.num_kv_head_replicas,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size)
|
||||
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None):
|
||||
param.load_qkv_weight(
|
||||
loaded_weight=loaded_weight,
|
||||
num_heads=self.num_kv_head_replicas,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size,
|
||||
)
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None,
|
||||
):
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
# Special case for AQLM codebooks.
|
||||
@@ -748,14 +827,20 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offsets = [
|
||||
# (shard_id, shard_offset, shard_size)
|
||||
("q", 0, self.total_num_heads * self.head_size),
|
||||
("k", self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
("v", (self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size, self.total_num_kv_heads * self.head_size),
|
||||
(
|
||||
"k",
|
||||
self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
(
|
||||
"v",
|
||||
(self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
]
|
||||
|
||||
for shard_id, shard_offset, shard_size in shard_offsets:
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(
|
||||
output_dim, shard_offset, shard_size)
|
||||
self.weight_loader(param, loaded_weight_shard, shard_id)
|
||||
@@ -843,16 +928,18 @@ class RowParallelLinear(LinearBase):
|
||||
quant_config: Quantization configure.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
input_is_parallel: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
reduce_results: bool = True,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
input_is_parallel: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
reduce_results: bool = True,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
# Divide the weight matrix along the first dimension.
|
||||
self.tp_rank = get_tp_rank()
|
||||
self.tp_size = get_tp_world_size()
|
||||
@@ -860,8 +947,14 @@ class RowParallelLinear(LinearBase):
|
||||
self.output_size_per_partition = output_size
|
||||
self.output_partition_sizes = [output_size]
|
||||
|
||||
super().__init__(input_size, output_size, skip_bias_add, params_dtype,
|
||||
quant_config, prefix)
|
||||
super().__init__(
|
||||
input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix,
|
||||
)
|
||||
|
||||
self.input_is_parallel = input_is_parallel
|
||||
self.reduce_results = reduce_results
|
||||
@@ -876,7 +969,8 @@ class RowParallelLinear(LinearBase):
|
||||
params_dtype=self.params_dtype,
|
||||
weight_loader=(
|
||||
self.weight_loader_v2 if self.quant_method.__class__.__name__
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader))
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader),
|
||||
)
|
||||
if not reduce_results and (bias and not skip_bias_add):
|
||||
raise ValueError("When not reduce the results, adding bias to the "
|
||||
"results can lead to incorrect results")
|
||||
@@ -884,10 +978,13 @@ class RowParallelLinear(LinearBase):
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(self.output_size, dtype=params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
set_weight_attrs(
|
||||
self.bias,
|
||||
{
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
},
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
@@ -916,7 +1013,6 @@ class RowParallelLinear(LinearBase):
|
||||
|
||||
def weight_loader_v2(self, param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor):
|
||||
|
||||
# Special case for loading scales off disk, which often do not
|
||||
# have a shape (such as in the case of AutoFP8).
|
||||
if len(loaded_weight.shape) == 0:
|
||||
|
||||
@@ -2,7 +2,7 @@ from typing import Literal, get_args
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig
|
||||
|
||||
QuantizationMethods = Literal[None]
|
||||
QuantizationMethods = Literal[None, "AbsMaxFP8"]
|
||||
|
||||
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
|
||||
|
||||
@@ -50,7 +50,12 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]:
|
||||
if quantization not in QUANTIZATION_METHODS:
|
||||
raise ValueError(f"Invalid quantization method: {quantization}")
|
||||
|
||||
method_to_config: dict[str, type[QuantizationConfig]] = {}
|
||||
# lazy import to avoid triggering `torch.compile` too early
|
||||
from .absmax_fp8 import AbsMaxFP8Config
|
||||
|
||||
method_to_config: dict[str, type[QuantizationConfig]] = {
|
||||
"AbsMaxFP8": AbsMaxFP8Config,
|
||||
}
|
||||
# Update the `method_to_config` with customized quantization methods.
|
||||
method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG)
|
||||
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
from typing import Any
|
||||
import torch
|
||||
from fastvideo.distributed.parallel_state import get_tp_world_size
|
||||
from fastvideo.layers.linear import (
|
||||
LinearBase,
|
||||
LinearMethodBase,
|
||||
MergedColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
)
|
||||
from fastvideo.layers.quantization import QuantizationMethods
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class AbsMaxFP8Config(QuantizationConfig):
|
||||
"""
|
||||
Config class for absmax float8_e4m3fn quantization.
|
||||
Currently only support per-tensor quantization.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "QuantizationConfig":
|
||||
return cls()
|
||||
|
||||
def get_name(self) -> QuantizationMethods:
|
||||
return "AbsMaxFP8"
|
||||
|
||||
def get_supported_act_dtypes(self) -> list[torch.dtype]:
|
||||
return [torch.bfloat16, torch.float16, torch.float32]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 75
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module,
|
||||
prefix: str) -> QuantizeMethodBase | None:
|
||||
if isinstance(layer, LinearBase):
|
||||
return AbsMaxFP8LinearMethod()
|
||||
return None
|
||||
|
||||
|
||||
class AbsMaxFP8Parameter(nn.Parameter):
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: nn.Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
_share_id: str | None = None,
|
||||
) -> None:
|
||||
if len(loaded_weight.shape) == 0:
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
|
||||
assert param.size() == loaded_weight.size(), (
|
||||
f"Tried to load weights of size {loaded_weight.size()}"
|
||||
f"to a parameter of size {param.size()}")
|
||||
param.data.copy_(loaded_weight)
|
||||
|
||||
|
||||
class AbsMaxFP8MergedParameter(nn.Parameter):
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: nn.Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
share_id: str | int | None = None,
|
||||
) -> None:
|
||||
# currently only support QKVParallelLinear and MergedColumnParallelLinear
|
||||
output_partition_sizes: list[int] = self.output_partition_sizes
|
||||
if share_id is None:
|
||||
share_id = 0
|
||||
if isinstance(share_id, str) and share_id in ["q", "k", "v"]:
|
||||
# QKVParallelLinear case
|
||||
share_idx = ["q", "k", "v"].index(share_id)
|
||||
start_idx = sum(output_partition_sizes[:share_idx])
|
||||
end_idx = start_idx + output_partition_sizes[share_idx]
|
||||
elif isinstance(share_id, int):
|
||||
# MergedColumnParallelLinear case
|
||||
tp_size = get_tp_world_size()
|
||||
if tp_size > 1:
|
||||
# TODO: support this case
|
||||
raise NotImplementedError(
|
||||
"AbsMaxFP8MergedParameter with integer share_id is not supported in tensor parallelism greater than 1 yet."
|
||||
)
|
||||
start_idx = sum(output_partition_sizes[:share_id])
|
||||
end_idx = start_idx + output_partition_sizes[share_id]
|
||||
else:
|
||||
raise ValueError(
|
||||
f"AbsMaxFP8MergedParameter requires share_id to be ['q', 'k', 'v'] or int, got {share_id}."
|
||||
)
|
||||
if len(loaded_weight.shape) == 0:
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
assert loaded_weight.numel() == 1
|
||||
# fill in the corresponding partition by repeating the val
|
||||
param.data[start_idx:end_idx].fill_(loaded_weight.item())
|
||||
|
||||
|
||||
class AbsMaxFP8LinearMethod(LinearMethodBase):
|
||||
"""Linear method with AbsMax FP8 quantization."""
|
||||
|
||||
@staticmethod
|
||||
def _convert_scale(scale: Any) -> torch.nn.Parameter:
|
||||
if scale is None:
|
||||
scale = torch.tensor([1.0], dtype=torch.float32)
|
||||
if not isinstance(scale, torch.Tensor):
|
||||
scale = torch.tensor([scale], dtype=torch.float32)
|
||||
if scale.dtype != torch.float32:
|
||||
raise NotImplementedError("Only float32 scale is supported")
|
||||
return AbsMaxFP8Parameter(scale, requires_grad=False)
|
||||
|
||||
@staticmethod
|
||||
def _merged_placeholder(
|
||||
output_partition_sizes: list[int], ) -> torch.nn.Parameter:
|
||||
scale = torch.ones(
|
||||
sum(output_partition_sizes),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
para = AbsMaxFP8MergedParameter(
|
||||
scale,
|
||||
False,
|
||||
)
|
||||
set_weight_attrs(
|
||||
para,
|
||||
{
|
||||
"output_partition_sizes": output_partition_sizes,
|
||||
},
|
||||
)
|
||||
return para
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
assert params_dtype in [
|
||||
torch.bfloat16, torch.float16, torch.float32
|
||||
], (f"AbsMaxFP8LinearMethod only supports bfloat16, float16, or float32 original dtype, got {params_dtype}."
|
||||
)
|
||||
weight = nn.Parameter(
|
||||
torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
if isinstance(layer, QKVParallelLinear | MergedColumnParallelLinear):
|
||||
scale_weight = self._merged_placeholder(output_partition_sizes, )
|
||||
else:
|
||||
scale_weight = self._convert_scale(
|
||||
extra_weight_attrs.get("scale_weight"))
|
||||
scale_input = self._convert_scale(extra_weight_attrs.get("scale_input"))
|
||||
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
layer.register_parameter("scale_weight", scale_weight)
|
||||
layer.register_parameter("scale_input", scale_input)
|
||||
set_weight_attrs(
|
||||
weight,
|
||||
{
|
||||
"output_dtype": params_dtype,
|
||||
},
|
||||
)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
weight_quant = layer.weight
|
||||
output_dtype: torch.dtype = weight_quant.output_dtype
|
||||
scale_weight: torch.Tensor = layer.scale_weight.data.to(output_dtype)
|
||||
scale_input: torch.Tensor = layer.scale_input.data.to(output_dtype)
|
||||
weight_output_type = weight_quant.to(dtype=output_dtype)
|
||||
weight_final = weight_output_type * scale_weight.unsqueeze(1)
|
||||
x_final = x.to(dtype=output_dtype) * scale_input
|
||||
|
||||
return nn.functional.linear(x_final, weight_final,
|
||||
bias=bias).to(dtype=output_dtype)
|
||||
@@ -0,0 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoder,
|
||||
LTX2AudioEncoder,
|
||||
LTX2Vocoder,
|
||||
)
|
||||
|
||||
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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:
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user