Compare commits

..
Author SHA1 Message Date
SolitaryThinker b30d98ca14 fp8 2025-12-31 00:19:52 +00:00
67 changed files with 1073 additions and 6008 deletions
+2 -2
View File
@@ -61,7 +61,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
command: "timeout 45m .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 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
+15 -21
View File
@@ -1,47 +1,41 @@
<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://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/c7g1qdD" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://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> |
</p>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
<div align="center">
<img src=assets/fastwan.png width="90%"/>
</div>
## 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 for bidirectional and autoregressive models:
- 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
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
- 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.
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
- State-of-the-art performance optimizations for inference
- 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.
- [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)
- 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:
+1 -8
View File
@@ -50,25 +50,21 @@ 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`](../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` (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
@@ -88,7 +84,6 @@ 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)**
@@ -116,7 +111,6 @@ endif()
```
### D. Expose in Python Ops
Update `fastvideo-kernel/python/fastvideo_kernel/ops.py` to make the function importable and handle fallbacks gracefully.
```python
@@ -139,7 +133,6 @@ 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
+14 -35
View File
@@ -4,26 +4,25 @@ This document outlines FastVideo's architecture for developers interested in fra
## Table of Contents - Directory Structure and Files
- [`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/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/utils.py` - Utility functions
- [`fastvideo/logger.py`](#logger) - Logging infrastructure
- [`fastvideo/logger.py`](#design-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)
@@ -35,14 +34,12 @@ 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`
@@ -93,9 +90,7 @@ 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
@@ -138,7 +133,6 @@ Transformer networks perform the actual denoising during diffusion:
- `HunyuanVideoTransformer3DModel`
Features include:
- Text/image conditioning
- Standardized interface for model-specific optimizations
@@ -167,7 +161,6 @@ 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
@@ -186,7 +179,6 @@ Encoders process conditioning inputs into embeddings:
- `CLIPVisionModel`
FastVideo implements optimizations such as:
- Vocab parallelism for distributed processing
- Caching for common prompts
- Precision-tuned computation
@@ -201,7 +193,6 @@ Schedulers manage the diffusion sampling process:
- `FlowMatchEulerDiscreteScheduler`
These components control:
- Diffusion timestep sequences
- Noise prediction to latent update conversions
- Quality/speed trade-offs
@@ -228,9 +219,7 @@ 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
@@ -251,9 +240,7 @@ self.attn = LocalAttention(
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
- Chunking, progressive computation, optimized masking
@@ -309,7 +296,6 @@ self.attn = DistributedAttention(
```
### Communication Primitives
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
Efficient communication primitives minimize distributed overhead:
@@ -328,7 +314,6 @@ 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
@@ -354,14 +339,12 @@ 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
@@ -376,13 +359,11 @@ 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
@@ -402,7 +383,6 @@ 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.
@@ -417,7 +397,6 @@ 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
+1 -1
View File
@@ -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](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
```bash
bash scripts/inference/v1_inference_wan_dmd.sh
+1 -1
View File
@@ -38,4 +38,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/examples_inference_index.md) - Explore example scripts and notebooks
- [Examples](../inference/examples/) - Explore example scripts and notebooks
+2 -2
View File
@@ -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](../../contributing/developer_env/docker.md)
[Docker Images](#docker)
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](../../contributing/overview.md)
[Contributor Guide](#developer-overview)
## Hardware Requirements
+1 -1
View File
@@ -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](../../contributing/overview.md)
[Contributor Guide](#developer-overview)
## Hardware Requirements
+1
View File
@@ -75,3 +75,4 @@ if __name__ == '__main__':
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/) - Explore more examples
- [Optimizations](../inference/optimizations.md) - Performance optimization tips
- [Low VRAM Inference](../inference/low_vram_inference.md) - Memory-saving settings (CPU offload, sharded loading, etc.)
-6
View File
@@ -45,7 +45,6 @@ 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/`
@@ -54,15 +53,12 @@ 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
@@ -95,7 +91,6 @@ self.out_proj = RowParallelLinear(
```
### Attention Layers
Replace standard attention with FastVideo's optimized attention:
```python
@@ -309,7 +304,6 @@ 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 -1
View File
@@ -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](examples/basic.md).
see the Python interface [here](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html).
## Basic Usage
+1 -1
View File
@@ -74,4 +74,4 @@ if __name__ == '__main__':
## Performance Optimization
For configuring optimizations, please see our [optimizations guide](optimizations.md)
For configuring optimizations, please see our [optimizations guide](#inference-optimizations)
+6 -15
View File
@@ -3,7 +3,6 @@
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
@@ -22,10 +21,9 @@ conda activate fastvideo
pip install fastvideo
```
For advanced installation options, see the [Installation Guide](../getting_started/installation.md).
For advanced installation options, see the [Installation Guide](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
@@ -62,10 +60,9 @@ 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.md) for the list of supported models and their available optimizations.
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
## Image-to-Video Generation
@@ -99,26 +96,20 @@ 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`
- Enable memory optimization with CPU-offload and sharded loading flags (see [Low VRAM Inference](low_vram_inference.md))
- Try a smaller model or use distilled versions
- Use `num_gpus` > 1 if multiple GPUs are available
### 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
@@ -126,8 +117,8 @@ If the generated video doesn't match your prompt:
## Next Steps
- Learn about [Advanced Inference Configurations](configuration.md)
- Learn about using [Optimizations](optimizations.md)
- See [Examples](examples/examples_inference_index.md) for more usage scenarios
- Learn about [Advanced Inference Configurations](#inference-configuration)
- Learn about using [Optimizations](#inference-optimizations)
- 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 -7
View File
@@ -7,13 +7,13 @@ This page describes the various options for speeding up generation times in Fast
- Optimized Attention Backends
- [Flash Attention](#flash-attention)
- [Sliding Tile Attention](#sliding-tile-attention)
- [Sage Attention](#sage-attention)
- [Sage Attention 3](#sage-attention-3)
- [Flash Attention](#optimizations-flash)
- [Sliding Tile Attention](#optimizations-sta)
- [Sage Attention](#optimizations-sage)
- [Sage Attention 3](#optimizations-sage3)
- Caching Techniques
- [TeaCache](#teacache)
- [TeaCache](#optimizations-teacache)
## Attention Backends
@@ -74,7 +74,7 @@ python setup.py install
pip install st_attn==0.0.4
```
Please see [this page](../attention/sta/index.md) for more installation instructions.
Please see [this page](#sta-installation) 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](../attention/vsa/index.md) for more installation instructions.
Please see [this page](#vsa-installation) for more installation instructions.
### Sage Attention
+14 -30
View File
@@ -40,26 +40,20 @@ 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 | 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 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| 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 | ❌ | ❌ | ✅ | ⭕ |
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
@@ -70,13 +64,3 @@ 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
+21 -106
View File
@@ -1,130 +1,45 @@
# 🧱 Data Preprocessing
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.
# 🧱 Data Preprocess
## Quick Start
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
Download the sample dataset and run preprocessing:
We provide a sample dataset to help you get started. Download the source media using the following command:
```bash
# 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
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=data/mini_i2v_dataset --repo_type=dataset
```
## Preprocessing Pipeline
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
The new preprocessing pipeline supports multiple dataset formats and video loaders:
```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
```
### Key Parameters
| 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:
To preprocess the dataset for fine-tuning or distillation, run:
```
your_dataset/
├── videos/
│ ├── video_001.mp4
│ ├── video_002.mp4
│ └── ...
└── videos2caption.json
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
```
The `videos2caption.json` maps video filenames to captions:
## Process your own dataset
```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:
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:
```
your_raw_data/
path_to_your_dataset_folder/
├── videos/
│ ├── 0.mp4
│ ├── 1.mp4
│ └── ...
├── videos.txt # list of video filenames
└── prompt.txt # corresponding captions (one per line)
├── videos.txt
└── prompt.txt
```
## Output Format
To generate the `videos2caption.json` and `merge.txt`, run
Preprocessing outputs Parquet files in the `combined_parquet_dataset/` subdirectory containing:
``` python
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
```
- `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
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
## Examples
```
bash scripts/preprocess/v1_preprocess_****.sh
```
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)**
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
+55 -153
View File
@@ -1,176 +1,78 @@
# 🧠 Finetuning
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.
# 🧠 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:
```bash
# Example: Wan2.1 T2V 1.3B full finetune (4 GPUs)
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
```
**Typical settings:**
Download the original model weights as specified in the [Distillation Section](../distillation/dmd.md):
- 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
Then you can run the finetune with:
## LoRA Finetuning
```
bash scripts/finetune/finetune_mochi.sh # for mochi
```
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
**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:
```bash
# Example: Wan2.1 T2V 1.3B LoRA finetune (1 GPU)
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v_lora.sh
bash scripts/finetune/finetune_v1_VSA.sh
```
Key differences from full finetune:
## ⚡ Lora Finetune
- 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:
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:
```bash
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
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
```
| 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) |
### 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.
### Merge LoRA Adapter
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 an adapter back into a base model:
### 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.
```bash
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
bash scripts/finetune/finetune_hunyuan.sh
bash scripts/finetune/finetune_mochi_lora_mix.sh
```
| 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
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
-67
View File
@@ -1,67 +0,0 @@
# 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)
-19
View File
@@ -1,19 +0,0 @@
# 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.
-1
View File
@@ -21,7 +21,6 @@ def main():
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
init_weights_from_safetensors="/mnt/weka/home/hao.zhang/wl/release/dmd_distill_1.3_4n_syn/checkpoint-900_weight_only/generator_inference_transformer"
)
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
@@ -1,98 +0,0 @@
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=True,
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())
@@ -1,57 +0,0 @@
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,
# TurboDiffusion uses a custom pipeline with RCM scheduler
override_pipeline_cls_name="TurboDiffusionPipeline",
)
# Generate videos with the same simple API, regardless of GPU count
# TurboDiffusion uses guidance_scale=1.0 (no CFG) and only 4 steps
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(
prompt,
output_path=OUTPUT_PATH,
save_video=True,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
# Generate another video with a different prompt, without reloading the model!
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,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
if __name__ == "__main__":
main()
@@ -1,55 +0,0 @@
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,
# TurboDiffusion uses a custom pipeline with RCM scheduler
override_pipeline_cls_name="TurboDiffusionPipeline",
)
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,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
# 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,
num_inference_steps=4,
seed=42,
guidance_scale=1.0,
)
if __name__ == "__main__":
main()
@@ -1,684 +0,0 @@
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()
@@ -1,46 +0,0 @@
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)
@@ -1,311 +0,0 @@
# 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
-587
View File
@@ -1,587 +0,0 @@
# 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
-4
View File
@@ -50,10 +50,6 @@ 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
+1 -2
View File
@@ -18,8 +18,7 @@ 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.SLA_ATTN,
AttentionBackendEnum.SAGE_SLA_ATTN)
AttentionBackendEnum.SAGE_ATTN_THREE)
hidden_size: int = 0
num_attention_heads: int = 0
+17 -1
View File
@@ -12,7 +12,8 @@ from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.utils import update_config_from_args
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean, shallow_asdict
from fastvideo.utils import (FlexibleArgumentParser, StoreBoolean,
PRECISION_TO_TYPE, shallow_asdict)
logger = init_logger(__name__)
@@ -311,6 +312,21 @@ class PipelineConfig:
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
unsupported_precisions = [
precision for precision in self.text_encoder_precisions
if precision not in PRECISION_TO_TYPE
and not precision.startswith("fp8")
]
if unsupported_precisions:
supported = ", ".join(PRECISION_TO_TYPE.keys())
logger.warning(
"Unsupported text encoder precision(s) detected in config: %s. "
"FastVideo will attempt to load them with transformers AutoModel when possible. "
"Supported fast paths: %s.",
unsupported_precisions,
supported,
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
@@ -1,288 +0,0 @@
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()
+38 -37
View File
@@ -12,7 +12,6 @@ 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
@@ -133,7 +132,6 @@ class FastVideoArgs:
# CPU offload parameters
dit_cpu_offload: bool = True
use_fsdp_inference: bool = True
dit_layerwise_offload: bool = False
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
vae_cpu_offload: bool = True
@@ -172,10 +170,9 @@ class FastVideoArgs:
"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
text_encoder_override: str | None = None
text_encoder_override_path: str | None = None
text_encoder_dtype: str | None = 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
@@ -422,11 +419,6 @@ 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,
@@ -439,6 +431,31 @@ class FastVideoArgs:
help=
"Use CPU offload for text encoder. Enable if run out of memory.",
)
parser.add_argument(
"--text-encoder-override",
type=str,
default=FastVideoArgs.text_encoder_override,
help=
("Load text encoder weights from a different local path or HF repo "
"instead of the main diffusers snapshot."),
)
parser.add_argument(
"--text-encoder-override-path",
type=str,
default=FastVideoArgs.text_encoder_override_path,
help=
("Optional relative path to the text encoder inside the override repository "
"(defaults to the module name)."),
)
parser.add_argument(
"--text-encoder-dtype",
type=str,
default=FastVideoArgs.text_encoder_dtype,
help=
("Torch dtype string for loading text encoders with transformers AutoModel "
"(e.g., fp16, bf16, fp32, fp8). If set, overrides pipeline-config precisions."
),
)
parser.add_argument(
"--image-encoder-cpu-offload",
action=StoreBoolean,
@@ -487,19 +504,6 @@ 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,
@@ -603,25 +607,22 @@ class FastVideoArgs:
kwargs['preprocess_config'] = PreprocessConfig.from_kwargs(kwargs)
return cls(**kwargs)
def get_component_override(
self, module_name: str) -> tuple[str | None, str | None]:
"""Return override repo/path for a given module if configured."""
if module_name.startswith(
"text_encoder") and self.text_encoder_override:
return self.text_encoder_override, self.text_encoder_override_path
return None, None
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
from fastvideo.platforms import current_platform
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(
+190 -286
View File
@@ -7,20 +7,13 @@ 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,
@@ -34,22 +27,12 @@ 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"
]
@@ -58,8 +41,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}
@@ -82,23 +65,18 @@ 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.
@@ -108,12 +86,10 @@ 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
@@ -122,37 +98,28 @@ 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
@@ -190,8 +157,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)
@@ -214,36 +181,29 @@ 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(
@@ -251,13 +211,10 @@ 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)
@@ -308,21 +265,19 @@ 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
@@ -335,14 +290,8 @@ 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
@@ -359,21 +308,17 @@ 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)
@@ -456,37 +401,32 @@ 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,
)
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:
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.
@@ -578,22 +518,20 @@ 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,
@@ -630,12 +568,10 @@ 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):
@@ -664,18 +600,16 @@ 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
@@ -692,31 +626,29 @@ 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)
@@ -731,7 +663,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.
@@ -742,39 +674,31 @@ 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)
@@ -791,20 +715,17 @@ 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,
)
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):
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.
@@ -827,20 +748,14 @@ 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)
@@ -928,18 +843,16 @@ 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()
@@ -947,14 +860,8 @@ 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
@@ -969,8 +876,7 @@ 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")
@@ -978,13 +884,10 @@ 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)
@@ -1013,6 +916,7 @@ 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
View File
@@ -2,7 +2,7 @@ from typing import Literal, get_args
from fastvideo.layers.quantization.base_config import QuantizationConfig
QuantizationMethods = Literal[None, "AbsMaxFP8"]
QuantizationMethods = Literal[None]
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
@@ -50,12 +50,7 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]:
if quantization not in QUANTIZATION_METHODS:
raise ValueError(f"Invalid quantization method: {quantization}")
# lazy import to avoid triggering `torch.compile` too early
from .absmax_fp8 import AbsMaxFP8Config
method_to_config: dict[str, type[QuantizationConfig]] = {
"AbsMaxFP8": AbsMaxFP8Config,
}
method_to_config: dict[str, type[QuantizationConfig]] = {}
# Update the `method_to_config` with customized quantization methods.
method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG)
-193
View File
@@ -1,193 +0,0 @@
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)
+9 -116
View File
@@ -1,13 +1,14 @@
from __future__ import annotations
import asyncio
import os
import random
# import cv2
import numpy as np
import torch
from diffusers.utils import export_to_video
from PIL import Image
from fastvideo.distributed.parallel_state import get_local_torch_device
from fastvideo.utils import logger
@@ -35,120 +36,6 @@ KEYBOARD_MAP_7 = { # templerun_distilled_model: still/w/s/left/right/a/d
}
KEYBOARD_MAP = KEYBOARD_MAP_4 # Default for backward compatibility
def expand_action_to_frames(action: dict, num_frames: int) -> tuple[torch.Tensor, torch.Tensor]:
result = {}
for key, tensor in action.items():
if tensor is not None:
# Expand to [num_frames, D] then unsqueeze to [1, num_frames, D]
result[key] = tensor.unsqueeze(0).repeat(num_frames, 1).unsqueeze(0)
else:
result[key] = None
if "mouse" not in result or result["mouse"] is None:
# keyboard device if available, otherwise default
device = result.get("keyboard", torch.tensor([])).device if result.get("keyboard") is not None else get_local_torch_device()
result["mouse"] = torch.zeros(1, num_frames, 2, device=device)
return result["keyboard"], result["mouse"]
def get_current_action(mode="universal"):
CAM_VALUE = 0.1
if mode == 'universal':
logger.info("")
logger.info('-'*30)
logger.info("PRESS [I, K, J, L, U] FOR CAMERA TRANSFORM\n (I: up, K: down, J: left, L: right, U: no move)")
logger.info("PRESS [W, S, A, D, Q] FOR MOVEMENT\n (W: forward, S: back, A: left, D: right, Q: no move)")
logger.info('-'*30)
CAMERA_VALUE_MAP = {
"i": [CAM_VALUE, 0],
"k": [-CAM_VALUE, 0],
"j": [0, -CAM_VALUE],
"l": [0, CAM_VALUE],
"u": [0, 0]
}
KEYBOARD_IDX = {
"w": [1, 0, 0, 0], "s": [0, 1, 0, 0], "a": [0, 0, 1, 0], "d": [0, 0, 0, 1],
"q": [0, 0, 0, 0]
}
flag = 0
while flag != 1:
try:
idx_mouse = input('Please input the mouse action (e.g. `U`):\n').strip().lower()
idx_keyboard = input('Please input the keyboard action (e.g. `W`):\n').strip().lower()
if idx_mouse in CAMERA_VALUE_MAP and idx_keyboard in KEYBOARD_IDX:
flag = 1
except Exception:
pass
mouse_cond = torch.tensor(CAMERA_VALUE_MAP[idx_mouse]).cuda()
keyboard_cond = torch.tensor(KEYBOARD_IDX[idx_keyboard]).cuda()
elif mode == 'gta_drive':
logger.info("")
logger.info('-'*30)
logger.info("PRESS [W, S, A, D, Q] FOR MOVEMENT\n (W: forward, S: back, A: left, D: right, Q: no move)")
logger.info('-'*30)
CAMERA_VALUE_MAP = {
"a": [0, -CAM_VALUE],
"d": [0, CAM_VALUE],
"q": [0, 0]
}
KEYBOARD_IDX = {
"w": [1, 0], "s": [0, 1],
"q": [0, 0]
}
flag = 0
while flag != 1:
try:
indexes = input('Please input the actions (split with ` `):\n(e.g. `W` for forward, `W A` for forward and left)\n').strip().lower().split(' ')
idx_mouse = []
idx_keyboard = []
for i in indexes:
if i in CAMERA_VALUE_MAP.keys():
idx_mouse += [i]
elif i in KEYBOARD_IDX.keys():
idx_keyboard += [i]
if len(idx_mouse) == 0:
idx_mouse += ['q']
if len(idx_keyboard) == 0:
idx_keyboard += ['q']
assert idx_mouse in [['a'], ['d'], ['q']] and idx_keyboard in [['q'], ['w'], ['s']]
flag = 1
except Exception:
pass
mouse_cond = torch.tensor(CAMERA_VALUE_MAP[idx_mouse[0]]).cuda()
keyboard_cond = torch.tensor(KEYBOARD_IDX[idx_keyboard[0]]).cuda()
elif mode == 'templerun':
logger.info("")
logger.info('-'*30)
logger.info("PRESS [W, S, A, D, Z, C, Q] FOR ACTIONS\n (W: jump, S: slide, A: left side, D: right side, Z: turn left, C: turn right, Q: no move)")
logger.info('-'*30)
KEYBOARD_IDX = {
"w": [0, 1, 0, 0, 0, 0, 0], "s": [0, 0, 1, 0, 0, 0, 0],
"a": [0, 0, 0, 0, 0, 1, 0], "d": [0, 0, 0, 0, 0, 0, 1],
"z": [0, 0, 0, 1, 0, 0, 0], "c": [0, 0, 0, 0, 1, 0, 0],
"q": [1, 0, 0, 0, 0, 0, 0]
}
flag = 0
while flag != 1:
try:
idx_keyboard = input('Please input the action: \n(e.g. `W` for forward, `Z` for turning left)\n').strip().lower()
if idx_keyboard in KEYBOARD_IDX.keys():
flag = 1
except Exception:
pass
keyboard_cond = torch.tensor(KEYBOARD_IDX[idx_keyboard]).cuda()
if mode != 'templerun':
return {
"mouse": mouse_cond,
"keyboard": keyboard_cond
}
return {
"keyboard": keyboard_cond
}
async def get_current_action_async(mode="universal"):
return await asyncio.to_thread(get_current_action, mode)
def load_initial_image(image_path: str = None) -> Image.Image:
if image_path and os.path.exists(image_path):
@@ -156,6 +43,7 @@ def load_initial_image(image_path: str = None) -> Image.Image:
logger.warning("No image provided, creating placeholder...")
return Image.new("RGB", (640, 352), (128, 128, 128))
def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = None):
if keyboard_dim not in (2, 4, 7):
raise ValueError(f"keyboard_dim must be 2, 4, or 7, got {keyboard_dim}")
@@ -257,6 +145,7 @@ def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = No
return {"keyboard": keyboard_condition, "mouse": mouse_condition}
def parse_config(config, mode="universal"):
assert mode in ['universal', 'gta_drive', 'templerun']
key_data = {}
@@ -294,6 +183,7 @@ def parse_config(config, mode="universal"):
)
return key_data, mouse_data
# NOTE: drawing functions are commented out to avoid cv2/libGL dependency.
#
# def draw_rounded_rectangle(image, top_left, bottom_right, color, radius=10, alpha=0.5):
@@ -309,6 +199,7 @@ def parse_config(config, mode="universal"):
# cv2.ellipse(overlay, (x2 - radius, y2 - radius), (radius, radius), 0, 0, 90, color, -1)
# cv2.addWeighted(overlay, alpha, image, 1 - alpha, 0, image)
#
#
# def draw_keys_on_frame(frame, keys, key_size=(80, 50), spacing=20, bottom_margin=30, mode='universal'):
# h, w, _ = frame.shape
# horison_shift = 90
@@ -353,6 +244,7 @@ def parse_config(config, mode="universal"):
# text_y = y + (key_size[1] + text_size[1]) // 2
# cv2.putText(frame, key_icon[key], (text_x, text_y), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 0, 0), 2)
#
#
# def overlay_icon(frame, icon, position, scale=1.0, rotation=0):
# x, y = position
# h, w, _ = icon.shape
@@ -390,6 +282,7 @@ def parse_config(config, mode="universal"):
# frame_region[:, :, c] = (1 - alpha) * frame_region[:, :, c] + alpha * icon_rgb[:, :, c]
# frame[top_left_y:bottom_right_y, top_left_x:bottom_right_x] = frame_region
#
#
# def process_video(input_video, output_video, config, mouse_icon_path,
# mouse_scale=1.0, mouse_rotation=0, process_icon=True, mode='universal'):
# key_data, mouse_data = parse_config(config, mode=mode)
+3 -14
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import math
from contextlib import nullcontext
from typing import Any
import numpy as np
@@ -734,19 +733,9 @@ class WanTransformer3DModel(CachableDiT):
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
else:
offload_mgr = getattr(self, "_layerwise_offload_manager", None)
use_offload = offload_mgr is not None and getattr(offload_mgr, "enabled", False)
for i, block in enumerate(self.blocks):
scope = offload_mgr.layer_scope(
prefetch_layer_idx=i + 1 if i + 1 < len(self.blocks) else None,
release_layer_idx=i,
non_blocking=True,
) if use_offload else nullcontext()
with scope:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
+172 -218
View File
@@ -31,11 +31,8 @@ from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.distributed import get_tp_rank, get_tp_world_size
from fastvideo.layers.activation import get_act_fn
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import (
MergedColumnParallelLinear,
QKVParallelLinear,
RowParallelLinear,
)
from fastvideo.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
@@ -47,7 +44,6 @@ class AttentionType:
Attention type.
Use string to be compatible with `torch.compile`.
"""
# Decoder attention between previous layer Q/K/V
DECODER = "decoder"
# Encoder attention between previous layer Q/K/V for encoder-decoder
@@ -64,16 +60,17 @@ class AttentionMetadata:
class T5DenseActDense(nn.Module):
def __init__(
self, config: T5Config, quant_config: QuantizationConfig | None = None
):
def __init__(self,
config: T5Config,
quant_config: QuantizationConfig | None = None):
super().__init__()
self.wi = MergedColumnParallelLinear(
config.d_model, [config.d_ff], bias=False
)
self.wo = RowParallelLinear(
config.d_ff, config.d_model, bias=False, quant_config=quant_config
)
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False)
self.wo = RowParallelLinear(config.d_ff,
config.d_model,
bias=False,
quant_config=quant_config)
self.act = get_act_fn(config.dense_act_fn)
def forward(self, hidden_states) -> torch.Tensor:
@@ -84,21 +81,23 @@ class T5DenseActDense(nn.Module):
class T5DenseGatedActDense(nn.Module):
def __init__(
self, config: T5Config, quant_config: QuantizationConfig | None = None
):
def __init__(self,
config: T5Config,
quant_config: QuantizationConfig | None = None):
super().__init__()
self.wi_0 = MergedColumnParallelLinear(
config.d_model, [config.d_ff], bias=False, quant_config=quant_config
)
self.wi_1 = MergedColumnParallelLinear(
config.d_model, [config.d_ff], bias=False, quant_config=quant_config
)
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False,
quant_config=quant_config)
self.wi_1 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False,
quant_config=quant_config)
# Should not run in fp16 unless mixed-precision is used,
# see https://github.com/huggingface/transformers/issues/20287.
self.wo = RowParallelLinear(
config.d_ff, config.d_model, bias=False, quant_config=quant_config
)
self.wo = RowParallelLinear(config.d_ff,
config.d_model,
bias=False,
quant_config=quant_config)
self.act = get_act_fn(config.dense_act_fn)
def forward(self, hidden_states) -> torch.Tensor:
@@ -110,18 +109,17 @@ class T5DenseGatedActDense(nn.Module):
class T5LayerFF(nn.Module):
def __init__(
self, config: T5Config, quant_config: QuantizationConfig | None = None
):
def __init__(self,
config: T5Config,
quant_config: QuantizationConfig | None = None):
super().__init__()
if config.is_gated_act:
self.DenseReluDense = T5DenseGatedActDense(
config, quant_config=quant_config
)
config, quant_config=quant_config)
else:
self.DenseReluDense = T5DenseActDense(
config, quant_config=quant_config
)
self.DenseReluDense = T5DenseActDense(config,
quant_config=quant_config)
self.layer_norm = RMSNorm(config.d_model, eps=config.layer_norm_epsilon)
@@ -134,41 +132,39 @@ class T5LayerFF(nn.Module):
# T5 has attn_bias and does not use softmax scaling
class T5MultiHeadAttention(nn.Module):
def __init__(self) -> None:
super().__init__()
def forward(self, q, k, v, attn_bias=None):
b, _, n, c = q.shape
attn = torch.einsum("binc,bjnc->bnij", q, k)
attn = torch.einsum('binc,bjnc->bnij', q, k)
if attn_bias is not None:
attn += attn_bias
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
x = torch.einsum("bnij,bjnc->binc", attn, v)
x = torch.einsum('bnij,bjnc->binc', attn, v)
x = x.reshape(b, -1, n * c)
return x
class T5Attention(nn.Module):
def __init__(
self,
config: T5Config,
attn_type: str,
has_relative_attention_bias=False,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
def __init__(self,
config: T5Config,
attn_type: str,
has_relative_attention_bias=False,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
self.attn_type = attn_type
# Cross-attention has no relative pos encoding anyway
self.is_decoder = attn_type == AttentionType.DECODER
self.has_relative_attention_bias = has_relative_attention_bias
self.relative_attention_num_buckets = (
self.relative_attention_num_buckets = \
config.relative_attention_num_buckets
)
self.relative_attention_max_distance = (
self.relative_attention_max_distance = \
config.relative_attention_max_distance
)
self.d_model = config.d_model
self.key_value_proj_dim = config.d_kv
self.total_num_heads = self.total_num_kv_heads = config.num_heads
@@ -195,13 +191,12 @@ class T5Attention(nn.Module):
self.attn = T5MultiHeadAttention()
if self.has_relative_attention_bias:
self.relative_attention_bias = VocabParallelEmbedding(
self.relative_attention_num_buckets,
self.total_num_heads,
org_num_embeddings=self.relative_attention_num_buckets,
padding_size=self.relative_attention_num_buckets,
quant_config=quant_config,
)
self.relative_attention_bias = \
VocabParallelEmbedding(self.relative_attention_num_buckets,
self.total_num_heads,
org_num_embeddings=self.relative_attention_num_buckets,
padding_size=self.relative_attention_num_buckets,
quant_config=quant_config)
self.o = RowParallelLinear(
self.total_num_heads * self.key_value_proj_dim,
self.d_model,
@@ -211,20 +206,21 @@ class T5Attention(nn.Module):
)
@staticmethod
def _relative_position_bucket(
relative_position, bidirectional=True, num_buckets=32, max_distance=128
) -> torch.Tensor:
def _relative_position_bucket(relative_position,
bidirectional=True,
num_buckets=32,
max_distance=128) -> torch.Tensor:
"""
Adapted from Mesh Tensorflow:
https://github.com/tensorflow/mesh/blob/0cb87fe07da627bf0b7e60475d59f95ed6b5be3d/mesh_tensorflow/transformer/transformer_layers.py#L593
Translate relative position to a bucket number for relative attention.
The relative position is defined as memory_position - query_position,
i.e. the distance in tokens from the attending position to the
attended-to position. If bidirectional=False, then positive relative
positions are invalid. We use smaller buckets for small absolute
relative_position and larger buckets for larger absolute
Translate relative position to a bucket number for relative attention.
The relative position is defined as memory_position - query_position,
i.e. the distance in tokens from the attending position to the
attended-to position. If bidirectional=False, then positive relative
positions are invalid. We use smaller buckets for small absolute
relative_position and larger buckets for larger absolute
relative_positions. All relative positions >=max_distance map to the
same bucket. All relative positions <=-max_distance map to the same
same bucket. All relative positions <=-max_distance map to the same
bucket. This should allow for more graceful generalization to longer
sequences than the model has been trained on
Args:
@@ -235,18 +231,16 @@ class T5Attention(nn.Module):
Returns:
a Tensor with the same shape as relative_position, containing int32
values in the range [0, num_buckets)
""" # noqa: E501
"""# noqa: E501
relative_buckets = 0
if bidirectional:
num_buckets //= 2
relative_buckets += (relative_position > 0).to(
torch.long
) * num_buckets
torch.long) * num_buckets
relative_position = torch.abs(relative_position)
else:
relative_position = -torch.min(
relative_position, torch.zeros_like(relative_position)
)
relative_position = -torch.min(relative_position,
torch.zeros_like(relative_position))
# now relative_position is in the range [0, inf)
# half of the buckets are for exact increments in positions
@@ -256,32 +250,30 @@ class T5Attention(nn.Module):
# The other half of the buckets are for logarithmically bigger bins
# in positions up to max_distance
relative_position_if_large = max_exact + (
torch.log(relative_position.float() / max_exact)
/ math.log(max_distance / max_exact)
* (num_buckets - max_exact)
).to(torch.long)
torch.log(relative_position.float() / max_exact) /
math.log(max_distance / max_exact) *
(num_buckets - max_exact)).to(torch.long)
relative_position_if_large = torch.min(
relative_position_if_large,
torch.full_like(relative_position_if_large, num_buckets - 1),
)
torch.full_like(relative_position_if_large, num_buckets - 1))
relative_buckets += torch.where(
is_small, relative_position, relative_position_if_large
)
relative_buckets += torch.where(is_small, relative_position,
relative_position_if_large)
return relative_buckets
def compute_bias(
self, query_length, key_length, device=None
) -> torch.Tensor:
def compute_bias(self,
query_length,
key_length,
device=None) -> torch.Tensor:
"""Compute binned relative position bias"""
if device is None:
device = self.relative_attention_bias.weight.device
context_position = torch.arange(
query_length, dtype=torch.long, device=device
)[:, None]
memory_position = torch.arange(
key_length, dtype=torch.long, device=device
)[None, :]
context_position = torch.arange(query_length,
dtype=torch.long,
device=device)[:, None]
memory_position = torch.arange(key_length,
dtype=torch.long,
device=device)[None, :]
# max_seq_len, nh
relative_position = memory_position - context_position
relative_position_bucket = self._relative_position_bucket(
@@ -294,8 +286,7 @@ class T5Attention(nn.Module):
relative_position_bucket
) # shape (query_length, key_length, num_heads)
x = values.permute([2, 0, 1]).unsqueeze(
0
) # shape (1, num_heads, query_length, key_length)
0) # shape (1, num_heads, query_length, key_length)
return x
def forward(
@@ -323,9 +314,8 @@ class T5Attention(nn.Module):
# The bias term is computed on longest sequence in batch. Biases
# for shorter sequences are slices of the longest.
assert self.attn_type == AttentionType.ENCODER
attn_bias = self.compute_bias(seq_len, seq_len).repeat(
num_seqs, 1, 1, 1
)
attn_bias = self.compute_bias(seq_len,
seq_len).repeat(num_seqs, 1, 1, 1)
attn_metadata.attn_bias = attn_bias
else:
# Encoder/Decoder Self-Attention Layer, attn bias already cached.
@@ -334,27 +324,24 @@ class T5Attention(nn.Module):
from fastvideo.platforms import current_platform
if attention_mask is not None:
attention_mask = (
attention_mask.view(bs, 1, 1, -1)
if attention_mask.ndim == 2
else attention_mask.unsqueeze(1)
)
mask_val = (
-1e4 if current_platform.is_mps() else torch.finfo(q.dtype).min
)
attention_mask = attention_mask.view(
bs, 1, 1,
-1) if attention_mask.ndim == 2 else attention_mask.unsqueeze(1)
mask_val = -1e4 if current_platform.is_mps() else torch.finfo(
q.dtype).min
attn_bias.masked_fill_(attention_mask == 0, mask_val)
if get_tp_world_size() > 1:
rank = get_tp_rank()
attn_bias = attn_bias[
:, rank * self.n_heads : (rank + 1) * self.n_heads, :, :
]
attn_bias = attn_bias[:, rank * self.n_heads:(rank + 1) *
self.n_heads, :, :]
attn_output = self.attn(q, k, v, attn_bias)
output, _ = self.o(attn_output)
return output
class T5LayerSelfAttention(nn.Module):
def __init__(
self,
config,
@@ -366,12 +353,10 @@ class T5LayerSelfAttention(nn.Module):
self.SelfAttention = T5Attention(
config,
AttentionType.DECODER
if "decoder" in prefix
else AttentionType.ENCODER,
if "decoder" in prefix else AttentionType.ENCODER,
has_relative_attention_bias=has_relative_attention_bias,
quant_config=quant_config,
prefix=f"{prefix}.SelfAttention",
)
prefix=f"{prefix}.SelfAttention")
self.layer_norm = RMSNorm(config.d_model, eps=config.layer_norm_epsilon)
def forward(
@@ -391,20 +376,17 @@ class T5LayerSelfAttention(nn.Module):
class T5LayerCrossAttention(nn.Module):
def __init__(
self,
config,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
def __init__(self,
config,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
self.EncDecAttention = T5Attention(
config,
AttentionType.ENCODER_DECODER,
has_relative_attention_bias=False,
quant_config=quant_config,
prefix=f"{prefix}.EncDecAttention",
)
self.EncDecAttention = T5Attention(config,
AttentionType.ENCODER_DECODER,
has_relative_attention_bias=False,
quant_config=quant_config,
prefix=f"{prefix}.EncDecAttention")
self.layer_norm = RMSNorm(config.d_model, eps=config.layer_norm_epsilon)
def forward(
@@ -422,14 +404,13 @@ class T5LayerCrossAttention(nn.Module):
class T5Block(nn.Module):
def __init__(
self,
config: T5Config,
is_decoder: bool,
has_relative_attention_bias=False,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
def __init__(self,
config: T5Config,
is_decoder: bool,
has_relative_attention_bias=False,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
self.is_decoder = is_decoder
self.layer = nn.ModuleList()
@@ -438,18 +419,13 @@ class T5Block(nn.Module):
config,
has_relative_attention_bias=has_relative_attention_bias,
quant_config=quant_config,
prefix=f"{prefix}.self_attn",
)
)
prefix=f"{prefix}.self_attn"))
if self.is_decoder:
self.layer.append(
T5LayerCrossAttention(
config,
quant_config=quant_config,
prefix=f"{prefix}.cross_attn",
)
)
T5LayerCrossAttention(config,
quant_config=quant_config,
prefix=f"{prefix}.cross_attn"))
self.layer.append(T5LayerFF(config, quant_config=quant_config))
@@ -459,15 +435,13 @@ class T5Block(nn.Module):
attention_mask: torch.Tensor,
attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor:
hidden_states = self.layer[0](
hidden_states=hidden_states,
attention_mask=attention_mask,
attn_metadata=attn_metadata,
)
hidden_states = self.layer[0](hidden_states=hidden_states,
attention_mask=attention_mask,
attn_metadata=attn_metadata)
if self.is_decoder:
hidden_states = self.layer[1](
hidden_states=hidden_states, attn_metadata=attn_metadata
)
hidden_states = self.layer[1](hidden_states=hidden_states,
attn_metadata=attn_metadata)
# Apply Feed Forward layer
hidden_states = self.layer[2](hidden_states)
@@ -477,49 +451,37 @@ class T5Block(nn.Module):
class T5Stack(nn.Module):
def __init__(
self,
config: T5Config,
is_decoder: bool,
n_layers: int,
embed_tokens=None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
is_umt5: bool = False,
):
def __init__(self,
config: T5Config,
is_decoder: bool,
n_layers: int,
embed_tokens=None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
is_umt5: bool = False):
super().__init__()
self.embed_tokens = embed_tokens
self.is_umt5 = is_umt5
if is_umt5:
self.block = nn.ModuleList(
[
T5Block(
config,
self.block = nn.ModuleList([
T5Block(config,
is_decoder=is_decoder,
has_relative_attention_bias=True,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{i}",
)
for i in range(n_layers)
]
)
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
])
else:
# Only the first block has relative positional encoding.
self.block = nn.ModuleList(
[
T5Block(
config,
self.block = nn.ModuleList([
T5Block(config,
is_decoder=is_decoder,
has_relative_attention_bias=i == 0,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{i}",
)
for i in range(n_layers)
]
)
self.final_layer_norm = RMSNorm(
config.d_model, eps=config.layer_norm_epsilon
)
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
])
self.final_layer_norm = RMSNorm(config.d_model,
eps=config.layer_norm_epsilon)
def forward(
self,
@@ -540,24 +502,24 @@ class T5Stack(nn.Module):
class T5EncoderModel(TextEncoder):
def __init__(self, config: T5Config, prefix: str = ""):
super().__init__(config)
quant_config = None
self.shared = VocabParallelEmbedding(
config.vocab_size,
config.d_model,
org_num_embeddings=config.vocab_size,
)
org_num_embeddings=config.vocab_size)
self.encoder = T5Stack(
config,
False,
config.num_layers,
self.shared,
quant_config=config.quant_config,
prefix=f"{prefix}.encoder",
is_umt5=False,
)
self.encoder = T5Stack(config,
False,
config.num_layers,
self.shared,
quant_config=quant_config,
prefix=f"{prefix}.encoder",
is_umt5=False)
def get_input_embeddings(self):
return self.shared
@@ -583,9 +545,8 @@ class T5EncoderModel(TextEncoder):
attention_mask=attention_mask,
)
def load_weights(
self, weights: Iterable[tuple[str, torch.Tensor]]
) -> set[str]:
def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"),
@@ -623,33 +584,32 @@ class T5EncoderModel(TextEncoder):
continue
param = params_dict[name]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
)
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
class UMT5EncoderModel(TextEncoder):
def __init__(self, config: T5Config, prefix: str = ""):
super().__init__(config)
quant_config = None
self.shared = VocabParallelEmbedding(
config.vocab_size,
config.d_model,
org_num_embeddings=config.vocab_size,
)
org_num_embeddings=config.vocab_size)
self.encoder = T5Stack(
config,
False,
config.num_layers,
self.shared,
quant_config=config.quant_config,
prefix=f"{prefix}.encoder",
is_umt5=True,
)
self.encoder = T5Stack(config,
False,
config.num_layers,
self.shared,
quant_config=quant_config,
prefix=f"{prefix}.encoder",
is_umt5=True)
def get_input_embeddings(self):
return self.shared
@@ -675,20 +635,15 @@ class UMT5EncoderModel(TextEncoder):
attention_mask=attention_mask,
)
def load_weights(
self, weights: Iterable[tuple[str, torch.Tensor]]
) -> set[str]:
def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
loaded = False
if "decoder" in name or "lm_head" in name:
continue
for (
param_name,
weight_name,
shard_id,
) in self.config.arch_config.stacked_params_mapping:
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
@@ -713,9 +668,8 @@ class UMT5EncoderModel(TextEncoder):
continue
param = params_dict[name]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
)
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
-224
View File
@@ -1,224 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Inspired by SGLang's layerwise offload implementation:
# https://github.com/sgl-project/sglang/pull/15511
#
# This implementation provides a lightweight layerwise CPU offload manager
# with async H2D prefetch using a dedicated CUDA stream, following SGLang's design.
import re
from contextlib import contextmanager
from typing import Dict, Set, Optional, Tuple
import torch
class LayerwiseOffloadManager:
"""A lightweight layerwise CPU offload manager.
Offloads per-layer parameters/buffers from GPU to CPU, and supports async H2D
prefetch using a dedicated CUDA stream.
"""
def __init__(
self,
model: torch.nn.Module,
*,
module_list_attr: str,
num_layers: int,
enabled: bool,
pin_cpu_memory: bool = True,
auto_initialize: bool = False,
) -> None:
self.model = model
self.module_list_attr = module_list_attr
self.num_layers = int(num_layers)
self.pin_cpu_memory = bool(pin_cpu_memory)
self.enabled = bool(enabled and torch.cuda.is_available())
self.device = (
torch.device("cuda", torch.cuda.current_device()) if self.enabled else None
)
self.copy_stream = torch.cuda.Stream() if self.enabled else None
self._layer_name_re = re.compile(
rf"(^|\.){re.escape(module_list_attr)}\.(\d+)(\.|$)"
)
self._cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
self._cpu_dtypes: Dict[int, Dict[str, torch.dtype]] = {}
self._gpu_layers: Dict[int, Set[str]] = {}
self._named_parameters: Dict[str, torch.nn.Parameter] = {}
self._named_buffers: Dict[str, torch.Tensor] = {}
self._meta: Dict[str, Tuple[int, torch.dtype]] = {}
if auto_initialize:
self.initialize()
def _match_layer_idx(self, name: str) -> Optional[int]:
m = self._layer_name_re.search(name)
if not m:
return None
try:
return int(m.group(2))
except Exception:
return None
def _record_meta(self, name: str, t: torch.Tensor) -> None:
if name not in self._meta:
self._meta[name] = (int(t.ndim), t.dtype)
def _make_placeholder(self, name: str) -> torch.Tensor:
"""Rank-preserving empty placeholder on GPU."""
assert self.device is not None
ndim, dtype = self._meta[name]
shape = (0,) if ndim <= 0 else (0,) * ndim
return torch.empty(shape, device=self.device, dtype=dtype)
def _get_target(self, name: str) -> torch.Tensor:
if name in self._named_parameters:
return self._named_parameters[name]
return self._named_buffers[name]
def _offload_tensor(self, name: str, tensor: torch.Tensor, layer_idx: int) -> None:
if layer_idx not in self._cpu_weights:
self._cpu_weights[layer_idx] = {}
self._cpu_dtypes[layer_idx] = {}
self._record_meta(name, tensor)
cpu_weight = tensor.detach().to("cpu")
if self.pin_cpu_memory:
cpu_weight = cpu_weight.pin_memory()
self._cpu_weights[layer_idx][name] = cpu_weight
self._cpu_dtypes[layer_idx][name] = tensor.dtype
if self.device is not None:
tensor.data = self._make_placeholder(name)
@torch.compiler.disable
def initialize(self) -> None:
"""Offload all matched layer tensors to CPU and prefetch layer 0 (sync)."""
if not self.enabled:
return
self._named_parameters = dict(self.model.named_parameters())
self._named_buffers = dict(self.model.named_buffers())
for name, param in self._named_parameters.items():
layer_idx = self._match_layer_idx(name)
if layer_idx is None or layer_idx >= self.num_layers:
continue
self._offload_tensor(name, param, layer_idx)
for name, buf in self._named_buffers.items():
layer_idx = self._match_layer_idx(name)
if layer_idx is None or layer_idx >= self.num_layers:
continue
self._offload_tensor(name, buf, layer_idx)
self.prefetch_layer(0, non_blocking=False)
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
@torch.compiler.disable
def prefetch_layer(self, layer_idx: int, non_blocking: bool = True) -> None:
"""Prefetch a layer's tensors from CPU to GPU (async on copy_stream)."""
if not self.enabled or self.device is None or self.copy_stream is None:
return
if layer_idx < 0 or layer_idx >= self.num_layers:
return
if layer_idx in self._gpu_layers:
return
if layer_idx not in self._cpu_weights:
return
self.copy_stream.wait_stream(torch.cuda.current_stream())
param_names: Set[str] = set()
with torch.cuda.stream(self.copy_stream):
for name, cpu_weight in self._cpu_weights[layer_idx].items():
target = self._get_target(name)
gpu_weight = torch.empty(
cpu_weight.shape,
dtype=self._cpu_dtypes[layer_idx][name],
device=self.device,
)
gpu_weight.copy_(cpu_weight, non_blocking=non_blocking)
target.data = gpu_weight
param_names.add(name)
self._gpu_layers[layer_idx] = param_names
@contextmanager
def layer_scope(
self,
*,
prefetch_layer_idx: Optional[int],
release_layer_idx: Optional[int],
non_blocking: bool = True,
):
if self.enabled and release_layer_idx is not None:
cur = release_layer_idx
if (
cur not in self._gpu_layers
and cur in self._cpu_weights
and self.device is not None
and self.copy_stream is not None
):
self.prefetch_layer(cur, non_blocking=False)
torch.cuda.current_stream().wait_stream(self.copy_stream)
if self.enabled and prefetch_layer_idx is not None:
self.prefetch_layer(prefetch_layer_idx, non_blocking=non_blocking)
try:
yield
finally:
if self.enabled and self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
if self.enabled and release_layer_idx is not None:
self.release_layer(release_layer_idx)
@torch.compiler.disable
def release_layer(self, layer_idx: int) -> None:
"""Release a layer's tensors back to placeholders (free VRAM)."""
if not self.enabled or self.device is None:
return
if layer_idx < 0:
return
param_names = self._gpu_layers.pop(layer_idx, None)
if not param_names:
return
for name in param_names:
target = self._get_target(name)
# Ensure meta exists even if something unexpected happened
self._record_meta(name, target)
target.data = self._make_placeholder(name)
@torch.compiler.disable
def release_all(self) -> None:
"""Release all currently-resident layers back to placeholders."""
if not self.enabled or self.device is None:
return
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
for layer_idx in list(self._gpu_layers.keys()):
param_names = self._gpu_layers.pop(layer_idx, None)
if not param_names:
continue
for name in param_names:
target = self._get_target(name)
self._record_meta(name, target)
target.data = self._make_placeholder(name)
+171 -246
View File
@@ -15,28 +15,22 @@ import torch.distributed as dist
import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file
from torch.distributed import init_device_mesh
from transformers import AutoImageProcessor, AutoTokenizer
from transformers import AutoImageProcessor, AutoModel, AutoTokenizer
from transformers import UMT5EncoderModel
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.configs.models import EncoderConfig
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.quantization import get_quantization_config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files,
filter_files_not_needed_for_inference,
pt_weights_iterator,
safetensors_weights_iterator,
)
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
pt_weights_iterator, safetensors_weights_iterator)
from fastvideo.models.registry import ModelRegistry
from fastvideo.utils import PRECISION_TO_TYPE
from fastvideo.models.layerwise_offload import LayerwiseOffloadManager
logger = init_logger(__name__)
@@ -51,27 +45,26 @@ class ComponentLoader(ABC):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""
Load the component based on the model path, architecture, and inference args.
Args:
model_path: Path to the component model
fastvideo_args: FastVideoArgs
Returns:
The loaded component
"""
raise NotImplementedError
@classmethod
def for_module_type(
cls, module_type: str, transformers_or_diffusers: str
) -> "ComponentLoader":
def for_module_type(cls, module_type: str,
transformers_or_diffusers: str) -> 'ComponentLoader':
"""
Factory method to create a component loader for a specific module type.
Args:
module_type: Type of module (e.g., "vae", "text_encoder", "transformer", "scheduler")
transformers_or_diffusers: Whether the module is from transformers or diffusers
Returns:
A component loader for the specified module type
"""
@@ -92,16 +85,13 @@ class ComponentLoader(ABC):
if module_type in module_loaders:
loader_cls, expected_library = module_loaders[module_type]
# Assert that the library matches what's expected for this module type
assert transformers_or_diffusers == expected_library, (
f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
)
assert transformers_or_diffusers == expected_library, f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
return loader_cls()
# For unknown module types, use a generic loader
logger.warning(
"No specific loader found for module type: %s. Using generic loader.",
module_type,
)
module_type)
return GenericComponentLoader(transformers_or_diffusers)
@@ -164,45 +154,36 @@ class TextEncoderLoader(ComponentLoader):
if use_safetensors:
hf_weights_files = filter_duplicate_safetensors_files(
hf_weights_files, hf_folder, index_file
)
hf_weights_files, hf_folder, index_file)
else:
hf_weights_files = filter_files_not_needed_for_inference(
hf_weights_files
)
hf_weights_files)
if len(hf_weights_files) == 0:
raise RuntimeError(
f"Cannot find any model weights with `{model_name_or_path}`"
)
f"Cannot find any model weights with `{model_name_or_path}`")
return hf_folder, hf_weights_files, use_safetensors
def _get_weights_iterator(
self, source: "Source", to_cpu: bool
) -> Generator[tuple[str, torch.Tensor], None, None]:
self, source: "Source",
to_cpu: bool) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Get an iterator for the model weights based on the load format."""
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
source.model_or_path,
source.fall_back_to_pt,
source.allow_patterns_overrides,
)
source.model_or_path, source.fall_back_to_pt,
source.allow_patterns_overrides)
if use_safetensors:
weights_iterator = safetensors_weights_iterator(
hf_weights_files, to_cpu=to_cpu
)
weights_iterator = safetensors_weights_iterator(hf_weights_files,
to_cpu=to_cpu)
else:
weights_iterator = pt_weights_iterator(
hf_weights_files, to_cpu=to_cpu
)
weights_iterator = pt_weights_iterator(hf_weights_files,
to_cpu=to_cpu)
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
# Apply the prefix.
return (
(source.prefix + name, tensor)
for (name, tensor) in weights_iterator
)
return ((source.prefix + name, tensor)
for (name, tensor) in weights_iterator)
def _get_all_weights(
self,
@@ -214,9 +195,8 @@ class TextEncoderLoader(ComponentLoader):
model_path,
prefix="",
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
allow_patterns_overrides=getattr(
model, "allow_patterns_overrides", None
),
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
None),
)
yield from self._get_weights_iterator(primary_weights, to_cpu)
@@ -245,94 +225,78 @@ class TextEncoderLoader(ComponentLoader):
# @TODO(Wei): Better way to handle this?
try:
encoder_config = (
fastvideo_args.pipeline_config.text_encoder_configs[0]
)
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
0]
encoder_config.update_model_arch(model_config)
encoder_precision = (
fastvideo_args.pipeline_config.text_encoder_precisions[0]
)
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
0]
except Exception:
encoder_config = (
fastvideo_args.pipeline_config.text_encoder_configs[1]
)
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
1]
encoder_config.update_model_arch(model_config)
encoder_precision = (
fastvideo_args.pipeline_config.text_encoder_precisions[1]
)
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
1]
requested_dtype = fastvideo_args.text_encoder_dtype or encoder_precision
target_device = get_local_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(
model_path,
encoder_config,
target_device,
fastvideo_args,
encoder_precision,
use_text_encoder_override=True,
)
if requested_dtype not in PRECISION_TO_TYPE:
logger.info(
"Loading text encoder via transformers AutoModel with dtype=%s",
requested_dtype,
)
return self._load_with_transformers(
model_path,
requested_dtype,
fastvideo_args,
target_device,
)
def load_model(
self,
model_path: str,
model_config: EncoderConfig,
target_device: torch.device,
fastvideo_args: FastVideoArgs,
dtype: str = "fp16",
use_text_encoder_override: bool = False, # prevent subclasses from misusing
):
use_cpu_offload = (
fastvideo_args.text_encoder_cpu_offload
and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0
)
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args, requested_dtype)
def load_model(self,
model_path: str,
model_config: EncoderConfig,
target_device: torch.device,
fastvideo_args: FastVideoArgs,
dtype: str = "fp16"):
use_cpu_offload = fastvideo_args.text_encoder_cpu_offload and len(
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
from fastvideo.platforms import current_platform
if fastvideo_args.text_encoder_cpu_offload:
target_device = (
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
target_device = torch.device(
"mps") if current_platform.is_mps() else torch.device("cpu")
if dtype not in PRECISION_TO_TYPE:
supported = ", ".join(PRECISION_TO_TYPE.keys())
raise ValueError(
f"Unsupported text encoder precision '{dtype}'. "
f"Supported precisions: {supported}. "
"FP8 checkpoints will currently be materialized in a supported dtype, "
"so they will not reduce VRAM usage."
)
# Set quantization config if specified
if use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError(
"override_text_encoder_quant is set but override_text_encoder_safetensors is None"
)
quant_cls = get_quantization_config(
fastvideo_args.override_text_encoder_quant
)
model_config.quant_config = quant_cls()
logger.info("Loading text encoder with precision=%s (%s)",
dtype, PRECISION_TO_TYPE[dtype])
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
with target_device:
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
model: TextEncoder = model_cls(model_config) # type: ignore
model = model_cls(model_config)
weights_to_load = {name for name, _ in model.named_parameters()}
if use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None:
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
to_cpu=use_cpu_offload,
)
) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
model, model_path, to_cpu=use_cpu_offload
)
) # type: ignore
loaded_weights = model.load_weights(
self._get_all_weights(model, model_path,
to_cpu=use_cpu_offload))
self.counter_after_loading_weights = time.perf_counter()
logger.info(
"Loading weights took %.2f seconds",
self.counter_after_loading_weights
- self.counter_before_loading_weights,
)
self.counter_after_loading_weights -
self.counter_before_loading_weights)
# Explicitly move model to target device after loading weights
model = model.to(target_device)
@@ -357,8 +321,7 @@ class TextEncoderLoader(ComponentLoader):
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
)
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
else:
mesh = init_device_mesh(
"cuda",
@@ -371,22 +334,65 @@ class TextEncoderLoader(ComponentLoader):
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
)
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
# We only enable strict check for non-quantized models
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded and model_config.quant_config is None:
if weights_not_loaded:
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
return model.eval()
def _resolve_torch_dtype(self, dtype: str) -> torch.dtype:
if dtype in PRECISION_TO_TYPE:
return PRECISION_TO_TYPE[dtype]
if dtype.startswith("fp8"):
torch_dtype = getattr(torch, "float8_e4m3fn", None)
if torch_dtype is None:
torch_dtype = getattr(torch, "float8_e4m3fnuz", None)
if torch_dtype is None:
raise ValueError(
"Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}"
"FP8 requested for text encoder loading, but the current "
"PyTorch build does not expose float8 dtypes. Upgrade PyTorch "
"or choose a supported dtype (fp16/bf16/fp32)."
)
return torch_dtype
raise ValueError(
f"Unsupported text encoder dtype '{dtype}'. "
"Pass a torch dtype string such as fp16, bf16, fp32, or fp8."
)
def _load_with_transformers(
self,
model_path: str,
dtype: str,
fastvideo_args: FastVideoArgs,
target_device: torch.device,
) -> nn.Module:
torch_dtype = self._resolve_torch_dtype(dtype)
device_map = "auto" if fastvideo_args.text_encoder_cpu_offload else None
model = AutoModel.from_pretrained(
model_path,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
torch_dtype=torch_dtype,
device_map=device_map,
)
if device_map is None:
model = model.to(target_device)
return model.eval()
class ImageEncoderLoader(TextEncoderLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, and inference args."""
# model_config: PretrainedConfig = get_hf_config(
@@ -409,21 +415,13 @@ class ImageEncoderLoader(TextEncoderLoader):
from fastvideo.platforms import current_platform
if fastvideo_args.image_encoder_cpu_offload:
target_device = (
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
)
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
target_device = get_local_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(
model_path,
encoder_config,
target_device,
fastvideo_args,
fastvideo_args.pipeline_config.image_encoder_precision,
)
model_path, encoder_config, target_device, fastvideo_args,
fastvideo_args.pipeline_config.image_encoder_precision)
class ImageProcessorLoader(ComponentLoader):
@@ -433,12 +431,9 @@ class ImageProcessorLoader(ComponentLoader):
"""Load the image processor based on the model path, and inference args."""
logger.info("Loading image processor from %s", model_path)
image_processor = AutoImageProcessor.from_pretrained(
model_path,
)
logger.info(
"Loaded image processor: %s", image_processor.__class__.__name__
)
image_processor = AutoImageProcessor.from_pretrained(model_path, )
logger.info("Loaded image processor: %s",
image_processor.__class__.__name__)
return image_processor
@@ -454,7 +449,7 @@ class TokenizerLoader(ComponentLoader):
# in v0, this was same string as encoder_name "ClipTextModel"
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
# other method of config?
padding_size="right",
padding_size='right',
)
logger.info("Loaded tokenizer: %s", tokenizer.__class__.__name__)
return tokenizer
@@ -467,9 +462,7 @@ class VAELoader(ComponentLoader):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, (
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
)
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
fastvideo_args.model_paths["vae"] = model_path
vae_config = fastvideo_args.pipeline_config.vae_config
@@ -478,32 +471,23 @@ class VAELoader(ComponentLoader):
from fastvideo.platforms import current_platform
if fastvideo_args.vae_cpu_offload:
target_device = (
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
)
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
target_device = get_local_torch_device()
with set_default_torch_dtype(
PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
if fastvideo_args.pipeline_config.vae_precision
else torch.bfloat16
):
with set_default_torch_dtype(PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision] if fastvideo_args.pipeline_config.vae_precision else torch.bfloat16):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
os.path.join(str(model_path), "*.safetensors"))
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
vae.load_state_dict(
loaded, strict=False
) # We might only load encoder or decoder
loaded, strict=False) # We might only load encoder or decoder
return vae.eval()
@@ -519,8 +503,7 @@ class TransformerLoader(ComponentLoader):
if cls_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported."
)
"Only diffusers format is supported.")
logger.info("transformer cls_name: %s", cls_name)
if fastvideo_args.override_transformer_cls_name is not None:
@@ -537,54 +520,40 @@ class TransformerLoader(ComponentLoader):
# Find all safetensors files
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Check if we should use custom initialization weights
custom_weights_path = getattr(
fastvideo_args, "init_weights_from_safetensors", None
)
use_custom_weights = (
custom_weights_path
and os.path.exists(custom_weights_path)
and not hasattr(fastvideo_args, "_loading_teacher_critic_model")
)
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
if use_custom_weights:
if "transformer_2" in model_path:
custom_weights_path = getattr(
fastvideo_args, "init_weights_from_safetensors_2", None
)
assert custom_weights_path is not None, (
"Custom initialization weights must be provided"
)
if 'transformer_2' in model_path:
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors_2', None)
assert custom_weights_path is not None, "Custom initialization weights must be provided"
if os.path.isdir(custom_weights_path):
safetensors_list = glob.glob(
os.path.join(str(custom_weights_path), "*.safetensors")
)
os.path.join(str(custom_weights_path), "*.safetensors"))
else:
assert custom_weights_path.endswith(".safetensors"), (
"Custom initialization weights must be a safetensors file"
)
assert custom_weights_path.endswith(".safetensors"), "Custom initialization weights must be a safetensors file"
safetensors_list = [custom_weights_path]
logger.info(
"Loading model from %s safetensors files: %s",
len(safetensors_list),
safetensors_list,
)
logger.info("Loading model from %s safetensors files: %s",
len(safetensors_list), safetensors_list)
default_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision
]
fastvideo_args.pipeline_config.dit_precision]
# Load the model using FSDP loader
assert fastvideo_args.hsdp_shard_dim is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={"config": dit_config, "hf_config": hf_config},
init_params={
"config": dit_config,
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
device=get_local_torch_device(),
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
@@ -599,47 +568,15 @@ class TransformerLoader(ComponentLoader):
output_dtype=None,
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
)
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
assert next(model.parameters()).dtype == default_dtype, (
"Model dtype does not match default dtype"
)
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
model = model.eval()
if fastvideo_args.dit_layerwise_offload and hasattr(model, "blocks"):
# Check if this is a Wan model (only Wan models support layerwise offload)
is_wan_model = "Wan" in cls_name
if not is_wan_model:
logger.warning(
"Layerwise offload is currently only supported for Wan models. "
"Model class '%s' does not support layerwise offload. "
"Disabling layerwise offload for this model.",
cls_name
)
else:
try:
num_layers = len(getattr(model, "blocks"))
except TypeError:
num_layers = None
if isinstance(num_layers, int) and num_layers > 0:
# Ensure model is on the correct device (CUDA) before initializing manager
# This ensures non-managed parameters (embeddings, final norms) are on GPU
model = model.to(get_local_torch_device())
mgr = LayerwiseOffloadManager(
model,
module_list_attr="blocks",
num_layers=num_layers,
enabled=True,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
auto_initialize=True,
)
setattr(model, "_layerwise_offload_manager", mgr)
return model
@@ -651,9 +588,7 @@ class SchedulerLoader(ComponentLoader):
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, (
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
)
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
@@ -662,8 +597,7 @@ class SchedulerLoader(ComponentLoader):
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
if fastvideo_args.pipeline_config.timesteps_scale is not None:
scheduler.set_timesteps_scale(
fastvideo_args.pipeline_config.timesteps_scale
)
fastvideo_args.pipeline_config.timesteps_scale)
return scheduler
@@ -676,11 +610,8 @@ class GenericComponentLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load a generic component based on the model path, and inference args."""
logger.warning(
"Using generic loader for %s with library %s",
model_path,
self.library,
)
logger.warning("Using generic loader for %s with library %s",
model_path, self.library)
if self.library == "transformers":
from transformers import AutoModel
@@ -690,10 +621,8 @@ class GenericComponentLoader(ComponentLoader):
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
)
logger.info(
"Loaded generic transformers model: %s",
model.__class__.__name__,
)
logger.info("Loaded generic transformers model: %s",
model.__class__.__name__)
return model
elif self.library == "diffusers":
logger.warning(
@@ -715,21 +644,18 @@ class PipelineComponentLoader:
"""
@staticmethod
def load_module(
module_name: str,
component_model_path: str,
transformers_or_diffusers: str,
fastvideo_args: FastVideoArgs,
):
def load_module(module_name: str, component_model_path: str,
transformers_or_diffusers: str,
fastvideo_args: FastVideoArgs):
"""
Load a pipeline module.
Args:
module_name: Name of the module (e.g., "vae", "text_encoder", "transformer", "scheduler")
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
pipeline_args: Inference arguments
Returns:
The loaded module
"""
@@ -741,9 +667,8 @@ class PipelineComponentLoader:
)
# Get the appropriate loader for this module type
loader = ComponentLoader.for_module_type(
module_name, transformers_or_diffusers
)
loader = ComponentLoader.for_module_type(module_name,
transformers_or_diffusers)
# Load the module
return loader.load(component_model_path, fastvideo_args)
+1 -1
View File
@@ -318,7 +318,7 @@ def load_model_from_full_model_state_dict(
unused_keys)
# List of allowed parameter name patterns
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress"] # Can be extended as needed
for new_param_name in unused_keys:
if not any(pattern in new_param_name
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
-2
View File
@@ -75,8 +75,6 @@ _SCHEDULERS = {
"SelfForcingFlowMatchScheduler":
("schedulers", "scheduling_self_forcing_flow_match",
"SelfForcingFlowMatchScheduler"),
"RCMScheduler":
("schedulers", "scheduling_rcm", "RCMScheduler"),
}
_FAST_VIDEO_MODELS = {
@@ -1,323 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# rCM (recurrent Consistency Model) Scheduler for TurboDiffusion support
#
# This scheduler implements the rCM sampling method from TurboDiffusion,
# enabling 1-4 step video generation with distilled checkpoints.
#
# Reference:
# TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times
# https://arxiv.org/pdf/2512.16093
import math
from dataclasses import dataclass
from typing import Any
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.base import BaseScheduler
logger = init_logger(__name__)
@dataclass
class RCMSchedulerOutput(BaseOutput):
"""
Output class for the RCM scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, ...)`):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used
as next model input in the denoising loop.
"""
prev_sample: torch.FloatTensor
class RCMScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
rCM (recurrent Consistency Model) scheduler for TurboDiffusion.
This scheduler implements the rCM sampling method which enables 1-4 step
video generation using distilled checkpoints. It uses:
1. TrigFlow → RectifiedFlow timestep conversion
2. SDE sampling formula: x = (1 - t_next) * (x - t_cur * v_pred) + t_next * noise
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps used to train the model.
sigma_max (`float`, defaults to 80.0):
The initial sigma value for rCM sampling. Controls the noise level
at the start of sampling.
mid_timesteps (`list[float]`, *optional*):
Custom intermediate timesteps. If None, uses optimized defaults
[1.5, 1.4, 1.0] for better visual quality.
"""
_compatibles: list[Any] = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
sigma_max: float = 80.0,
mid_timesteps: list[float] | None = None,
):
# Default mid timesteps optimized for visual quality
if mid_timesteps is None:
mid_timesteps = [1.5, 1.4, 1.0]
self._mid_timesteps = mid_timesteps
self.num_train_timesteps = num_train_timesteps
self.sigma_max = sigma_max
# Initialize with default timesteps (will be set properly via set_timesteps)
self.timesteps = torch.tensor([1.0, 0.0], dtype=torch.float64)
self.sigmas = self.timesteps.clone()
self._step_index: int | None = None
self._begin_index: int | None = None
BaseScheduler.__init__(self)
@property
def step_index(self) -> int | None:
"""
The index counter for current timestep. Increases by 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self) -> int | None:
"""
The index for the first timestep.
"""
return self._begin_index
@property
def init_noise_sigma(self) -> float:
"""
Initial noise sigma for scaling latents.
In rCM, initial noise is scaled by the first sigma value:
x_0 = noise * sigmas[0]
This property is used by LatentPreparationStage to scale initial latents.
"""
return float(self.sigmas[0])
def set_begin_index(self, begin_index: int = 0) -> None:
"""
Sets the begin index for the scheduler.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def set_shift(self, shift: float) -> None:
"""
rCM doesn't use shift parameter, but required by BaseScheduler.
"""
pass
def set_timesteps(
self,
num_inference_steps: int,
device: str | torch.device | None = None,
sigma_max: float | None = None,
) -> None:
"""
Sets the discrete timesteps used for the rCM sampling process.
The timesteps are computed using TrigFlow → RectifiedFlow conversion:
1. Start with atan(sigma_max) and intermediate values
2. Convert via: t = sin(t) / (cos(t) + sin(t))
Args:
num_inference_steps (`int`):
The number of diffusion steps (1-4 for rCM).
device (`str` or `torch.device`, *optional*):
The device to move timesteps to.
sigma_max (`float`, *optional*):
Override the initial sigma value.
"""
if num_inference_steps < 1 or num_inference_steps > 4:
logger.warning(
"rCM is optimized for 1-4 steps, got %d steps. "
"Performance may be suboptimal.", num_inference_steps
)
self.num_inference_steps = num_inference_steps
if sigma_max is not None:
self.sigma_max = sigma_max
# Build timestep schedule
mid_t = self._mid_timesteps[:num_inference_steps - 1]
# TrigFlow timesteps: [atan(sigma_max), mid_t..., 0]
t_steps = torch.tensor(
[math.atan(self.sigma_max), *mid_t, 0],
dtype=torch.float64,
device=device,
)
# Convert TrigFlow → RectifiedFlow: t = sin(t) / (cos(t) + sin(t))
t_steps = torch.sin(t_steps) / (torch.cos(t_steps) + torch.sin(t_steps))
# Store raw sigmas for use in step() formula
self.sigmas = t_steps.clone()
# Scale timesteps by 1000 for model input (as per TurboDiffusion)
self.timesteps = t_steps * 1000
self._step_index = None
self._begin_index = None
logger.debug("rCM timesteps (scaled): %s", self.timesteps.tolist())
logger.debug("rCM sigmas (raw): %s", self.sigmas.tolist())
def _init_step_index(self, timestep: torch.FloatTensor | None = None) -> None:
"""Initialize step index at the beginning of sampling."""
if self._begin_index is None:
self._step_index = 0
else:
self._step_index = self._begin_index
def scale_model_input(
self,
sample: torch.Tensor,
timestep: int | None = None,
) -> torch.Tensor:
"""
rCM doesn't scale model input, returns sample as-is.
"""
return sample
def scale_noise(
self,
sample: torch.FloatTensor,
timestep: torch.FloatTensor | None = None,
noise: torch.FloatTensor | None = None,
) -> torch.FloatTensor:
"""
Scale initial noise for rCM sampling.
In rCM, initial noise is scaled by the first timestep (raw sigma):
x_0 = noise * t_steps[0]
Args:
sample: Not used (for API compatibility)
timestep: Not used (for API compatibility)
noise: The noise tensor to scale
Returns:
Scaled noise tensor ready for sampling
"""
if noise is None:
raise ValueError("noise must be provided for rCM scale_noise")
# Use raw sigma (not scaled timestep) for initial noise scaling
t_initial = self.sigmas[0]
return noise.to(torch.float64) * t_initial
def step(
self,
model_output: torch.FloatTensor,
timestep: int | torch.Tensor,
sample: torch.FloatTensor,
generator: torch.Generator | None = None,
return_dict: bool = True,
) -> RCMSchedulerOutput | tuple[torch.FloatTensor, ...]:
"""
Predict the sample from the previous timestep using rCM update rule.
The rCM update formula is:
x_{t+1} = (1 - t_next) * (x_t - t_cur * v_pred) + t_next * noise
Args:
model_output (`torch.FloatTensor`):
The velocity prediction from the model (v_pred).
timestep (`int` or `torch.Tensor`):
The current timestep index (not the actual timestep value).
For rCM, this should be the index into self.timesteps.
sample (`torch.FloatTensor`):
Current sample x_t.
generator (`torch.Generator`, *optional*):
Random number generator for noise.
return_dict (`bool`):
Whether to return RCMSchedulerOutput or tuple.
Returns:
`RCMSchedulerOutput` or `tuple`:
The denoised sample for the next step.
"""
if self._step_index is None:
self._init_step_index()
assert self._step_index is not None
# Get current and next sigma values (raw, unscaled) for rCM formula
# Note: self.timesteps is scaled by 1000 for model input,
# but we need raw values for the step formula
t_cur = self.sigmas[self._step_index]
# On the final step, t_next should be 0 (fully denoised)
if self._step_index + 1 < len(self.sigmas):
t_next = self.sigmas[self._step_index + 1]
else:
t_next = torch.tensor(0.0, device=sample.device, dtype=torch.float64)
# Ensure we're working in float64 for precision
sample = sample.to(torch.float64)
model_output = model_output.to(torch.float64)
# rCM update: x = (1 - t_next) * (x - t_cur * v_pred) + t_next * noise
x_denoised = sample - t_cur * model_output
# Generate noise for SDE sampling
if isinstance(generator, list):
generator = generator[0]
noise = torch.randn(
sample.shape,
dtype=torch.float32,
device="cpu",
generator=generator,
).to(sample.device).to(torch.float64)
prev_sample = (1 - t_next) * x_denoised + t_next * noise
# Increment step counter
self._step_index += 1
# Cast back to model output dtype
prev_sample = prev_sample.to(model_output.dtype)
if not return_dict:
return (prev_sample,)
return RCMSchedulerOutput(prev_sample=prev_sample)
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
"""
Add noise to samples (forward diffusion process).
Not typically used for rCM inference, but provided for API compatibility.
"""
raise NotImplementedError(
"add_noise is not implemented for RCMScheduler. "
"Use scale_noise for initializing the sampling process."
)
def __len__(self) -> int:
return self.config.num_train_timesteps
-47
View File
@@ -1252,52 +1252,6 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
dec = dec[:, :, start_frame_idx:]
return dec
def get_streaming_cache(self) -> list[torch.Tensor | None]:
def _count_conv3d(model) -> int:
count = 0
for m in model.modules():
if isinstance(m, WanCausalConv3d):
count += 1
return count
conv_num = _count_conv3d(self.decoder)
return [None] * conv_num
def streaming_decode(
self,
z: torch.Tensor,
cache: list[torch.Tensor | None],
is_first_chunk: bool = False,
) -> tuple[torch.Tensor, list[torch.Tensor | None]]:
"""
Args:
z (`torch.Tensor`): Latent tensor of shape [B, C, T, H, W].
cache (`list[torch.Tensor | None]`): The VAE cache.
is_first_chunk (`bool`): Whether this is the first chunk in the sequence.
Returns:
A tuple of (decoded_frames, updated_cache).
"""
iter_ = z.shape[2]
x = self.post_quant_conv(z)
with forward_context(feat_cache_arg=cache, feat_idx_arg=0):
outputs = []
for i in range(iter_):
feat_idx.set(0)
first_chunk.set(is_first_chunk and i == 0)
decoded_chunk = self.decoder(x[:, :, i:i + 1, :, :])
outputs.append(decoded_chunk)
out = torch.cat(outputs, dim=2)
if self.config.patch_size is not None:
out = unpatchify(out, patch_size=self.config.patch_size)
out = out.float()
out = torch.clamp(out, min=-1.0, max=1.0)
return out, cache
def forward(
self,
sample: torch.Tensor,
@@ -1318,4 +1272,3 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
z = posterior.mode()
dec = self.decode(z)
return dec
@@ -2,9 +2,8 @@
"""Matrix-Game causal DMD pipeline implementation."""
from fastvideo.fastvideo_args import FastVideoArgs
import torch
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, ForwardBatch, LoRAPipeline
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
InputValidationStage,
@@ -70,61 +69,5 @@ class MatrixGameCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
logger.info(
"MatrixGameCausalDMDPipeline initialized with action support")
@torch.no_grad()
def streaming_reset(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs):
if not self.post_init_called:
self.post_init()
# 1. Run Pre-processing stages
stages_to_run = [
"input_validation_stage", "prompt_encoding_stage",
"image_encoding_stage", "conditioning_stage",
"latent_preparation_stage", "image_latent_preparation_stage"
]
for stage_name in stages_to_run:
if stage_name in self._stage_name_mapping:
batch = self._stage_name_mapping[stage_name].forward(
batch, fastvideo_args)
# 2. Reset Denoising Stage
denoiser = self._stage_name_mapping["denoising_stage"]
denoiser.streaming_reset(batch, fastvideo_args)
# 3. Initialize VAE cache
self._vae_cache = None
def streaming_step(self, keyboard_action, mouse_action) -> ForwardBatch:
denoiser = self._stage_name_mapping["denoising_stage"]
ctx = denoiser._streaming_ctx
assert ctx is not None, "streaming_ctx must be set"
start_idx = ctx.start_index
batch = denoiser.streaming_step(keyboard_action, mouse_action)
end_idx = ctx.start_index
# Decode only the new generated block
if end_idx > start_idx:
current_latents = batch.latents[:, :, start_idx:end_idx, :, :]
args = ctx.fastvideo_args
decoder = self._stage_name_mapping["decoding_stage"]
decoded_frames, self._vae_cache = decoder.streaming_decode(
current_latents,
args,
cache=self._vae_cache,
is_first_chunk=(start_idx == 0))
batch.output = decoded_frames
else:
batch.output = None
return batch
def streaming_clear(self) -> None:
denoiser = self._stage_name_mapping.get("denoising_stage")
if denoiser is not None and hasattr(denoiser, "streaming_clear"):
denoiser.streaming_clear()
self._vae_cache = None
EntryClass = [MatrixGameCausalDMDPipeline]
@@ -1,81 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
TurboDiffusion Video Pipeline Implementation.
This module contains an implementation of the TurboDiffusion video diffusion pipeline
for 1-4 step video generation using rCM (recurrent Consistency Model) sampling
with SLA (Sparse-Linear Attention).
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_rcm import RCMScheduler
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class TurboDiffusionPipeline(LoRAPipeline, ComposedPipelineBase):
"""
TurboDiffusion video pipeline for 1-4 step generation.
Uses RCM scheduler and SLA attention for fast, high-quality video generation.
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Use RCM scheduler for TurboDiffusion
logger.info("Initializing RCM scheduler for TurboDiffusion")
self.modules["scheduler"] = RCMScheduler(sigma_max=80.0)
# Store checkpoint path for later loading
self._turbodiffusion_checkpoint = getattr(fastvideo_args,
'turbodiffusion_checkpoint',
None)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae"),
pipeline=self))
EntryClass = TurboDiffusionPipeline
+27 -2
View File
@@ -352,8 +352,8 @@ class ComposedPipelineBase(ABC):
else:
load_module_name = module_name
component_model_path = os.path.join(self.model_path,
load_module_name)
component_model_path = self._resolve_component_model_path(
fastvideo_args, module_name, load_module_name)
module = PipelineComponentLoader.load_module(
module_name=load_module_name,
component_model_path=component_model_path,
@@ -376,6 +376,31 @@ class ComposedPipelineBase(ABC):
return modules
def _resolve_component_model_path(
self,
fastvideo_args: FastVideoArgs,
module_name: str,
load_module_name: str,
) -> str:
"""Resolve the on-disk path for a given module, respecting overrides."""
override_repo, override_subpath = fastvideo_args.get_component_override(
module_name)
if override_repo is not None:
override_root = maybe_download_model(override_repo)
module_subpath = override_subpath or load_module_name
override_path = os.path.join(override_root, module_subpath)
if not os.path.exists(override_path):
raise FileNotFoundError(
f"Override path for {module_name} not found: {override_path}"
)
logger.info("Using override for %s from %s", module_name,
override_path)
return override_path
return os.path.join(self.model_path, load_module_name)
def add_stage(self, stage_name: str, stage: PipelineStage):
assert self.modules is not None, "No modules are registered"
self._stages.append(stage)
-1
View File
@@ -23,7 +23,6 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanImageToVideoPipeline": "wan",
"WanVideoToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"TurboDiffusionPipeline": "turbodiffusion",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
+26 -81
View File
@@ -50,7 +50,32 @@ class DecodingStage(PipelineStage):
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
return result
def _denormalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
@torch.no_grad()
def decode(self, latents: torch.Tensor,
fastvideo_args: FastVideoArgs) -> torch.Tensor:
"""
Decode latent representations into pixel space using VAE.
Args:
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
fastvideo_args: Configuration containing:
- disable_autocast: Whether to disable automatic mixed precision (default: False)
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
Returns:
Decoded video tensor with shape (batch, channels, frames, height, width),
normalized to [0, 1] range and moved to CPU as float32
"""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# denormalization for MatrixGame VAE
# z = z * std + mean during decode
if (hasattr(self.vae.config, 'latents_mean')
@@ -84,35 +109,6 @@ class DecodingStage(PipelineStage):
latents.dtype)
else:
latents += self.vae.shift_factor
return latents
@torch.no_grad()
def decode(self, latents: torch.Tensor,
fastvideo_args: FastVideoArgs) -> torch.Tensor:
"""
Decode latent representations into pixel space using VAE.
Args:
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
fastvideo_args: Configuration containing:
- disable_autocast: Whether to disable automatic mixed precision (default: False)
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
Returns:
Decoded video tensor with shape (batch, channels, frames, height, width),
normalized to [0, 1] range and moved to CPU as float32
"""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
latents = self._denormalize_latents(latents)
# Decode latents
with torch.autocast(device_type="cuda",
@@ -130,57 +126,6 @@ class DecodingStage(PipelineStage):
image = (image / 2 + 0.5).clamp(0, 1)
return image
@torch.no_grad()
def streaming_decode(
self,
latents: torch.Tensor,
fastvideo_args: FastVideoArgs,
cache: list[torch.Tensor | None] | None = None,
is_first_chunk: bool = False,
) -> tuple[torch.Tensor, list[torch.Tensor | None]]:
"""
Decode latent representations into pixel space using VAE with streaming cache.
Args:
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
fastvideo_args: Configuration object.
cache: VAE cache from previous call, or None to initialize a new cache.
is_first_chunk: Whether this is the first chunk.
Returns:
A tuple of (decoded_frames, updated_cache).
"""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
latents = self._denormalize_latents(latents)
# Initialize cache if needed
if cache is None:
cache = self.vae.get_streaming_cache()
# Decode latents with streaming
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image, cache = self.vae.streaming_decode(latents, cache,
is_first_chunk)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
assert cache is not None, "cache should not be None after streaming_decode"
return image, cache
@torch.no_grad()
def forward(
self,
+5 -32
View File
@@ -253,37 +253,19 @@ class DenoisingStage(PipelineStage):
if boundary_timestep is None or t >= boundary_timestep:
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and self.transformer_2 is not None and next(
self.transformer_2.parameters()).device.type
== 'cuda'):
self.transformer_2.to('cpu')
current_model = self.transformer
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and not fastvideo_args.use_fsdp_inference
and current_model is not None):
transformer_device = next(
current_model.parameters()).device.type
if transformer_device == 'cpu':
current_model.to(get_local_torch_device())
current_guidance_scale = batch.guidance_scale
else:
# low-noise stage in wan2.2
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and next(self.transformer.parameters()).device.type
== 'cuda'):
if fastvideo_args.dit_cpu_offload and next(
self.transformer.parameters(
)).device.type == 'cuda':
self.transformer.to('cpu')
current_model = self.transformer_2
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and not fastvideo_args.use_fsdp_inference
and current_model is not None):
transformer_2_device = next(
current_model.parameters()).device.type
if transformer_2_device == 'cpu':
current_model.to(get_local_torch_device())
current_guidance_scale = batch.guidance_scale_2
assert current_model is not None, "current_model is None"
@@ -460,6 +442,7 @@ class DenoisingStage(PipelineStage):
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
@@ -476,16 +459,6 @@ class DenoisingStage(PipelineStage):
# Update batch with final latents
batch.latents = latents
if fastvideo_args.dit_layerwise_offload:
mgr = getattr(self.transformer, "_layerwise_offload_manager", None)
if mgr is not None and getattr(mgr, "enabled", False):
mgr.release_all()
if self.transformer_2 is not None:
mgr2 = getattr(self.transformer_2, "_layerwise_offload_manager",
None)
if mgr2 is not None and getattr(mgr2, "enabled", False):
mgr2.release_all()
# Save STA mask search results if needed
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
self.save_sta_search_results(batch)
@@ -1212,4 +1185,4 @@ class DmdDenoisingStage(DenoisingStage):
# Update batch with final latents
batch.latents = latents
return batch
return batch
+205 -526
View File
@@ -1,6 +1,4 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
import torch # type: ignore
@@ -34,45 +32,6 @@ except ImportError:
logger = init_logger(__name__)
@dataclass
class BlockProcessingContext:
"""Dataclass contains for block processing."""
batch: ForwardBatch
block_idx: int
start_index: int
kv_cache1: list[dict[Any, Any]]
kv_cache2: list[dict[Any, Any]] | None
kv_cache_mouse: list[dict[Any, Any]] | None
kv_cache_keyboard: list[dict[Any, Any]] | None
crossattn_cache: list[dict[Any, Any]]
timesteps: torch.Tensor
block_sizes: list[int]
noise_pool: list[torch.Tensor] | None
fastvideo_args: FastVideoArgs
target_dtype: torch.dtype
autocast_enabled: bool
boundary_timestep: float | None
high_noise_timesteps: torch.Tensor | None
context_noise: float
image_kwargs: dict[str, Any]
pos_cond_kwargs: dict[str, Any]
def get_kv_cache(self, timestep_val: float) -> list[dict[Any, Any]]:
if self.boundary_timestep is not None:
if timestep_val >= self.boundary_timestep:
return self.kv_cache1
else:
assert self.kv_cache2 is not None, "kv_cache2 is not initialized"
return self.kv_cache2
return self.kv_cache1
class MatrixGameCausalDenoisingStage(DenoisingStage):
def __init__(self,
@@ -121,9 +80,6 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
self.action_config = getattr(self.transformer, 'action_config', {})
self.use_action_module = len(self.action_config) > 0
self._streaming_initialized: bool = False
self._streaming_ctx: BlockProcessingContext | None = None
def forward(
self,
batch: ForwardBatch,
@@ -138,6 +94,8 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
independent_first_frame = getattr(self.transformer,
'independent_first_frame', False)
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
@@ -194,12 +152,23 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=target_dtype,
device=latents.device)
def _get_kv_cache(timestep_val: float) -> list[dict]:
if boundary_timestep is not None:
if timestep_val >= boundary_timestep:
return kv_cache1
else:
assert kv_cache2 is not None, "kv_cache2 is not initialized"
return kv_cache2
return kv_cache1
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257, # 1 CLS + 256 patch tokens
dtype=target_dtype,
device=latents.device)
pos_start_base = 0
if t % self.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for causal denoising"
@@ -215,67 +184,211 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
# The first frame information is already encoded in batch.image_latent (cond_concat)
# and will be used by the model via channel concatenation: torch.cat([x, cond_concat], dim=1)
ctx = BlockProcessingContext(
batch=batch,
block_idx=0,
start_index=0,
kv_cache1=kv_cache1,
kv_cache2=kv_cache2,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
timesteps=timesteps,
block_sizes=block_sizes,
noise_pool=None,
fastvideo_args=fastvideo_args,
target_dtype=target_dtype,
autocast_enabled=autocast_enabled,
boundary_timestep=boundary_timestep,
high_noise_timesteps=high_noise_timesteps,
context_noise=getattr(fastvideo_args.pipeline_config,
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
)
context_noise = getattr(fastvideo_args.pipeline_config, "context_noise",
0)
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for block_idx, current_num_frames in enumerate(block_sizes):
ctx.block_idx = block_idx
ctx.start_index = start_index
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
_video_raw_latent_shape = noise_latents_btchw.shape # noqa: F841
# NOTE: crossattn_cache should NOT be reset between blocks!
action_kwargs = self._prepare_action_kwargs(
batch, start_index, current_num_frames)
current_latents = self._process_single_block(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timesteps=timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
progress_bar=progress_bar,
)
for i, t_cur in enumerate(timesteps):
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
else:
current_model = self.transformer
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
if batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(target_dtype)
],
dim=2)
# t_expand needs to be [batch * frames] to match flattened pred_noise/noise_latents
t_expand = t_cur.repeat(latent_model_input.shape[0] *
current_num_frames)
if vsa_available and self.attn_backend == VideoSparseAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
raw_latent_shape=(current_num_frames, h, w),
patch_size=fastvideo_args.pipeline_config.
dit_config.patch_size,
STA_param=batch.STA_param,
VSA_sparsity=fastvideo_args.VSA_sparsity,
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Expand timestep to per-frame format [batch, num_frames] for causal model
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
model_kwargs = {
"kv_cache": _get_kv_cache(t_cur),
"crossattn_cache": crossattn_cache,
"current_start": (pos_start_base + start_index) *
self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module and current_model == self.transformer:
model_kwargs.update({
"kv_cache_mouse":
kv_cache_mouse,
"kv_cache_keyboard":
kv_cache_keyboard,
})
model_kwargs.update(action_kwargs)
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
**image_kwargs,
**pos_cond_kwargs,
**model_kwargs,
).permute(0, 2, 1, 3, 4)
if boundary_timestep is not None and t_cur >= boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1],
dtype=torch.long,
device=pred_video_btchw.device)
noise = torch.randn(
pred_video_btchw.shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else
batch.generator)).to(
pred_video_btchw.device)
noise_btchw = noise
if boundary_timestep is not None and high_noise_timesteps is not None and i < len(
high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_timestep,
torch.ones_like(next_timestep) *
boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif boundary_timestep is not None and high_noise_timesteps is not None and i == len(
high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(
0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(
0, 2, 1, 3, 4)
if progress_bar is not None:
progress_bar.update()
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Update KV caches with clean context
self._update_context_cache(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
context_noise=context_noise,
)
context_noise = getattr(fastvideo_args.pipeline_config,
"context_noise", 0)
# Expand context timestep to per-frame format [batch, num_frames] for causal model
t_context = torch.ones([latents.shape[0], current_num_frames],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
context_model_kwargs = {
"kv_cache": kv_cache1,
"crossattn_cache": crossattn_cache,
"current_start":
(pos_start_base + start_index) * self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module:
context_model_kwargs.update({
"kv_cache_mouse":
kv_cache_mouse,
"kv_cache_keyboard":
kv_cache_keyboard,
})
context_model_kwargs.update(action_kwargs)
if boundary_timestep is not None and self.transformer_2 is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_context,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
self.transformer(
context_bcthw,
prompt_embeds,
t_context,
**image_kwargs,
**pos_cond_kwargs,
**context_model_kwargs,
)
start_index += current_num_frames
@@ -429,440 +542,6 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
})
return crossattn_cache
def _process_single_block(
self,
current_latents: torch.Tensor,
batch: ForwardBatch,
start_index: int,
current_num_frames: int,
timesteps: torch.Tensor,
ctx: BlockProcessingContext,
action_kwargs: dict[str, Any],
noise_generator: Callable[[tuple, torch.dtype, int], torch.Tensor]
| None = None,
progress_bar: Any | None = None,
) -> torch.Tensor:
prompt_embeds = batch.prompt_embeds
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
for i, t_cur in enumerate(timesteps):
if ctx.boundary_timestep is not None and t_cur < ctx.boundary_timestep:
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
else:
current_model = self.transformer
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(ctx.target_dtype)
independent_first_frame = getattr(self.transformer,
'independent_first_frame', False)
if batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(ctx.target_dtype)
],
dim=2)
# t_expand needs to be [batch * frames] to match flattened pred_noise/noise_latents
t_expand = t_cur.repeat(latent_model_input.shape[0] *
current_num_frames)
# Build attention metadata if VSA is available
if vsa_available and self.attn_backend == VideoSparseAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
h, w = current_latents.shape[-2:]
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
raw_latent_shape=(current_num_frames, h, w),
patch_size=ctx.fastvideo_args.pipeline_config.
dit_config.patch_size,
STA_param=batch.STA_param,
VSA_sparsity=ctx.fastvideo_args.VSA_sparsity,
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
enabled=ctx.autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Expand timestep to per-frame format [batch, num_frames] for causal model
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
model_kwargs = {
"kv_cache": ctx.get_kv_cache(t_cur),
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module and current_model == self.transformer:
model_kwargs.update({
"kv_cache_mouse":
ctx.kv_cache_mouse,
"kv_cache_keyboard":
ctx.kv_cache_keyboard,
})
model_kwargs.update(action_kwargs)
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**model_kwargs,
).permute(0, 2, 1, 3, 4)
if ctx.boundary_timestep is not None and t_cur >= ctx.boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
ctx.boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video_btchw.device)
# Use custom noise generator if provided (for streaming), else generate
if noise_generator is not None:
noise = noise_generator(pred_video_btchw.shape,
pred_video_btchw.dtype, i)
else:
noise = torch.randn(
pred_video_btchw.shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else batch.generator)).to(
pred_video_btchw.device)
noise_btchw = noise
if ctx.boundary_timestep is not None and ctx.high_noise_timesteps is not None and i < len(
ctx.high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_timestep,
torch.ones_like(next_timestep) *
ctx.boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif ctx.boundary_timestep is not None and ctx.high_noise_timesteps is not None and i == len(
ctx.high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0,
1), noise_btchw.flatten(0, 1),
next_timestep).unflatten(0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(0, 2, 1, 3, 4)
if progress_bar is not None:
progress_bar.update()
return current_latents
def _update_context_cache(
self,
current_latents: torch.Tensor,
batch: ForwardBatch,
start_index: int,
current_num_frames: int,
ctx: BlockProcessingContext,
action_kwargs: dict[str, Any],
context_noise: float,
) -> None:
prompt_embeds = batch.prompt_embeds
latents_device = current_latents.device
# Expand context timestep to per-frame format [batch, num_frames] for causal model
t_context = torch.ones([current_latents.shape[0], current_num_frames],
device=latents_device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(ctx.target_dtype)
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
enabled=ctx.autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
context_model_kwargs = {
"kv_cache": ctx.kv_cache1,
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module:
context_model_kwargs.update({
"kv_cache_mouse":
ctx.kv_cache_mouse,
"kv_cache_keyboard":
ctx.kv_cache_keyboard,
})
context_model_kwargs.update(action_kwargs)
if ctx.boundary_timestep is not None and self.transformer_2 is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_context,
kv_cache=ctx.kv_cache2,
crossattn_cache=ctx.crossattn_cache,
current_start=start_index * self.frame_seq_length,
start_frame=start_index,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
)
self.transformer(
context_bcthw,
prompt_embeds,
t_context,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**context_model_kwargs,
)
def streaming_reset(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_size = self.transformer.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
'boundary_ratio', None)
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# directly set the kwarg.
image_kwargs = {"encoder_hidden_states_image": image_embeds}
pos_cond_kwargs: dict[str, Any] = {}
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize caches
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache_mouse = None
kv_cache_keyboard = None
if self.use_action_module:
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257, # 1 CLS + 256 patch tokens
dtype=target_dtype,
device=latents.device)
# Calculate block sizes
if t % self.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for causal denoising"
)
num_blocks = t // self.num_frame_per_block
block_sizes = [self.num_frame_per_block] * num_blocks
if boundary_timestep is not None:
block_sizes[0] = 1
# Pre-allocate noise pool
num_denoising_steps = len(timesteps)
noise_shape = (b, self.num_frame_per_block, c, h, w)
noise_pool = [
torch.randn(
noise_shape,
dtype=target_dtype,
device=latents.device,
) for _ in range(num_denoising_steps - 1)
]
# Create and store context
self._streaming_ctx = BlockProcessingContext(
batch=batch,
block_idx=0,
start_index=0,
kv_cache1=kv_cache1,
kv_cache2=kv_cache2,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
timesteps=timesteps,
block_sizes=block_sizes,
noise_pool=noise_pool,
fastvideo_args=fastvideo_args,
target_dtype=target_dtype,
autocast_enabled=autocast_enabled,
boundary_timestep=boundary_timestep,
high_noise_timesteps=high_noise_timesteps,
context_noise=getattr(fastvideo_args.pipeline_config,
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
)
self._streaming_initialized = True
return batch
def streaming_step(
self,
keyboard_action: torch.Tensor | None = None,
mouse_action: torch.Tensor | None = None) -> ForwardBatch:
if not self._streaming_initialized or self._streaming_ctx is None:
raise RuntimeError(
"Streaming not initialized! Call streaming_reset first.")
ctx = self._streaming_ctx
if ctx.block_idx >= len(ctx.block_sizes):
return ctx.batch
batch = ctx.batch
latents = batch.latents
assert latents is not None, "latents must be set in batch"
current_num_frames = ctx.block_sizes[ctx.block_idx]
start_index = ctx.start_index
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# Update batch with new actions for this block
if keyboard_action is not None or mouse_action is not None:
vae_ratio = 4
start_frame = 0 if start_index == 0 else 1 + vae_ratio * (
start_index - 1)
if keyboard_action is not None:
n = keyboard_action.shape[1]
batch.keyboard_cond[:, start_frame:start_frame +
n] = keyboard_action.to(
batch.keyboard_cond.device)
if mouse_action is not None:
n = mouse_action.shape[1]
batch.mouse_cond[:, start_frame:start_frame +
n] = mouse_action.to(batch.mouse_cond.device)
action_kwargs = self._prepare_action_kwargs(batch, start_index,
current_num_frames)
# Create noise generator that uses pre-allocated noise pool
def streaming_noise_generator(shape: tuple, dtype: torch.dtype,
step_idx: int) -> torch.Tensor:
if ctx.noise_pool is not None and step_idx < len(ctx.noise_pool):
return ctx.noise_pool[step_idx][:, :shape[1], :, :, :].to(
latents.device)
else:
# Fallback to dynamic allocation if pool not available
return torch.randn(
shape,
dtype=dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else batch.generator)).to(
latents.device)
current_latents = self._process_single_block(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timesteps=ctx.timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
noise_generator=streaming_noise_generator,
)
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Update KV caches with clean context
self._update_context_cache(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
context_noise=ctx.context_noise,
)
# Advance streaming state
ctx.start_index += current_num_frames
ctx.block_idx += 1
return batch
def streaming_clear(self) -> None:
self._streaming_initialized = False
self._streaming_ctx = None
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
-28
View File
@@ -199,34 +199,6 @@ class CudaPlatformBase(Platform):
"Failed to import Video MoBA Attention backend: %s", str(e))
raise ImportError(
"Video MoBA Attention backend is not installed. ") from e
elif selected_backend == AttentionBackendEnum.SLA_ATTN:
try:
from fastvideo.attention.backends.sla import ( # noqa: F401
SLAAttentionBackend)
logger.info("Using SLA (Sparse-Linear Attention) backend.")
return "fastvideo.attention.backends.sla.SLAAttentionBackend"
except ImportError as e:
logger.error("Failed to import SLA Attention backend: %s",
str(e))
raise ImportError(
"SLA Attention backend is not available. ") from e
elif selected_backend == AttentionBackendEnum.SAGE_SLA_ATTN:
try:
from fastvideo.attention.backends.sla import ( # noqa: F401
SageSLAAttentionBackend)
logger.info(
"Using SageSLA (Quantized Sparse-Linear Attention) backend."
)
return "fastvideo.attention.backends.sla.SageSLAAttentionBackend"
except ImportError as e:
logger.error("Failed to import SageSLA Attention backend: %s",
str(e))
raise ImportError(
"SageSLA Attention backend requires spas_sage_attn. "
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git"
) from e
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.attention.backends.sdpa.SDPABackend"
-2
View File
@@ -18,8 +18,6 @@ class AttentionBackendEnum(enum.Enum):
SAGE_ATTN_THREE = enum.auto()
VIDEO_SPARSE_ATTN = enum.auto()
VMOBA_ATTN = enum.auto()
SLA_ATTN = enum.auto()
SAGE_SLA_ATTN = enum.auto()
NO_ATTENTION = enum.auto()
+2 -2
View File
@@ -82,7 +82,7 @@ def run_transformer_tests():
@app.function(
gpu="L40S:4",
image=image,
timeout=3600,
timeout=2700,
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol}
)
@@ -122,7 +122,7 @@ def run_kernel_tests():
def run_inference_tests_vmoba():
run_test('python fastvideo/tests/inference/vmoba/test_vmoba_inference.py')
@app.function(gpu="L40S:1", image=image, timeout=1200)
@app.function(gpu="L40S:1", image=image, timeout=3600)
def run_inference_lora_tests():
run_test("pytest ./fastvideo/tests/inference/lora/test_lora_inference_similarity.py -vs")
@@ -1,94 +0,0 @@
import unittest
import torch
import torch.nn as nn
from torch.testing import assert_close
from fastvideo.layers.quantization.absmax_fp8 import (
AbsMaxFP8LinearMethod,
AbsMaxFP8MergedParameter,
AbsMaxFP8Parameter,
)
from fastvideo.models.utils import set_weight_attrs
class TestAbsMaxFP8LinearMethod(unittest.TestCase):
def test_convert_scale_none(self):
method = AbsMaxFP8LinearMethod()
scale = method._convert_scale(None)
self.assertIsInstance(scale, AbsMaxFP8Parameter)
self.assertEqual(scale.dtype, torch.float32)
assert_close(scale, torch.tensor([1.0], dtype=torch.float32))
def test_convert_scale_scalar(self):
method = AbsMaxFP8LinearMethod()
scale = method._convert_scale(2.5)
self.assertIsInstance(scale, AbsMaxFP8Parameter)
self.assertEqual(scale.dtype, torch.float32)
assert_close(scale, torch.tensor([2.5], dtype=torch.float32))
def test_convert_scale_rejects_non_float32(self):
method = AbsMaxFP8LinearMethod()
scale = torch.tensor([1.0], dtype=torch.float16)
with self.assertRaisesRegex(NotImplementedError, "float32"):
method._convert_scale(scale)
def test_create_weights_rejects_invalid_dtype(self):
method = AbsMaxFP8LinearMethod()
layer = nn.Module()
with self.assertRaisesRegex(AssertionError, "only supports"):
method.create_weights(
layer=layer,
input_size_per_partition=2,
output_partition_sizes=[3],
input_size=2,
output_size=3,
params_dtype=torch.float32,
)
def test_absmax_fp8_parameter_weight_loader(self):
param = AbsMaxFP8Parameter(torch.zeros(1), requires_grad=False)
param.weight_loader(param, torch.tensor(3.0))
assert_close(param, torch.tensor([3.0]))
def test_absmax_fp8_merged_parameter_weight_loader(self):
method = AbsMaxFP8LinearMethod()
output_partition_sizes = [2, 3, 4]
param = method._merged_placeholder(output_partition_sizes)
self.assertIsInstance(param, AbsMaxFP8MergedParameter)
param.weight_loader(param, torch.tensor(7.0), share_id="k")
expected = torch.ones(sum(output_partition_sizes), dtype=torch.float32)
expected[2:5] = 7.0
assert_close(param, expected)
def test_absmax_fp8_merged_parameter_rejects_invalid_share_id(self):
method = AbsMaxFP8LinearMethod()
param = method._merged_placeholder([2, 2, 2])
with self.assertRaisesRegex(ValueError, "requires share_id"):
param.weight_loader(param, torch.tensor(1.0), share_id="bad")
def test_apply_matches_linear(self):
method = AbsMaxFP8LinearMethod()
layer = nn.Module()
method.create_weights(
layer=layer,
input_size_per_partition=3,
output_partition_sizes=[2],
input_size=3,
output_size=2,
params_dtype=torch.float16,
)
weight_fp16 = torch.tensor(
[[1.0, -2.0, 3.0], [4.0, 0.5, -1.5]], dtype=torch.float16
)
layer.weight.data = weight_fp16.to(dtype=torch.float8_e4m3fn)
layer.scale_weight.data = torch.tensor([2.0, 3.0], dtype=torch.float32)
layer.scale_input.data = torch.tensor([4.0], dtype=torch.float32)
x = torch.tensor([[1.0, 2.0, -1.0]], dtype=torch.float16)
expected = torch.nn.functional.linear(
x * layer.scale_input.data.to(dtype=torch.float16),
weight_fp16
* layer.scale_weight.data.to(dtype=torch.float16).unsqueeze(1),
).to(dtype=torch.float16)
output = method.apply(layer, x, bias=None)
assert_close(output, expected)
@@ -1,11 +0,0 @@
{
"mean_ssim": 0.9917645101194028,
"min_ssim": 0.9908460974693298,
"max_ssim": 0.9925968050956726,
"reference_video": "/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
"generated_video": "/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
"parameters": {
"num_inference_steps": 4,
"prompt": "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
}
}
@@ -1,159 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
SSIM-based similarity test for TurboDiffusion inference.
TurboDiffusion uses the SLA (Sparse-Linear Attention) backend and RCM scheduler
for 1-4 step video generation.
"""
import os
import torch
import pytest
from fastvideo import VideoGenerator
from fastvideo.logger import init_logger
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
from fastvideo.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = '_reference_videos'
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}, using L40S references")
raise ValueError(f"Unsupported device for ssim tests: {device_name}")
# TurboDiffusion parameters (1-4 step generation with RCM scheduler + SLA attention)
TURBODIFFUSION_PARAMS = {
"num_gpus": 2,
"model_path": "loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
"height": 480,
"width": 832,
"num_frames": 81,
"num_inference_steps": 4, # TurboDiffusion uses 1-4 steps
"guidance_scale": 1.0, # No CFG for TurboDiffusion
"seed": 42,
"sp_size": 2,
"tp_size": 1,
"fps": 24,
}
TURBODIFFUSION_MODEL_TO_PARAMS = {
"TurboWan2.1-T2V-1.3B-Diffusers": TURBODIFFUSION_PARAMS,
}
TURBODIFFUSION_TEST_PROMPTS = [
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
]
@pytest.mark.parametrize("prompt", TURBODIFFUSION_TEST_PROMPTS)
@pytest.mark.parametrize("model_id", list(TURBODIFFUSION_MODEL_TO_PARAMS.keys()))
def test_turbodiffusion_inference_similarity(prompt, model_id):
"""
Test that runs TurboDiffusion inference with SLA attention and RCM scheduler,
then compares the output to reference videos using SSIM.
"""
# TurboDiffusion requires SLA attention backend
ATTENTION_BACKEND = "SLA_ATTN"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
BASE_PARAMS = TURBODIFFUSION_MODEL_TO_PARAMS[model_id]
num_inference_steps = BASE_PARAMS["num_inference_steps"]
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
"override_pipeline_cls_name": "TurboDiffusionPipeline",
}
generation_kwargs = {
"num_inference_steps": num_inference_steps,
"output_path": output_dir,
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"guidance_scale": BASE_PARAMS["guidance_scale"],
"seed": BASE_PARAMS["seed"],
"fps": BASE_PARAMS["fps"],
}
generator = VideoGenerator.from_pretrained(
model_path=BASE_PARAMS["model_path"],
**init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
if isinstance(generator.executor, MultiprocExecutor):
generator.executor.shutdown()
assert os.path.exists(output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
# Find the matching reference video based on the prompt
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith('.mp4') and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
logger.error(
f"Reference video not found for prompt: {prompt} with backend: {ATTENTION_BACKEND}"
)
raise FileNotFoundError(f"Reference video missing")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name)
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
)
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
logger.info(f"Writing SSIM results to directory: {output_dir}")
success = write_ssim_results(
output_dir, ssim_values, reference_video_path,
generated_video_path, num_inference_steps, prompt
)
if not success:
logger.error("Failed to write SSIM results to file")
# TurboDiffusion uses fewer steps, may have slightly lower SSIM
min_acceptable_ssim = 0.90
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}"
)
@@ -121,4 +121,4 @@ def test_wan_transformer():
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
+1 -5
View File
@@ -21,12 +21,8 @@ class ModelWrapper(torch.distributed.checkpoint.stateful.Stateful):
state_dict = get_model_state_dict(
self.model) # type: ignore[no-any-return]
# filter out non-trainable parameters
# Note: activation checkpointing adds ._checkpoint_wrapped_module. prefix
# to parameter names, but get_model_state_dict returns normalized keys.
# We need to normalize the parameter names for proper matching.
param_requires_grad = set([
k.replace("._checkpoint_wrapped_module.", ".")
for k, v in dict(self.model.named_parameters()).items()
k for k, v in dict(self.model.named_parameters()).items()
if v.requires_grad
])
state_dict = {
+56 -110
View File
@@ -116,44 +116,6 @@ def get_sigmas(noise_scheduler,
return sigma
def _is_lora_training(transformer: torch.nn.Module) -> bool:
"""Best-effort check for LoRA fine-tuning.
We treat a run as LoRA training when the model already contains LoRA
layers with trainable parameters. This avoids touching consolidated
full-model weights which are not updated during LoRA-only training.
"""
state_keys = transformer.state_dict().keys()
return any("lora_" in name for name in state_keys)
def _save_lora_adapter(cpu_state: dict[str, Any], adapter_path: str) -> None:
"""Persist only the LoRA adapter weights for LoRA training runs.
Args:
cpu_state: Pre-gathered state dict from gather_state_dict_on_cpu_rank0.
Using the pre-gathered state dict avoids issues with invalid
tensor storage pointers in FSDP-wrapped models.
adapter_path: Path to save the LoRA adapter safetensors file.
"""
lora_state = {
name: tensor.detach().clone().cpu().contiguous()
for name, tensor in cpu_state.items() if "lora_" in name
}
if len(lora_state) == 0:
logger.warning(
"LoRA training detected but no LoRA parameters found to save.")
return
os.makedirs(os.path.dirname(adapter_path), exist_ok=True)
save_file(lora_state, adapter_path)
logger.info("Saved LoRA adapter with %d tensors to %s", len(lora_state),
adapter_path)
def save_checkpoint(transformer,
rank,
output_dir,
@@ -199,43 +161,35 @@ def save_checkpoint(transformer,
cpu_state = gather_state_dict_on_cpu_rank0(transformer, device=None)
if rank == 0:
if _is_lora_training(transformer):
adapter_path = os.path.join(save_dir, "lora_adapter.safetensors")
_save_lora_adapter(cpu_state, adapter_path)
logger.info(
"LoRA training detected; saved adapter instead of consolidated weights."
)
else:
# Save model weights (consolidated)
transformer_save_dir = os.path.join(save_dir, "transformer")
os.makedirs(transformer_save_dir, exist_ok=True)
weight_path = os.path.join(transformer_save_dir,
"diffusion_pytorch_model.safetensors")
logger.info("rank: %s, saving consolidated checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
# Save model weights (consolidated)
transformer_save_dir = os.path.join(save_dir, "transformer")
os.makedirs(transformer_save_dir, exist_ok=True)
weight_path = os.path.join(transformer_save_dir,
"diffusion_pytorch_model.safetensors")
logger.info("rank: %s, saving consolidated checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, transformer.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
# Convert training format to diffusers format and save
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, transformer.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
logger.info("rank: %s, consolidated checkpoint saved to %s",
rank,
weight_path,
local_main_process_only=False)
logger.info("rank: %s, consolidated checkpoint saved to %s",
rank,
weight_path,
local_main_process_only=False)
# Save model config
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(transformer_save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> checkpoint saved at step %s to %s", step,
weight_path)
# Save model config
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(transformer_save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
def save_distillation_checkpoint(
@@ -467,44 +421,37 @@ def save_distillation_checkpoint(
device=None)
if rank == 0:
if _is_lora_training(generator_transformer):
adapter_path = os.path.join(save_dir, "lora_adapter.safetensors")
_save_lora_adapter(cpu_state, adapter_path)
logger.info(
"LoRA training detected; saved adapter instead of consolidated generator weights."
)
else:
# Save generator model weights (consolidated) for inference
os.makedirs(inference_save_dir, exist_ok=True)
weight_path = os.path.join(inference_save_dir,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator inference checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
# Save generator model weights (consolidated) for inference
os.makedirs(inference_save_dir, exist_ok=True)
weight_path = os.path.join(inference_save_dir,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator inference checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, generator_transformer.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
# Convert training format to diffusers format and save
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, generator_transformer.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
logger.info(
"rank: %s, consolidated generator inference checkpoint saved to %s",
rank,
weight_path,
local_main_process_only=False)
logger.info(
"rank: %s, consolidated generator inference checkpoint saved to %s",
rank,
weight_path,
local_main_process_only=False)
# Save model config
config_dict = generator_transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(inference_save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> distillation checkpoint saved at step %s to %s",
step, weight_path)
# Save model config
config_dict = generator_transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(inference_save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
@@ -652,9 +599,8 @@ def load_distillation_checkpoint(
checkpoint_path)
return 0
# Extract step number from checkpoint path (normpath handles trailing slashes)
step = int(
os.path.basename(os.path.normpath(checkpoint_path)).split('-')[-1])
# Extract step number from checkpoint path
step = int(os.path.basename(checkpoint_path).split('-')[-1])
if rank == 0:
logger.info("Loading distillation checkpoint from step %s", step)
-14
View File
@@ -114,20 +114,6 @@ class Worker:
return {"status": "lora_adapter_unmerged"}
return {"status": "failed: pipeline is not a LoRAPipeline"}
def execute_streaming_reset(
self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
self.pipeline.streaming_reset(forward_batch, self.fastvideo_args)
return {"status": "reset_complete"}
def execute_streaming_step(self, keyboard_action: torch.Tensor,
mouse_action: torch.Tensor) -> ForwardBatch:
return self.pipeline.streaming_step(keyboard_action, mouse_action)
def execute_streaming_clear(self) -> dict[str, Any]:
self.pipeline.streaming_clear()
return {"status": "cleared"}
def merge_lora_weights(self) -> dict[str, Any]:
if isinstance(self.pipeline, LoRAPipeline):
self.pipeline.merge_lora_weights()
+2 -229
View File
@@ -1,17 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import asyncio
import atexit
import contextlib
from dataclasses import dataclass
from enum import Enum
import faulthandler
import multiprocessing as mp
from multiprocessing.connection import Connection
from multiprocessing.queues import Queue
import os
import queue
import signal
import time
from collections.abc import Callable
@@ -19,55 +13,19 @@ from multiprocessing.process import BaseProcess
from typing import Any, cast
import psutil
import torch
from fastvideo.distributed.parallel_state import get_dp_group, get_tp_group
import fastvideo.envs as envs
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.utils import (decorate_logs, get_distributed_init_method,
get_exception_traceback, get_loopback_ip,
get_mp_context, get_open_port,
kill_itself_when_parent_died, force_spawn)
from fastvideo.utils import decorate_logs, get_distributed_init_method, get_exception_traceback, get_loopback_ip, get_mp_context, get_open_port, kill_itself_when_parent_died, force_spawn
from fastvideo.worker.executor import Executor
from fastvideo.worker.worker_base import WorkerWrapperBase
logger = init_logger(__name__)
class StreamingTaskType(str, Enum):
"""
Enumeration for different streaming task types.
Inherits from str to allow string comparison for backward compatibility.
"""
RESET = "reset"
STEP = "step"
CLEAR = "clear"
EXIT = "exit"
@dataclass
class StreamingTask:
"""Task submitted to worker via input queue."""
task_type: StreamingTaskType
# For STEP tasks:
keyboard_action: torch.Tensor | None = None
mouse_action: torch.Tensor | None = None
# For RESET tasks:
batch: ForwardBatch | None = None
fastvideo_args: FastVideoArgs | None = None
@dataclass
class StreamingResult:
"""Result returned from worker via output queue."""
task_type: StreamingTaskType
output_batch: ForwardBatch | None = None
error: Exception | None = None
class MultiprocExecutor(Executor):
def _init_executor(self) -> None:
@@ -82,12 +40,6 @@ class MultiprocExecutor(Executor):
get_loopback_ip(), master_port)
logger.info("Use master port: %s", master_port)
# Create streaming queues BEFORE spawning workers
ctx = get_mp_context()
self._streaming_input_queue: Queue | None = ctx.Queue()
self._streaming_output_queue: Queue | None = ctx.Queue()
self._streaming_enabled = False
unready_workers: list[UnreadyWorkerProcHandle] = []
success = False
try:
@@ -98,8 +50,6 @@ class MultiprocExecutor(Executor):
local_rank=rank,
rank=rank,
distributed_init_method=distributed_init_method,
streaming_input_queue=self._streaming_input_queue,
streaming_output_queue=self._streaming_output_queue,
))
# Workers must be created before wait_for_ready to avoid
@@ -138,108 +88,6 @@ class MultiprocExecutor(Executor):
return result_batch
def execute_streaming_reset(
self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
responses = self.collective_rpc("execute_streaming_reset",
kwargs={
"forward_batch": forward_batch,
"fastvideo_args": fastvideo_args,
})
return responses[0]
def execute_streaming_step(self, keyboard_action: Any,
mouse_action: Any) -> ForwardBatch:
responses = self.collective_rpc("execute_streaming_step",
kwargs={
"keyboard_action": keyboard_action,
"mouse_action": mouse_action,
})
return responses[0]
async def execute_streaming_step_async(self, keyboard_action: Any,
mouse_action: Any) -> ForwardBatch:
responses = await self.collective_rpc_async("execute_streaming_step",
kwargs={
"keyboard_action":
keyboard_action,
"mouse_action":
mouse_action,
})
return responses[0]
def execute_streaming_clear(self) -> dict[str, Any]:
responses = self.collective_rpc("execute_streaming_clear")
return responses[0]
def enable_streaming(self) -> None:
if self._streaming_enabled:
return
self.collective_rpc("start_streaming_queue_loop")
self._streaming_enabled = True
def disable_streaming(self) -> None:
if not self._streaming_enabled:
return
if self._streaming_input_queue is not None:
self._streaming_input_queue.put(
StreamingTask(task_type=StreamingTaskType.EXIT))
self._streaming_enabled = False
self._streaming_input_queue = None
self._streaming_output_queue = None
def submit_reset(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> None:
if not self._streaming_enabled:
self.enable_streaming()
self._streaming_input_queue.put(
StreamingTask(
task_type=StreamingTaskType.RESET,
batch=forward_batch,
fastvideo_args=fastvideo_args,
))
def submit_step(self, keyboard_action: torch.Tensor | None,
mouse_action: torch.Tensor | None) -> None:
if not self._streaming_enabled:
raise RuntimeError(
"Streaming mode not enabled. Call enable_streaming() first.")
self._streaming_input_queue.put(
StreamingTask(
task_type=StreamingTaskType.STEP,
keyboard_action=keyboard_action,
mouse_action=mouse_action,
))
def submit_clear(self) -> None:
if self._streaming_enabled and self._streaming_input_queue is not None:
self._streaming_input_queue.put(
StreamingTask(task_type=StreamingTaskType.CLEAR))
def get_result(self,
timeout: float | None = None) -> StreamingResult | None:
if not self._streaming_enabled or self._streaming_output_queue is None:
return None
try:
if timeout == 0:
return self._streaming_output_queue.get_nowait()
else:
return self._streaming_output_queue.get(timeout=timeout)
except queue.Empty:
return None
def wait_result(self) -> StreamingResult:
if not self._streaming_enabled or self._streaming_output_queue is None:
raise RuntimeError("Streaming mode not enabled.")
return self._streaming_output_queue.get()
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
@@ -299,24 +147,6 @@ class MultiprocExecutor(Executor):
except Exception as e:
raise e
async def collective_rpc_async(self,
method: str | Callable,
timeout: float | None = None,
args: tuple = (),
kwargs: dict | None = None) -> list[Any]:
kwargs = kwargs or {}
loop = asyncio.get_running_loop()
for worker in self.workers:
worker.pipe.send({"method": method, "args": args, "kwargs": kwargs})
async def recv_from_worker(worker: WorkerProcHandle) -> Any:
return await loop.run_in_executor(None, worker.pipe.recv)
responses = await asyncio.gather(
*[recv_from_worker(worker) for worker in self.workers])
return list(responses)
def shutdown(self) -> None:
"""Properly shut down the executor and its workers"""
if hasattr(self, 'shutting_down') and self.shutting_down:
@@ -436,7 +266,7 @@ class WorkerProcHandle:
@classmethod
def from_unready_handle(
cls, unready_handle: UnreadyWorkerProcHandle) -> WorkerProcHandle:
cls, unready_handle: UnreadyWorkerProcHandle) -> "WorkerProcHandle":
return cls(
proc=unready_handle.proc,
rank=unready_handle.rank,
@@ -456,13 +286,9 @@ class WorkerMultiprocProc:
rank: int,
distributed_init_method: str,
pipe: Connection,
streaming_input_queue: Queue | None = None,
streaming_output_queue: Queue | None = None,
):
self.rank = rank
self.pipe = pipe
self.streaming_input_queue = streaming_input_queue
self.streaming_output_queue = streaming_output_queue
wrapper = WorkerWrapperBase(fastvideo_args=fastvideo_args,
rpc_rank=rank)
@@ -488,8 +314,6 @@ class WorkerMultiprocProc:
local_rank: int,
rank: int,
distributed_init_method: str,
streaming_input_queue: Queue | None = None,
streaming_output_queue: Queue | None = None,
) -> UnreadyWorkerProcHandle:
context = get_mp_context()
executor_pipe, worker_pipe = context.Pipe(duplex=True)
@@ -502,8 +326,6 @@ class WorkerMultiprocProc:
"distributed_init_method": distributed_init_method,
"pipe": worker_pipe,
"ready_pipe": writer,
"streaming_input_queue": streaming_input_queue,
"streaming_output_queue": streaming_output_queue,
}
# Run EngineCore busy loop in background process.
proc = context.Process(target=WorkerMultiprocProc.worker_main,
@@ -633,11 +455,6 @@ class WorkerMultiprocProc:
with contextlib.suppress(Exception):
self.pipe.send(response)
break
if method == "start_streaming_queue_loop":
self.pipe.send(
{"status": "streaming_queue_loop_started"})
self.streaming_queue_loop()
continue
if method == 'execute_forward':
forward_batch = kwargs['forward_batch']
fastvideo_args = kwargs['fastvideo_args']
@@ -671,50 +488,6 @@ class WorkerMultiprocProc:
self.rank, str(e))
continue
def streaming_queue_loop(self) -> None:
if self.streaming_input_queue is None or self.streaming_output_queue is None:
logger.error("Worker %d: streaming queues not initialized",
self.rank)
return
while True:
try:
task: StreamingTask = self.streaming_input_queue.get()
if task.task_type == StreamingTaskType.EXIT:
break
elif task.task_type == StreamingTaskType.RESET:
try:
self.worker.execute_streaming_reset(
task.batch, task.fastvideo_args)
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.RESET))
except Exception as e:
logger.error("Worker %d reset error: %s", self.rank, e)
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.RESET,
error=e))
elif task.task_type == StreamingTaskType.STEP:
try:
batch = self.worker.execute_streaming_step(
task.keyboard_action, task.mouse_action)
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.STEP,
output_batch=batch))
except Exception as e:
logger.error("Worker %d step error: %s", self.rank, e)
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.STEP,
error=e))
elif task.task_type == StreamingTaskType.CLEAR:
self.worker.execute_streaming_clear()
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.CLEAR))
except Exception as e:
logger.error("Worker %d queue loop error: %s", self.rank, e)
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.STEP, error=e))
@staticmethod
def setup_proc_title_and_log_prefix() -> None:
dp_size = get_dp_group().world_size
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/vllm-project/vllm/blob/releases/v0.11.0/vllm/executor/ray_distributed_executor.py
import asyncio
from collections import defaultdict
import os
import cloudpickle
@@ -270,47 +269,6 @@ class RayDistributedExecutor(Executor):
else:
self.non_driver_workers.append(worker)
def execute_streaming_reset(
self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
responses: list[dict[str, Any]] = self.collective_rpc(
"execute_streaming_reset",
kwargs={
"forward_batch": forward_batch,
"fastvideo_args": fastvideo_args,
},
)
return responses[0]
def execute_streaming_step(self,
keyboard_action=None,
mouse_action=None) -> ForwardBatch:
responses: list[ForwardBatch] = self.collective_rpc(
"execute_streaming_step",
kwargs={
"keyboard_action": keyboard_action,
"mouse_action": mouse_action,
},
)
return responses[0]
async def execute_streaming_step_async(self,
keyboard_action=None,
mouse_action=None) -> ForwardBatch:
kwargs = {
"keyboard_action": keyboard_action,
"mouse_action": mouse_action,
}
futures = [
w.execute_method.remote("execute_streaming_step", **kwargs)
for w in self.workers
]
responses = await asyncio.gather(*futures)
return responses[0]
def execute_streaming_clear(self) -> None:
self.collective_rpc("execute_streaming_clear")
def execute_forward(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
responses: list[ForwardBatch] = self.collective_rpc(
-1
View File
@@ -135,7 +135,6 @@ nav:
- Optimizations: inference/examples/optimizations.md
- STA Mask Search: inference/examples/sta_mask_search.md
- Training:
- Overview: training/overview.md
- Data Preprocessing: training/data_preprocess.md
- Fine-tuning: training/finetune.md
- Examples:
@@ -1,205 +0,0 @@
#!/usr/bin/env python3
"""
Convert TurboDiffusion .pth checkpoint to Diffusers safetensors format.
This script:
1. Loads the TurboDiffusion .pth checkpoint
2. Applies weight key renaming to match FastVideo/Diffusers format
3. Reshapes patch_embedding from [D, C*P] to [D, C, P_t, P_h, P_w]
4. Saves as sharded safetensors files compatible with diffusers
Usage:
python convert_turbodiffusion_to_diffusers.py \
--input_path /path/to/TurboWan2.1-T2V-1.3B-480P.pth \
--output_dir /path/to/output/transformer \
--reference_repo Wan-AI/Wan2.1-T2V-1.3B-Diffusers
"""
import argparse
import os
import re
import json
import torch
import shutil
import glob
from safetensors import safe_open
from safetensors.torch import save_file
from huggingface_hub import snapshot_download
# Weight mapping from TurboDiffusion -> Diffusers/FastVideo format
TURBODIFFUSION_WEIGHT_MAPPING = {
# Self attention
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
# Cross attention
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
# Norms and FFN
r"^blocks\.(\d+)\.norm1\.(.*)$": r"blocks.\1.norm1.\2",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.norm3.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
# Embeddings - DON'T add .proj here! WanVideoArchConfig's param_names_mapping will add it
# patch_embedding.weight stays as patch_embedding.weight (HF format needs this)
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
# Head
r"^head\.head\.(.*)$": r"proj_out.\1",
r"^head\.norm\.(.*)$": r"norm_out.\1",
r"^head\.modulation$": r"scale_shift_table",
# SLA proj_l weights - include them! They're the distilled attention weights
r"^blocks\.(\d+)\.self_attn\.attn_op\.local_attn\.proj_l\.(.*)$": r"blocks.\1.attn1.attn_impl.proj_l.\2",
}
# No keys to skip - we want all weights including proj_l
SKIP_PATTERNS = []
def should_skip_key(key: str) -> bool:
"""Check if a key should be skipped (SLA-specific weights)."""
for pattern in SKIP_PATTERNS:
if re.match(pattern, key):
return True
return False
def convert_key(turbo_key: str) -> str:
"""Convert TurboDiffusion key to Diffusers format."""
for pattern, replacement in TURBODIFFUSION_WEIGHT_MAPPING.items():
if re.match(pattern, turbo_key):
return re.sub(pattern, replacement, turbo_key)
return turbo_key # Return unchanged if no pattern matches
def reshape_patch_embedding(tensor: torch.Tensor, target_shape: tuple) -> torch.Tensor:
"""Reshape patch_embedding from [D, C*P_t*P_h*P_w] to [D, C, P_t, P_h, P_w]."""
if len(tensor.shape) == 2 and len(target_shape) == 5:
return tensor.view(target_shape)
return tensor
def get_reference_shapes(reference_repo: str) -> dict:
"""Download reference model and get expected shapes for patch_embedding."""
print(f"Downloading reference model shapes from {reference_repo}...")
# Download just the transformer config and a weight file to get shapes
local_dir = snapshot_download(
repo_id=reference_repo,
allow_patterns=["transformer/config.json", "transformer/diffusion_pytorch_model*.safetensors"],
local_dir_use_symlinks=False
)
# Load the first safetensors file to get shapes
weight_files = glob.glob(os.path.join(local_dir, "transformer", "*.safetensors"))
shapes = {}
for wf in weight_files:
with safe_open(wf, framework="pt") as f:
for key in f.keys():
shapes[key] = f.get_tensor(key).shape
return shapes
def main():
parser = argparse.ArgumentParser(description="Convert TurboDiffusion checkpoint to Diffusers format")
parser.add_argument("--input_path", type=str, required=True,
help="Path to TurboDiffusion .pth checkpoint")
parser.add_argument("--output_dir", type=str, required=True,
help="Output directory for converted safetensors")
parser.add_argument("--reference_repo", type=str, default="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
help="Reference HF repo to get expected tensor shapes")
parser.add_argument("--skip_sla_weights", action="store_true", default=False,
help="Skip SLA-specific weights (proj_l) that aren't in base model")
args = parser.parse_args()
# Load TurboDiffusion checkpoint
print(f"Loading TurboDiffusion checkpoint from {args.input_path}...")
turbo_state_dict = torch.load(args.input_path, map_location="cpu", weights_only=True)
print(f"Loaded {len(turbo_state_dict)} keys")
# Get reference shapes for reshaping
ref_shapes = get_reference_shapes(args.reference_repo)
print(f"Got {len(ref_shapes)} reference shapes")
# Convert keys and reshape tensors
converted_state_dict = {}
skipped_keys = []
for turbo_key, tensor in turbo_state_dict.items():
# Skip SLA-specific weights if requested
if args.skip_sla_weights and should_skip_key(turbo_key):
skipped_keys.append(turbo_key)
continue
# Convert key name
new_key = convert_key(turbo_key)
# Reshape patch_embedding if needed
if "patch_embedding" in new_key and new_key in ref_shapes:
target_shape = ref_shapes[new_key]
if tensor.shape != target_shape:
print(f"Reshaping {new_key}: {tensor.shape} -> {target_shape}")
tensor = reshape_patch_embedding(tensor, target_shape)
# Verify shape matches reference if available
if new_key in ref_shapes:
if tensor.shape != ref_shapes[new_key]:
print(f"WARNING: Shape mismatch for {new_key}: got {tensor.shape}, expected {ref_shapes[new_key]}")
converted_state_dict[new_key] = tensor
print(f"\nConversion summary:")
print(f" Converted: {len(converted_state_dict)} keys")
print(f" Skipped (SLA): {len(skipped_keys)} keys")
if skipped_keys:
print(f"\nSkipped SLA keys (first 5):")
for k in skipped_keys[:5]:
print(f" - {k}")
# Create output directory
os.makedirs(args.output_dir, exist_ok=True)
# Save as safetensors
output_path = os.path.join(args.output_dir, "diffusion_pytorch_model.safetensors")
print(f"\nSaving to {output_path}...")
save_file(converted_state_dict, output_path)
# Copy config.json from reference
ref_local = snapshot_download(
repo_id=args.reference_repo,
allow_patterns=["transformer/config.json"],
local_dir_use_symlinks=False
)
src_config = os.path.join(ref_local, "transformer", "config.json")
dst_config = os.path.join(args.output_dir, "config.json")
shutil.copy(src_config, dst_config)
print(f"Copied config.json")
print(f"\nDone! Converted weights saved to: {args.output_dir}")
print(f"\nNext steps:")
print(f" 1. Use create_hf_repo.py to create a complete diffusers repo:")
print(f" python scripts/checkpoint_conversion/create_hf_repo.py \\")
print(f" --repo_id Wan-AI/Wan2.1-T2V-1.3B-Diffusers \\")
print(f" --local_dir /tmp/turbodiffusion-wan \\")
print(f" --checkpoint_dir {args.output_dir} \\")
print(f" --push_to_hub --upload_repo_id YOUR_USERNAME/TurboWan-Diffusers")
if __name__ == "__main__":
main()