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

|
||||
|
||||
### Attention Patterns
|
||||
|
||||
Supports various patterns with memory optimization techniques:
|
||||
|
||||
- **Cross/Self/Temporal/Global-Local Attention**
|
||||
- Chunking, progressive computation, optimized masking
|
||||
|
||||
@@ -296,6 +309,7 @@ self.attn = DistributedAttention(
|
||||
```
|
||||
|
||||
### Communication Primitives
|
||||
|
||||
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
|
||||
|
||||
Efficient communication primitives minimize distributed overhead:
|
||||
@@ -314,6 +328,7 @@ Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-sp
|
||||
- **Profiling Data**: Potential hooks for performance metrics collection
|
||||
|
||||
This context-based approach enables:
|
||||
|
||||
- Dynamic optimization based on execution state (e.g., attention backend selection)
|
||||
- Step-specific customizations within model components
|
||||
|
||||
@@ -339,12 +354,14 @@ FastVideo implements a flexible execution model for distributed processing:
|
||||
- **GPU Workers**: Handle actual model execution on individual GPUs
|
||||
|
||||
The MultiProcExecutor implementation:
|
||||
|
||||
1. Spawns worker processes for each GPU
|
||||
2. Establishes communication channels via pipes
|
||||
3. Coordinates distributed operations across workers
|
||||
4. Handles graceful startup and shutdown of the process group
|
||||
|
||||
Each GPU worker:
|
||||
|
||||
1. Initializes the distributed environment
|
||||
2. Builds the pipeline for the specified model
|
||||
3. Executes requested operations on its assigned GPU
|
||||
@@ -359,11 +376,13 @@ The `fastvideo/platforms/` directory provides hardware platform abstractions tha
|
||||
### Platform Abstraction
|
||||
|
||||
FastVideo's platform abstraction layer enables:
|
||||
|
||||
- **Hardware Detection**: Automatic detection of available hardware
|
||||
- **Backend Selection**: Appropriate selection of compute kernels
|
||||
- **Memory Management**: Efficient utilization of hardware-specific memory features
|
||||
|
||||
The primary components include:
|
||||
|
||||
- **Platform Interface**: Defines the common API for all platform implementations
|
||||
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
|
||||
- **Backend Enum**: Used throughout the codebase for feature selection
|
||||
@@ -383,6 +402,7 @@ else:
|
||||
The platform system is designed to be extensible for future hardware targets.
|
||||
|
||||
## Logger
|
||||
|
||||
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
|
||||
@@ -397,6 +417,7 @@ If you're a new contributor, here are some common areas to explore:
|
||||
4. **Hardware support**: Extend the `platforms` module for new hardware targets
|
||||
|
||||
When adding code, follow these practices:
|
||||
|
||||
- Use type hints for better code readability
|
||||
- Add appropriate docstrings
|
||||
- Maintain the separation between model components and execution logic
|
||||
|
||||
@@ -13,7 +13,7 @@ We provide two distilled models:
|
||||
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
|
||||
|
||||
## ⚙️ Inference
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
|
||||
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
|
||||
@@ -11,15 +11,13 @@ FastVideo supports the following hardware platforms:
|
||||
### Using pip
|
||||
|
||||
```bash
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
### Using conda
|
||||
|
||||
```bash
|
||||
conda install -c conda-forge fastvideo
|
||||
```
|
||||
|
||||
### From source
|
||||
|
||||
```bash
|
||||
@@ -28,6 +26,12 @@ cd FastVideo
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
- **NVIDIA GPUs**: CUDA 11.8+ with compute capability 7.0+
|
||||
@@ -38,4 +42,4 @@ pip install -e .
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/) - Explore example scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
|
||||
@@ -84,12 +84,12 @@ pip install flash-attn --no-build-isolation
|
||||
|
||||
## Set up using Docker
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
[Docker Images](../../contributing/developer_env/docker.md)
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
[Contributor Guide](../../contributing/overview.md)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
|
||||
@@ -78,7 +78,7 @@ uv pip install -e .
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
[Contributor Guide](../../contributing/overview.md)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
|
||||
@@ -15,6 +15,12 @@ conda activate fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
|
||||
### Text-to-Video Generation
|
||||
|
||||
@@ -45,6 +45,7 @@ FastVideo uses the Hugging Face Diffusers format for model organization:
|
||||
### Implementing Modules
|
||||
|
||||
Place new modules in the appropriate directories:
|
||||
|
||||
- Encoders: `fastvideo/models/encoders/`
|
||||
- VAEs: `fastvideo/models/vaes/`
|
||||
- Transformer models: `fastvideo/models/dits/`
|
||||
@@ -53,12 +54,15 @@ Place new modules in the appropriate directories:
|
||||
### Adapting Model Layers
|
||||
|
||||
#### Layer Replacements
|
||||
|
||||
Replace standard PyTorch layers with FastVideo optimized versions:
|
||||
|
||||
- nn.LayerNorm → fastvideo.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.layers.activation
|
||||
|
||||
#### Distributed Linear Layers
|
||||
|
||||
Use appropriate parallel layers for distribution:
|
||||
|
||||
```python
|
||||
@@ -91,6 +95,7 @@ self.out_proj = RowParallelLinear(
|
||||
```
|
||||
|
||||
### Attention Layers
|
||||
|
||||
Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
@@ -304,6 +309,7 @@ EntryClass = [MyCustomPipeline, MyOtherPipeline]
|
||||
```
|
||||
|
||||
The registry will automatically:
|
||||
|
||||
1. Scan all packages under `fastvideo/pipelines/`
|
||||
2. Look for `EntryClass` variables
|
||||
3. Register pipelines using their class names as identifiers
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# FastVideo CLI Inference
|
||||
|
||||
The FastVideo CLI provides a quick way to access the FastVideo inference pipeline for video generation. For more advanced usage,
|
||||
see the Python interface [here](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html).
|
||||
see the Python interface [here](examples/basic.md).
|
||||
|
||||
## Basic Usage
|
||||
|
||||
|
||||
@@ -74,4 +74,4 @@ if __name__ == '__main__':
|
||||
|
||||
## Performance Optimization
|
||||
|
||||
For configuring optimizations, please see our [optimizations guide](#inference-optimizations)
|
||||
For configuring optimizations, please see our [optimizations guide](optimizations.md)
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
This page contains step-by-step instructions to get you quickly started with video generation using FastVideo.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS**: Linux (Tested on Ubuntu 22.04+)
|
||||
- **Python**: 3.10-3.12
|
||||
- **CUDA**: 12.8
|
||||
@@ -21,9 +22,10 @@ conda activate fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
For advanced installation options, see the [Installation Guide](installation.md).
|
||||
For advanced installation options, see the [Installation Guide](../getting_started/installation.md).
|
||||
|
||||
## Generating Your First Video
|
||||
|
||||
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
@@ -60,9 +62,10 @@ python example.py
|
||||
The generated video will be saved in the current directory under `my_videos/`
|
||||
|
||||
More inference example scripts can be found in `scripts/inference/`
|
||||
|
||||
## Available Models
|
||||
|
||||
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
|
||||
Please see the [support matrix](support_matrix.md) for the list of supported models and their available optimizations.
|
||||
|
||||
## Image-to-Video Generation
|
||||
|
||||
@@ -96,20 +99,26 @@ if __name__ == '__main__':
|
||||
Common issues and their solutions:
|
||||
|
||||
### Out of Memory Errors
|
||||
|
||||
If you encounter CUDA out of memory errors:
|
||||
|
||||
- Reduce `num_frames` or video resolution
|
||||
- Enable memory optimization with `enable_model_cpu_offload`
|
||||
- Try a smaller model or use distilled versions
|
||||
- Use `num_gpus` > 1 if multiple GPUs are available
|
||||
|
||||
### Slow Generation
|
||||
|
||||
To speed up generation:
|
||||
|
||||
- Reduce `num_inference_steps` (20-30 is usually sufficient)
|
||||
- Use half precision (`fp16`) for the VAE
|
||||
- Use multiple GPUs if available
|
||||
|
||||
### Unexpected Results
|
||||
|
||||
If the generated video doesn't match your prompt:
|
||||
|
||||
- Try increasing `guidance_scale` (7.0-9.0 works well)
|
||||
- Make your prompt more detailed and specific
|
||||
- Experiment with different random seeds
|
||||
@@ -117,8 +126,8 @@ If the generated video doesn't match your prompt:
|
||||
|
||||
## Next Steps
|
||||
|
||||
- Learn about [Advanced Inference Configurations](#inference-configuration)
|
||||
- Learn about using [Optimizations](#inference-optimizations)
|
||||
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
|
||||
- Learn about [Advanced Inference Configurations](configuration.md)
|
||||
- Learn about using [Optimizations](optimizations.md)
|
||||
- See [Examples](examples/examples_inference_index.md) for more usage scenarios
|
||||
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
@@ -7,13 +7,13 @@ This page describes the various options for speeding up generation times in Fast
|
||||
|
||||
- Optimized Attention Backends
|
||||
|
||||
- [Flash Attention](#optimizations-flash)
|
||||
- [Sliding Tile Attention](#optimizations-sta)
|
||||
- [Sage Attention](#optimizations-sage)
|
||||
- [Sage Attention 3](#optimizations-sage3)
|
||||
- [Flash Attention](#flash-attention)
|
||||
- [Sliding Tile Attention](#sliding-tile-attention)
|
||||
- [Sage Attention](#sage-attention)
|
||||
- [Sage Attention 3](#sage-attention-3)
|
||||
|
||||
- Caching Techniques
|
||||
- [TeaCache](#optimizations-teacache)
|
||||
- [TeaCache](#teacache)
|
||||
|
||||
## Attention Backends
|
||||
|
||||
@@ -74,7 +74,7 @@ python setup.py install
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
Please see [this page](../attention/sta/index.md) for more installation instructions.
|
||||
|
||||
### Video Sparse Attention
|
||||
|
||||
@@ -85,7 +85,7 @@ git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
Please see [this page](#vsa-installation) for more installation instructions.
|
||||
Please see [this page](../attention/vsa/index.md) for more installation instructions.
|
||||
|
||||
### Sage Attention
|
||||
|
||||
|
||||
@@ -40,20 +40,26 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
|
||||
}
|
||||
</style>
|
||||
|
||||
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA |
|
||||
|------------|---------------------|-------------|----------|-------------------|-----------|-----|
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
|
||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
|
||||
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
|
||||
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ |
|
||||
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ |
|
||||
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA | BSA |
|
||||
|------------|---------------------|-------------|----------|-------------------|-----------|-----|-----|
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| TurboWan2.1 T2V 1.3B | `loayrashid/TurboWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| TurboWan2.1 T2V 14B | `loayrashid/TurboWan2.1-T2V-14B-Diffusers` | 480P, 720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| LongCat T2V 13.6B | See note** | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
|
||||
| Matrix Game 2.0 Base | `FastVideo/Matrix-Game-2.0-Base-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 GTA | `FastVideo/Matrix-Game-2.0-GTA-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 TempleRun | `FastVideo/Matrix-Game-2.0-TempleRun-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
@@ -64,3 +70,13 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
|
||||
|
||||
### Sliding Tile Attention
|
||||
- Currently only Hopper GPUs (H100s) are supported.
|
||||
|
||||
### TurboWan2.1 (TurboDiffusion)
|
||||
- Uses TurboDiffusionPipeline with RCM scheduler for 1-4 step generation
|
||||
- Requires SLA attention backend: `export FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
|
||||
- Uses `guidance_scale=1.0` (no classifier-free guidance)
|
||||
|
||||
### Matrix Game 2.0
|
||||
- Image-to-video game world models with keyboard/mouse control input
|
||||
- Three variants available: Base (universal), GTA, and TempleRun
|
||||
- Each variant has different keyboard dimensions for control inputs
|
||||
|
||||
@@ -1,45 +1,130 @@
|
||||
# 🧱 Data Preprocessing
|
||||
|
||||
# 🧱 Data Preprocess
|
||||
To save GPU memory during training, FastVideo precomputes text embeddings and VAE latents. This eliminates the need to load the text encoder and VAE during training.
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
## Quick Start
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
Download the sample dataset and run preprocessing:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=data/mini_i2v_dataset --repo_type=dataset
|
||||
# Download the crush-smol dataset
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id "wlsaidhi/crush-smol-merged" \
|
||||
--local_dir "data/crush-smol" \
|
||||
--repo_type "dataset"
|
||||
|
||||
# Run preprocessing
|
||||
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v_new.sh
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
## Preprocessing Pipeline
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
The new preprocessing pipeline supports multiple dataset formats and video loaders:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
```bash
|
||||
GPU_NUM=2
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATASET_PATH="data/crush-smol/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path $MODEL_PATH \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 77 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.samples_per_file 8 \
|
||||
--preprocess.flush_frequency 8 \
|
||||
--preprocess.video_length_tolerance_range 5
|
||||
```
|
||||
|
||||
## Process your own dataset
|
||||
### Key Parameters
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
| Parameter | Description |
|
||||
|-----------|-------------|
|
||||
| `--workload_type` | Task type: `t2v` (text-to-video) or `i2v` (image-to-video) |
|
||||
| `--preprocess.dataset_type` | Input format: `hf` (HuggingFace) or `merged` (local folder) |
|
||||
| `--preprocess.dataset_path` | Path to dataset (HF repo ID or local folder) |
|
||||
| `--preprocess.dataset_output_dir` | Output directory for Parquet files |
|
||||
| `--preprocess.video_loader_type` | Video decoder: `torchcodec` or `torchvision` |
|
||||
| `--preprocess.max_height` / `max_width` | Target resolution for videos |
|
||||
| `--preprocess.num_frames` | Number of frames to extract per video |
|
||||
| `--preprocess.train_fps` | Target FPS for frame extraction |
|
||||
|
||||
## Dataset Formats
|
||||
|
||||
### Merged Dataset (Local Folder)
|
||||
|
||||
Structure your dataset as follows:
|
||||
|
||||
```
|
||||
path_to_your_dataset_folder/
|
||||
your_dataset/
|
||||
├── videos/
|
||||
│ ├── video_001.mp4
|
||||
│ ├── video_002.mp4
|
||||
│ └── ...
|
||||
└── videos2caption.json
|
||||
```
|
||||
|
||||
The `videos2caption.json` maps video filenames to captions:
|
||||
|
||||
```json
|
||||
[
|
||||
{"path": "video_001.mp4", "cap": "A cat playing with yarn..."},
|
||||
{"path": "video_002.mp4", "cap": "Ocean waves at sunset..."}
|
||||
]
|
||||
```
|
||||
|
||||
### HuggingFace Dataset
|
||||
|
||||
Use `--preprocess.dataset_type hf` and point `--preprocess.dataset_path` to a HuggingFace dataset with `video` and `caption` columns.
|
||||
|
||||
## Creating Your Own Dataset
|
||||
|
||||
If you have raw videos and captions in separate files, generate the `videos2caption.json`:
|
||||
|
||||
```bash
|
||||
python scripts/dataset_preparation/prepare_json_file.py \
|
||||
--data_folder path/to/your_raw_data/ \
|
||||
--output path/to/output_folder
|
||||
```
|
||||
|
||||
Your raw data folder should contain:
|
||||
|
||||
```
|
||||
your_raw_data/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
│ ├── 1.mp4
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
│ └── ...
|
||||
├── videos.txt # list of video filenames
|
||||
└── prompt.txt # corresponding captions (one per line)
|
||||
```
|
||||
|
||||
To generate the `videos2caption.json` and `merge.txt`, run
|
||||
## Output Format
|
||||
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
Preprocessing outputs Parquet files in the `combined_parquet_dataset/` subdirectory containing:
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
- `vae_latent_bytes` — VAE-encoded video latent
|
||||
- `text_embedding_bytes` — text encoder output
|
||||
- `clip_feature_bytes` — CLIP image features (I2V only)
|
||||
- `first_frame_latent_bytes` — first frame latent (I2V only)
|
||||
- Metadata: shapes, dtypes, and sample identifiers
|
||||
|
||||
```
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
```
|
||||
## Examples
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
See ready-to-run preprocessing scripts in the training examples:
|
||||
|
||||
- **T2V**: `examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v_new.sh`
|
||||
- **I2V**: `examples/training/finetune/wan_i2v_14B_480p/crush_smol/preprocess_wan_data_i2v_new.sh`
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
+153
-55
@@ -1,78 +1,176 @@
|
||||
# 🧠 Finetuning
|
||||
|
||||
# 🧠 Finetune
|
||||
## ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
This guide covers finetuning video diffusion models with FastVideo, including full finetuning and LoRA.
|
||||
|
||||
## Training Arguments
|
||||
|
||||
FastVideo training scripts use several argument groups:
|
||||
|
||||
### Training Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--max_train_steps` | Total training steps |
|
||||
| `--train_batch_size` | Batch size per GPU |
|
||||
| `--gradient_accumulation_steps` | Steps to accumulate before optimizer update |
|
||||
| `--num_latent_t` | Temporal latent dimension (reduce to save memory) |
|
||||
| `--num_height` / `--num_width` | Video resolution |
|
||||
| `--num_frames` | Number of frames per video |
|
||||
| `--output_dir` | Directory for checkpoints |
|
||||
|
||||
### Parallelism Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--num_gpus` | Total number of GPUs |
|
||||
| `--sp_size` | Sequence parallel size (increase to reduce memory per GPU) |
|
||||
| `--tp_size` | Tensor parallel size |
|
||||
| `--hsdp_replicate_dim` | HSDP replication dimension |
|
||||
| `--hsdp_shard_dim` | HSDP sharding dimension |
|
||||
|
||||
### Optimizer Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--learning_rate` | Base learning rate |
|
||||
| `--mixed_precision` | Precision mode (`bf16` recommended) |
|
||||
| `--weight_decay` | Weight decay for regularization |
|
||||
| `--max_grad_norm` | Gradient clipping threshold |
|
||||
|
||||
### Validation Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--log_validation` | Enable validation logging |
|
||||
| `--validation_dataset_file` | JSON file with validation prompts |
|
||||
| `--validation_steps` | Run validation every N steps |
|
||||
| `--validation_sampling_steps` | Inference steps for validation |
|
||||
| `--validation_guidance_scale` | CFG scale for validation |
|
||||
|
||||
## Full Finetuning
|
||||
|
||||
Full finetuning updates all model weights. This provides the best quality but requires more GPU memory.
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
# Example: Wan2.1 T2V 1.3B full finetune (4 GPUs)
|
||||
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
|
||||
```
|
||||
|
||||
Download the original model weights as specified in the [Distillation Section](../distillation/dmd.md):
|
||||
**Typical settings:**
|
||||
|
||||
Then you can run the finetune with:
|
||||
- Learning rate: `1e-5` to `5e-5`
|
||||
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
|
||||
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
## LoRA Finetuning
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
|
||||
|
||||
### LoRA-Specific Arguments
|
||||
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--lora_training True` | Enable LoRA mode |
|
||||
| `--lora_rank` | Rank of LoRA adapters (16, 32, 64, 128) |
|
||||
|
||||
### Learning Rate for LoRA
|
||||
|
||||
**Important:** LoRA typically requires a **10–20× higher learning rate** than full finetuning because only the low-rank adapters are being trained while the base model is frozen.
|
||||
|
||||
| Training Mode | Recommended Learning Rate |
|
||||
|---------------|---------------------------|
|
||||
| Full finetune | `1e-5` to `5e-5` |
|
||||
| LoRA | `1e-4` to `2e-4` |
|
||||
|
||||
### Example LoRA Training
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
# Example: Wan2.1 T2V 1.3B LoRA finetune (1 GPU)
|
||||
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v_lora.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
Key differences from full finetune:
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
- Add `--lora_training True --lora_rank 32`
|
||||
- Use higher learning rate (10–20× full finetune)
|
||||
- Can run on fewer GPUs (even single GPU)
|
||||
- Outputs adapter weights instead of full model
|
||||
|
||||
## LoRA Extraction and Merging
|
||||
|
||||
FastVideo provides tools to extract LoRA adapters from finetuned models and merge them back.
|
||||
|
||||
### Extract LoRA Adapter
|
||||
|
||||
Extract a LoRA adapter by comparing a finetuned model to its base:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
python scripts/lora_extraction/extract_lora.py \
|
||||
--base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--ft path/to/your/finetuned_model \
|
||||
--out adapter_r32.safetensors \
|
||||
--rank 32
|
||||
```
|
||||
|
||||
### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--base` | Base model (HuggingFace ID or local path) |
|
||||
| `--ft` | Finetuned model path |
|
||||
| `--out` | Output adapter file (.safetensors) |
|
||||
| `--rank` | LoRA rank (16, 32, 64, 128) |
|
||||
| `--full-rank` | Extract full-rank adapter (optional) |
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
### Merge LoRA Adapter
|
||||
|
||||
### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
Merge an adapter back into a base model:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
python scripts/lora_extraction/merge_lora.py \
|
||||
--base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--adapter adapter_r32.safetensors \
|
||||
--ft path/to/your/finetuned_model \
|
||||
--output merged_model
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
| Argument | Description |
|
||||
|----------|-------------|
|
||||
| `--base` | Base model path |
|
||||
| `--adapter` | LoRA adapter file |
|
||||
| `--ft` | Finetuned model (for config reference) |
|
||||
| `--output` | Output directory for merged model |
|
||||
|
||||
### Validate Merged Model
|
||||
|
||||
Compare the merged model against the original finetuned model:
|
||||
|
||||
```bash
|
||||
python scripts/lora_extraction/lora_inference_comparison.py \
|
||||
--base merged_model \
|
||||
--ft path/to/your/finetuned_model \
|
||||
--adapter NONE \
|
||||
--output-dir results \
|
||||
--prompt "A cat sitting on a windowsill" \
|
||||
--compute-ssim \
|
||||
--compute-lpips
|
||||
```
|
||||
|
||||
## Training Examples
|
||||
|
||||
Ready-to-run training scripts are available for multiple models:
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
| Model | Type | Example |
|
||||
|-------|------|---------|
|
||||
| Wan2.1 T2V 1.3B | T2V | `examples/training/finetune/wan_t2v_1.3B/crush_smol/` |
|
||||
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
|
||||
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
|
||||
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
- `preprocess_*.sh` — run preprocessing
|
||||
- `finetune_*.sh` — full finetune launcher
|
||||
- `finetune_*_lora.sh` — LoRA finetune launcher
|
||||
- `validation.json` — validation prompts
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# Training Overview
|
||||
|
||||
FastVideo supports finetuning video diffusion models on custom datasets. This page explains what data you need and how to get started.
|
||||
|
||||
## Data Requirements
|
||||
|
||||
To save GPU memory during training, FastVideo precomputes embeddings and latents ahead of time. This eliminates the need to load the text encoder and VAE during training, significantly reducing memory usage.
|
||||
|
||||
### Text-to-Video (T2V) Finetuning
|
||||
|
||||
For T2V models, you need:
|
||||
|
||||
| Component | Description |
|
||||
|-----------|-------------|
|
||||
| **Text embeddings** | Precomputed embeddings from the model's text encoder (e.g., T5 or LLaMA). Stored as numpy arrays in Parquet files. |
|
||||
| **Video latents** | VAE-encoded representations of your training videos. Each video is encoded into a compressed latent tensor. |
|
||||
|
||||
### Image-to-Video (I2V) Finetuning
|
||||
|
||||
For I2V models, you need everything from T2V plus additional image conditioning. Note that not all I2V architectures require encoded images—this depends on how the model conditions on the input frame. Wan2.1 and Wan2.2 A14B I2V models do require these additional components:
|
||||
|
||||
| Component | Description |
|
||||
|-----------|-------------|
|
||||
| **Text embeddings** | Same as T2V—precomputed from the text encoder. |
|
||||
| **Video latents** | Same as T2V—VAE-encoded video representations. |
|
||||
| **First frame latent** | VAE-encoded representation of the first frame, used as the conditioning image. |
|
||||
| **CLIP features** | Image embeddings from a CLIP vision encoder for the conditioning frame. |
|
||||
|
||||
## Preprocessing
|
||||
|
||||
Before training, you need to preprocess your raw videos and captions into Parquet files containing precomputed latents and embeddings.
|
||||
|
||||
FastVideo supports two input formats:
|
||||
|
||||
- **HuggingFace datasets** — load directly from HF Hub or local HF datasets
|
||||
- **Merged datasets** — local folder with videos and a `videos2caption.json` metadata file
|
||||
|
||||
**→ See [Data Preprocessing](data_preprocess.md) for full details and examples.**
|
||||
|
||||
## Training Examples
|
||||
|
||||
Ready-to-run examples with preprocessing scripts, training launchers, and validation configs are available for multiple models and datasets:
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
- `preprocess_*.sh` — run preprocessing
|
||||
- `finetune_*.sh` — launch training (full finetune or LoRA)
|
||||
- `validation.json` — validation prompts for checkpoints
|
||||
|
||||
## Training Methods
|
||||
|
||||
FastVideo supports several training approaches:
|
||||
|
||||
| Method | Use Case |
|
||||
|--------|----------|
|
||||
| **Full finetune** | Adapt entire model to a new domain or style |
|
||||
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
|
||||
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
|
||||
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
|
||||
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
@@ -0,0 +1,19 @@
|
||||
# Self-Forcing Distillation for SFWan2.1 T2V 1.3B
|
||||
|
||||
These scripts demonstrate self-forcing distillation (SFwan) for the causal Wan2.1 T2V 1.3B model. The workflow mirrors DMD2 while injecting self-forcing blocks so the student can autoregressively refine later frames.
|
||||
|
||||
## Run the recipe
|
||||
1. Download the preprocessed text-video dataset:
|
||||
```bash
|
||||
bash examples/distill/SFWan2.1-T2V/download_dataset.sh
|
||||
```
|
||||
2. (Optional) Regenerate parquet shards locally:
|
||||
```bash
|
||||
bash examples/distill/SFWan2.1-T2V/preprocess_data.sh
|
||||
```
|
||||
3. Launch self-forcing distillation with your cluster settings:
|
||||
```bash
|
||||
sbatch examples/distill/SFWan2.1-T2V/distill_dmd_t2v_1.3B.sh
|
||||
```
|
||||
|
||||
Update the dataset paths and wandb credentials inside the script before running on your environment.
|
||||
@@ -0,0 +1,98 @@
|
||||
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
|
||||
from fastvideo.models.dits.matrix_game.utils import get_current_action_async, expand_action_to_frames
|
||||
|
||||
import torch
|
||||
import asyncio
|
||||
|
||||
# Available variants: "base_distilled_model", "gta_distilled_model", "templerun_distilled_model"
|
||||
# Each variant has different keyboard_dim:
|
||||
# - base_distilled_model: keyboard_dim=4
|
||||
# - gta_distilled_model: keyboard_dim=2
|
||||
# - templerun_distilled_model: keyboard_dim=7 (keyboard only, no mouse)
|
||||
MODEL_VARIANT = "base_distilled_model"
|
||||
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"mode": "universal",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"mode": "gta_drive",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"mode": "templerun",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
async def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
config = VARIANT_CONFIG[MODEL_VARIANT]
|
||||
|
||||
generator = StreamingVideoGenerator.from_pretrained(
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=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())
|
||||
@@ -0,0 +1,60 @@
|
||||
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",
|
||||
|
||||
# set to false if using RTX 4090
|
||||
# pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
# TurboDiffusion uses guidance_scale=1.0 (no CFG) and only 4 steps
|
||||
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()
|
||||
@@ -0,0 +1,55 @@
|
||||
import os
|
||||
|
||||
# Set SLA attention backend BEFORE fastvideo imports
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_turbodiffusion_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()
|
||||
@@ -0,0 +1,41 @@
|
||||
import os
|
||||
|
||||
# Set SLA attention backend BEFORE fastvideo imports
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Use local model path
|
||||
MODEL_PATH = "loayrashid/TurboWan2.2-I2V-A14B-Diffusers"
|
||||
OUTPUT_PATH = "video_samples_turbodiffusion_i2v"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# TurboDiffusion I2V: 1-4 step image-to-video generation
|
||||
# Uses RCM scheduler with sigma_max=200 for I2V
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_PATH,
|
||||
num_gpus=2,
|
||||
override_pipeline_cls_name="TurboDiffusionI2VPipeline",
|
||||
)
|
||||
|
||||
# Example prompt and image for I2V
|
||||
prompt = ("Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside.")
|
||||
|
||||
# Use an example image path
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
image_path=image_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_inference_steps=4,
|
||||
seed=42,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,684 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Request, HTTPException
|
||||
from fastapi.responses import HTMLResponse, FileResponse
|
||||
|
||||
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
|
||||
from fastvideo.models.dits.matrix_game.utils import expand_action_to_frames
|
||||
|
||||
|
||||
VARIANT_CONFIG = {
|
||||
"Matrix-Game-2.0-Base": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"mode": "universal",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"Matrix-Game-2.0-GTA": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"mode": "gta_drive",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"Matrix-Game-2.0-TempleRun": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"mode": "templerun",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
MODEL_PATH_MAPPING = {
|
||||
name: config["model_path"] for name, config in VARIANT_CONFIG.items()
|
||||
}
|
||||
|
||||
|
||||
CAM_VALUE = 0.1
|
||||
KEYBOARD_MAP_UNIVERSAL = {
|
||||
"W (Forward)": [1, 0, 0, 0],
|
||||
"S (Back)": [0, 1, 0, 0],
|
||||
"A (Left)": [0, 0, 1, 0],
|
||||
"D (Right)": [0, 0, 0, 1],
|
||||
"Q (Stop)": [0, 0, 0, 0],
|
||||
}
|
||||
KEYBOARD_MAP_GTA = {
|
||||
"W (Forward)": [1, 0],
|
||||
"S (Back)": [0, 1],
|
||||
"Q (Stop)": [0, 0],
|
||||
}
|
||||
KEYBOARD_MAP_TEMPLERUN = {
|
||||
"Q (Run)": [1, 0, 0, 0, 0, 0, 0],
|
||||
"W (Jump)": [0, 1, 0, 0, 0, 0, 0],
|
||||
"S (Slide)": [0, 0, 1, 0, 0, 0, 0],
|
||||
"Z (Turn Left)": [0, 0, 0, 1, 0, 0, 0],
|
||||
"C (Turn Right)": [0, 0, 0, 0, 1, 0, 0],
|
||||
"A (Left)": [0, 0, 0, 0, 0, 1, 0],
|
||||
"D (Right)": [0, 0, 0, 0, 0, 0, 1],
|
||||
}
|
||||
|
||||
|
||||
CAMERA_MAP_UNIVERSAL = {
|
||||
"U (Center)": [0, 0],
|
||||
"I (Up)": [CAM_VALUE, 0],
|
||||
"K (Down)": [-CAM_VALUE, 0],
|
||||
"J (Left)": [0, -CAM_VALUE],
|
||||
"L (Right)": [0, CAM_VALUE],
|
||||
}
|
||||
CAMERA_MAP_GTA = {
|
||||
"Q (Straight)": [0, 0],
|
||||
"A (Steer Left)": [0, -CAM_VALUE],
|
||||
"D (Steer Right)": [0, CAM_VALUE],
|
||||
}
|
||||
|
||||
def setup_model_environment(model_path: str) -> None:
|
||||
# if "fullattn" in model_path.lower():
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
# else:
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
|
||||
|
||||
def create_timing_display(inference_time, total_time, stage_execution_times, num_frames):
|
||||
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
|
||||
|
||||
timing_html = f"""
|
||||
<div style="margin: 10px 0;">
|
||||
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
|
||||
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
|
||||
<div class="timing-card timing-card-highlight">
|
||||
<div style="font-size: 20px;">🚀</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
|
||||
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🧠</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
|
||||
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🎬</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
|
||||
<div style="font-size: 16px; color: #dc2626;">N/A</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🌐</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
|
||||
<div style="font-size: 16px; color: #059669;">N/A</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">📊</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
|
||||
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
|
||||
</div>
|
||||
</div>"""
|
||||
|
||||
if inference_time > 0:
|
||||
fps = num_frames / inference_time
|
||||
timing_html += f"""
|
||||
<div class="performance-card" style="margin-top: 15px;">
|
||||
<span style="font-weight: bold;">Generation Speed: </span>
|
||||
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
|
||||
</div>"""
|
||||
|
||||
return timing_html + "</div>"
|
||||
|
||||
def get_action_tensors(mode: str, keyboard_key: str, mouse_key: str | None):
|
||||
if mode == "universal":
|
||||
keyboard = torch.tensor(KEYBOARD_MAP_UNIVERSAL.get(keyboard_key, [0, 0, 0, 0])).cuda()
|
||||
mouse = torch.tensor(CAMERA_MAP_UNIVERSAL.get(mouse_key, [0, 0])).cuda()
|
||||
elif mode == "gta_drive":
|
||||
keyboard = torch.tensor(KEYBOARD_MAP_GTA.get(keyboard_key, [0, 0])).cuda()
|
||||
mouse = torch.tensor(CAMERA_MAP_GTA.get(mouse_key, [0, 0])).cuda()
|
||||
elif mode == "templerun":
|
||||
keyboard = torch.tensor(KEYBOARD_MAP_TEMPLERUN.get(keyboard_key, [1, 0, 0, 0, 0, 0, 0])).cuda()
|
||||
mouse = None
|
||||
else:
|
||||
raise ValueError(f"Unknown mode: {mode}")
|
||||
|
||||
return {"keyboard": keyboard, "mouse": mouse}
|
||||
|
||||
def create_gradio_interface(generators: dict[str, StreamingVideoGenerator], loaded_model_name: str):
|
||||
initial_config = VARIANT_CONFIG.get(loaded_model_name, VARIANT_CONFIG["Matrix-Game-2.0-Base"])
|
||||
initial_mode = initial_config["mode"]
|
||||
|
||||
if initial_mode == "universal":
|
||||
initial_kb_choices = list(KEYBOARD_MAP_UNIVERSAL.keys())
|
||||
initial_mouse_choices = list(CAMERA_MAP_UNIVERSAL.keys())
|
||||
initial_mouse_visible = True
|
||||
elif initial_mode == "gta_drive":
|
||||
initial_kb_choices = list(KEYBOARD_MAP_GTA.keys())
|
||||
initial_mouse_choices = list(CAMERA_MAP_GTA.keys())
|
||||
initial_mouse_visible = True
|
||||
else: # templerun
|
||||
initial_kb_choices = list(KEYBOARD_MAP_TEMPLERUN.keys())
|
||||
initial_mouse_choices = []
|
||||
initial_mouse_visible = False
|
||||
|
||||
theme = gr.themes.Base().set(
|
||||
button_primary_background_fill="#2563eb",
|
||||
button_primary_background_fill_hover="#1d4ed8",
|
||||
button_primary_text_color="white",
|
||||
slider_color="#2563eb",
|
||||
checkbox_background_color_selected="#2563eb",
|
||||
)
|
||||
|
||||
with gr.Blocks(title="FastVideo - Matrix Game 2.0", theme=theme) as demo:
|
||||
game_state = gr.State({
|
||||
"initialized": False,
|
||||
"current_model": None,
|
||||
"block_idx": 0,
|
||||
"max_blocks": 50,
|
||||
})
|
||||
|
||||
# Header
|
||||
gr.Image("assets/full.svg", show_label=False, container=False, height=80)
|
||||
|
||||
gr.HTML("""
|
||||
<div style="text-align: center; margin-bottom: 10px;">
|
||||
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
|
||||
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
with gr.Accordion("🎥 What Is FastVideo?", open=False):
|
||||
gr.HTML("""
|
||||
<div style="padding: 20px; line-height: 1.6;">
|
||||
<p style="font-size: 16px; margin-bottom: 15px;">
|
||||
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
</p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
# Model Selection
|
||||
with gr.Row():
|
||||
model_selection = gr.Dropdown(
|
||||
choices=[loaded_model_name],
|
||||
value=loaded_model_name,
|
||||
label="Select Model",
|
||||
interactive=False
|
||||
)
|
||||
|
||||
|
||||
# Main Layout
|
||||
with gr.Row(equal_height=True, elem_classes="main-content-row"):
|
||||
with gr.Column(scale=1, elem_classes="advanced-options-column"):
|
||||
with gr.Group():
|
||||
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Game Controls</div>")
|
||||
|
||||
with gr.Group():
|
||||
gr.HTML("<div style='font-size: 14px; margin-bottom: 5px; font-weight: bold;'>🎮 Keyboard Control</div>")
|
||||
keyboard_action = gr.Radio(
|
||||
choices=initial_kb_choices,
|
||||
value=initial_kb_choices[0] if initial_kb_choices else None,
|
||||
label="Movement",
|
||||
show_label=False,
|
||||
interactive=True
|
||||
)
|
||||
|
||||
with gr.Group(visible=initial_mouse_visible) as mouse_group:
|
||||
gr.HTML("<div style='font-size: 14px; margin-bottom: 5px; font-weight: bold;'>🖱️ Mouse/Camera Control</div>")
|
||||
mouse_action = gr.Radio(
|
||||
choices=initial_mouse_choices if initial_mouse_visible else [],
|
||||
value=initial_mouse_choices[0] if initial_mouse_choices else None,
|
||||
label="Camera",
|
||||
show_label=False,
|
||||
interactive=True
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
action_btn = gr.Button("Start", variant="primary")
|
||||
stop_btn = gr.Button("Stop", variant="stop")
|
||||
|
||||
gr.HTML("<div style='margin-top: 15px;'></div>")
|
||||
|
||||
seed = gr.Slider(
|
||||
label="Seed",
|
||||
minimum=0,
|
||||
maximum=1000000,
|
||||
step=1,
|
||||
value=1024,
|
||||
)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
block_counter = gr.Textbox(label="Progress", value="Block: 0 / 50", interactive=False, lines=1)
|
||||
|
||||
|
||||
# Right Column: Video Output
|
||||
with gr.Column(scale=1, elem_classes="video-column"):
|
||||
video_output = gr.Video(
|
||||
label="Generated Video",
|
||||
show_label=True,
|
||||
height=466,
|
||||
width=600,
|
||||
container=True,
|
||||
elem_classes="video-component",
|
||||
autoplay=True
|
||||
)
|
||||
|
||||
# Styles
|
||||
gr.HTML("""
|
||||
<style>
|
||||
.center-button {
|
||||
display: flex !important;
|
||||
justify-content: center !important;
|
||||
height: 100% !important;
|
||||
padding-top: 1.4em !important;
|
||||
}
|
||||
|
||||
.gradio-container {
|
||||
max-width: 1200px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.main {
|
||||
max-width: 1200px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.gr-form, .gr-box, .gr-group {
|
||||
max-width: 1200px !important;
|
||||
}
|
||||
|
||||
.gr-video {
|
||||
max-width: 500px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.main-content-row {
|
||||
display: flex !important;
|
||||
align-items: flex-start !important;
|
||||
min-height: 500px !important;
|
||||
gap: 20px !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
display: flex !important;
|
||||
flex-direction: column !important;
|
||||
flex: 1 !important;
|
||||
min-height: 400px !important;
|
||||
align-items: stretch !important;
|
||||
}
|
||||
|
||||
.video-column > * {
|
||||
margin-top: 0 !important;
|
||||
}
|
||||
|
||||
.video-column .gr-video,
|
||||
.video-component {
|
||||
margin-top: 0 !important;
|
||||
padding-top: 0 !important;
|
||||
}
|
||||
|
||||
.video-column .gr-video .gr-form {
|
||||
margin-top: 0 !important;
|
||||
}
|
||||
|
||||
.advanced-options-column .gr-group,
|
||||
.video-column .gr-video {
|
||||
margin-top: 0 !important;
|
||||
vertical-align: top !important;
|
||||
}
|
||||
|
||||
.advanced-options-column > *:last-child,
|
||||
.video-column > *:last-child {
|
||||
flex-grow: 0 !important;
|
||||
}
|
||||
|
||||
@media (max-width: 1400px) {
|
||||
.main-content-row {
|
||||
min-height: 600px !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
min-height: 600px !important;
|
||||
}
|
||||
}
|
||||
|
||||
@media (max-width: 1200px) {
|
||||
.main-content-row {
|
||||
flex-direction: column !important;
|
||||
align-items: stretch !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
min-height: auto !important;
|
||||
width: 100% !important;
|
||||
}
|
||||
}
|
||||
|
||||
.timing-card {
|
||||
background: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color) !important;
|
||||
padding: 10px;
|
||||
border-radius: 8px;
|
||||
text-align: center;
|
||||
min-height: 80px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.timing-card-highlight {
|
||||
background: var(--background-fill-primary) !important;
|
||||
border: 2px solid var(--color-accent) !important;
|
||||
}
|
||||
|
||||
.performance-card {
|
||||
background: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color) !important;
|
||||
padding: 10px;
|
||||
border-radius: 6px;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
.gr-number input[readonly] {
|
||||
background-color: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color-subdued) !important;
|
||||
cursor: default !important;
|
||||
text-align: center !important;
|
||||
font-weight: 500 !important;
|
||||
}
|
||||
</style>
|
||||
""")
|
||||
|
||||
# UI update based on model selection
|
||||
def on_model_change(model_name):
|
||||
config = VARIANT_CONFIG.get(model_name, VARIANT_CONFIG["Matrix-Game-2.0-Base"])
|
||||
mode = config["mode"]
|
||||
|
||||
if mode == "universal":
|
||||
kb_choices = list(KEYBOARD_MAP_UNIVERSAL.keys())
|
||||
mouse_choices = list(CAMERA_MAP_UNIVERSAL.keys())
|
||||
mouse_visible = True
|
||||
elif mode == "gta_drive":
|
||||
kb_choices = list(KEYBOARD_MAP_GTA.keys())
|
||||
mouse_choices = list(CAMERA_MAP_GTA.keys())
|
||||
mouse_visible = True
|
||||
else: # templerun
|
||||
kb_choices = list(KEYBOARD_MAP_TEMPLERUN.keys())
|
||||
mouse_choices = []
|
||||
mouse_visible = False
|
||||
|
||||
return (
|
||||
gr.update(choices=kb_choices, value=kb_choices[0] if kb_choices else None),
|
||||
gr.update(choices=mouse_choices, value=mouse_choices[0] if mouse_choices else None, visible=mouse_visible),
|
||||
gr.update(visible=mouse_visible),
|
||||
)
|
||||
|
||||
model_selection.change(
|
||||
fn=on_model_change,
|
||||
inputs=model_selection,
|
||||
outputs=[keyboard_action, mouse_action, mouse_group]
|
||||
)
|
||||
|
||||
def start_game(model_name, seed_val, randomize, state):
|
||||
if randomize:
|
||||
seed_val = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
config = VARIANT_CONFIG.get(model_name)
|
||||
if not config:
|
||||
return state, seed_val, "Block: 0 / 50", None, "", gr.update(), gr.update()
|
||||
|
||||
generator = generators.get(config["model_path"])
|
||||
if not generator:
|
||||
return state, seed_val, "Block: 0 / 50", None, "", gr.update(), gr.update()
|
||||
|
||||
# If already initialized, clean up first
|
||||
if state.get("initialized"):
|
||||
try:
|
||||
# Clear accumulated frames without saving
|
||||
generator.accumulated_frames = []
|
||||
generator.executor.execute_streaming_clear()
|
||||
except Exception as e:
|
||||
print(f"Warning: cleanup error: {e}")
|
||||
|
||||
# Streaming parameters
|
||||
num_latent_frames_per_block = 3
|
||||
max_blocks = 50
|
||||
total_latent_frames = num_latent_frames_per_block * max_blocks
|
||||
num_frames = (total_latent_frames - 1) * 4 + 1
|
||||
|
||||
actions = {
|
||||
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
|
||||
"mouse": torch.zeros((num_frames, 2))
|
||||
}
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
|
||||
output_dir = os.path.abspath("outputs/matrixgame")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
video_path = os.path.join(output_dir, f"video_{int(time.time())}.mp4")
|
||||
|
||||
generator.reset(
|
||||
prompt="",
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
output_path=video_path,
|
||||
)
|
||||
|
||||
new_state = {
|
||||
"initialized": True,
|
||||
"current_model": model_name,
|
||||
"block_idx": 0,
|
||||
"max_blocks": max_blocks,
|
||||
"video_path": video_path,
|
||||
"frames_per_block": num_latent_frames_per_block * 4,
|
||||
"mode": config["mode"],
|
||||
"seed": seed_val,
|
||||
}
|
||||
|
||||
return new_state, seed_val, "Block: 0 / 50", None, gr.update(value="Step"), gr.update(interactive=True)
|
||||
|
||||
async def step_game(keyboard_key, mouse_key, model_name, state):
|
||||
if not state.get("initialized"):
|
||||
return state, state.get("seed", 0), "Block: 0 / 50", None, gr.update(), gr.update()
|
||||
|
||||
# total_start_time = time.time()
|
||||
config = VARIANT_CONFIG.get(model_name)
|
||||
generator = generators.get(config["model_path"])
|
||||
mode = state["mode"]
|
||||
frames_per_block = state["frames_per_block"]
|
||||
|
||||
# Parse inputs to tensors
|
||||
action = get_action_tensors(mode, keyboard_key, mouse_key)
|
||||
keyboard_cond, mouse_cond = expand_action_to_frames(action, frames_per_block)
|
||||
|
||||
# run step async
|
||||
# inference_start_time = time.time()
|
||||
frames, block_future = await generator.step_async(keyboard_cond, mouse_cond)
|
||||
# inference_time = time.time() - inference_start_time
|
||||
|
||||
# wait for block file to be written
|
||||
block_path = await asyncio.to_thread(block_future.result) if block_future else None
|
||||
state["block_idx"] = generator.block_idx
|
||||
block_str = f"Block: {state['block_idx']} / {state['max_blocks']}"
|
||||
|
||||
# total_time = time.time() - total_start_time
|
||||
|
||||
# Timing breakdown
|
||||
# timing_html = create_timing_display(inference_time, total_time, [], frames_per_block)
|
||||
|
||||
return state, state.get("seed", 0), block_str, block_path, gr.update(), gr.update()
|
||||
|
||||
def stop_game(model_name, state):
|
||||
if not state.get("initialized"):
|
||||
return {"initialized": False}, 0, "Block: 0 / 50", None, gr.update(value="Start"), gr.update(interactive=False)
|
||||
|
||||
config = VARIANT_CONFIG.get(model_name)
|
||||
generator = generators.get(config["model_path"])
|
||||
|
||||
final_path = state.get("video_path")
|
||||
generator.finalize(final_path)
|
||||
|
||||
return {"initialized": False}, state.get("seed", 0), "Block: 0 / 50", final_path, gr.update(value="Start"), gr.update(interactive=False)
|
||||
|
||||
async def handle_action(keyboard_key, mouse_key, model_name, seed_val, randomize, state):
|
||||
if not state.get("initialized"):
|
||||
return start_game(model_name, seed_val, randomize, state)
|
||||
else:
|
||||
return await step_game(keyboard_key, mouse_key, model_name, state)
|
||||
|
||||
action_btn.click(
|
||||
fn=handle_action,
|
||||
inputs=[keyboard_action, mouse_action, model_selection, seed, randomize_seed, game_state],
|
||||
outputs=[game_state, seed_output, block_counter, video_output, action_btn, stop_btn]
|
||||
)
|
||||
|
||||
stop_btn.click(
|
||||
fn=stop_game,
|
||||
inputs=[model_selection, game_state],
|
||||
outputs=[game_state, seed_output, block_counter, video_output, action_btn, stop_btn]
|
||||
)
|
||||
|
||||
gr.HTML("""
|
||||
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
|
||||
<p style="font-size: 16px; margin: 0;">Note that this demo is meant to showcase Matrix Game's quality and that under a large number of requests, generation speed may be affected.</p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Matrix Game Gradio Demo")
|
||||
parser.add_argument("--model", type=str, default="Matrix-Game-2.0-Base",
|
||||
choices=list(VARIANT_CONFIG.keys()),
|
||||
help="Model variant to load")
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0")
|
||||
parser.add_argument("--port", type=int, default=7860)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Load the selected model
|
||||
config = VARIANT_CONFIG[args.model]
|
||||
model_path = config["model_path"]
|
||||
|
||||
print(f"Loading model: {model_path}")
|
||||
setup_model_environment(model_path)
|
||||
generator = StreamingVideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
generators = {model_path: generator}
|
||||
|
||||
demo = create_gradio_interface(generators, args.model)
|
||||
|
||||
print(f"Starting Gradio at http://{args.host}:{args.port}")
|
||||
|
||||
# FastAPI Wrapper
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/logo.png")
|
||||
def get_logo():
|
||||
return FileResponse(
|
||||
"assets/full.svg",
|
||||
media_type="image/svg+xml",
|
||||
headers={
|
||||
"Cache-Control": "public, max-age=3600",
|
||||
"Access-Control-Allow-Origin": "*"
|
||||
}
|
||||
)
|
||||
|
||||
@app.get("/favicon.ico")
|
||||
def get_favicon():
|
||||
favicon_path = "assets/icon-simple.svg"
|
||||
|
||||
if os.path.exists(favicon_path):
|
||||
return FileResponse(
|
||||
favicon_path,
|
||||
media_type="image/svg+xml",
|
||||
headers={
|
||||
"Cache-Control": "public, max-age=3600",
|
||||
"Access-Control-Allow-Origin": "*"
|
||||
}
|
||||
)
|
||||
else:
|
||||
raise HTTPException(status_code=404, detail="Favicon not found")
|
||||
|
||||
@app.get("/", response_class=HTMLResponse)
|
||||
def index(request: Request):
|
||||
base_url = str(request.base_url).rstrip('/')
|
||||
return f"""
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
|
||||
<title>FastVideo - Matrix Game 2.0</title>
|
||||
<meta name="title" content="MatrixGame2.0">
|
||||
<meta name="description" content="Make video generation go blurrrrrrr">
|
||||
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, Matrix Game 2.0">
|
||||
|
||||
<meta property="og:type" content="website">
|
||||
<meta property="og:url" content="{base_url}/">
|
||||
<meta property="og:title" content="FastVideo - Matrix Game 2.0">
|
||||
<meta property="og:description" content="Make video generation go blurrrrrrr">
|
||||
<meta property="og:image" content="{base_url}/logo.png">
|
||||
<meta property="og:image:width" content="1200">
|
||||
<meta property="og:image:height" content="630">
|
||||
<meta property="og:site_name" content="MatrixGame2.0">
|
||||
|
||||
<meta property="twitter:card" content="summary_large_image">
|
||||
<meta property="twitter:url" content="{base_url}/">
|
||||
<meta property="twitter:title" content="MatrixGame2.0">
|
||||
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
|
||||
<meta property="twitter:image" content="{base_url}/logo.png">
|
||||
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
|
||||
<link rel="icon" type="image/png" sizes="16x16" href="/favicon.ico">
|
||||
<link rel="apple-touch-icon" href="/favicon.ico">
|
||||
<style>
|
||||
body, html {{
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
height: 100%;
|
||||
overflow: hidden;
|
||||
}}
|
||||
iframe {{
|
||||
width: 100%;
|
||||
height: 100vh;
|
||||
border: none;
|
||||
}}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
|
||||
</body>
|
||||
</html>
|
||||
"""
|
||||
|
||||
app = gr.mount_gradio_app(
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
|
||||
)
|
||||
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,46 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import argparse
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
|
||||
|
||||
|
||||
def main(text_encoder_path: str):
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
# AbsMaxFP8 is the quantization method used by ComfyUI;
|
||||
# check fastvideo/layers/quantization/* for more quantization methods
|
||||
override_text_encoder_quant="AbsMaxFP8",
|
||||
# for Wan 2.2, this is the path to "umt5_xxl_fp8_e4m3fn_scaled.safetensors"
|
||||
override_text_encoder_safetensors=text_encoder_path,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
)
|
||||
|
||||
# I2V is triggered just by passing in an image_path argument
|
||||
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
video = generator.generate_video(
|
||||
prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--text_encoder_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the quantized text encoder safetensors file.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args.text_encoder_path)
|
||||
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.2.1"
|
||||
version = "0.2.2"
|
||||
description = "Unified CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -0,0 +1,311 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from TurboDiffusion SLA implementation
|
||||
# Copyright (c) 2025 by SLA team.
|
||||
#
|
||||
# Citation:
|
||||
# @article{zhang2025sla,
|
||||
# title={SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse-Linear Attention},
|
||||
# author={Jintao Zhang and Haoxu Wang and Kai Jiang and Shuo Yang and Kaiwen Zheng and
|
||||
# Haocheng Xi and Ziteng Wang and Hongzhou Zhu and Min Zhao and Ion Stoica and
|
||||
# Joseph E. Gonzalez and Jun Zhu and Jianfei Chen},
|
||||
# journal={arXiv preprint arXiv:2509.24006},
|
||||
# year={2025}
|
||||
# }
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd(
|
||||
Q, K, V,
|
||||
qk_scale: tl.constexpr,
|
||||
topk: tl.constexpr,
|
||||
LUT, LSE, OS,
|
||||
L: tl.constexpr,
|
||||
M_BLOCKS: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
idx_m = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
qkv_offset = idx_bh * L * D
|
||||
lut_offset = (idx_bh * M_BLOCKS + idx_m) * topk
|
||||
lse_offset = idx_bh * L
|
||||
offs_m = idx_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
Q_ptrs = Q + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
K_ptrs = K + qkv_offset + offs_n[None, :] * D + offs_d[:, None]
|
||||
V_ptrs = V + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
OS_ptrs = OS + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
LUT_ptr = LUT + lut_offset
|
||||
LSE_ptrs = LSE + lse_offset + offs_m
|
||||
|
||||
m_i = tl.full([BLOCK_M], -float('inf'), dtype=tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
|
||||
o_s = tl.zeros([BLOCK_M, D], dtype=tl.float32)
|
||||
|
||||
q = tl.load(Q_ptrs, mask=offs_m[:, None] < L)
|
||||
for block_idx in tl.range(topk):
|
||||
idx_n = tl.load(LUT_ptr + block_idx)
|
||||
n_mask = offs_n < L - idx_n * BLOCK_N
|
||||
|
||||
k = tl.load(K_ptrs + idx_n * BLOCK_N * D, mask=n_mask[None, :])
|
||||
qk = tl.dot(q, k) * (qk_scale * 1.4426950408889634) # = 1 / ln(2)
|
||||
if L - idx_n * BLOCK_N < BLOCK_N:
|
||||
qk = tl.where(n_mask[None, :], qk, float("-inf"))
|
||||
|
||||
v = tl.load(V_ptrs + idx_n * BLOCK_N * D, mask=n_mask[:, None])
|
||||
local_m = tl.max(qk, 1)
|
||||
new_m = tl.maximum(m_i, local_m)
|
||||
qk = qk - new_m[:, None]
|
||||
|
||||
p = tl.math.exp2(qk)
|
||||
l_ij = tl.sum(p, 1)
|
||||
alpha = tl.math.exp2(m_i - new_m)
|
||||
o_s = o_s * alpha[:, None]
|
||||
o_s += tl.dot(p.to(v.dtype), v)
|
||||
|
||||
l_i = l_i * alpha + l_ij
|
||||
m_i = new_m
|
||||
|
||||
o_s = o_s / l_i[:, None]
|
||||
tl.store(OS_ptrs, o_s.to(OS.type.element_ty), mask=offs_m[:, None] < L)
|
||||
|
||||
m_i += tl.math.log2(l_i)
|
||||
tl.store(LSE_ptrs, m_i, mask=offs_m < L)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(
|
||||
OS, DOS, DELTAS,
|
||||
L,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
):
|
||||
idx_m = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
OS += idx_bh * L * D
|
||||
DOS += idx_bh * L * D
|
||||
DELTAS += idx_bh * L
|
||||
|
||||
offs_m = idx_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
o_s = tl.load(OS + offs_m[:, None] * D + offs_d[None, :], mask=offs_m[:, None] < L)
|
||||
do_s = tl.load(DOS + offs_m[:, None] * D + offs_d[None, :], mask=offs_m[:, None] < L)
|
||||
|
||||
delta_s = tl.sum(o_s * do_s, axis=1).to(DELTAS.type.element_ty)
|
||||
tl.store(DELTAS + offs_m, delta_s, mask=offs_m < L)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(
|
||||
Q, K, V, LSE, DELTAS,
|
||||
DOS, DQ, LUT,
|
||||
qk_scale: tl.constexpr,
|
||||
topk: tl.constexpr,
|
||||
L: tl.constexpr,
|
||||
M_BLOCKS: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
idx_m = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
offs_m = idx_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
qkv_offset = idx_bh * L * D
|
||||
lse_offset = idx_bh * L
|
||||
lut_offset = (idx_bh * M_BLOCKS + idx_m) * topk
|
||||
|
||||
Q_ptrs = Q + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
K_ptrs = K + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
V_ptrs = V + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
DQ_ptrs = DQ + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
DOS_ptrs = DOS + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
LSE_ptrs = LSE + lse_offset + offs_m
|
||||
DELTAS_ptrs = DELTAS + lse_offset + offs_m
|
||||
LUT_ptr = LUT + lut_offset
|
||||
|
||||
q = tl.load(Q_ptrs, mask=offs_m[:, None] < L)
|
||||
do_s = tl.load(DOS_ptrs, mask=offs_m[:, None] < L)
|
||||
delta_s = tl.load(DELTAS_ptrs, mask=offs_m < L)
|
||||
lse = tl.load(LSE_ptrs, mask=offs_m < L, other=float("inf"))
|
||||
|
||||
dq = tl.zeros([BLOCK_M, D], dtype=tl.float32)
|
||||
for block_idx in tl.range(topk, num_stages=2):
|
||||
idx_n = tl.load(LUT_ptr + block_idx)
|
||||
n_mask = offs_n < L - idx_n * BLOCK_N
|
||||
|
||||
k = tl.load(K_ptrs + idx_n * BLOCK_N * D, mask=n_mask[:, None])
|
||||
v = tl.load(V_ptrs + idx_n * BLOCK_N * D, mask=n_mask[:, None])
|
||||
qk = tl.dot(q, k.T) * (qk_scale * 1.4426950408889634)
|
||||
p = tl.math.exp2(qk - lse[:, None])
|
||||
p = tl.where(n_mask[None, :], p, 0.0)
|
||||
|
||||
dp = tl.dot(do_s, v.T).to(tl.float32)
|
||||
ds = p * (dp - delta_s[:, None])
|
||||
dq += tl.dot(ds.to(k.dtype), k)
|
||||
tl.store(DQ_ptrs, dq * qk_scale, mask=offs_m[:, None] < L)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(
|
||||
Q, K, V, DOS, DK, DV,
|
||||
qk_scale, KBID, LSE, DELTAS,
|
||||
L: tl.constexpr,
|
||||
M_BLOCKS: tl.constexpr,
|
||||
N_BLOCKS: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
BLOCK_SLICE_FACTOR: tl.constexpr,
|
||||
):
|
||||
BLOCK_M2: tl.constexpr = BLOCK_M // BLOCK_SLICE_FACTOR
|
||||
|
||||
idx_n = tl.program_id(0).to(tl.int64)
|
||||
idx_bh = tl.program_id(1).to(tl.int64)
|
||||
|
||||
offs_n = idx_n * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
offs_m = tl.arange(0, BLOCK_M2)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
qkv_offset = idx_bh * L * D
|
||||
kbid_offset = idx_bh * M_BLOCKS * N_BLOCKS
|
||||
lse_offset = idx_bh * L
|
||||
|
||||
Q_ptrs = Q + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
K_ptrs = K + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
V_ptrs = V + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
DOS_ptrs = DOS + qkv_offset + offs_m[:, None] * D + offs_d[None, :]
|
||||
DK_ptrs = DK + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
DV_ptrs = DV + qkv_offset + offs_n[:, None] * D + offs_d[None, :]
|
||||
LSE_ptrs = LSE + lse_offset + offs_m
|
||||
DELTAS_ptrs = DELTAS + lse_offset + offs_m
|
||||
KBID_ptr = KBID + kbid_offset + idx_n
|
||||
|
||||
k = tl.load(K_ptrs, mask=offs_n[:, None] < L)
|
||||
v = tl.load(V_ptrs, mask=offs_n[:, None] < L)
|
||||
|
||||
dk = tl.zeros([BLOCK_N, D], dtype=tl.float32)
|
||||
dv = tl.zeros([BLOCK_N, D], dtype=tl.float32)
|
||||
for idx_m in tl.range(0, L, BLOCK_M2):
|
||||
kbid = tl.load(KBID_ptr)
|
||||
if kbid == 1:
|
||||
m_mask = offs_m < L - idx_m
|
||||
q = tl.load(Q_ptrs, mask=m_mask[:, None])
|
||||
lse = tl.load(LSE_ptrs, mask=m_mask, other=float("inf"))
|
||||
qkT = tl.dot(k, q.T) * (qk_scale * 1.4426950408889634)
|
||||
pT = tl.math.exp2(qkT - lse[None, :])
|
||||
pT = tl.where(offs_n[:, None] < L, pT, 0.0)
|
||||
|
||||
do = tl.load(DOS_ptrs, mask=m_mask[:, None])
|
||||
dv += tl.dot(pT.to(do.dtype), do)
|
||||
delta = tl.load(DELTAS_ptrs, mask=m_mask)
|
||||
dpT = tl.dot(v, tl.trans(do))
|
||||
dsT = pT * (dpT - delta[None, :])
|
||||
dk += tl.dot(dsT.to(q.dtype), q)
|
||||
|
||||
Q_ptrs += BLOCK_M2 * D
|
||||
DOS_ptrs += BLOCK_M2 * D
|
||||
LSE_ptrs += BLOCK_M2
|
||||
DELTAS_ptrs += BLOCK_M2
|
||||
if (idx_m + BLOCK_M2) % BLOCK_M == 0:
|
||||
KBID_ptr += N_BLOCKS
|
||||
|
||||
tl.store(DK_ptrs, dk * qk_scale, mask=offs_n[:, None] < L)
|
||||
tl.store(DV_ptrs, dv, mask=offs_n[:, None] < L)
|
||||
|
||||
|
||||
class _attention(torch.autograd.Function):
|
||||
"""Sparse attention forward/backward with autograd support."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, k_block_id, lut, topk, BLOCK_M, BLOCK_N, qk_scale=None):
|
||||
assert q.is_contiguous() and k.is_contiguous() and v.is_contiguous()
|
||||
assert k_block_id.is_contiguous() and lut.is_contiguous()
|
||||
|
||||
assert BLOCK_M == 64 or BLOCK_M == 128
|
||||
assert BLOCK_N == 64
|
||||
|
||||
B, H, L, D = q.shape
|
||||
if qk_scale is None:
|
||||
qk_scale = D**-0.5
|
||||
|
||||
M_BLOCKS = triton.cdiv(L, BLOCK_M)
|
||||
|
||||
o_s = torch.empty_like(v)
|
||||
lse = torch.empty(q.shape[:-1], device=q.device, dtype=torch.float32)
|
||||
|
||||
grid = (M_BLOCKS, B * H)
|
||||
_attn_fwd[grid](
|
||||
q, k, v, qk_scale, topk,
|
||||
lut, lse, o_s,
|
||||
L, M_BLOCKS,
|
||||
D, BLOCK_M, BLOCK_N,
|
||||
num_warps=4 if q.shape[-1] == 64 else 8,
|
||||
num_stages=3
|
||||
)
|
||||
|
||||
ctx.save_for_backward(q, k, v, k_block_id, lut, lse, o_s)
|
||||
ctx.qk_scale = qk_scale
|
||||
ctx.topk = topk
|
||||
ctx.BLOCK_M = BLOCK_M
|
||||
ctx.BLOCK_N = BLOCK_N
|
||||
return o_s
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, do_s):
|
||||
q, k, v, k_block_id, lut, lse, o_s = ctx.saved_tensors
|
||||
do_s = do_s.contiguous()
|
||||
|
||||
BLOCK_M, BLOCK_N = ctx.BLOCK_M, ctx.BLOCK_N
|
||||
B, H, L, D = q.shape
|
||||
|
||||
M_BLOCKS = triton.cdiv(L, BLOCK_M)
|
||||
N_BLOCKS = triton.cdiv(L, BLOCK_N)
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
delta_s = torch.empty_like(lse)
|
||||
|
||||
grid = (M_BLOCKS, B * H)
|
||||
_attn_bwd_preprocess[grid](
|
||||
o_s, do_s, delta_s,
|
||||
L, D, BLOCK_M,
|
||||
)
|
||||
|
||||
grid = (M_BLOCKS, B * H)
|
||||
_attn_bwd_dq[grid](
|
||||
q, k, v, lse, delta_s,
|
||||
do_s, dq, lut,
|
||||
ctx.qk_scale, ctx.topk,
|
||||
L, M_BLOCKS,
|
||||
D, BLOCK_M, BLOCK_N,
|
||||
num_warps=4 if q.shape[-1] == 64 else 8,
|
||||
num_stages=4 if q.shape[-1] == 64 else 5
|
||||
)
|
||||
|
||||
grid = (N_BLOCKS, B * H)
|
||||
_attn_bwd_dkdv[grid](
|
||||
q, k, v, do_s, dk, dv,
|
||||
ctx.qk_scale, k_block_id, lse, delta_s,
|
||||
L, M_BLOCKS, N_BLOCKS,
|
||||
D, BLOCK_M, BLOCK_N,
|
||||
BLOCK_SLICE_FACTOR=BLOCK_M // 64,
|
||||
num_warps=4 if q.shape[-1] == 64 else 8,
|
||||
num_stages=4 if q.shape[-1] == 64 else 5
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, None, None, None, None, None
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.2.1"
|
||||
__version__ = "0.2.2"
|
||||
|
||||
@@ -0,0 +1,587 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SLA (Sparse-Linear Attention) backend for FastVideo
|
||||
# Adapted from TurboDiffusion SLA implementation
|
||||
#
|
||||
# Copyright (c) 2025 by SLA team.
|
||||
# Citation:
|
||||
# @article{zhang2025sla,
|
||||
# title={SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse-Linear Attention},
|
||||
# author={Jintao Zhang and Haoxu Wang and Kai Jiang and Shuo Yang and Kaiwen Zheng and
|
||||
# Haocheng Xi and Ziteng Wang and Hongzhou Zhu and Min Zhao and Ion Stoica and
|
||||
# Joseph E. Gonzalez and Jun Zhu and Jianfei Chen},
|
||||
# journal={arXiv preprint arXiv:2509.24006},
|
||||
# year={2025}
|
||||
# }
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo_kernel.triton_kernels.sla_triton import _attention
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# ============================================================================
|
||||
# SLA Utility functions (moved from sla_kernels/utils.py)
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@triton.jit
|
||||
def compress_kernel(
|
||||
X,
|
||||
XM,
|
||||
L: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BLOCK_L: tl.constexpr,
|
||||
):
|
||||
idx_l = tl.program_id(0)
|
||||
idx_bh = tl.program_id(1)
|
||||
|
||||
offs_l = idx_l * BLOCK_L + tl.arange(0, BLOCK_L)
|
||||
offs_d = tl.arange(0, D)
|
||||
|
||||
x_offset = idx_bh * L * D
|
||||
xm_offset = idx_bh * ((L + BLOCK_L - 1) // BLOCK_L) * D
|
||||
x = tl.load(X + x_offset + offs_l[:, None] * D + offs_d[None, :],
|
||||
mask=offs_l[:, None] < L)
|
||||
|
||||
nx = min(BLOCK_L, L - idx_l * BLOCK_L)
|
||||
x_mean = tl.sum(x, axis=0, dtype=tl.float32) / nx
|
||||
tl.store(XM + xm_offset + idx_l * D + offs_d,
|
||||
x_mean.to(XM.dtype.element_ty))
|
||||
|
||||
|
||||
def mean_pool(x: torch.Tensor, BLK: int) -> torch.Tensor:
|
||||
"""Mean pool tensor along sequence dimension with block size BLK."""
|
||||
assert x.is_contiguous()
|
||||
|
||||
B, H, L, D = x.shape
|
||||
L_BLOCKS = (L + BLK - 1) // BLK
|
||||
x_mean = torch.empty((B, H, L_BLOCKS, D), device=x.device, dtype=x.dtype)
|
||||
|
||||
grid = (L_BLOCKS, B * H)
|
||||
compress_kernel[grid](x, x_mean, L, D, BLK)
|
||||
return x_mean
|
||||
|
||||
|
||||
def get_block_map(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
topk_ratio: float,
|
||||
BLKQ: int = 64,
|
||||
BLKK: int = 64,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""Compute sparse block map for attention based on QK similarity.
|
||||
|
||||
Args:
|
||||
q: Query tensor of shape (B, H, L, D)
|
||||
k: Key tensor of shape (B, H, L, D)
|
||||
topk_ratio: Ratio of key blocks to attend to (0-1)
|
||||
BLKQ: Query block size
|
||||
BLKK: Key block size
|
||||
|
||||
Returns:
|
||||
sparse_map: Binary mask of shape (B, H, num_q_blocks, num_k_blocks)
|
||||
lut: Top-k indices of shape (B, H, num_q_blocks, topk)
|
||||
topk: Number of key blocks selected
|
||||
"""
|
||||
arg_k = k - torch.mean(
|
||||
k, dim=-2, keepdim=True) # smooth-k technique from SageAttention
|
||||
pooled_qblocks = mean_pool(q, BLKQ)
|
||||
pooled_kblocks = mean_pool(arg_k, BLKK)
|
||||
pooled_score = pooled_qblocks @ pooled_kblocks.transpose(-1, -2)
|
||||
|
||||
K = pooled_score.shape[-1]
|
||||
topk = min(K, int(topk_ratio * K))
|
||||
lut = torch.topk(pooled_score, topk, dim=-1, sorted=False).indices
|
||||
|
||||
sparse_map = torch.zeros_like(pooled_score, dtype=torch.int8)
|
||||
sparse_map.scatter_(-1, lut, 1)
|
||||
return sparse_map, lut, topk
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# SLA Backend classes
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class SLAAttentionBackend(AttentionBackend):
|
||||
"""Sparse-Linear Attention backend."""
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SLAAttentionImpl"]:
|
||||
return SLAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["SLAAttentionMetadata"]:
|
||||
return SLAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["SLAAttentionMetadataBuilder"]:
|
||||
return SLAAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class SLAAttentionMetadata(AttentionMetadata):
|
||||
"""Metadata for SLA attention."""
|
||||
current_timestep: int
|
||||
topk_ratio: float = 0.5 # Ratio of key blocks to attend to
|
||||
|
||||
|
||||
class SLAAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
"""Builder for SLA attention metadata."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
topk_ratio: float = 0.5,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> SLAAttentionMetadata:
|
||||
return SLAAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
topk_ratio=topk_ratio,
|
||||
)
|
||||
|
||||
|
||||
class SLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
"""SLA attention implementation with learnable linear projection.
|
||||
|
||||
This implementation combines sparse attention with linear attention,
|
||||
using a learnable projection to blend the outputs. The sparse attention
|
||||
uses a block-sparse pattern determined by QK similarity.
|
||||
|
||||
Args:
|
||||
num_heads: Number of attention heads
|
||||
head_size: Dimension of each head
|
||||
topk_ratio: Ratio of key blocks to attend to (0-1), default 0.5
|
||||
feature_map: Feature map for linear attention ('softmax', 'elu', 'relu')
|
||||
BLKQ: Query block size for sparse attention
|
||||
BLKK: Key block size for sparse attention
|
||||
use_bf16: Whether to use bfloat16 for computation
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool = False,
|
||||
softmax_scale: float | None = None,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
# SLA-specific parameters - matched to TurboDiffusion defaults
|
||||
topk_ratio: float = 0.1, # TurboDiffusion uses topk=0.1
|
||||
feature_map: str = "softmax",
|
||||
BLKQ: int = 128, # TurboDiffusion uses BLKQ=128
|
||||
BLKK: int = 64, # TurboDiffusion uses BLKK=64
|
||||
use_bf16: bool = True,
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
nn.Module.__init__(self)
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.softmax_scale = softmax_scale if softmax_scale else head_size**-0.5
|
||||
self.causal = causal
|
||||
self.prefix = prefix
|
||||
|
||||
# SLA-specific config
|
||||
self.topk_ratio = topk_ratio
|
||||
self.BLKQ = BLKQ
|
||||
self.BLKK = BLKK
|
||||
self.dtype = torch.bfloat16 if use_bf16 else torch.float16
|
||||
|
||||
# Learnable linear projection for combining sparse + linear attention
|
||||
self.proj_l = nn.Linear(head_size, head_size, dtype=torch.float32)
|
||||
|
||||
# Feature map for linear attention
|
||||
# Type annotation for callables
|
||||
self.feature_map_q: Callable[[torch.Tensor], torch.Tensor]
|
||||
self.feature_map_k: Callable[[torch.Tensor], torch.Tensor]
|
||||
if feature_map == "elu":
|
||||
self.feature_map_q = lambda x: F.elu(x) + 1
|
||||
self.feature_map_k = lambda x: F.elu(x) + 1
|
||||
elif feature_map == "relu":
|
||||
self.feature_map_q = F.relu
|
||||
self.feature_map_k = F.relu
|
||||
elif feature_map == "softmax":
|
||||
self.feature_map_q = lambda x: F.softmax(x, dim=-1)
|
||||
self.feature_map_k = lambda x: F.softmax(x, dim=-1)
|
||||
else:
|
||||
raise ValueError(f"Unknown feature map: {feature_map}")
|
||||
|
||||
self._init_weights()
|
||||
|
||||
def _init_weights(self) -> None:
|
||||
"""Initialize projection weights to zero for residual-like behavior."""
|
||||
with torch.no_grad():
|
||||
nn.init.zeros_(self.proj_l.weight)
|
||||
nn.init.zeros_(self.proj_l.bias) # type: ignore[arg-type]
|
||||
|
||||
def _calc_linear_attention(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Compute linear attention: (Q @ K^T @ V) / normalizer.
|
||||
|
||||
Args:
|
||||
q: Query tensor (B, H, L, D) after feature map
|
||||
k: Key tensor (B, H, L, D) after feature map
|
||||
v: Value tensor (B, H, L, D)
|
||||
|
||||
Returns:
|
||||
Linear attention output (B, H, L, D)
|
||||
"""
|
||||
kvsum = k.transpose(-1, -2) @ v # (B, H, D, D)
|
||||
ksum = torch.sum(k, dim=-2, keepdim=True) # (B, H, 1, D)
|
||||
return (q @ kvsum) / (1e-5 + (q * ksum).sum(dim=-1, keepdim=True))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass for SLA attention.
|
||||
|
||||
Input tensors are in FastVideo format: (B, L, H, D)
|
||||
Internally converted to SLA format: (B, H, L, D)
|
||||
|
||||
Args:
|
||||
query: Query tensor (B, L, H, D)
|
||||
key: Key tensor (B, L, H, D)
|
||||
value: Value tensor (B, L, H, D)
|
||||
attn_metadata: Attention metadata
|
||||
|
||||
Returns:
|
||||
Output tensor (B, L, H, D)
|
||||
"""
|
||||
original_dtype = query.dtype
|
||||
|
||||
# Convert from FastVideo format (B, L, H, D) to SLA format (B, H, L, D)
|
||||
q = query.transpose(1, 2).contiguous()
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Get topk ratio from metadata if available
|
||||
topk_ratio = self.topk_ratio
|
||||
if hasattr(attn_metadata, 'topk_ratio'):
|
||||
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
|
||||
|
||||
# Compute block-sparse attention pattern
|
||||
sparse_map, lut, real_topk = get_block_map(q,
|
||||
k,
|
||||
topk_ratio=topk_ratio,
|
||||
BLKQ=self.BLKQ,
|
||||
BLKK=self.BLKK)
|
||||
|
||||
# Convert to compute dtype
|
||||
q = q.to(self.dtype)
|
||||
k = k.to(self.dtype)
|
||||
v = v.to(self.dtype)
|
||||
|
||||
# Sparse attention
|
||||
o_s = _attention.apply(q, k, v, sparse_map, lut, real_topk, self.BLKQ,
|
||||
self.BLKK)
|
||||
|
||||
# Linear attention with feature maps
|
||||
q_linear = self.feature_map_q(q).contiguous().to(self.dtype)
|
||||
k_linear = self.feature_map_k(k).contiguous().to(self.dtype)
|
||||
o_l = self._calc_linear_attention(q_linear, k_linear, v)
|
||||
|
||||
# Project linear attention output and combine
|
||||
with torch.amp.autocast('cuda', dtype=self.dtype):
|
||||
o_l = self.proj_l(o_l)
|
||||
|
||||
# Combine sparse and linear outputs
|
||||
output = (o_s + o_l).to(original_dtype)
|
||||
|
||||
# Convert back to FastVideo format (B, L, H, D)
|
||||
output = output.transpose(1, 2)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
# Check if spas_sage_attn is available for SageSLA
|
||||
SAGESLA_ENABLED = True
|
||||
try:
|
||||
import spas_sage_attn._qattn as qattn
|
||||
import spas_sage_attn._fused as fused
|
||||
from spas_sage_attn.utils import get_vanilla_qk_quant, block_map_lut_triton
|
||||
except ImportError:
|
||||
SAGESLA_ENABLED = False
|
||||
|
||||
SAGE2PP_ENABLED = True
|
||||
try:
|
||||
from spas_sage_attn._qattn import qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold
|
||||
except ImportError:
|
||||
SAGE2PP_ENABLED = False
|
||||
|
||||
|
||||
class SageSLAAttentionBackend(AttentionBackend):
|
||||
"""Quantized Sparse-Linear Attention backend using SageAttention kernels."""
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_SLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SageSLAAttentionImpl"]:
|
||||
return SageSLAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["SLAAttentionMetadata"]:
|
||||
return SLAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["SLAAttentionMetadataBuilder"]:
|
||||
return SLAAttentionMetadataBuilder
|
||||
|
||||
|
||||
def _get_cuda_arch(device_index: int) -> str:
|
||||
"""Get CUDA architecture string for the given device."""
|
||||
major, minor = torch.cuda.get_device_capability(device_index)
|
||||
return f"sm{major}{minor}"
|
||||
|
||||
|
||||
class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
"""SageSLA attention implementation using quantized SageAttention kernels.
|
||||
|
||||
This uses INT8 quantization for Q/K and FP8 for V to achieve better performance
|
||||
while maintaining accuracy. Requires spas_sage_attn package.
|
||||
|
||||
Args:
|
||||
num_heads: Number of attention heads
|
||||
head_size: Dimension of each head (must be 64 or 128)
|
||||
topk_ratio: Ratio of key blocks to attend to (0-1), default 0.5
|
||||
feature_map: Feature map for linear attention ('softmax', 'elu', 'relu')
|
||||
use_bf16: Whether to use bfloat16 for computation
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool = False,
|
||||
softmax_scale: float | None = None,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
# SageSLA-specific parameters
|
||||
topk_ratio: float = 0.5,
|
||||
feature_map: str = "softmax",
|
||||
use_bf16: bool = True,
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
nn.Module.__init__(self)
|
||||
|
||||
if not SAGESLA_ENABLED:
|
||||
raise ImportError(
|
||||
"SageSLA requires spas_sage_attn. "
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git"
|
||||
)
|
||||
|
||||
assert head_size in [
|
||||
64, 128
|
||||
], f"SageSLA requires head_size in [64, 128], got {head_size}"
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.softmax_scale = softmax_scale if softmax_scale else head_size**-0.5
|
||||
self.causal = causal
|
||||
self.prefix = prefix
|
||||
|
||||
# SageSLA-specific config
|
||||
self.topk_ratio = topk_ratio
|
||||
self.dtype = torch.bfloat16 if use_bf16 else torch.float16
|
||||
|
||||
# Learnable linear projection for combining sparse + linear attention
|
||||
self.proj_l = nn.Linear(head_size, head_size, dtype=torch.float32)
|
||||
|
||||
# Feature map for linear attention
|
||||
# Type annotation for callables
|
||||
self.feature_map_q: Callable[[torch.Tensor], torch.Tensor]
|
||||
self.feature_map_k: Callable[[torch.Tensor], torch.Tensor]
|
||||
if feature_map == "elu":
|
||||
self.feature_map_q = lambda x: F.elu(x) + 1
|
||||
self.feature_map_k = lambda x: F.elu(x) + 1
|
||||
elif feature_map == "relu":
|
||||
self.feature_map_q = F.relu
|
||||
self.feature_map_k = F.relu
|
||||
elif feature_map == "softmax":
|
||||
self.feature_map_q = lambda x: F.softmax(x, dim=-1)
|
||||
self.feature_map_k = lambda x: F.softmax(x, dim=-1)
|
||||
else:
|
||||
raise ValueError(f"Unknown feature map: {feature_map}")
|
||||
|
||||
self._init_weights()
|
||||
|
||||
def _init_weights(self) -> None:
|
||||
"""Initialize projection weights to zero for residual-like behavior."""
|
||||
with torch.no_grad():
|
||||
nn.init.zeros_(self.proj_l.weight)
|
||||
nn.init.zeros_(self.proj_l.bias) # type: ignore[arg-type]
|
||||
|
||||
def _calc_linear_attention(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Compute linear attention: (Q @ K^T @ V) / normalizer."""
|
||||
kvsum = k.transpose(-1, -2) @ v
|
||||
ksum = torch.sum(k, dim=-2, keepdim=True)
|
||||
return (q @ kvsum) / (1e-5 + (q * ksum).sum(dim=-1, keepdim=True))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass for SageSLA attention with quantized kernels.
|
||||
|
||||
Input tensors are in FastVideo format: (B, L, H, D)
|
||||
|
||||
Args:
|
||||
query: Query tensor (B, L, H, D)
|
||||
key: Key tensor (B, L, H, D)
|
||||
value: Value tensor (B, L, H, D)
|
||||
attn_metadata: Attention metadata
|
||||
|
||||
Returns:
|
||||
Output tensor (B, L, H, D)
|
||||
"""
|
||||
original_dtype = query.dtype
|
||||
|
||||
# Convert from FastVideo format (B, L, H, D) to SLA format (B, H, L, D)
|
||||
q = query.transpose(1, 2).contiguous()
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Get topk ratio from metadata if available
|
||||
topk_ratio = self.topk_ratio
|
||||
if hasattr(attn_metadata, 'topk_ratio'):
|
||||
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
|
||||
|
||||
# Determine block sizes based on GPU architecture
|
||||
arch = _get_cuda_arch(q.device.index)
|
||||
if arch == "sm90":
|
||||
BLKQ, BLKK = 64, 128
|
||||
else:
|
||||
BLKQ, BLKK = 128, 64
|
||||
|
||||
# Compute block-sparse attention pattern
|
||||
sparse_map, lut, real_topk = get_block_map(q,
|
||||
k,
|
||||
topk_ratio=topk_ratio,
|
||||
BLKQ=BLKQ,
|
||||
BLKK=BLKK)
|
||||
|
||||
# Convert to compute dtype
|
||||
q = q.to(self.dtype)
|
||||
k = k.to(self.dtype)
|
||||
v = v.to(self.dtype)
|
||||
|
||||
# ========== SPARGE QUANTIZED ATTENTION ==========
|
||||
km = k.mean(dim=-2, keepdim=True)
|
||||
headdim = q.size(-1)
|
||||
scale = 1.0 / (headdim**0.5)
|
||||
|
||||
# Quantize Q, K to INT8
|
||||
q_int8, q_scale, k_int8, k_scale = get_vanilla_qk_quant(
|
||||
q, k, km, BLKQ, BLKK)
|
||||
lut_triton, valid_block_num = block_map_lut_triton(sparse_map)
|
||||
|
||||
# Quantize V to FP8
|
||||
b, h_kv, kv_len, head_dim = v.shape
|
||||
padded_len = (kv_len + 127) // 128 * 128
|
||||
v_transposed_permutted = torch.empty((b, h_kv, head_dim, padded_len),
|
||||
dtype=v.dtype,
|
||||
device=v.device)
|
||||
fused.transpose_pad_permute_cuda(v, v_transposed_permutted, 1)
|
||||
v_fp8 = torch.empty(v_transposed_permutted.shape,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=v.device)
|
||||
v_scale = torch.empty((b, h_kv, head_dim),
|
||||
dtype=torch.float32,
|
||||
device=v.device)
|
||||
fused.scale_fuse_quant_cuda(v_transposed_permutted, v_fp8, v_scale,
|
||||
kv_len, 2.25, 1)
|
||||
|
||||
# Sparse attention with quantized kernels
|
||||
o_s = torch.empty_like(q)
|
||||
if arch == "sm90":
|
||||
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_sm90(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
q_scale, k_scale, v_scale, 1, False, 1, scale)
|
||||
else:
|
||||
pvthreshold = torch.full((q.shape[-3], ),
|
||||
1e6,
|
||||
dtype=torch.float32,
|
||||
device=q.device)
|
||||
if SAGE2PP_ENABLED:
|
||||
qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
pvthreshold, q_scale, k_scale, v_scale, 1, False, 1, scale,
|
||||
0)
|
||||
else:
|
||||
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
pvthreshold, q_scale, k_scale, v_scale, 1, False, 1, scale,
|
||||
0)
|
||||
# ========== END SPARGE ==========
|
||||
|
||||
# Linear attention with feature maps
|
||||
q_linear = self.feature_map_q(q).contiguous().to(self.dtype)
|
||||
k_linear = self.feature_map_k(k).contiguous().to(self.dtype)
|
||||
o_l = self._calc_linear_attention(q_linear, k_linear, v)
|
||||
|
||||
# Project linear attention output and combine
|
||||
with torch.amp.autocast('cuda', dtype=self.dtype):
|
||||
o_l = self.proj_l(o_l)
|
||||
|
||||
# Combine sparse and linear outputs
|
||||
output = (o_s + o_l).to(original_dtype)
|
||||
|
||||
# Convert back to FastVideo format (B, L, H, D)
|
||||
output = output.transpose(1, 2)
|
||||
|
||||
return output
|
||||
@@ -50,6 +50,10 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
# Register attn_impl as submodule if it has learnable parameters (e.g., SLA's proj_l)
|
||||
# This ensures its parameters are included in state_dict() for saving/loading
|
||||
if isinstance(self.attn_impl, nn.Module):
|
||||
self.add_module('attn_impl', self.attn_impl)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
|
||||
@@ -18,7 +18,8 @@ class DiTArchConfig(ArchConfig):
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE)
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE, AttentionBackendEnum.SLA_ATTN,
|
||||
AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -26,7 +26,6 @@ class LongCatVideoArchConfig(DiTArchConfig):
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Parameter name mapping for weight conversion
|
||||
# Maps original LongCat third_party names -> native FastVideo names
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Embedders
|
||||
|
||||
@@ -17,11 +17,7 @@ from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields.
|
||||
|
||||
NOTE: This is for Phase 1 wrapper compatibility. For native model (Phase 2),
|
||||
use LongCatVideoConfig from fastvideo.configs.models.dits.longcat instead.
|
||||
"""
|
||||
"""Extended DiTArchConfig with LongCat-specific fields."""
|
||||
# LongCat-specific architecture parameters
|
||||
adaln_tembed_dim: int = 512
|
||||
caption_channels: int = 4096
|
||||
@@ -88,20 +84,16 @@ def umt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
|
||||
@dataclass
|
||||
class LongCatT2V480PConfig(PipelineConfig):
|
||||
"""Configuration for LongCat pipeline (480p) aligned to LongCat-Video modules.
|
||||
"""Configuration for LongCat pipeline (480p).
|
||||
|
||||
Components expected by loaders:
|
||||
- tokenizer: AutoTokenizer
|
||||
- text_encoder: UMT5EncoderModel
|
||||
- transformer: LongCatVideoTransformer3DModel (Phase 1 wrapper)
|
||||
OR LongCatTransformer3DModel (Phase 2 native)
|
||||
- transformer: LongCatTransformer3DModel
|
||||
- vae: AutoencoderKLWan (Wan VAE, 4x8 compression)
|
||||
- scheduler: FlowMatchEulerDiscreteScheduler
|
||||
"""
|
||||
|
||||
# DiT config with LongCat-specific arch_config
|
||||
# NOTE: For Phase 1 wrapper, uses LongCatDiTArchConfig
|
||||
# For Phase 2 native model, can use LongCatVideoConfig directly
|
||||
dit_config: DiTConfig = field(
|
||||
default_factory=lambda: DiTConfig(arch_config=LongCatDiTArchConfig()))
|
||||
|
||||
|
||||
@@ -55,11 +55,21 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGameI2V480PConfig,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGameI2V480PConfig,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGameI2V480PConfig,
|
||||
# LongCat Video models
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
|
||||
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"longcatimagetovideo":
|
||||
lambda id: "longcatimagetovideo" in id.lower(),
|
||||
"longcatvideocontinuation":
|
||||
lambda id: "longcatvideocontinuation" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
"hunyuan":
|
||||
lambda id: "hunyuan" in id.lower(),
|
||||
"hunyuan15":
|
||||
@@ -78,13 +88,15 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
lambda id: "stepvideo" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
"turbodiffusion":
|
||||
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"longcatimagetovideo": LongCatT2V480PConfig,
|
||||
"longcatvideocontinuation": LongCatT2V480PConfig,
|
||||
"longcat": LongCatT2V480PConfig,
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
@@ -96,7 +108,8 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
"stepvideo": StepVideoT2VConfig,
|
||||
"turbodiffusion": Wan2_2_I2V_A14B_Config,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -223,7 +223,7 @@ class SamplingParam:
|
||||
help="Path to input image for image-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path",
|
||||
"--video-path",
|
||||
type=str,
|
||||
default=SamplingParam.video_path,
|
||||
help="Path to input video for video-to-video generation",
|
||||
|
||||
@@ -0,0 +1,288 @@
|
||||
import asyncio
|
||||
import os
|
||||
from concurrent.futures import Future, ThreadPoolExecutor
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ForwardBatch
|
||||
from fastvideo.utils import align_to, shallow_asdict
|
||||
from fastvideo.worker.executor import Executor
|
||||
from fastvideo.worker.multiproc_executor import MultiprocExecutor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class IncrementalVideoWriter:
|
||||
|
||||
def __init__(self, path: str, fps: int = 24, block_dir: str | None = None):
|
||||
self._executor = ThreadPoolExecutor(max_workers=2,
|
||||
thread_name_prefix="video_write_")
|
||||
self._path = path
|
||||
self._writer = imageio.get_writer(path, fps=fps, format="mp4")
|
||||
self._pending_main: Future | None = None
|
||||
self._block_dir = block_dir
|
||||
self._block_idx = 0
|
||||
self._fps = fps
|
||||
|
||||
@property
|
||||
def path(self) -> str:
|
||||
return self._path
|
||||
|
||||
def add_frames(self, frames: list[np.ndarray]) -> Future | None:
|
||||
# Wait for previous main video write to complete
|
||||
if self._pending_main is not None:
|
||||
self._pending_main.result()
|
||||
|
||||
# Copy frames to avoid race conditions
|
||||
frames_copy = [f.copy() for f in frames]
|
||||
self._pending_main = self._executor.submit(self._write_frames,
|
||||
frames_copy)
|
||||
|
||||
# Write block file if block_dir is set
|
||||
block_future = None
|
||||
if self._block_dir:
|
||||
self._block_idx += 1
|
||||
block_path = os.path.join(self._block_dir,
|
||||
f"b{self._block_idx}.mp4")
|
||||
block_future = self._executor.submit(self._write_block, frames_copy,
|
||||
block_path)
|
||||
return block_future
|
||||
|
||||
def _write_frames(self, frames: list[np.ndarray]) -> None:
|
||||
for frame in frames:
|
||||
self._writer.append_data(frame)
|
||||
|
||||
def _write_block(self, frames: list[np.ndarray], path: str) -> str:
|
||||
imageio.mimsave(path, frames, fps=self._fps)
|
||||
return path
|
||||
|
||||
def close(self) -> None:
|
||||
if self._pending_main is not None:
|
||||
self._pending_main.result()
|
||||
self._pending_main = None
|
||||
if self._writer:
|
||||
self._writer.close()
|
||||
self._writer = None
|
||||
self._executor.shutdown(wait=True)
|
||||
|
||||
|
||||
class StreamingVideoGenerator(VideoGenerator):
|
||||
"""
|
||||
This class extends VideoGenerator with streaming capabilities,
|
||||
allowing incremental video generation with step-by-step control.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
executor_class: type[Executor],
|
||||
log_stats: bool,
|
||||
use_queue_mode: bool = True):
|
||||
super().__init__(fastvideo_args, executor_class, log_stats)
|
||||
self.accumulated_frames: list[np.ndarray] = []
|
||||
self.sampling_param: SamplingParam | None = None
|
||||
self.batch: ForwardBatch | None = None
|
||||
self._use_queue_mode = use_queue_mode and isinstance(
|
||||
self.executor, MultiprocExecutor)
|
||||
self.writer: IncrementalVideoWriter | None = None
|
||||
self.block_dir: str | None = None
|
||||
self.block_idx: int = 0
|
||||
|
||||
@classmethod
|
||||
def from_fastvideo_args(
|
||||
cls, fastvideo_args: FastVideoArgs) -> "StreamingVideoGenerator":
|
||||
executor_class = Executor.get_class(fastvideo_args)
|
||||
return cls(
|
||||
fastvideo_args=fastvideo_args,
|
||||
executor_class=executor_class,
|
||||
log_stats=False,
|
||||
)
|
||||
|
||||
def reset(
|
||||
self,
|
||||
prompt: str = "A gameplay video of a cyberpunk city",
|
||||
image_path: str | None = None,
|
||||
num_frames: int = 120, # Default max frames
|
||||
**kwargs):
|
||||
self.accumulated_frames = []
|
||||
self.block_idx = 0
|
||||
self.block_dir = None
|
||||
if self.writer:
|
||||
self.writer.close()
|
||||
self.writer = None
|
||||
self.executor.execute_streaming_clear()
|
||||
|
||||
# Handle batch processing from text file
|
||||
if self.sampling_param is None:
|
||||
self.sampling_param = SamplingParam.from_pretrained(
|
||||
self.fastvideo_args.model_path)
|
||||
|
||||
self.sampling_param.update(kwargs)
|
||||
self.sampling_param.prompt = prompt
|
||||
if image_path:
|
||||
self.sampling_param.image_path = image_path
|
||||
self.sampling_param.num_frames = num_frames
|
||||
|
||||
if "output_path" in kwargs:
|
||||
output_path = self._prepare_output_path(kwargs["output_path"],
|
||||
prompt)
|
||||
# Create block directory for individual block files
|
||||
block_dir = output_path.replace(".mp4", "")
|
||||
os.makedirs(block_dir, exist_ok=True)
|
||||
self.block_dir = block_dir
|
||||
self.writer = IncrementalVideoWriter(output_path,
|
||||
fps=24,
|
||||
block_dir=block_dir)
|
||||
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
self.sampling_param.height = align_to(self.sampling_param.height, 16)
|
||||
self.sampling_param.width = align_to(self.sampling_param.width, 16)
|
||||
|
||||
latents_size = [(self.sampling_param.num_frames - 1) // 4 + 1,
|
||||
self.sampling_param.height // 8,
|
||||
self.sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
self.sampling_param.return_frames = True
|
||||
self.sampling_param.save_video = False
|
||||
|
||||
self.batch = ForwardBatch(
|
||||
**shallow_asdict(self.sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
)
|
||||
|
||||
if self._use_queue_mode:
|
||||
self.executor.submit_reset(self.batch, fastvideo_args)
|
||||
result = self.executor.wait_result()
|
||||
if result.error:
|
||||
raise result.error
|
||||
else:
|
||||
self.executor.execute_streaming_reset(self.batch, fastvideo_args)
|
||||
|
||||
def step(
|
||||
self, keyboard_cond: torch.Tensor,
|
||||
mouse_cond: torch.Tensor) -> tuple[list[np.ndarray], Future | None]:
|
||||
if self.batch is None:
|
||||
raise RuntimeError("Call reset() before step()")
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.submit_step(keyboard_cond, mouse_cond)
|
||||
result = self.executor.wait_result()
|
||||
if result.error:
|
||||
raise result.error
|
||||
output_batch = result.output_batch
|
||||
else:
|
||||
# Fallback to RPC-based
|
||||
output_batch = self.executor.execute_streaming_step(
|
||||
keyboard_action=keyboard_cond, mouse_action=mouse_cond)
|
||||
|
||||
frames = self._process_output_batch(output_batch)
|
||||
block_future = None
|
||||
if len(frames) > 0:
|
||||
self.accumulated_frames.extend(frames)
|
||||
self.block_idx += 1
|
||||
if self.writer:
|
||||
# Returns Future for block file, or None if no block_dir
|
||||
block_future = self.writer.add_frames(frames)
|
||||
|
||||
return frames, block_future
|
||||
|
||||
async def step_async(
|
||||
self, keyboard_cond: torch.Tensor,
|
||||
mouse_cond: torch.Tensor) -> tuple[list[np.ndarray], Future | None]:
|
||||
if self.batch is None:
|
||||
raise RuntimeError("Call reset() before step_async()")
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.submit_step(keyboard_cond, mouse_cond)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
result = await loop.run_in_executor(None, self.executor.wait_result)
|
||||
|
||||
if result.error:
|
||||
raise result.error
|
||||
output_batch = result.output_batch
|
||||
else:
|
||||
# Fallback to RPC-based
|
||||
output_batch = await self.executor.execute_streaming_step_async(
|
||||
keyboard_action=keyboard_cond,
|
||||
mouse_action=mouse_cond,
|
||||
)
|
||||
|
||||
frames = self._process_output_batch(output_batch)
|
||||
block_future = None
|
||||
if len(frames) > 0:
|
||||
self.accumulated_frames.extend(frames)
|
||||
self.block_idx += 1
|
||||
if self.writer:
|
||||
block_future = self.writer.add_frames(frames)
|
||||
|
||||
return frames, block_future
|
||||
|
||||
def finalize(self,
|
||||
output_path: str = "streaming_output.mp4",
|
||||
fps: int = 24) -> str:
|
||||
if not self.accumulated_frames:
|
||||
logger.warning("No frames to save.")
|
||||
return ""
|
||||
|
||||
if self.writer:
|
||||
output_path = self.writer.path
|
||||
self.writer.close()
|
||||
self.writer = None
|
||||
logger.info("Saved video to %s", output_path)
|
||||
else:
|
||||
imageio.mimsave(output_path,
|
||||
self.accumulated_frames,
|
||||
fps=fps,
|
||||
format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.submit_clear()
|
||||
else:
|
||||
self.executor.execute_streaming_clear()
|
||||
self.accumulated_frames = []
|
||||
return output_path
|
||||
|
||||
def _process_output_batch(self,
|
||||
output_batch: ForwardBatch) -> list[np.ndarray]:
|
||||
if output_batch.output is None:
|
||||
return []
|
||||
|
||||
samples = output_batch.output
|
||||
# [B, C, T, H, W] or [1, C, T, H, W]
|
||||
if len(samples.shape) == 5:
|
||||
# Rearrange to [T, B, C, H, W] for processing loop
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
else:
|
||||
logger.warning("Unexpected output shape: %s", samples.shape)
|
||||
return []
|
||||
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=1)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).cpu().numpy().astype(np.uint8))
|
||||
|
||||
return frames
|
||||
|
||||
def shutdown(self):
|
||||
if self.writer:
|
||||
self.writer.close()
|
||||
self.writer = None
|
||||
|
||||
if self._use_queue_mode and self.executor._streaming_enabled:
|
||||
self.executor.disable_streaming()
|
||||
|
||||
super().shutdown()
|
||||
@@ -12,6 +12,7 @@ from typing import Any, TYPE_CHECKING
|
||||
from fastvideo.configs.configs import PreprocessConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.configs.utils import clean_cli_args
|
||||
from fastvideo.layers.quantization import QUANTIZATION_METHODS, QuantizationMethods
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
@@ -132,6 +133,7 @@ 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
|
||||
@@ -170,6 +172,10 @@ 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
|
||||
|
||||
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
|
||||
@@ -416,6 +422,11 @@ 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,
|
||||
@@ -476,6 +487,19 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-text-encoder-safetensors",
|
||||
type=str,
|
||||
default=FastVideoArgs.override_text_encoder_safetensors,
|
||||
help="Path to safetensors file for text encoder override",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-text-encoder-quant",
|
||||
type=str,
|
||||
choices=QUANTIZATION_METHODS,
|
||||
default=FastVideoArgs.override_text_encoder_quant,
|
||||
help="Quantization method for text encoder override",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-transformer-cls-name",
|
||||
type=str,
|
||||
@@ -585,6 +609,19 @@ class FastVideoArgs:
|
||||
|
||||
if current_platform.is_mps():
|
||||
self.use_fsdp_inference = False
|
||||
self.dit_layerwise_offload = False
|
||||
|
||||
if self.dit_layerwise_offload:
|
||||
if self.use_fsdp_inference:
|
||||
logger.warning(
|
||||
"dit_layerwise_offload is enabled, automatically disabling use_fsdp_inference."
|
||||
)
|
||||
self.use_fsdp_inference = False
|
||||
if self.dit_cpu_offload:
|
||||
logger.warning(
|
||||
"dit_layerwise_offload is enabled, automatically disabling dit_cpu_offload."
|
||||
)
|
||||
self.dit_cpu_offload = False
|
||||
|
||||
# Validate mode and inference_mode consistency
|
||||
assert isinstance(
|
||||
|
||||
+286
-190
@@ -7,13 +7,20 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.distributed import (divide, get_tp_rank, get_tp_world_size,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from fastvideo.layers.quantization.base_config import (QuantizationConfig,
|
||||
QuantizeMethodBase)
|
||||
from fastvideo.distributed import (
|
||||
divide,
|
||||
get_tp_rank,
|
||||
get_tp_world_size,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
# yapf: disable
|
||||
from fastvideo.models.parameter import (BasevLLMParameter,
|
||||
BlockQuantScaleParameter,
|
||||
@@ -27,12 +34,22 @@ from fastvideo.models.utils import set_weight_attrs
|
||||
logger = init_logger(__name__)
|
||||
|
||||
WEIGHT_LOADER_V2_SUPPORTED = [
|
||||
"CompressedTensorsLinearMethod", "AWQMarlinLinearMethod", "AWQLinearMethod",
|
||||
"GPTQMarlinLinearMethod", "Fp8LinearMethod", "MarlinLinearMethod",
|
||||
"QQQLinearMethod", "GPTQMarlin24LinearMethod", "TPUInt8LinearMethod",
|
||||
"GPTQLinearMethod", "FBGEMMFp8LinearMethod", "ModelOptFp8LinearMethod",
|
||||
"IPEXAWQLinearMethod", "IPEXGPTQLinearMethod", "HQQMarlinMethod",
|
||||
"QuarkLinearMethod"
|
||||
"CompressedTensorsLinearMethod",
|
||||
"AWQMarlinLinearMethod",
|
||||
"AWQLinearMethod",
|
||||
"GPTQMarlinLinearMethod",
|
||||
"Fp8LinearMethod",
|
||||
"MarlinLinearMethod",
|
||||
"QQQLinearMethod",
|
||||
"GPTQMarlin24LinearMethod",
|
||||
"TPUInt8LinearMethod",
|
||||
"GPTQLinearMethod",
|
||||
"FBGEMMFp8LinearMethod",
|
||||
"ModelOptFp8LinearMethod",
|
||||
"IPEXAWQLinearMethod",
|
||||
"IPEXGPTQLinearMethod",
|
||||
"HQQMarlinMethod",
|
||||
"QuarkLinearMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -41,8 +58,8 @@ def adjust_scalar_to_fused_array(
|
||||
shard_id: str | int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""For fused modules (QKV and MLP) we have an array of length
|
||||
N that holds 1 scale for each "logical" matrix. So the param
|
||||
is an array of length N. The loaded_weight corresponds to
|
||||
one of the shards on disk. Here, we slice the param based on
|
||||
is an array of length N. The loaded_weight corresponds to
|
||||
one of the shards on disk. Here, we slice the param based on
|
||||
the shard_id for loading.
|
||||
"""
|
||||
qkv_idxs = {"q": 0, "k": 1, "v": 2}
|
||||
@@ -65,18 +82,23 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
"""Base class for different (maybe quantized) linear methods."""
|
||||
|
||||
@abstractmethod
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
"""Create weights for a linear layer.
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
"""Create weights for a linear layer.
|
||||
The weights will be set as attributes of the layer.
|
||||
|
||||
Args:
|
||||
layer: The layer that is using the LinearMethodBase factory.
|
||||
input_size_per_partition: Size of the weight input dim on rank X.
|
||||
output_partition_sizes: Sizes of the output dim of each logical
|
||||
output_partition_sizes: Sizes of the output dim of each logical
|
||||
weight on rank X. E.g., output_partition_sizes for QKVLinear
|
||||
is a list contains the width of Wq, Wk, Wv on rank X.
|
||||
input_size: Size of the input dim of the weight across all ranks.
|
||||
@@ -86,10 +108,12 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply the weights in layer to the input tensor.
|
||||
Expects create_weights to have been called before on the layer."""
|
||||
raise NotImplementedError
|
||||
@@ -98,28 +122,37 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
class UnquantizedLinearMethod(LinearMethodBase):
|
||||
"""Linear method without quantization."""
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False)
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
weight = Parameter(
|
||||
torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
output = F.linear(x, layer.weight, bias) if torch.cuda.is_available(
|
||||
) or bias is None else F.linear(
|
||||
x, layer.weight, bias.to(x.dtype)
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
output = (
|
||||
F.linear(x, layer.weight, bias) if torch.cuda.is_available()
|
||||
or bias is None else F.linear(x, layer.weight, bias.to(x.dtype))
|
||||
) # NOTE: this line assumes that we are using amp when using cuda and is needed to account for the fact that amp isn't supported in mps
|
||||
return output
|
||||
|
||||
@@ -157,8 +190,8 @@ class LinearBase(torch.nn.Module):
|
||||
self.quant_config = quant_config
|
||||
self.prefix = prefix
|
||||
if quant_config is None:
|
||||
self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
|
||||
)
|
||||
self.quant_method: QuantizeMethodBase | None = (
|
||||
UnquantizedLinearMethod())
|
||||
else:
|
||||
self.quant_method = quant_config.get_quant_method(self,
|
||||
prefix=prefix)
|
||||
@@ -181,29 +214,36 @@ class ReplicatedLinear(LinearBase):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
super().__init__(input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix=prefix)
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__(
|
||||
input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
# All the linear layer supports quant method.
|
||||
assert self.quant_method is not None
|
||||
self.quant_method.create_weights(self,
|
||||
self.input_size, [self.output_size],
|
||||
self.input_size,
|
||||
self.output_size,
|
||||
self.params_dtype,
|
||||
weight_loader=self.weight_loader)
|
||||
self.quant_method.create_weights(
|
||||
self,
|
||||
self.input_size,
|
||||
[self.output_size],
|
||||
self.input_size,
|
||||
self.output_size,
|
||||
self.params_dtype,
|
||||
weight_loader=self.weight_loader,
|
||||
)
|
||||
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
@@ -211,10 +251,13 @@ class ReplicatedLinear(LinearBase):
|
||||
self.output_size,
|
||||
dtype=self.params_dtype,
|
||||
))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
set_weight_attrs(
|
||||
self.bias,
|
||||
{
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
},
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
@@ -265,19 +308,21 @@ class ColumnParallelLinear(LinearBase):
|
||||
output_sizes: list of output sizes packed into one output, like for QKV
|
||||
the list would be size 3.
|
||||
prefix: The name of the layer in the state dict, including all parents
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
output_sizes: list[int] | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
output_sizes: list[int] | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
# Divide the weight matrix along the last dimension.
|
||||
self.tp_size = get_tp_world_size()
|
||||
self.input_size_per_partition = input_size
|
||||
@@ -290,8 +335,14 @@ class ColumnParallelLinear(LinearBase):
|
||||
for output_size in self.output_sizes
|
||||
]
|
||||
|
||||
super().__init__(input_size, output_size, skip_bias_add, params_dtype,
|
||||
quant_config, prefix)
|
||||
super().__init__(
|
||||
input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix,
|
||||
)
|
||||
|
||||
self.gather_output = gather_output
|
||||
|
||||
@@ -308,17 +359,21 @@ class ColumnParallelLinear(LinearBase):
|
||||
params_dtype=self.params_dtype,
|
||||
weight_loader=(
|
||||
self.weight_loader_v2 if self.quant_method.__class__.__name__
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader))
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader),
|
||||
)
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(
|
||||
self.output_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
set_weight_attrs(
|
||||
self.bias,
|
||||
{
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
},
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
@@ -401,32 +456,37 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_sizes: list[int],
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_sizes: list[int],
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
self.output_sizes = output_sizes
|
||||
tp_size = get_tp_world_size()
|
||||
assert all(output_size % tp_size == 0 for output_size in output_sizes)
|
||||
super().__init__(input_size=input_size,
|
||||
output_size=sum(output_sizes),
|
||||
bias=bias,
|
||||
gather_output=gather_output,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None) -> None:
|
||||
super().__init__(
|
||||
input_size=input_size,
|
||||
output_size=sum(output_sizes),
|
||||
bias=bias,
|
||||
gather_output=gather_output,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None,
|
||||
) -> None:
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
# Special case for AQLM codebooks.
|
||||
@@ -518,20 +578,22 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
if (isinstance(param, PackedColumnParameter | PackedvLLMParameter)
|
||||
and param.packed_dim == param.output_dim):
|
||||
shard_size, shard_offset = (
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
shard_size=shard_size, shard_offset=shard_offset))
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(param.output_dim,
|
||||
shard_offset, shard_size)
|
||||
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
|
||||
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None) -> None:
|
||||
def weight_loader_v2(
|
||||
self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None,
|
||||
) -> None:
|
||||
if loaded_shard_id is None:
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight,
|
||||
@@ -568,10 +630,12 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
|
||||
shard_size = self.output_sizes[loaded_shard_id] // tp_size
|
||||
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size)
|
||||
param.load_merged_column_weight(
|
||||
loaded_weight=loaded_weight,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size,
|
||||
)
|
||||
|
||||
|
||||
class QKVParallelLinear(ColumnParallelLinear):
|
||||
@@ -600,16 +664,18 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
hidden_size: int,
|
||||
head_size: int,
|
||||
total_num_heads: int,
|
||||
total_num_kv_heads: int | None = None,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
head_size: int,
|
||||
total_num_heads: int,
|
||||
total_num_kv_heads: int | None = None,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
self.hidden_size = hidden_size
|
||||
self.head_size = head_size
|
||||
self.total_num_heads = total_num_heads
|
||||
@@ -626,29 +692,31 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
self.num_kv_heads = divide(self.total_num_kv_heads, tp_size)
|
||||
self.num_kv_head_replicas = 1
|
||||
input_size = self.hidden_size
|
||||
output_size = (self.num_heads +
|
||||
2 * self.num_kv_heads) * tp_size * self.head_size
|
||||
output_size = ((self.num_heads + 2 * self.num_kv_heads) * tp_size *
|
||||
self.head_size)
|
||||
self.output_sizes = [
|
||||
self.num_heads * self.head_size * tp_size, # q_proj
|
||||
self.num_kv_heads * self.head_size * tp_size, # k_proj
|
||||
self.num_kv_heads * self.head_size * tp_size, # v_proj
|
||||
self.num_kv_heads * self.head_size * tp_size, # v_proj
|
||||
]
|
||||
|
||||
super().__init__(input_size=input_size,
|
||||
output_size=output_size,
|
||||
bias=bias,
|
||||
gather_output=False,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
super().__init__(
|
||||
input_size=input_size,
|
||||
output_size=output_size,
|
||||
bias=bias,
|
||||
gather_output=False,
|
||||
skip_bias_add=skip_bias_add,
|
||||
params_dtype=params_dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
|
||||
shard_offset_mapping = {
|
||||
"q": 0,
|
||||
"k": self.num_heads * self.head_size,
|
||||
"v": (self.num_heads + self.num_kv_heads) * self.head_size,
|
||||
"total": (self.num_heads + 2 * self.num_kv_heads) * self.head_size
|
||||
"total": (self.num_heads + 2 * self.num_kv_heads) * self.head_size,
|
||||
}
|
||||
return shard_offset_mapping.get(loaded_shard_id)
|
||||
|
||||
@@ -663,7 +731,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
def _load_fused_module_from_checkpoint(self, param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor):
|
||||
"""
|
||||
Handle special case for models where QKV layers are already
|
||||
Handle special case for models where QKV layers are already
|
||||
fused on disk. In this case, we have no shard id. This function
|
||||
determmines the shard id by splitting these layers and then calls
|
||||
the weight loader using the shard id.
|
||||
@@ -674,31 +742,39 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offsets = [
|
||||
# (shard_id, shard_offset, shard_size)
|
||||
("q", 0, self.total_num_heads * self.head_size),
|
||||
("k", self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
("v",
|
||||
(self.total_num_heads + self.total_num_kv_heads) * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
(
|
||||
"k",
|
||||
self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
(
|
||||
"v",
|
||||
(self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
]
|
||||
|
||||
for shard_id, shard_offset, shard_size in shard_offsets:
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
if (isinstance(param, PackedColumnParameter | PackedvLLMParameter)
|
||||
and param.packed_dim == param.output_dim):
|
||||
shard_size, shard_offset = (
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
shard_size=shard_size, shard_offset=shard_offset))
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(param.output_dim,
|
||||
shard_offset, shard_size)
|
||||
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
|
||||
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None):
|
||||
def weight_loader_v2(
|
||||
self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None,
|
||||
):
|
||||
if loaded_shard_id is None: # special case for certain models
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0)
|
||||
@@ -715,17 +791,20 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offset = self._get_shard_offset_mapping(loaded_shard_id)
|
||||
shard_size = self._get_shard_size_mapping(loaded_shard_id)
|
||||
|
||||
param.load_qkv_weight(loaded_weight=loaded_weight,
|
||||
num_heads=self.num_kv_head_replicas,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size)
|
||||
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None):
|
||||
param.load_qkv_weight(
|
||||
loaded_weight=loaded_weight,
|
||||
num_heads=self.num_kv_head_replicas,
|
||||
shard_id=loaded_shard_id,
|
||||
shard_offset=shard_offset,
|
||||
shard_size=shard_size,
|
||||
)
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None,
|
||||
):
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
# Special case for AQLM codebooks.
|
||||
@@ -748,14 +827,20 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offsets = [
|
||||
# (shard_id, shard_offset, shard_size)
|
||||
("q", 0, self.total_num_heads * self.head_size),
|
||||
("k", self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
("v", (self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size, self.total_num_kv_heads * self.head_size),
|
||||
(
|
||||
"k",
|
||||
self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
(
|
||||
"v",
|
||||
(self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size,
|
||||
self.total_num_kv_heads * self.head_size,
|
||||
),
|
||||
]
|
||||
|
||||
for shard_id, shard_offset, shard_size in shard_offsets:
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(
|
||||
output_dim, shard_offset, shard_size)
|
||||
self.weight_loader(param, loaded_weight_shard, shard_id)
|
||||
@@ -843,16 +928,18 @@ class RowParallelLinear(LinearBase):
|
||||
quant_config: Quantization configure.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
input_is_parallel: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
reduce_results: bool = True,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
input_is_parallel: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
reduce_results: bool = True,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
# Divide the weight matrix along the first dimension.
|
||||
self.tp_rank = get_tp_rank()
|
||||
self.tp_size = get_tp_world_size()
|
||||
@@ -860,8 +947,14 @@ class RowParallelLinear(LinearBase):
|
||||
self.output_size_per_partition = output_size
|
||||
self.output_partition_sizes = [output_size]
|
||||
|
||||
super().__init__(input_size, output_size, skip_bias_add, params_dtype,
|
||||
quant_config, prefix)
|
||||
super().__init__(
|
||||
input_size,
|
||||
output_size,
|
||||
skip_bias_add,
|
||||
params_dtype,
|
||||
quant_config,
|
||||
prefix,
|
||||
)
|
||||
|
||||
self.input_is_parallel = input_is_parallel
|
||||
self.reduce_results = reduce_results
|
||||
@@ -876,7 +969,8 @@ class RowParallelLinear(LinearBase):
|
||||
params_dtype=self.params_dtype,
|
||||
weight_loader=(
|
||||
self.weight_loader_v2 if self.quant_method.__class__.__name__
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader))
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader),
|
||||
)
|
||||
if not reduce_results and (bias and not skip_bias_add):
|
||||
raise ValueError("When not reduce the results, adding bias to the "
|
||||
"results can lead to incorrect results")
|
||||
@@ -884,10 +978,13 @@ class RowParallelLinear(LinearBase):
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(self.output_size, dtype=params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
set_weight_attrs(
|
||||
self.bias,
|
||||
{
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
},
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
@@ -916,7 +1013,6 @@ class RowParallelLinear(LinearBase):
|
||||
|
||||
def weight_loader_v2(self, param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor):
|
||||
|
||||
# Special case for loading scales off disk, which often do not
|
||||
# have a shape (such as in the case of AutoFP8).
|
||||
if len(loaded_weight.shape) == 0:
|
||||
|
||||
@@ -2,7 +2,7 @@ from typing import Literal, get_args
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig
|
||||
|
||||
QuantizationMethods = Literal[None]
|
||||
QuantizationMethods = Literal[None, "AbsMaxFP8"]
|
||||
|
||||
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
|
||||
|
||||
@@ -50,7 +50,12 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]:
|
||||
if quantization not in QUANTIZATION_METHODS:
|
||||
raise ValueError(f"Invalid quantization method: {quantization}")
|
||||
|
||||
method_to_config: dict[str, type[QuantizationConfig]] = {}
|
||||
# lazy import to avoid triggering `torch.compile` too early
|
||||
from .absmax_fp8 import AbsMaxFP8Config
|
||||
|
||||
method_to_config: dict[str, type[QuantizationConfig]] = {
|
||||
"AbsMaxFP8": AbsMaxFP8Config,
|
||||
}
|
||||
# Update the `method_to_config` with customized quantization methods.
|
||||
method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG)
|
||||
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
from typing import Any
|
||||
import torch
|
||||
from fastvideo.distributed.parallel_state import get_tp_world_size
|
||||
from fastvideo.layers.linear import (
|
||||
LinearBase,
|
||||
LinearMethodBase,
|
||||
MergedColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
)
|
||||
from fastvideo.layers.quantization import QuantizationMethods
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class AbsMaxFP8Config(QuantizationConfig):
|
||||
"""
|
||||
Config class for absmax float8_e4m3fn quantization.
|
||||
Currently only support per-tensor quantization.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "QuantizationConfig":
|
||||
return cls()
|
||||
|
||||
def get_name(self) -> QuantizationMethods:
|
||||
return "AbsMaxFP8"
|
||||
|
||||
def get_supported_act_dtypes(self) -> list[torch.dtype]:
|
||||
return [torch.bfloat16, torch.float16, torch.float32]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 75
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module,
|
||||
prefix: str) -> QuantizeMethodBase | None:
|
||||
if isinstance(layer, LinearBase):
|
||||
return AbsMaxFP8LinearMethod()
|
||||
return None
|
||||
|
||||
|
||||
class AbsMaxFP8Parameter(nn.Parameter):
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: nn.Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
_share_id: str | None = None,
|
||||
) -> None:
|
||||
if len(loaded_weight.shape) == 0:
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
|
||||
assert param.size() == loaded_weight.size(), (
|
||||
f"Tried to load weights of size {loaded_weight.size()}"
|
||||
f"to a parameter of size {param.size()}")
|
||||
param.data.copy_(loaded_weight)
|
||||
|
||||
|
||||
class AbsMaxFP8MergedParameter(nn.Parameter):
|
||||
|
||||
def weight_loader(
|
||||
self,
|
||||
param: nn.Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
share_id: str | int | None = None,
|
||||
) -> None:
|
||||
# currently only support QKVParallelLinear and MergedColumnParallelLinear
|
||||
output_partition_sizes: list[int] = self.output_partition_sizes
|
||||
if share_id is None:
|
||||
share_id = 0
|
||||
if isinstance(share_id, str) and share_id in ["q", "k", "v"]:
|
||||
# QKVParallelLinear case
|
||||
share_idx = ["q", "k", "v"].index(share_id)
|
||||
start_idx = sum(output_partition_sizes[:share_idx])
|
||||
end_idx = start_idx + output_partition_sizes[share_idx]
|
||||
elif isinstance(share_id, int):
|
||||
# MergedColumnParallelLinear case
|
||||
tp_size = get_tp_world_size()
|
||||
if tp_size > 1:
|
||||
# TODO: support this case
|
||||
raise NotImplementedError(
|
||||
"AbsMaxFP8MergedParameter with integer share_id is not supported in tensor parallelism greater than 1 yet."
|
||||
)
|
||||
start_idx = sum(output_partition_sizes[:share_id])
|
||||
end_idx = start_idx + output_partition_sizes[share_id]
|
||||
else:
|
||||
raise ValueError(
|
||||
f"AbsMaxFP8MergedParameter requires share_id to be ['q', 'k', 'v'] or int, got {share_id}."
|
||||
)
|
||||
if len(loaded_weight.shape) == 0:
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
assert loaded_weight.numel() == 1
|
||||
# fill in the corresponding partition by repeating the val
|
||||
param.data[start_idx:end_idx].fill_(loaded_weight.item())
|
||||
|
||||
|
||||
class AbsMaxFP8LinearMethod(LinearMethodBase):
|
||||
"""Linear method with AbsMax FP8 quantization."""
|
||||
|
||||
@staticmethod
|
||||
def _convert_scale(scale: Any) -> torch.nn.Parameter:
|
||||
if scale is None:
|
||||
scale = torch.tensor([1.0], dtype=torch.float32)
|
||||
if not isinstance(scale, torch.Tensor):
|
||||
scale = torch.tensor([scale], dtype=torch.float32)
|
||||
if scale.dtype != torch.float32:
|
||||
raise NotImplementedError("Only float32 scale is supported")
|
||||
return AbsMaxFP8Parameter(scale, requires_grad=False)
|
||||
|
||||
@staticmethod
|
||||
def _merged_placeholder(
|
||||
output_partition_sizes: list[int], ) -> torch.nn.Parameter:
|
||||
scale = torch.ones(
|
||||
sum(output_partition_sizes),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
para = AbsMaxFP8MergedParameter(
|
||||
scale,
|
||||
False,
|
||||
)
|
||||
set_weight_attrs(
|
||||
para,
|
||||
{
|
||||
"output_partition_sizes": output_partition_sizes,
|
||||
},
|
||||
)
|
||||
return para
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
assert params_dtype in [
|
||||
torch.bfloat16, torch.float16, torch.float32
|
||||
], (f"AbsMaxFP8LinearMethod only supports bfloat16, float16, or float32 original dtype, got {params_dtype}."
|
||||
)
|
||||
weight = nn.Parameter(
|
||||
torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
if isinstance(layer, QKVParallelLinear | MergedColumnParallelLinear):
|
||||
scale_weight = self._merged_placeholder(output_partition_sizes, )
|
||||
else:
|
||||
scale_weight = self._convert_scale(
|
||||
extra_weight_attrs.get("scale_weight"))
|
||||
scale_input = self._convert_scale(extra_weight_attrs.get("scale_input"))
|
||||
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
layer.register_parameter("scale_weight", scale_weight)
|
||||
layer.register_parameter("scale_input", scale_input)
|
||||
set_weight_attrs(
|
||||
weight,
|
||||
{
|
||||
"output_dtype": params_dtype,
|
||||
},
|
||||
)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
weight_quant = layer.weight
|
||||
output_dtype: torch.dtype = weight_quant.output_dtype
|
||||
scale_weight: torch.Tensor = layer.scale_weight.data.to(output_dtype)
|
||||
scale_input: torch.Tensor = layer.scale_input.data.to(output_dtype)
|
||||
weight_output_type = weight_quant.to(dtype=output_dtype)
|
||||
weight_final = weight_output_type * scale_weight.unsqueeze(1)
|
||||
x_final = x.to(dtype=output_dtype) * scale_input
|
||||
|
||||
return nn.functional.linear(x_final, weight_final,
|
||||
bias=bias).to(dtype=output_dtype)
|
||||
@@ -1,9 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Native LongCat Video DiT implementation using FastVideo conventions.
|
||||
|
||||
This is a Phase 2 reimplementation that replaces the third_party wrapper
|
||||
with native FastVideo layers for better performance and integration.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
@@ -129,7 +126,7 @@ class TimestepEmbedder(nn.Module):
|
||||
# Sinusoidal embedding in FP32
|
||||
t_freq = self.timestep_embedding(t.flatten(), self.frequency_embedding_size)
|
||||
|
||||
# Cast to model dtype before MLP
|
||||
# Cast to model dtype before MLP (matching original LongCat)
|
||||
# Handle LoRA wrapper if present
|
||||
linear_layer = self.linear_1.base_layer if hasattr(self.linear_1, 'base_layer') else self.linear_1
|
||||
target_dtype = linear_layer.weight.dtype
|
||||
@@ -166,13 +163,14 @@ class CaptionEmbedder(nn.Module):
|
||||
self.text_tokens_zero_pad = text_tokens_zero_pad
|
||||
|
||||
# Two-layer MLP using ReplicatedLinear
|
||||
# CRITICAL: Original LongCat uses GELU(approximate="tanh"), NOT SiLU!
|
||||
self.linear_1 = ReplicatedLinear(
|
||||
caption_channels,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
self.act = nn.SiLU()
|
||||
self.act = nn.GELU(approximate="tanh") # Match original LongCat
|
||||
self.linear_2 = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
@@ -268,10 +266,19 @@ class LongCatSelfAttention(nn.Module):
|
||||
self,
|
||||
x: torch.Tensor, # [B, N, C]
|
||||
latent_shape: tuple, # (T, H, W)
|
||||
num_cond_latents: int = 0, # Number of conditioning latent frames (for I2V)
|
||||
return_kv: bool = False, # Return K/V for caching
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
) -> torch.Tensor | tuple:
|
||||
"""
|
||||
Forward pass with 3D RoPE and optional BSA.
|
||||
|
||||
For I2V mode (num_cond_latents > 0):
|
||||
- Conditioned tokens only attend to themselves
|
||||
- Noise tokens attend to ALL tokens (cond + noise)
|
||||
|
||||
Args:
|
||||
return_kv: If True, return (output, (k_cache, v_cache)) for KV caching
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T, H, W = latent_shape
|
||||
@@ -290,6 +297,12 @@ class LongCatSelfAttention(nn.Module):
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
# Save pre-RoPE K/V for cache if requested (before RoPE is applied)
|
||||
if return_kv:
|
||||
# [B, N, num_heads, head_dim] -> [B, num_heads, N, head_dim]
|
||||
k_cache = k.transpose(1, 2).clone()
|
||||
v_cache = v.transpose(1, 2).clone()
|
||||
|
||||
# For RoPE: need [B, num_heads, N, head_dim]
|
||||
q_rope = q.transpose(1, 2)
|
||||
k_rope = k.transpose(1, 2)
|
||||
@@ -301,6 +314,50 @@ class LongCatSelfAttention(nn.Module):
|
||||
q = q_rope.transpose(1, 2)
|
||||
k = k_rope.transpose(1, 2)
|
||||
|
||||
# === I2V Split Attention ===
|
||||
# For I2V, conditioned tokens and noise tokens are processed separately
|
||||
if num_cond_latents > 0:
|
||||
# Calculate number of conditioned tokens (cond_latents * spatial_tokens_per_frame)
|
||||
num_cond_tokens = num_cond_latents * (N // T)
|
||||
|
||||
# Conditioned tokens: only attend to themselves (same seq length, use self.attn)
|
||||
q_cond = q[:, :num_cond_tokens].contiguous()
|
||||
k_cond = k[:, :num_cond_tokens].contiguous()
|
||||
v_cond = v[:, :num_cond_tokens].contiguous()
|
||||
out_cond, _ = self.attn(q_cond, k_cond, v_cond)
|
||||
|
||||
# Noise tokens: attend to ALL tokens (different seq lengths!)
|
||||
# Need to use flash attention directly since q has different length than k/v
|
||||
q_noise = q[:, num_cond_tokens:].contiguous() # [B, N_noise, num_heads, head_dim]
|
||||
# k, v are full: [B, N, num_heads, head_dim]
|
||||
|
||||
# Transpose for flash attention: [B, num_heads, seq, head_dim]
|
||||
q_noise_t = q_noise.transpose(1, 2)
|
||||
k_t = k.transpose(1, 2)
|
||||
v_t = v.transpose(1, 2)
|
||||
|
||||
# Use scaled dot product attention (handles different q/kv lengths)
|
||||
out_noise_t = torch.nn.functional.scaled_dot_product_attention(
|
||||
q_noise_t, k_t, v_t,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
is_causal=False
|
||||
) # [B, num_heads, N_noise, head_dim]
|
||||
|
||||
# Transpose back: [B, N_noise, num_heads, head_dim]
|
||||
out_noise = out_noise_t.transpose(1, 2)
|
||||
|
||||
# Merge conditioned and noise outputs
|
||||
out = torch.cat([out_cond, out_noise], dim=1)
|
||||
|
||||
# Reshape and project out
|
||||
out = out.reshape(B, N, C)
|
||||
out, _ = self.to_out(out)
|
||||
|
||||
if return_kv:
|
||||
return out, (k_cache, v_cache)
|
||||
return out
|
||||
|
||||
# === Attention: BSA or standard ===
|
||||
if self.enable_bsa and T > 1: # Only use BSA for multi-frame videos
|
||||
# BSA expects [B, H, S, D] format
|
||||
@@ -348,6 +405,96 @@ class LongCatSelfAttention(nn.Module):
|
||||
out = out.reshape(B, N, C)
|
||||
out, _ = self.to_out(out)
|
||||
|
||||
if return_kv:
|
||||
return out, (k_cache, v_cache)
|
||||
return out
|
||||
|
||||
def forward_with_kv_cache(
|
||||
self,
|
||||
x: torch.Tensor, # [B, N_noise, C] - only noise tokens
|
||||
latent_shape: tuple, # (T_noise, H, W) - shape for noise only
|
||||
num_cond_latents: int, # Number of conditioning latent frames
|
||||
kv_cache: tuple, # (k_cond, v_cond) - [B, heads, N_cond, head_dim]
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward using cached K/V from conditioning frames.
|
||||
|
||||
x contains only NOISE tokens.
|
||||
kv_cache contains pre-computed K/V for CONDITIONING tokens.
|
||||
|
||||
CRITICAL: RoPE positions for noise tokens must start AFTER conditioning.
|
||||
We achieve this by padding Q with dummy tokens for conditioning positions,
|
||||
applying RoPE to the full sequence, then extracting only noise token Q.
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T, H, W = latent_shape
|
||||
|
||||
k_cache, v_cache = kv_cache
|
||||
|
||||
# Handle batch size mismatch (cache might be smaller for CFG)
|
||||
# When using CFG, latent_model_input is doubled [neg, pos], but cache is for original batch
|
||||
if k_cache.shape[0] != B:
|
||||
# Expand cache to match input batch size
|
||||
# For CFG: repeat the cache for both negative and positive branches
|
||||
repeat_factor = B // k_cache.shape[0]
|
||||
k_cache = k_cache.repeat(repeat_factor, 1, 1, 1)
|
||||
v_cache = v_cache.repeat(repeat_factor, 1, 1, 1)
|
||||
|
||||
# Project to Q/K/V for noise tokens
|
||||
q, _ = self.to_q(x)
|
||||
k, _ = self.to_k(x)
|
||||
v, _ = self.to_v(x)
|
||||
|
||||
# Reshape to heads: [B, N, num_heads, head_dim]
|
||||
q = q.view(B, N, self.num_heads, self.head_dim)
|
||||
k = k.view(B, N, self.num_heads, self.head_dim)
|
||||
v = v.view(B, N, self.num_heads, self.head_dim)
|
||||
|
||||
# Per-head RMS normalization
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
# Transpose for RoPE: [B, heads, N, head_dim]
|
||||
q_rope = q.transpose(1, 2)
|
||||
k_rope = k.transpose(1, 2)
|
||||
v = v.transpose(1, 2)
|
||||
|
||||
# CRITICAL: Apply RoPE with correct positional offset
|
||||
# Noise frame queries need positions starting from num_cond_latents
|
||||
# Following the original LongCat approach:
|
||||
# 1. Pad Q with dummy tokens matching k_cache shape
|
||||
# 2. Apply RoPE to full sequence (T_cond + T_noise)
|
||||
# 3. Extract only the noise portion of Q
|
||||
|
||||
# Create dummy Q padding to fill conditioning positions
|
||||
# k_cache shape: [B, heads, N_cond, head_dim]
|
||||
q_padding = torch.cat([torch.empty_like(k_cache), q_rope], dim=2).contiguous()
|
||||
|
||||
# Concatenate cached K with noise K for RoPE
|
||||
k_full = torch.cat([k_cache, k_rope], dim=2)
|
||||
v_full = torch.cat([v_cache, v], dim=2)
|
||||
|
||||
# Apply RoPE to full sequence (includes both cond and noise positions)
|
||||
# Grid size: (T_cond + T_noise, H, W)
|
||||
full_T = num_cond_latents + T
|
||||
q_padding, k_full = self.rope_3d(q_padding, k_full, grid_size=(full_T, H, W))
|
||||
|
||||
# Extract only the noise portion of Q (last N tokens)
|
||||
q_rope = q_padding[:, :, -N:].contiguous()
|
||||
|
||||
# Run attention: Q_noise attends to full K/V (cond + noise)
|
||||
out = torch.nn.functional.scaled_dot_product_attention(
|
||||
q_rope, k_full, v_full,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
is_causal=False
|
||||
) # [B, heads, N_noise, head_dim]
|
||||
|
||||
# Transpose back: [B, N_noise, heads, head_dim]
|
||||
out = out.transpose(1, 2)
|
||||
out = out.reshape(B, N, C)
|
||||
out, _ = self.to_out(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
@@ -394,6 +541,8 @@ class LongCatCrossAttention(nn.Module):
|
||||
self,
|
||||
x: torch.Tensor, # [B, N_img, C]
|
||||
context: torch.Tensor, # [B, N_text, C]
|
||||
latent_shape: tuple = None, # (T, H, W) - needed for I2V
|
||||
num_cond_latents: int = 0, # Number of conditioning latent frames (for I2V)
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
@@ -402,9 +551,57 @@ class LongCatCrossAttention(nn.Module):
|
||||
Args:
|
||||
x: Image tokens [B, N_img, C]
|
||||
context: Text tokens [B, N_text, C] (standard padded format)
|
||||
latent_shape: (T, H, W) - needed for calculating num_cond_tokens
|
||||
num_cond_latents: Number of conditioning latent frames (for I2V)
|
||||
|
||||
For I2V mode (num_cond_latents > 0):
|
||||
- Conditioned tokens get ZERO cross-attention output
|
||||
- Only noise tokens get cross-attention with text
|
||||
"""
|
||||
B, N_img, C = x.shape
|
||||
|
||||
# === I2V: Only noise tokens get cross-attention ===
|
||||
if num_cond_latents > 0 and latent_shape is not None:
|
||||
T, H, W = latent_shape
|
||||
num_cond_tokens = num_cond_latents * (N_img // T)
|
||||
|
||||
# Only process noise tokens
|
||||
x_noise = x[:, num_cond_tokens:] # [B, N_noise, C]
|
||||
|
||||
# Project Q, K, V for noise tokens only
|
||||
q, _ = self.to_q(x_noise)
|
||||
k, _ = self.to_k(context)
|
||||
v, _ = self.to_v(context)
|
||||
|
||||
N_text = context.shape[1]
|
||||
N_noise = x_noise.shape[1]
|
||||
|
||||
# Reshape to heads
|
||||
q = q.view(B, N_noise, self.num_heads, self.head_dim)
|
||||
k = k.view(B, N_text, self.num_heads, self.head_dim)
|
||||
v = v.view(B, N_text, self.num_heads, self.head_dim)
|
||||
|
||||
# Per-head RMS normalization
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
# Run cross-attention
|
||||
out_noise = self.attn(q, k, v) # [B, N_noise, num_heads, head_dim]
|
||||
out_noise = out_noise.reshape(B, N_noise, C)
|
||||
out_noise, _ = self.to_out(out_noise)
|
||||
|
||||
# Conditioned tokens get zero output
|
||||
out_cond = torch.zeros(
|
||||
(B, num_cond_tokens, C),
|
||||
dtype=out_noise.dtype,
|
||||
device=out_noise.device
|
||||
)
|
||||
|
||||
# Merge
|
||||
out = torch.cat([out_cond, out_noise], dim=1)
|
||||
return out
|
||||
|
||||
# === Standard cross-attention ===
|
||||
# Project Q, K, V (standard cross-attention like WanVideo/StepVideo/Cosmos)
|
||||
q, _ = self.to_q(x)
|
||||
k, _ = self.to_k(context)
|
||||
@@ -475,19 +672,19 @@ class LongCatSwiGLUFFN(nn.Module):
|
||||
|
||||
def modulate_fp32(norm: nn.Module, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply modulation in FP32 for numerical stability (matching original LongCat).
|
||||
Apply modulation in FP32 for numerical stability.
|
||||
|
||||
shift and scale should already be FP32 from torch.amp.autocast context.
|
||||
Converts inputs to FP32 for the modulation operation, then casts back.
|
||||
"""
|
||||
# Ensure modulation params are FP32 (should be from autocast)
|
||||
assert shift.dtype == torch.float32 and scale.dtype == torch.float32, \
|
||||
f"shift and scale must be FP32, got {shift.dtype} and {scale.dtype}"
|
||||
|
||||
orig_dtype = x.dtype
|
||||
|
||||
# Convert to FP32 for numerical stability
|
||||
shift_fp32 = shift.float()
|
||||
scale_fp32 = scale.float()
|
||||
|
||||
# Normalize and modulate in FP32
|
||||
x_norm = norm(x.to(torch.float32))
|
||||
x_mod = x_norm * (scale + 1) + shift
|
||||
x_mod = x_norm * (scale_fp32 + 1) + shift_fp32
|
||||
|
||||
return x_mod.to(orig_dtype)
|
||||
|
||||
@@ -568,10 +765,21 @@ class LongCatTransformerBlock(nn.Module):
|
||||
context: torch.Tensor, # [B, N_text, C]
|
||||
t: torch.Tensor, # [B, T, C_t]
|
||||
latent_shape: tuple, # (T, H, W)
|
||||
num_cond_latents: int = 0, # Number of conditioning latent frames (for I2V)
|
||||
return_kv: bool = False, # Return K/V for caching
|
||||
kv_cache: tuple | None = None, # Pre-computed K/V cache
|
||||
skip_crs_attn: bool = False, # Skip cross-attention (for cache init)
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
) -> torch.Tensor | tuple:
|
||||
"""
|
||||
Forward pass with AdaLN modulation.
|
||||
|
||||
Args:
|
||||
num_cond_latents: For I2V, number of conditioning latent frames.
|
||||
These frames use split attention behavior.
|
||||
return_kv: If True, return (x, (k_cache, v_cache))
|
||||
kv_cache: Pre-computed K/V from conditioning frames
|
||||
skip_crs_attn: If True, skip cross-attention (used during cache init)
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T, H, W = latent_shape
|
||||
@@ -592,17 +800,47 @@ class LongCatTransformerBlock(nn.Module):
|
||||
x_norm = modulate_fp32(self.norm_attn, x.view(B, T, -1, C), shift_msa, scale_msa)
|
||||
x_norm = x_norm.view(B, N, C)
|
||||
|
||||
attn_out = self.self_attn(x_norm, latent_shape=latent_shape)
|
||||
# Handle KV cache
|
||||
if kv_cache is not None:
|
||||
# Move cache to device if offloaded
|
||||
kv_cache = (kv_cache[0].to(x.device), kv_cache[1].to(x.device))
|
||||
attn_out = self.self_attn.forward_with_kv_cache(
|
||||
x_norm,
|
||||
latent_shape=latent_shape,
|
||||
num_cond_latents=num_cond_latents,
|
||||
kv_cache=kv_cache,
|
||||
)
|
||||
kv_cache_new = None # Don't return cache when using cache
|
||||
else:
|
||||
attn_result = self.self_attn(
|
||||
x_norm,
|
||||
latent_shape=latent_shape,
|
||||
num_cond_latents=num_cond_latents,
|
||||
return_kv=return_kv,
|
||||
)
|
||||
if return_kv:
|
||||
attn_out, kv_cache_new = attn_result
|
||||
else:
|
||||
attn_out = attn_result
|
||||
kv_cache_new = None
|
||||
|
||||
# Residual with gating (CRITICAL: FP32 like original, then cast back)
|
||||
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
|
||||
x = x + (gate_msa * attn_out.view(B, T, -1, C)).view(B, N, C)
|
||||
x = x.to(x_orig_dtype)
|
||||
|
||||
# === Cross-Attention ===
|
||||
x_norm_cross = self.norm_cross(x)
|
||||
cross_out = self.cross_attn(x_norm_cross, context)
|
||||
x = x + cross_out
|
||||
# === Cross-Attention (skip if requested) ===
|
||||
if not skip_crs_attn:
|
||||
x_norm_cross = self.norm_cross(x)
|
||||
# When using KV cache, no need for num_cond_latents in cross-attn
|
||||
cross_num_cond = 0 if kv_cache is not None else num_cond_latents
|
||||
cross_out = self.cross_attn(
|
||||
x_norm_cross,
|
||||
context,
|
||||
latent_shape=latent_shape,
|
||||
num_cond_latents=cross_num_cond
|
||||
)
|
||||
x = x + cross_out
|
||||
|
||||
# === FFN ===
|
||||
x_norm_ffn = modulate_fp32(self.norm_ffn, x.view(B, T, -1, C), shift_mlp, scale_mlp)
|
||||
@@ -615,6 +853,8 @@ class LongCatTransformerBlock(nn.Module):
|
||||
x = x + (gate_mlp * ffn_out.view(B, T, -1, C)).view(B, N, C)
|
||||
x = x.to(x_orig_dtype)
|
||||
|
||||
if return_kv:
|
||||
return x, kv_cache_new
|
||||
return x
|
||||
|
||||
|
||||
@@ -670,16 +910,12 @@ class FinalLayer(nn.Module):
|
||||
B, N, C = x.shape
|
||||
T, _, _ = latent_shape
|
||||
|
||||
# AdaLN modulation (FP32 for stability like original)
|
||||
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
|
||||
t_mod = self.adaln_act(t)
|
||||
mod_params, _ = self.adaln_linear(t_mod)
|
||||
# Ensure FP32 output (needed when LoRA is applied)
|
||||
if mod_params.dtype != torch.float32:
|
||||
mod_params = mod_params.float()
|
||||
shift, scale = mod_params.unsqueeze(2).chunk(2, dim=-1)
|
||||
# AdaLN modulation
|
||||
t_mod = self.adaln_act(t)
|
||||
mod_params, _ = self.adaln_linear(t_mod)
|
||||
shift, scale = mod_params.unsqueeze(2).chunk(2, dim=-1)
|
||||
|
||||
# Modulate
|
||||
# Modulate (converts to FP32 internally for stability)
|
||||
x = modulate_fp32(self.norm, x.view(B, T, -1, C), shift, scale)
|
||||
x = x.reshape(B, N, C)
|
||||
|
||||
@@ -696,8 +932,6 @@ class FinalLayer(nn.Module):
|
||||
class LongCatTransformer3DModel(CachableDiT):
|
||||
"""
|
||||
Native LongCat Video Transformer using FastVideo layers.
|
||||
|
||||
This is a Phase 2 implementation that replaces third_party dependencies.
|
||||
"""
|
||||
|
||||
# FSDP sharding: shard at each transformer block
|
||||
@@ -789,13 +1023,28 @@ class LongCatTransformer3DModel(CachableDiT):
|
||||
encoder_attention_mask: torch.Tensor | None = None, # [B, N_text]
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
guidance: float | None = None, # Unused, for API compatibility
|
||||
num_cond_latents: int = 0, # For I2V: number of conditioning latent frames
|
||||
# === KV Cache Parameters ===
|
||||
return_kv: bool = False, # If True, return (output, kv_cache_dict)
|
||||
kv_cache_dict: dict | None = None, # Pre-computed {block_idx: (k, v)}
|
||||
skip_crs_attn: bool = False, # Skip cross-attention (for cache init)
|
||||
offload_kv_cache: bool = False, # Move cache to CPU after compute
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
) -> torch.Tensor | tuple[torch.Tensor, dict]:
|
||||
"""
|
||||
Forward pass with FastVideo parameter ordering.
|
||||
|
||||
NOTE: This follows FastVideo convention:
|
||||
(hidden_states, encoder_hidden_states, timestep)
|
||||
|
||||
Args:
|
||||
num_cond_latents: For I2V, number of conditioning latent frames.
|
||||
These frames are treated as "clean" (timestep=0)
|
||||
and use split attention behavior.
|
||||
return_kv: If True, return (output, kv_cache_dict)
|
||||
kv_cache_dict: Pre-computed K/V cache {block_idx: (k, v)}
|
||||
skip_crs_attn: If True, skip cross-attention (for cache init)
|
||||
offload_kv_cache: If True, move cache to CPU after compute
|
||||
"""
|
||||
B, _, T, H, W = hidden_states.shape
|
||||
|
||||
@@ -825,12 +1074,31 @@ class LongCatTransformer3DModel(CachableDiT):
|
||||
encoder_attention_mask=encoder_attention_mask
|
||||
) # [B, N_text, C]
|
||||
|
||||
# 4. Transformer blocks
|
||||
# 4. Transformer blocks with optional KV cache
|
||||
kv_cache_dict_ret = {} if return_kv else None
|
||||
|
||||
for i, block in enumerate(self.blocks):
|
||||
x = block(
|
||||
# Get cache for this block if available
|
||||
block_kv_cache = kv_cache_dict.get(i, None) if kv_cache_dict else None
|
||||
|
||||
block_out = block(
|
||||
x, context, t,
|
||||
latent_shape=(N_t, N_h, N_w)
|
||||
latent_shape=(N_t, N_h, N_w),
|
||||
num_cond_latents=num_cond_latents,
|
||||
return_kv=return_kv,
|
||||
kv_cache=block_kv_cache,
|
||||
skip_crs_attn=skip_crs_attn,
|
||||
)
|
||||
|
||||
if return_kv:
|
||||
x, kv_cache = block_out
|
||||
# Store cache
|
||||
if offload_kv_cache:
|
||||
kv_cache_dict_ret[i] = (kv_cache[0].cpu(), kv_cache[1].cpu())
|
||||
else:
|
||||
kv_cache_dict_ret[i] = (kv_cache[0].contiguous(), kv_cache[1].contiguous())
|
||||
else:
|
||||
x = block_out
|
||||
|
||||
# 5. Output projection
|
||||
output = self.final_layer(x, t, latent_shape=(N_t, N_h, N_w))
|
||||
@@ -841,6 +1109,8 @@ class LongCatTransformer3DModel(CachableDiT):
|
||||
# Cast to float32 for better accuracy (as per original)
|
||||
output = output.to(torch.float32)
|
||||
|
||||
if return_kv:
|
||||
return output, kv_cache_dict_ret
|
||||
return output
|
||||
|
||||
def unpatchify(self, x: torch.Tensor, N_t: int, N_h: int, N_w: int) -> torch.Tensor:
|
||||
|
||||
@@ -1,14 +1,13 @@
|
||||
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
|
||||
|
||||
|
||||
@@ -36,6 +35,120 @@ 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):
|
||||
@@ -43,7 +156,6 @@ 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}")
|
||||
@@ -145,7 +257,6 @@ 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 = {}
|
||||
@@ -183,7 +294,6 @@ 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):
|
||||
@@ -199,7 +309,6 @@ 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
|
||||
@@ -244,7 +353,6 @@ 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
|
||||
@@ -282,7 +390,6 @@ 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)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from contextlib import nullcontext
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -733,9 +734,19 @@ class WanTransformer3DModel(CachableDiT):
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, attention_mask)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, attention_mask)
|
||||
offload_mgr = getattr(self, "_layerwise_offload_manager", None)
|
||||
use_offload = offload_mgr is not None and getattr(offload_mgr, "enabled", False)
|
||||
|
||||
for i, block in enumerate(self.blocks):
|
||||
scope = offload_mgr.layer_scope(
|
||||
prefetch_layer_idx=i + 1 if i + 1 < len(self.blocks) else None,
|
||||
release_layer_idx=i,
|
||||
non_blocking=True,
|
||||
) if use_offload else nullcontext()
|
||||
|
||||
with scope:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, attention_mask)
|
||||
# if teacache is enabled, we need to cache the original hidden states
|
||||
|
||||
if enable_teacache:
|
||||
|
||||
+218
-172
@@ -31,8 +31,11 @@ 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
|
||||
@@ -44,6 +47,7 @@ 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
|
||||
@@ -60,17 +64,16 @@ 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:
|
||||
@@ -81,23 +84,21 @@ 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:
|
||||
@@ -109,17 +110,18 @@ 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)
|
||||
|
||||
@@ -132,39 +134,41 @@ 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
|
||||
@@ -191,12 +195,13 @@ 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,
|
||||
@@ -206,21 +211,20 @@ 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:
|
||||
@@ -231,16 +235,18 @@ 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
|
||||
@@ -250,30 +256,32 @@ 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(
|
||||
@@ -286,7 +294,8 @@ 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(
|
||||
@@ -314,8 +323,9 @@ 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.
|
||||
@@ -324,24 +334,27 @@ 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,
|
||||
@@ -353,10 +366,12 @@ 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(
|
||||
@@ -376,17 +391,20 @@ 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(
|
||||
@@ -404,13 +422,14 @@ 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()
|
||||
@@ -419,13 +438,18 @@ 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))
|
||||
|
||||
@@ -435,13 +459,15 @@ 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)
|
||||
@@ -451,37 +477,49 @@ 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,
|
||||
@@ -502,24 +540,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=quant_config,
|
||||
prefix=f"{prefix}.encoder",
|
||||
is_umt5=False)
|
||||
self.encoder = T5Stack(
|
||||
config,
|
||||
False,
|
||||
config.num_layers,
|
||||
self.shared,
|
||||
quant_config=config.quant_config,
|
||||
prefix=f"{prefix}.encoder",
|
||||
is_umt5=False,
|
||||
)
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.shared
|
||||
@@ -545,8 +583,9 @@ 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"),
|
||||
@@ -584,32 +623,33 @@ 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=quant_config,
|
||||
prefix=f"{prefix}.encoder",
|
||||
is_umt5=True)
|
||||
self.encoder = T5Stack(
|
||||
config,
|
||||
False,
|
||||
config.num_layers,
|
||||
self.shared,
|
||||
quant_config=config.quant_config,
|
||||
prefix=f"{prefix}.encoder",
|
||||
is_umt5=True,
|
||||
)
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.shared
|
||||
@@ -635,15 +675,20 @@ 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)
|
||||
@@ -668,8 +713,9 @@ 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
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
# 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)
|
||||
@@ -22,15 +22,21 @@ 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__)
|
||||
|
||||
@@ -45,26 +51,27 @@ 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
|
||||
"""
|
||||
@@ -85,13 +92,16 @@ 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)
|
||||
|
||||
|
||||
@@ -154,36 +164,45 @@ 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,
|
||||
@@ -195,8 +214,9 @@ 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)
|
||||
|
||||
@@ -225,53 +245,94 @@ 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]
|
||||
)
|
||||
|
||||
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)
|
||||
return self.load_model(
|
||||
model_path,
|
||||
encoder_config,
|
||||
target_device,
|
||||
fastvideo_args,
|
||||
encoder_precision,
|
||||
use_text_encoder_override=True,
|
||||
)
|
||||
|
||||
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
|
||||
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
|
||||
)
|
||||
|
||||
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")
|
||||
)
|
||||
|
||||
# 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()
|
||||
|
||||
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 = model_cls(model_config)
|
||||
model: TextEncoder = model_cls(model_config) # type: ignore
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
loaded_weights = model.load_weights(
|
||||
self._get_all_weights(model, model_path,
|
||||
to_cpu=use_cpu_offload))
|
||||
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
|
||||
|
||||
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)
|
||||
@@ -296,7 +357,8 @@ 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",
|
||||
@@ -309,20 +371,22 @@ 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:
|
||||
raise ValueError("Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
if weights_not_loaded and model_config.quant_config is None:
|
||||
raise ValueError(
|
||||
"Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}"
|
||||
)
|
||||
|
||||
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(
|
||||
@@ -345,13 +409,21 @@ 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):
|
||||
@@ -361,9 +433,12 @@ 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
|
||||
|
||||
|
||||
@@ -379,7 +454,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
|
||||
@@ -392,7 +467,9 @@ 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
|
||||
@@ -401,23 +478,32 @@ 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()
|
||||
|
||||
@@ -433,7 +519,8 @@ 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:
|
||||
@@ -450,40 +537,54 @@ 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,
|
||||
@@ -498,15 +599,47 @@ 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
|
||||
|
||||
|
||||
@@ -518,7 +651,9 @@ 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)
|
||||
|
||||
@@ -527,7 +662,8 @@ 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
|
||||
|
||||
|
||||
@@ -540,8 +676,11 @@ 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
|
||||
@@ -551,8 +690,10 @@ 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(
|
||||
@@ -574,18 +715,21 @@ 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
|
||||
"""
|
||||
@@ -597,8 +741,9 @@ 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)
|
||||
|
||||
@@ -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"] # Can be extended as needed
|
||||
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # 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):
|
||||
|
||||
@@ -30,8 +30,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"),
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -75,6 +75,8 @@ _SCHEDULERS = {
|
||||
"SelfForcingFlowMatchScheduler":
|
||||
("schedulers", "scheduling_self_forcing_flow_match",
|
||||
"SelfForcingFlowMatchScheduler"),
|
||||
"RCMScheduler":
|
||||
("schedulers", "scheduling_rcm", "RCMScheduler"),
|
||||
}
|
||||
|
||||
_FAST_VIDEO_MODELS = {
|
||||
|
||||
@@ -0,0 +1,323 @@
|
||||
# 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
|
||||
@@ -1252,6 +1252,52 @@ 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,
|
||||
@@ -1272,3 +1318,4 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z)
|
||||
return dec
|
||||
|
||||
|
||||
@@ -2,5 +2,10 @@
|
||||
"""LongCat pipeline module."""
|
||||
|
||||
from fastvideo.pipelines.basic.longcat.longcat_pipeline import LongCatPipeline
|
||||
from fastvideo.pipelines.basic.longcat.longcat_i2v_pipeline import LongCatImageToVideoPipeline
|
||||
from fastvideo.pipelines.basic.longcat.longcat_vc_pipeline import LongCatVideoContinuationPipeline
|
||||
|
||||
__all__ = ["LongCatPipeline"]
|
||||
__all__ = [
|
||||
"LongCatPipeline", "LongCatImageToVideoPipeline",
|
||||
"LongCatVideoContinuationPipeline"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Image-to-Video pipeline implementation.
|
||||
|
||||
This module implements I2V (Image-to-Video) generation for LongCat using Tier 3
|
||||
conditioning with timestep masking, num_cond_latents support, and RoPE skipping.
|
||||
|
||||
Supports:
|
||||
- Basic I2V (50 steps, guidance_scale=4.0)
|
||||
- Distilled I2V with LoRA (16 steps, guidance_scale=1.0)
|
||||
- Refinement I2V for 720p upscaling (with refinement LoRA + BSA)
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
InputValidationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.longcat_image_vae_encoding import LongCatImageVAEEncodingStage
|
||||
from fastvideo.pipelines.stages.longcat_i2v_latent_preparation import LongCatI2VLatentPreparationStage
|
||||
from fastvideo.pipelines.stages.longcat_i2v_denoising import LongCatI2VDenoisingStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
LongCat Image-to-Video pipeline.
|
||||
|
||||
Generates video from a single input image using Tier 3 I2V conditioning:
|
||||
- Per-frame timestep masking (timestep[:, 0] = 0)
|
||||
- num_cond_latents parameter to transformer
|
||||
- RoPE skipping for conditioning frames
|
||||
- Selective denoising (skip first frame in scheduler)
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize LongCat-specific components."""
|
||||
# Same BSA initialization as base LongCat pipeline
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
transformer = self.get_module("transformer", None)
|
||||
if transformer is None:
|
||||
return
|
||||
|
||||
# Enable BSA if configured
|
||||
if pipeline_config.enable_bsa:
|
||||
bsa_params_cfg = getattr(pipeline_config, 'bsa_params', None) or {}
|
||||
sparsity = getattr(pipeline_config, 'bsa_sparsity', None)
|
||||
cdf_threshold = getattr(pipeline_config, 'bsa_cdf_threshold', None)
|
||||
chunk_q = getattr(pipeline_config, 'bsa_chunk_q', None)
|
||||
chunk_k = getattr(pipeline_config, 'bsa_chunk_k', None)
|
||||
|
||||
effective_bsa_params = dict(bsa_params_cfg) if isinstance(
|
||||
bsa_params_cfg, dict) else {}
|
||||
if sparsity is not None:
|
||||
effective_bsa_params['sparsity'] = sparsity
|
||||
if cdf_threshold is not None:
|
||||
effective_bsa_params['cdf_threshold'] = cdf_threshold
|
||||
if chunk_q is not None:
|
||||
effective_bsa_params['chunk_3d_shape_q'] = chunk_q
|
||||
if chunk_k is not None:
|
||||
effective_bsa_params['chunk_3d_shape_k'] = chunk_k
|
||||
|
||||
# Provide defaults
|
||||
effective_bsa_params.setdefault('sparsity', 0.9375)
|
||||
effective_bsa_params.setdefault('chunk_3d_shape_q', [4, 4, 4])
|
||||
effective_bsa_params.setdefault('chunk_3d_shape_k', [4, 4, 4])
|
||||
|
||||
if hasattr(transformer, 'enable_bsa'):
|
||||
logger.info("Enabling BSA for LongCat I2V transformer")
|
||||
transformer.enable_bsa()
|
||||
if hasattr(transformer, 'blocks'):
|
||||
try:
|
||||
for blk in transformer.blocks:
|
||||
if hasattr(blk, 'self_attn'):
|
||||
blk.self_attn.bsa_params = effective_bsa_params
|
||||
except Exception as e:
|
||||
logger.warning("Failed to set BSA params: %s", e)
|
||||
logger.info("BSA parameters: %s", effective_bsa_params)
|
||||
else:
|
||||
if hasattr(transformer, 'disable_bsa'):
|
||||
transformer.disable_bsa()
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up I2V-specific pipeline stages."""
|
||||
|
||||
# 1. Input validation
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
# 2. Text encoding (same as T2V)
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
# 3. Image VAE encoding (for I2V - skipped in refinement mode)
|
||||
self.add_stage(
|
||||
stage_name="image_vae_encoding_stage",
|
||||
stage=LongCatImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
# 4. Refinement initialization (skipped if not refining)
|
||||
self.add_stage(stage_name="longcat_refine_init_stage",
|
||||
stage=LongCatRefineInitStage(vae=self.get_module("vae")))
|
||||
|
||||
# 5. Timestep preparation (generic)
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
# 6. Refinement timestep override (skipped if not refining)
|
||||
self.add_stage(stage_name="longcat_refine_timestep_stage",
|
||||
stage=LongCatRefineTimestepStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
# 7. Latent preparation with I2V conditioning
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LongCatI2VLatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
# 8. Denoising with I2V support
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=LongCatI2VDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
# 9. Decoding
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
|
||||
EntryClass = LongCatImageToVideoPipeline
|
||||
@@ -1,9 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat video diffusion pipeline implementation (Phase 1: Wrapper).
|
||||
LongCat video diffusion pipeline implementation.
|
||||
|
||||
This module contains a wrapper implementation of the LongCat video diffusion pipeline
|
||||
using FastVideo's modular pipeline architecture with the original LongCat modules.
|
||||
This module implements the LongCat video diffusion pipeline using FastVideo's
|
||||
modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -26,9 +26,6 @@ logger = init_logger(__name__)
|
||||
class LongCatPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
LongCat video diffusion pipeline with LoRA support.
|
||||
|
||||
Phase 1 implementation using wrapper modules from third_party/longcat_video.
|
||||
This validates the pipeline infrastructure before full FastVideo integration.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Video Continuation (VC) pipeline implementation.
|
||||
|
||||
This module implements VC (Video Continuation) generation for LongCat with
|
||||
KV cache optimization for 2-3x speedup.
|
||||
|
||||
Supports:
|
||||
- Basic VC (50 steps, guidance_scale=4.0)
|
||||
- Distilled VC with LoRA (16 steps, guidance_scale=1.0)
|
||||
- KV cache for conditioning frames
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
InputValidationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
|
||||
from fastvideo.pipelines.stages.longcat_i2v_latent_preparation import LongCatI2VLatentPreparationStage
|
||||
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
|
||||
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatVideoContinuationPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
LongCat Video Continuation pipeline.
|
||||
|
||||
Generates video continuation from multiple conditioning frames using
|
||||
optional KV cache for 2-3x speedup.
|
||||
|
||||
Key features:
|
||||
- Takes video input (13+ frames typically)
|
||||
- Encodes conditioning frames via VAE
|
||||
- Optionally pre-computes KV cache for conditioning
|
||||
- Uses cached K/V during denoising for speedup
|
||||
- Concatenates conditioning back after denoising
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize LongCat-specific components."""
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
transformer = self.get_module("transformer", None)
|
||||
if transformer is None:
|
||||
return
|
||||
|
||||
# Enable BSA if configured (for VC, BSA may not be needed)
|
||||
if getattr(pipeline_config, 'enable_bsa', False):
|
||||
bsa_params_cfg = getattr(pipeline_config, 'bsa_params', None) or {}
|
||||
sparsity = getattr(pipeline_config, 'bsa_sparsity', None)
|
||||
cdf_threshold = getattr(pipeline_config, 'bsa_cdf_threshold', None)
|
||||
chunk_q = getattr(pipeline_config, 'bsa_chunk_q', None)
|
||||
chunk_k = getattr(pipeline_config, 'bsa_chunk_k', None)
|
||||
|
||||
effective_bsa_params = dict(bsa_params_cfg) if isinstance(
|
||||
bsa_params_cfg, dict) else {}
|
||||
if sparsity is not None:
|
||||
effective_bsa_params['sparsity'] = sparsity
|
||||
if cdf_threshold is not None:
|
||||
effective_bsa_params['cdf_threshold'] = cdf_threshold
|
||||
if chunk_q is not None:
|
||||
effective_bsa_params['chunk_3d_shape_q'] = chunk_q
|
||||
if chunk_k is not None:
|
||||
effective_bsa_params['chunk_3d_shape_k'] = chunk_k
|
||||
|
||||
# Provide defaults
|
||||
effective_bsa_params.setdefault('sparsity', 0.9375)
|
||||
effective_bsa_params.setdefault('chunk_3d_shape_q', [4, 4, 4])
|
||||
effective_bsa_params.setdefault('chunk_3d_shape_k', [4, 4, 4])
|
||||
|
||||
if hasattr(transformer, 'enable_bsa'):
|
||||
logger.info("Enabling BSA for LongCat VC transformer")
|
||||
transformer.enable_bsa()
|
||||
if hasattr(transformer, 'blocks'):
|
||||
try:
|
||||
for blk in transformer.blocks:
|
||||
if hasattr(blk, 'self_attn'):
|
||||
blk.self_attn.bsa_params = effective_bsa_params
|
||||
except Exception as e:
|
||||
logger.warning("Failed to set BSA params: %s", e)
|
||||
logger.info("BSA parameters: %s", effective_bsa_params)
|
||||
else:
|
||||
if hasattr(transformer, 'disable_bsa'):
|
||||
transformer.disable_bsa()
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up VC-specific pipeline stages."""
|
||||
|
||||
# 1. Input validation
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
# 2. Text encoding
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
# 3. Video VAE encoding (encodes conditioning frames)
|
||||
self.add_stage(
|
||||
stage_name="video_vae_encoding_stage",
|
||||
stage=LongCatVideoVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
# 4. Timestep preparation
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
# 5. Latent preparation (reuse I2V stage - it handles video_latent too)
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LongCatVCLatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
# 6. KV cache initialization (optional, based on config)
|
||||
# This is always added but will skip if use_kv_cache=False
|
||||
self.add_stage(stage_name="kv_cache_init_stage",
|
||||
stage=LongCatKVCacheInitStage(
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
# 7. Denoising with VC and KV cache support
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=LongCatVCDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
# 8. Decoding
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
|
||||
class LongCatVCLatentPreparationStage(LongCatI2VLatentPreparationStage):
|
||||
"""
|
||||
Prepare latents with video conditioning for first N frames.
|
||||
|
||||
Extends I2V latent preparation to handle video_latent (multiple frames)
|
||||
instead of image_latent (single frame).
|
||||
"""
|
||||
|
||||
def forward(self, batch, fastvideo_args):
|
||||
"""Prepare latents with VC conditioning."""
|
||||
|
||||
# Check if we have video_latent (from VC encoding stage)
|
||||
video_latent = getattr(batch, 'video_latent', None)
|
||||
if video_latent is not None:
|
||||
# Set image_latent to video_latent for parent class compatibility
|
||||
batch.image_latent = video_latent
|
||||
|
||||
# Call parent class forward
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
EntryClass = LongCatVideoContinuationPipeline
|
||||
@@ -2,8 +2,9 @@
|
||||
"""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, LoRAPipeline
|
||||
from fastvideo.pipelines import ComposedPipelineBase, ForwardBatch, LoRAPipeline
|
||||
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
InputValidationStage,
|
||||
@@ -69,5 +70,61 @@ 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]
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from fastvideo.pipelines.basic.turbodiffusion.turbodiffusion_pipeline import (
|
||||
TurboDiffusionPipeline, )
|
||||
from fastvideo.pipelines.basic.turbodiffusion.turbodiffusion_i2v_pipeline import (
|
||||
TurboDiffusionI2VPipeline, )
|
||||
|
||||
__all__ = ["TurboDiffusionPipeline", "TurboDiffusionI2VPipeline"]
|
||||
@@ -0,0 +1,89 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
TurboDiffusion I2V (Image-to-Video) Pipeline Implementation.
|
||||
|
||||
This module contains an implementation of the TurboDiffusion I2V pipeline
|
||||
for 1-4 step image-to-video generation using rCM (recurrent Consistency Model)
|
||||
sampling with SLA (Sparse-Linear Attention).
|
||||
|
||||
Key differences from T2V:
|
||||
- Uses dual models (high/low noise) with boundary switching
|
||||
- sigma_max=200 (vs 80 for T2V)
|
||||
- Mask conditioning with encoded first frame
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_rcm import RCMScheduler
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, ImageVAEEncodingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class TurboDiffusionI2VPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
TurboDiffusion I2V pipeline for 1-4 step image-to-video generation.
|
||||
|
||||
Uses RCM scheduler, SLA attention, and dual model switching for
|
||||
high-quality I2V generation.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "transformer_2",
|
||||
"scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Use RCM scheduler with higher sigma_max for I2V
|
||||
logger.info(
|
||||
"Initializing RCM scheduler for TurboDiffusion I2V (sigma_max=200)")
|
||||
self.modules["scheduler"] = RCMScheduler(sigma_max=200.0)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
# I2V: Encode initial image to latent space
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
|
||||
EntryClass = TurboDiffusionI2VPipeline
|
||||
@@ -0,0 +1,76 @@
|
||||
# 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)
|
||||
|
||||
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
|
||||
@@ -23,6 +23,8 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanImageToVideoPipeline": "wan",
|
||||
"WanVideoToVideoPipeline": "wan",
|
||||
"WanCausalDMDPipeline": "wan",
|
||||
"TurboDiffusionPipeline": "turbodiffusion",
|
||||
"TurboDiffusionI2VPipeline": "turbodiffusion",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
"HunyuanVideo15Pipeline": "hunyuan15",
|
||||
@@ -30,6 +32,8 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"MatrixGamePipeline": "matrixgame",
|
||||
"MatrixGameCausalDMDPipeline": "matrixgame",
|
||||
"LongCatPipeline": "longcat",
|
||||
"LongCatImageToVideoPipeline": "longcat",
|
||||
"LongCatVideoContinuationPipeline": "longcat",
|
||||
}
|
||||
|
||||
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
|
||||
|
||||
@@ -28,6 +28,11 @@ from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
from fastvideo.pipelines.stages.timestep_preparation import (
|
||||
TimestepPreparationStage)
|
||||
|
||||
# LongCat stages
|
||||
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
|
||||
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
|
||||
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
|
||||
|
||||
__all__ = [
|
||||
"PipelineStage",
|
||||
"InputValidationStage",
|
||||
@@ -50,4 +55,8 @@ __all__ = [
|
||||
"VideoVAEEncodingStage",
|
||||
"TextEncodingStage",
|
||||
"StepvideoPromptEncodingStage",
|
||||
# LongCat stages
|
||||
"LongCatVideoVAEEncodingStage",
|
||||
"LongCatKVCacheInitStage",
|
||||
"LongCatVCDenoisingStage",
|
||||
]
|
||||
|
||||
@@ -50,32 +50,7 @@ class DecodingStage(PipelineStage):
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@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
|
||||
|
||||
def _denormalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
# denormalization for MatrixGame VAE
|
||||
# z = z * std + mean during decode
|
||||
if (hasattr(self.vae.config, 'latents_mean')
|
||||
@@ -109,6 +84,35 @@ 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",
|
||||
@@ -126,6 +130,57 @@ 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,
|
||||
|
||||
@@ -253,19 +253,37 @@ 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 next(
|
||||
self.transformer.parameters(
|
||||
)).device.type == 'cuda':
|
||||
if (fastvideo_args.dit_cpu_offload
|
||||
and not fastvideo_args.dit_layerwise_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"
|
||||
|
||||
@@ -442,7 +460,6 @@ 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)
|
||||
@@ -459,6 +476,16 @@ 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)
|
||||
@@ -1185,4 +1212,4 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
return batch
|
||||
return batch
|
||||
|
||||
@@ -64,7 +64,7 @@ class LongCatDenoisingStage(DenoisingStage):
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
from fastvideo.models.model_loader import TransformerLoader
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat I2V Denoising Stage with conditioning support.
|
||||
|
||||
This stage implements Tier 3 I2V denoising:
|
||||
1. Per-frame timestep masking (timestep[:, :num_cond_latents] = 0)
|
||||
2. Passes num_cond_latents to transformer (for RoPE skipping)
|
||||
3. Selective denoising (only updates non-conditioned frames)
|
||||
4. CFG-zero optimized guidance
|
||||
"""
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatI2VDenoisingStage(LongCatDenoisingStage):
|
||||
"""
|
||||
LongCat denoising with I2V conditioning support.
|
||||
|
||||
Key modifications from base LongCat denoising:
|
||||
1. Sets timestep=0 for conditioning frames
|
||||
2. Passes num_cond_latents to transformer
|
||||
3. Only applies scheduler step to non-conditioned frames
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run denoising loop with I2V conditioning."""
|
||||
|
||||
# Load transformer if needed
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Setup
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get num_cond_latents from batch
|
||||
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
|
||||
|
||||
if num_cond_latents > 0:
|
||||
logger.info("I2V Denoising: num_cond_latents=%s, latent_shape=%s",
|
||||
num_cond_latents, latents.shape)
|
||||
|
||||
# Prepare negative prompts for CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="I2V Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
|
||||
# 1. Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# 2. Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# 3. CRITICAL: Expand timestep to temporal dimension
|
||||
# and set conditioning frames to timestep=0
|
||||
timestep = timestep.unsqueeze(-1).repeat(
|
||||
1, latent_model_input.shape[2])
|
||||
|
||||
# Mark conditioning frames as clean (timestep=0)
|
||||
if num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
# 4. Run transformer with num_cond_latents
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
num_cond_latents=num_cond_latents,
|
||||
)
|
||||
|
||||
# 5. Apply CFG with optimized scale
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
# CFG-zero optimized scale
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# 6. CRITICAL: Negate for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# 7. CRITICAL: Only update non-conditioned frames
|
||||
# The conditioning frames stay FIXED throughout denoising
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
# No conditioning, update all frames
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,105 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat I2V Latent Preparation Stage.
|
||||
|
||||
This stage prepares latents with image conditioning for the first frame.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatI2VLatentPreparationStage(LatentPreparationStage):
|
||||
"""
|
||||
Prepare latents with image conditioning for first frame.
|
||||
|
||||
This stage:
|
||||
1. Generates random noise for all frames
|
||||
2. Replaces first latent frame with encoded image latent
|
||||
3. Marks conditioning information in batch
|
||||
"""
|
||||
|
||||
# Uses parent __init__ - no need for additional constructor
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Prepare latents with I2V conditioning."""
|
||||
|
||||
# IMPORTANT: Skip if latents already prepared (e.g., by refinement init stage)
|
||||
# The refine_init stage encodes stage1 video and mixes with noise - don't overwrite!
|
||||
if batch.latents is not None:
|
||||
logger.info(
|
||||
"I2V Latent Prep: Skipping - latents already prepared "
|
||||
"(shape=%s), likely from refinement stage", batch.latents.shape)
|
||||
return batch
|
||||
|
||||
# 1. Calculate dimensions
|
||||
num_frames = batch.num_frames
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
# Get VAE compression factors
|
||||
# IMPORTANT: Use VAE's temporal compression (4), NOT transformer's patch_size[0] (1)
|
||||
vae_temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
|
||||
vae_spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
|
||||
num_latent_frames = (num_frames - 1) // vae_temporal_scale + 1
|
||||
latent_height = height // vae_spatial_scale
|
||||
latent_width = width // vae_spatial_scale
|
||||
|
||||
num_channels = self.transformer.config.in_channels
|
||||
|
||||
logger.info(
|
||||
"I2V Latent Prep: num_frames=%s, num_latent_frames=%s "
|
||||
"(vae_temporal_scale=%s), latent_shape=(%s, %s)", num_frames,
|
||||
num_latent_frames, vae_temporal_scale, latent_height, latent_width)
|
||||
|
||||
# 2. Generate random noise for all frames
|
||||
# batch_size might not be set, default to 1
|
||||
batch_size = batch.batch_size if batch.batch_size is not None else 1
|
||||
shape = (batch_size, num_channels, num_latent_frames, latent_height,
|
||||
latent_width)
|
||||
|
||||
# Handle generator - may be a list for batch handling
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list):
|
||||
generator = generator[0] if generator else None
|
||||
|
||||
# torch.randn requires specific argument order: size, generator, dtype
|
||||
latents = torch.randn(*shape,
|
||||
generator=generator).to(get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Replace first frame with conditioned image latent
|
||||
if batch.image_latent is not None:
|
||||
num_cond_latents = batch.num_cond_latents
|
||||
latents[:, :, :
|
||||
num_cond_latents] = batch.image_latent[:, :, :
|
||||
num_cond_latents]
|
||||
|
||||
logger.info(
|
||||
"I2V: Replaced first %s latent frame(s) with image conditioning",
|
||||
num_cond_latents)
|
||||
else:
|
||||
logger.warning(
|
||||
"No image_latent found in batch, proceeding without conditioning"
|
||||
)
|
||||
|
||||
# 4. Store in batch
|
||||
batch.latents = latents
|
||||
|
||||
# Required by base class output validator
|
||||
batch.raw_latent_shape = (num_latent_frames, latent_height,
|
||||
latent_width)
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,162 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Image VAE Encoding Stage for I2V generation.
|
||||
|
||||
This stage handles encoding a single input image to latent space with
|
||||
LongCat-specific normalization for I2V conditioning.
|
||||
"""
|
||||
|
||||
import PIL
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vision_utils import (normalize, numpy_to_pt, pil_to_numpy,
|
||||
resize)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatImageVAEEncodingStage(PipelineStage):
|
||||
"""
|
||||
Encode input image to latent space for I2V conditioning.
|
||||
|
||||
This stage:
|
||||
1. Preprocesses image to match target dimensions
|
||||
2. Encodes via VAE to latent space
|
||||
3. Applies LongCat-specific normalization
|
||||
4. Stores latent and calculates num_cond_latents
|
||||
"""
|
||||
|
||||
def __init__(self, vae):
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Encode image to latent for I2V conditioning."""
|
||||
|
||||
# Skip image encoding for refinement tasks - we're refining an existing video
|
||||
if getattr(batch, 'stage1_video', None) is not None or getattr(
|
||||
batch, 'refine_from', None) is not None:
|
||||
logger.info(
|
||||
"Skipping image encoding - refinement mode (using stage1_video)"
|
||||
)
|
||||
return batch
|
||||
|
||||
# 1. Get image from batch
|
||||
image = batch.pil_image # PIL.Image
|
||||
if image is None:
|
||||
raise ValueError("pil_image must be provided for I2V")
|
||||
|
||||
if not isinstance(image, PIL.Image.Image):
|
||||
raise TypeError(f"pil_image must be PIL.Image, got {type(image)}")
|
||||
|
||||
# 2. Get target dimensions
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
if height is None or width is None:
|
||||
raise ValueError("height and width must be set for I2V")
|
||||
|
||||
# 3. Preprocess image
|
||||
image = resize(image, height, width, resize_mode="default")
|
||||
image = pil_to_numpy(image)
|
||||
image = numpy_to_pt(image)
|
||||
image = normalize(image) # to [-1, 1]
|
||||
|
||||
# 4. Add temporal dimension
|
||||
# After numpy_to_pt: [1, C, H, W] (batch already added by pil_to_numpy)
|
||||
# Add T dimension: [1, C, H, W] -> [1, C, 1, H, W] = [B, C, T, H, W]
|
||||
image = image.unsqueeze(2)
|
||||
image = image.to(get_local_torch_device(), dtype=torch.float32)
|
||||
|
||||
# 5. Encode via VAE
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
|
||||
if not vae_autocast_enabled:
|
||||
image = image.to(vae_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
encoder_output = self.vae.encode(image)
|
||||
latent = self.retrieve_latents(encoder_output, batch.generator)
|
||||
|
||||
# 6. Apply LongCat-specific normalization
|
||||
# Formula: (latents - mean) / std
|
||||
latent = self.normalize_latents(latent)
|
||||
|
||||
# 7. Calculate num_cond_latents
|
||||
# Formula: 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
# For single image (num_cond_frames=1): always 1 latent frame
|
||||
num_cond_frames = 1 # Single image
|
||||
vae_temporal_scale = self.vae.config.scale_factor_temporal
|
||||
batch.num_cond_latents = 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
|
||||
# 8. Store in batch
|
||||
batch.image_latent = latent
|
||||
batch.num_cond_frames = 1
|
||||
|
||||
logger.info(
|
||||
"I2V: Encoded image to latent shape %s, num_cond_latents=%s",
|
||||
latent.shape, batch.num_cond_latents)
|
||||
|
||||
# Offload VAE if needed
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self, encoder_output: object,
|
||||
generator: torch.Generator | None) -> torch.Tensor:
|
||||
"""Sample from VAE posterior."""
|
||||
# WAN VAE returns an object with .sample() method
|
||||
if hasattr(encoder_output, 'sample'):
|
||||
return encoder_output.sample(generator)
|
||||
elif hasattr(encoder_output, 'latent_dist'):
|
||||
return encoder_output.latent_dist.sample(generator)
|
||||
elif hasattr(encoder_output, 'latents'):
|
||||
return encoder_output.latents
|
||||
else:
|
||||
raise AttributeError("Could not access latents from encoder output")
|
||||
|
||||
def normalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply LongCat-specific latent normalization.
|
||||
|
||||
Formula: (latents - mean) / std
|
||||
|
||||
This matches the original LongCat implementation and is DIFFERENT
|
||||
from standard VAE scaling (which uses scaling_factor).
|
||||
"""
|
||||
if not hasattr(self.vae.config, 'latents_mean') or not hasattr(
|
||||
self.vae.config, 'latents_std'):
|
||||
raise ValueError(
|
||||
"VAE config must have 'latents_mean' and 'latents_std' "
|
||||
"for LongCat normalization")
|
||||
|
||||
latents_mean = torch.tensor(self.vae.config.latents_mean).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
latents_std = torch.tensor(self.vae.config.latents_std).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
return (latents - latents_mean) / latents_std
|
||||
@@ -0,0 +1,123 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat KV Cache Initialization Stage for Video Continuation (VC).
|
||||
|
||||
This stage pre-computes K/V cache for conditioning frames.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatKVCacheInitStage(PipelineStage):
|
||||
"""
|
||||
Pre-compute KV cache for conditioning frames.
|
||||
|
||||
After this stage:
|
||||
- batch.kv_cache_dict contains {block_idx: (k, v)}
|
||||
- batch.cond_latents contains the conditioning latents
|
||||
- batch.latents contains ONLY noise latents
|
||||
"""
|
||||
|
||||
def __init__(self, transformer):
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Initialize KV cache from conditioning latents."""
|
||||
|
||||
# Check if KV cache is enabled
|
||||
use_kv_cache = getattr(fastvideo_args.pipeline_config, 'use_kv_cache',
|
||||
True)
|
||||
if not use_kv_cache:
|
||||
batch.kv_cache_dict = {}
|
||||
batch.use_kv_cache = False
|
||||
logger.info("KV cache disabled, skipping initialization")
|
||||
return batch
|
||||
|
||||
batch.use_kv_cache = True
|
||||
offload_kv_cache = getattr(fastvideo_args.pipeline_config,
|
||||
'offload_kv_cache', False)
|
||||
|
||||
# Get conditioning latents
|
||||
num_cond_latents = batch.num_cond_latents
|
||||
if num_cond_latents <= 0:
|
||||
batch.kv_cache_dict = {}
|
||||
logger.warning("num_cond_latents <= 0, skipping KV cache init")
|
||||
return batch
|
||||
|
||||
# Extract conditioning latents
|
||||
cond_latents = batch.latents[:, :, :num_cond_latents].clone()
|
||||
|
||||
logger.info(
|
||||
"Initializing KV cache for %d conditioning latents, shape: %s",
|
||||
num_cond_latents, cond_latents.shape)
|
||||
|
||||
# Timestep = 0 for conditioning (they are "clean")
|
||||
B = cond_latents.shape[0]
|
||||
T_cond = cond_latents.shape[2]
|
||||
timestep = torch.zeros(B,
|
||||
T_cond,
|
||||
device=cond_latents.device,
|
||||
dtype=cond_latents.dtype)
|
||||
|
||||
# Empty prompt embeddings (cross-attn will be skipped)
|
||||
max_seq_len = 512
|
||||
# Get caption dimension from transformer config
|
||||
caption_dim = self.transformer.config.caption_channels
|
||||
empty_embeds = torch.zeros(B,
|
||||
max_seq_len,
|
||||
caption_dim,
|
||||
device=cond_latents.device,
|
||||
dtype=cond_latents.dtype)
|
||||
|
||||
# Get transformer dtype
|
||||
if hasattr(self.transformer, 'module'):
|
||||
transformer_dtype = next(self.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
|
||||
# Run transformer with return_kv=True, skip_crs_attn=True
|
||||
with (
|
||||
torch.no_grad(),
|
||||
set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
),
|
||||
torch.autocast(device_type='cuda', dtype=transformer_dtype),
|
||||
):
|
||||
_, kv_cache_dict = self.transformer(
|
||||
hidden_states=cond_latents.to(transformer_dtype),
|
||||
encoder_hidden_states=empty_embeds.to(transformer_dtype),
|
||||
timestep=timestep.to(transformer_dtype),
|
||||
return_kv=True,
|
||||
skip_crs_attn=True,
|
||||
offload_kv_cache=offload_kv_cache,
|
||||
)
|
||||
|
||||
# Store cache and save cond_latents for later concatenation
|
||||
batch.kv_cache_dict = kv_cache_dict
|
||||
batch.cond_latents = cond_latents
|
||||
|
||||
# Remove conditioning latents from main latents
|
||||
# After this, batch.latents contains ONLY noise frames
|
||||
batch.latents = batch.latents[:, :, num_cond_latents:]
|
||||
|
||||
logger.info(
|
||||
"KV cache initialized: %d blocks, offload=%s, remaining latents shape: %s",
|
||||
len(kv_cache_dict), offload_kv_cache, batch.latents.shape)
|
||||
|
||||
return batch
|
||||
@@ -257,12 +257,16 @@ class LongCatRefineInitStage(PipelineStage):
|
||||
num_cond_frames_added, num_noise_frames_added,
|
||||
new_num_frames)
|
||||
|
||||
# VAE encode
|
||||
logger.info("Encoding stage1 video with VAE...")
|
||||
# VAE encode with tiling for memory efficiency
|
||||
logger.info("Encoding stage1 video with VAE (tiling enabled)...")
|
||||
vae_dtype = next(self.vae.parameters()).dtype
|
||||
vae_device = next(self.vae.parameters()).device
|
||||
video_up = video_up.to(dtype=vae_dtype, device=vae_device)
|
||||
|
||||
# Enable tiling for large video encoding
|
||||
if hasattr(self.vae, 'enable_tiling'):
|
||||
self.vae.enable_tiling()
|
||||
|
||||
with torch.no_grad():
|
||||
latent_dist = self.vae.encode(video_up)
|
||||
# Extract tensor from latent distribution
|
||||
@@ -301,10 +305,14 @@ class LongCatRefineInitStage(PipelineStage):
|
||||
|
||||
logger.info("Applied t_thresh=%s noise mixing", t_thresh)
|
||||
|
||||
# Store in batch
|
||||
batch.latents = latent_up.to(dtype)
|
||||
# Store in batch - ensure correct dtype and device
|
||||
# The latents need to be on the same device as the transformer (CUDA)
|
||||
target_device = batch.prompt_embeds[0].device
|
||||
batch.latents = latent_up.to(device=target_device, dtype=dtype)
|
||||
batch.raw_latent_shape = latent_up.shape
|
||||
|
||||
logger.info("Latents device: %s, dtype: %s", batch.latents.device,
|
||||
batch.latents.dtype)
|
||||
logger.info("LongCat refinement initialization complete")
|
||||
|
||||
return batch
|
||||
|
||||
@@ -0,0 +1,217 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat VC Denoising Stage with KV cache support.
|
||||
|
||||
This stage extends the I2V denoising stage to support:
|
||||
1. KV cache for conditioning frames
|
||||
2. Video continuation with multiple conditioning frames
|
||||
"""
|
||||
|
||||
import time
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatVCDenoisingStage(LongCatDenoisingStage):
|
||||
"""
|
||||
LongCat denoising with Video Continuation and KV cache support.
|
||||
|
||||
Key differences from I2V denoising:
|
||||
- Supports KV cache (reuses cached K/V from conditioning frames)
|
||||
- Handles larger num_cond_latents
|
||||
- Concatenates conditioning latents back after denoising
|
||||
|
||||
When use_kv_cache=True:
|
||||
- batch.latents contains ONLY noise frames (cond removed by KV cache init)
|
||||
- batch.kv_cache_dict contains cached K/V
|
||||
- batch.cond_latents contains conditioning latents for post-concat
|
||||
|
||||
When use_kv_cache=False:
|
||||
- batch.latents contains ALL frames (cond + noise)
|
||||
- Timestep masking: timestep[:, :num_cond_latents] = 0
|
||||
- Selective denoising: only update noise frames
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run denoising loop with VC conditioning and optional KV cache."""
|
||||
|
||||
# Load transformer if needed
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Setup
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_attention_mask = (batch.prompt_attention_mask[0]
|
||||
if batch.prompt_attention_mask else None)
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get VC-specific parameters
|
||||
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
|
||||
use_kv_cache = getattr(batch, 'use_kv_cache', False)
|
||||
kv_cache_dict = getattr(batch, 'kv_cache_dict', {})
|
||||
|
||||
logger.info(
|
||||
"VC Denoising: num_cond_latents=%d, use_kv_cache=%s, latent_shape=%s",
|
||||
num_cond_latents, use_kv_cache, latents.shape)
|
||||
|
||||
# Prepare negative prompts for CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
step_times = []
|
||||
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="VC Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
step_start = time.time()
|
||||
|
||||
# 1. Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# 2. Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# 3. Expand timestep to temporal dimension
|
||||
timestep = timestep.unsqueeze(-1).repeat(
|
||||
1, latent_model_input.shape[2])
|
||||
|
||||
# 4. Timestep masking (only when NOT using KV cache)
|
||||
if not use_kv_cache and num_cond_latents > 0:
|
||||
timestep[:, :num_cond_latents] = 0
|
||||
|
||||
# 5. Prepare transformer kwargs
|
||||
# IMPORTANT: num_cond_latents is ALWAYS passed - needed for RoPE position offset
|
||||
transformer_kwargs = {
|
||||
'num_cond_latents': num_cond_latents,
|
||||
}
|
||||
if use_kv_cache:
|
||||
transformer_kwargs['kv_cache_dict'] = kv_cache_dict
|
||||
|
||||
# 6. Run transformer
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
**transformer_kwargs,
|
||||
)
|
||||
|
||||
# 7. Apply CFG with optimized scale (CFG-zero)
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# 8. Negate for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# 9. Scheduler step
|
||||
if use_kv_cache:
|
||||
# All latents are noise frames (conditioning is in cache)
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
else:
|
||||
# Only update noise frames (skip conditioning)
|
||||
if num_cond_latents > 0:
|
||||
latents[:, :, num_cond_latents:] = self.scheduler.step(
|
||||
noise_pred[:, :, num_cond_latents:],
|
||||
t,
|
||||
latents[:, :, num_cond_latents:],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
else:
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
step_time = time.time() - step_start
|
||||
step_times.append(step_time)
|
||||
|
||||
# Log timing for first few steps
|
||||
if i < 3:
|
||||
logger.info("Step %d: %.2fs", i, step_time)
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# 10. If using KV cache, concatenate conditioning latents back
|
||||
if use_kv_cache and hasattr(
|
||||
batch, 'cond_latents') and batch.cond_latents is not None:
|
||||
latents = torch.cat([batch.cond_latents, latents], dim=2)
|
||||
logger.info(
|
||||
"Concatenated conditioning latents back, final shape: %s",
|
||||
latents.shape)
|
||||
|
||||
# Log average timing
|
||||
avg_time = sum(step_times) / len(step_times)
|
||||
logger.info("Average step time: %.2fs (total: %.1fs)", avg_time,
|
||||
sum(step_times))
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Video VAE Encoding Stage for Video Continuation (VC) generation.
|
||||
|
||||
This stage handles encoding multiple video frames to latent space with
|
||||
LongCat-specific normalization for VC conditioning.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vision_utils import normalize, numpy_to_pt, pil_to_numpy, resize
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatVideoVAEEncodingStage(PipelineStage):
|
||||
"""
|
||||
Encode video frames to latent space for VC conditioning.
|
||||
|
||||
This stage:
|
||||
1. Loads video frames from path or uses provided frames
|
||||
2. Takes the last num_cond_frames from the video
|
||||
3. Preprocesses and stacks frames
|
||||
4. Encodes via VAE to latent space
|
||||
5. Applies LongCat-specific normalization
|
||||
6. Calculates num_cond_latents
|
||||
"""
|
||||
|
||||
def __init__(self, vae):
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Encode video frames to latent for VC conditioning."""
|
||||
|
||||
# Get video from batch - can be path, list of PIL images, or already loaded
|
||||
video = getattr(batch, 'video_frames', None) or getattr(
|
||||
batch, 'video_path', None)
|
||||
num_cond_frames = getattr(batch, 'num_cond_frames',
|
||||
13) # Default 13 for VC
|
||||
|
||||
if video is None:
|
||||
raise ValueError(
|
||||
"video_frames or video_path must be provided for VC")
|
||||
|
||||
# Load video if path
|
||||
if isinstance(video, str):
|
||||
from diffusers.utils import load_video
|
||||
video = load_video(video)
|
||||
logger.info("Loaded video from path: %d frames", len(video))
|
||||
|
||||
# Take last num_cond_frames
|
||||
if len(video) > num_cond_frames:
|
||||
video = video[-num_cond_frames:]
|
||||
logger.info("Using last %d frames for conditioning",
|
||||
num_cond_frames)
|
||||
elif len(video) < num_cond_frames:
|
||||
logger.warning(
|
||||
"Video has only %d frames, less than num_cond_frames=%d",
|
||||
len(video), num_cond_frames)
|
||||
num_cond_frames = len(video)
|
||||
|
||||
# Get target dimensions
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
if height is None or width is None:
|
||||
raise ValueError("height and width must be set for VC")
|
||||
|
||||
# Preprocess and stack frames
|
||||
processed_frames = []
|
||||
for frame in video:
|
||||
if not isinstance(frame, PIL.Image.Image):
|
||||
raise TypeError(f"Frame must be PIL.Image, got {type(frame)}")
|
||||
|
||||
frame = resize(frame, height, width, resize_mode="default")
|
||||
frame = pil_to_numpy(frame) # Returns [1, H, W, C] then converted
|
||||
frame = numpy_to_pt(frame) # Returns [1, C, H, W]
|
||||
frame = normalize(frame) # to [-1, 1]
|
||||
processed_frames.append(frame)
|
||||
|
||||
# Stack frames: [num_frames, C, H, W] -> [1, C, T, H, W]
|
||||
video_tensor = torch.cat(processed_frames, dim=0) # [T, C, H, W]
|
||||
video_tensor = video_tensor.permute(1, 0, 2,
|
||||
3).unsqueeze(0) # [1, C, T, H, W]
|
||||
video_tensor = video_tensor.to(get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
logger.info("VC: Preprocessed video tensor shape: %s",
|
||||
video_tensor.shape)
|
||||
|
||||
# Encode via VAE
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
|
||||
if not vae_autocast_enabled:
|
||||
video_tensor = video_tensor.to(vae_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
encoder_output = self.vae.encode(video_tensor)
|
||||
latent = self.retrieve_latents(encoder_output, batch.generator)
|
||||
|
||||
# Apply LongCat-specific normalization
|
||||
latent = self.normalize_latents(latent)
|
||||
|
||||
# Calculate num_cond_latents
|
||||
# Formula: 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
vae_temporal_scale = self.vae.config.scale_factor_temporal
|
||||
num_cond_latents = 1 + (num_cond_frames - 1) // vae_temporal_scale
|
||||
|
||||
# Store in batch
|
||||
batch.video_latent = latent
|
||||
batch.num_cond_frames = num_cond_frames
|
||||
batch.num_cond_latents = num_cond_latents
|
||||
|
||||
logger.info(
|
||||
"VC: Encoded %d frames to latent shape %s, num_cond_latents=%d",
|
||||
num_cond_frames, latent.shape, num_cond_latents)
|
||||
|
||||
# Offload VAE if needed
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self, encoder_output: Any,
|
||||
generator: torch.Generator | None) -> torch.Tensor:
|
||||
"""Sample from VAE posterior."""
|
||||
if hasattr(encoder_output, 'sample'):
|
||||
return encoder_output.sample(generator)
|
||||
elif hasattr(encoder_output, 'latent_dist'):
|
||||
return encoder_output.latent_dist.sample(generator)
|
||||
elif hasattr(encoder_output, 'latents'):
|
||||
return encoder_output.latents
|
||||
else:
|
||||
raise AttributeError("Could not access latents from encoder output")
|
||||
|
||||
def normalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply LongCat-specific latent normalization.
|
||||
|
||||
Formula: (latents - mean) / std
|
||||
"""
|
||||
if not hasattr(self.vae.config, 'latents_mean') or not hasattr(
|
||||
self.vae.config, 'latents_std'):
|
||||
raise ValueError(
|
||||
"VAE config must have 'latents_mean' and 'latents_std' "
|
||||
"for LongCat normalization")
|
||||
|
||||
latents_mean = torch.tensor(self.vae.config.latents_mean).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
latents_std = torch.tensor(self.vae.config.latents_std).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
|
||||
return (latents - latents_mean) / latents_std
|
||||
@@ -1,4 +1,6 @@
|
||||
from __future__ import annotations
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch # type: ignore
|
||||
@@ -32,6 +34,45 @@ 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,
|
||||
@@ -80,6 +121,9 @@ 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,
|
||||
@@ -94,8 +138,6 @@ 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()
|
||||
@@ -152,23 +194,12 @@ 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"
|
||||
@@ -184,211 +215,67 @@ 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)
|
||||
|
||||
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()
|
||||
current_latents = self._process_single_block(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
timesteps=timesteps,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
progress_bar=progress_bar,
|
||||
)
|
||||
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = current_latents
|
||||
|
||||
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,
|
||||
)
|
||||
# Update KV caches with clean context
|
||||
self._update_context_cache(
|
||||
current_latents=current_latents,
|
||||
batch=batch,
|
||||
start_index=start_index,
|
||||
current_num_frames=current_num_frames,
|
||||
ctx=ctx,
|
||||
action_kwargs=action_kwargs,
|
||||
context_noise=context_noise,
|
||||
)
|
||||
|
||||
start_index += current_num_frames
|
||||
|
||||
@@ -542,6 +429,440 @@ 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()
|
||||
|
||||
@@ -199,6 +199,34 @@ 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"
|
||||
|
||||
@@ -18,6 +18,8 @@ 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()
|
||||
|
||||
|
||||
|
||||
@@ -82,7 +82,7 @@ def run_transformer_tests():
|
||||
@app.function(
|
||||
gpu="L40S:4",
|
||||
image=image,
|
||||
timeout=2700,
|
||||
timeout=3600,
|
||||
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=3600)
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1200)
|
||||
def run_inference_lora_tests():
|
||||
run_test("pytest ./fastvideo/tests/inference/lora/test_lora_inference_similarity.py -vs")
|
||||
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
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)
|
||||
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"mean_ssim": 1.0,
|
||||
"min_ssim": 1.0,
|
||||
"max_ssim": 1.0,
|
||||
"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."
|
||||
}
|
||||
}
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"mean_ssim": 1.0,
|
||||
"min_ssim": 1.0,
|
||||
"max_ssim": 1.0,
|
||||
"reference_video": "/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.2-I2V-A14B-Diffusers/SLA_ATTN/An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space reali.mp4",
|
||||
"generated_video": "/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.2-I2V-A14B-Diffusers/SLA_ATTN/An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space reali.mp4",
|
||||
"parameters": {
|
||||
"num_inference_steps": 4,
|
||||
"prompt": "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,297 @@
|
||||
# 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.95
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
|
||||
f"for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
)
|
||||
|
||||
|
||||
# TurboDiffusion I2V parameters (dual-model with RCM scheduler + SLA attention)
|
||||
TURBODIFFUSION_I2V_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 4, # TurboDiffusion uses 1-4 steps
|
||||
"guidance_scale": 1.0, # No CFG for TurboDiffusion
|
||||
"seed": 42,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
TURBODIFFUSION_I2V_MODEL_TO_PARAMS = {
|
||||
"TurboWan2.2-I2V-A14B-Diffusers": TURBODIFFUSION_I2V_PARAMS,
|
||||
}
|
||||
|
||||
TURBODIFFUSION_I2V_TEST_PROMPTS = [
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot.",
|
||||
]
|
||||
|
||||
TURBODIFFUSION_I2V_IMAGE_PATHS = [
|
||||
"https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", TURBODIFFUSION_I2V_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("model_id", list(TURBODIFFUSION_I2V_MODEL_TO_PARAMS.keys()))
|
||||
def test_turbodiffusion_i2v_inference_similarity(prompt, model_id):
|
||||
"""
|
||||
Test that runs TurboDiffusion I2V inference with dual-model switching,
|
||||
then compares the output to reference videos using SSIM.
|
||||
"""
|
||||
# TurboDiffusion requires SLA attention backend
|
||||
ATTENTION_BACKEND = "SLA_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
assert len(TURBODIFFUSION_I2V_TEST_PROMPTS) == len(TURBODIFFUSION_I2V_IMAGE_PATHS), \
|
||||
"Expect number of prompts equal to number of images"
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
|
||||
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
BASE_PARAMS = TURBODIFFUSION_I2V_MODEL_TO_PARAMS[model_id]
|
||||
num_inference_steps = BASE_PARAMS["num_inference_steps"]
|
||||
image_path = TURBODIFFUSION_I2V_IMAGE_PATHS[TURBODIFFUSION_I2V_TEST_PROMPTS.index(prompt)]
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"override_pipeline_cls_name": "TurboDiffusionI2VPipeline",
|
||||
}
|
||||
|
||||
generation_kwargs = {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"output_path": output_dir,
|
||||
"image_path": image_path,
|
||||
"height": BASE_PARAMS["height"],
|
||||
"width": BASE_PARAMS["width"],
|
||||
"num_frames": BASE_PARAMS["num_frames"],
|
||||
"guidance_scale": BASE_PARAMS["guidance_scale"],
|
||||
"seed": BASE_PARAMS["seed"],
|
||||
"fps": BASE_PARAMS["fps"],
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=BASE_PARAMS["model_path"],
|
||||
**init_kwargs
|
||||
)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
if isinstance(generator.executor, MultiprocExecutor):
|
||||
generator.executor.shutdown()
|
||||
|
||||
assert os.path.exists(output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(
|
||||
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
|
||||
)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}"
|
||||
)
|
||||
|
||||
# Find the matching reference video based on the prompt
|
||||
reference_video_name = None
|
||||
|
||||
for filename in os.listdir(reference_folder):
|
||||
if filename.endswith('.mp4') and prompt[:100].strip() in filename:
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
if not reference_video_name:
|
||||
logger.error(
|
||||
f"Reference video not found for prompt: {prompt} with backend: {ATTENTION_BACKEND}"
|
||||
)
|
||||
raise FileNotFoundError(f"Reference video missing")
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
|
||||
logger.info(
|
||||
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
|
||||
)
|
||||
ssim_values = compute_video_ssim_torchvision(
|
||||
reference_video_path, generated_video_path, use_ms_ssim=True
|
||||
)
|
||||
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info(f"SSIM mean value: {mean_ssim}")
|
||||
logger.info(f"Writing SSIM results to directory: {output_dir}")
|
||||
|
||||
success = write_ssim_results(
|
||||
output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path, num_inference_steps, prompt
|
||||
)
|
||||
|
||||
if not success:
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
# TurboDiffusion I2V uses fewer steps, may have slightly lower SSIM
|
||||
min_acceptable_ssim = 0.95
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
|
||||
f"for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
)
|
||||
@@ -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 +1 @@
|
||||
__version__ = "0.1.6"
|
||||
__version__ = "0.1.7"
|
||||
|
||||
@@ -114,6 +114,20 @@ 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()
|
||||
|
||||
@@ -1,11 +1,17 @@
|
||||
# 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
|
||||
@@ -13,19 +19,55 @@ 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:
|
||||
@@ -40,6 +82,12 @@ 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:
|
||||
@@ -50,6 +98,8 @@ 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
|
||||
@@ -88,6 +138,108 @@ 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:
|
||||
@@ -147,6 +299,24 @@ 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:
|
||||
@@ -266,7 +436,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,
|
||||
@@ -286,9 +456,13 @@ 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)
|
||||
|
||||
@@ -314,6 +488,8 @@ 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)
|
||||
@@ -326,6 +502,8 @@ 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,
|
||||
@@ -455,6 +633,11 @@ 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']
|
||||
@@ -488,6 +671,50 @@ 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,6 +1,7 @@
|
||||
# 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
|
||||
@@ -269,6 +270,47 @@ 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(
|
||||
|
||||
@@ -135,6 +135,7 @@ 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:
|
||||
|
||||
+2
-1
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.6"
|
||||
version = "0.1.7"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
@@ -63,6 +63,7 @@ dependencies = [
|
||||
"remote-pdb",
|
||||
|
||||
# Kernel & Packaging
|
||||
"fastvideo-kernel==0.2.2",
|
||||
"wheel",
|
||||
|
||||
# Training Dependencies
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.6"
|
||||
version = "0.1.7"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
@@ -63,6 +63,7 @@ dependencies = [
|
||||
"remote-pdb",
|
||||
|
||||
# Kernel & Packaging
|
||||
"fastvideo-kernel==0.2.2",
|
||||
"wheel",
|
||||
|
||||
# Training Dependencies
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Convert TurboDiffusion I2V .pth checkpoint to Diffusers safetensors format.
|
||||
|
||||
TurboDiffusion I2V uses two models: high-noise and low-noise.
|
||||
This script converts both checkpoints to Diffusers format.
|
||||
|
||||
Usage:
|
||||
python convert_turbodiffusion_i2v_to_diffusers.py \
|
||||
--high_noise_path /path/to/TurboWan2.2-I2V-A14B-high-720P.pth \
|
||||
--low_noise_path /path/to/TurboWan2.2-I2V-A14B-low-720P.pth \
|
||||
--output_dir /path/to/output \
|
||||
--reference_repo Wan-AI/Wan2.1-I2V-14B-720P-Diffusers
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
import torch
|
||||
import shutil
|
||||
import glob
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import save_file
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
# Weight mapping from TurboDiffusion -> Diffusers/FastVideo format
|
||||
# Same as T2V but may need additional I2V-specific mappings
|
||||
TURBODIFFUSION_WEIGHT_MAPPING = {
|
||||
# Self attention
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
|
||||
# Cross attention
|
||||
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
|
||||
# I2V-specific cross attention (add_k/v_proj for image context)
|
||||
r"^blocks\.(\d+)\.cross_attn\.add_k\.(.*)$": r"blocks.\1.attn2.add_k_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.add_v\.(.*)$": r"blocks.\1.attn2.add_v_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_add_k\.(.*)$": r"blocks.\1.attn2.norm_added_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_add_q\.(.*)$": r"blocks.\1.attn2.norm_added_q.\2",
|
||||
# Norms and FFN
|
||||
r"^blocks\.(\d+)\.norm1\.(.*)$": r"blocks.\1.norm1.\2",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.norm3.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
|
||||
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
|
||||
# Embeddings
|
||||
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
# Head
|
||||
r"^head\.head\.(.*)$": r"proj_out.\1",
|
||||
r"^head\.norm\.(.*)$": r"norm_out.\1",
|
||||
r"^head\.modulation$": r"scale_shift_table",
|
||||
# SLA proj_l weights
|
||||
r"^blocks\.(\d+)\.self_attn\.attn_op\.local_attn\.proj_l\.(.*)$": r"blocks.\1.attn1.attn_impl.proj_l.\2",
|
||||
}
|
||||
|
||||
SKIP_PATTERNS = []
|
||||
|
||||
|
||||
def should_skip_key(key: str) -> bool:
|
||||
for pattern in SKIP_PATTERNS:
|
||||
if re.match(pattern, key):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def convert_key(turbo_key: str) -> str:
|
||||
for pattern, replacement in TURBODIFFUSION_WEIGHT_MAPPING.items():
|
||||
if re.match(pattern, turbo_key):
|
||||
return re.sub(pattern, replacement, turbo_key)
|
||||
return turbo_key
|
||||
|
||||
|
||||
def reshape_patch_embedding(tensor: torch.Tensor, target_shape: tuple) -> torch.Tensor:
|
||||
if len(tensor.shape) == 2 and len(target_shape) == 5:
|
||||
return tensor.view(target_shape)
|
||||
return tensor
|
||||
|
||||
|
||||
def get_reference_shapes(reference_repo: str) -> dict:
|
||||
print(f"Downloading reference model shapes from {reference_repo}...")
|
||||
|
||||
local_dir = snapshot_download(
|
||||
repo_id=reference_repo,
|
||||
allow_patterns=["transformer/config.json", "transformer/diffusion_pytorch_model*.safetensors"],
|
||||
local_dir_use_symlinks=False
|
||||
)
|
||||
|
||||
weight_files = glob.glob(os.path.join(local_dir, "transformer", "*.safetensors"))
|
||||
shapes = {}
|
||||
|
||||
for wf in weight_files:
|
||||
with safe_open(wf, framework="pt") as f:
|
||||
for key in f.keys():
|
||||
shapes[key] = f.get_tensor(key).shape
|
||||
|
||||
return shapes, local_dir
|
||||
|
||||
|
||||
def convert_checkpoint(input_path: str, output_dir: str, ref_shapes: dict, model_name: str) -> None:
|
||||
"""Convert a single TurboDiffusion checkpoint to Diffusers format."""
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Converting {model_name}: {input_path}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
turbo_state_dict = torch.load(input_path, map_location="cpu", weights_only=True)
|
||||
print(f"Loaded {len(turbo_state_dict)} keys")
|
||||
|
||||
converted_state_dict = {}
|
||||
skipped_keys = []
|
||||
|
||||
for turbo_key, tensor in turbo_state_dict.items():
|
||||
if should_skip_key(turbo_key):
|
||||
skipped_keys.append(turbo_key)
|
||||
continue
|
||||
|
||||
new_key = convert_key(turbo_key)
|
||||
|
||||
if "patch_embedding" in new_key and new_key in ref_shapes:
|
||||
target_shape = ref_shapes[new_key]
|
||||
if tensor.shape != target_shape:
|
||||
print(f"Reshaping {new_key}: {tensor.shape} -> {target_shape}")
|
||||
tensor = reshape_patch_embedding(tensor, target_shape)
|
||||
|
||||
if new_key in ref_shapes:
|
||||
if tensor.shape != ref_shapes[new_key]:
|
||||
print(f"WARNING: Shape mismatch for {new_key}: got {tensor.shape}, expected {ref_shapes[new_key]}")
|
||||
|
||||
converted_state_dict[new_key] = tensor
|
||||
|
||||
print(f"Converted: {len(converted_state_dict)} keys, Skipped: {len(skipped_keys)} keys")
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
output_path = os.path.join(output_dir, "diffusion_pytorch_model.safetensors")
|
||||
print(f"Saving to {output_path}...")
|
||||
save_file(converted_state_dict, output_path)
|
||||
|
||||
return converted_state_dict
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Convert TurboDiffusion I2V checkpoints to Diffusers format")
|
||||
parser.add_argument("--high_noise_path", type=str, required=True,
|
||||
help="Path to high-noise TurboDiffusion .pth checkpoint")
|
||||
parser.add_argument("--low_noise_path", type=str, required=True,
|
||||
help="Path to low-noise TurboDiffusion .pth checkpoint")
|
||||
parser.add_argument("--output_dir", type=str, required=True,
|
||||
help="Output directory for converted safetensors")
|
||||
parser.add_argument("--reference_repo", type=str, default="Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
help="Reference HF repo to get expected tensor shapes")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Get reference shapes
|
||||
ref_shapes, ref_local_dir = get_reference_shapes(args.reference_repo)
|
||||
print(f"Got {len(ref_shapes)} reference shapes")
|
||||
|
||||
# Convert high-noise model
|
||||
high_noise_output = os.path.join(args.output_dir, "transformer_high")
|
||||
convert_checkpoint(args.high_noise_path, high_noise_output, ref_shapes, "high-noise")
|
||||
|
||||
# Convert low-noise model
|
||||
low_noise_output = os.path.join(args.output_dir, "transformer_low")
|
||||
convert_checkpoint(args.low_noise_path, low_noise_output, ref_shapes, "low-noise")
|
||||
|
||||
# Copy config.json to both
|
||||
src_config = os.path.join(ref_local_dir, "transformer", "config.json")
|
||||
shutil.copy(src_config, os.path.join(high_noise_output, "config.json"))
|
||||
shutil.copy(src_config, os.path.join(low_noise_output, "config.json"))
|
||||
print("Copied config.json to both transformer directories")
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Conversion complete!")
|
||||
print(f"High-noise model: {high_noise_output}")
|
||||
print(f"Low-noise model: {low_noise_output}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,205 @@
|
||||
#!/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()
|
||||
@@ -1,14 +1,30 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=2
|
||||
# LongCat Text-to-Video (T2V) Inference Script
|
||||
#
|
||||
# This script runs LongCat T2V inference using the fastvideo CLI.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# For longcat, we must first convert the official weights to FastVideo native format
|
||||
# Model path options:
|
||||
# Option 1: HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-T2V-Diffusers
|
||||
|
||||
# Option 2: Local weights (uncomment if you have local weights)
|
||||
# For local weights, convert the official weights to FastVideo native format
|
||||
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
|
||||
# --source /path/to/LongCat-Video/weights/LongCat-Video
|
||||
# --output weights/longcat-native
|
||||
export MODEL_BASE=weights/longcat-native
|
||||
# export MODEL_BASE=weights/longcat-native
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
@@ -26,7 +42,7 @@ fastvideo generate \
|
||||
--num-inference-steps 50 \
|
||||
--fps 15 \
|
||||
--guidance-scale 4.0 \
|
||||
--prompt-txt assets/prompt.txt \
|
||||
--prompt "In a realistic photography style, a white boy around seven or eight years old sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene." \
|
||||
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 42 \
|
||||
--output-path outputs_video/longcat_480p
|
||||
--output-path outputs_video/longcat_t2v
|
||||
|
||||
@@ -1,12 +1,23 @@
|
||||
#!/bin/bash
|
||||
|
||||
# LongCat T2V Distilled Inference Script
|
||||
#
|
||||
# This script runs LongCat T2V with distillation LoRA (16 steps instead of 50).
|
||||
# Uses the distilled LoRA for faster generation.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_distill.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
# For longcat, we must first convert the official weights to FastVideo native format
|
||||
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
|
||||
# --source /path/to/LongCat-Video/weights/LongCat-Video
|
||||
# --output weights/longcat-native
|
||||
export MODEL_BASE=weights/longcat-native
|
||||
|
||||
# Model path - HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-T2V-Diffusers
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
@@ -14,11 +25,11 @@ fastvideo generate \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload False \
|
||||
--vae-cpu-offload False \
|
||||
--text-encoder-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa False \
|
||||
--lora-path "$MODEL_BASE/lora/distilled" \
|
||||
--lora-path "FastVideo/LongCat-Video-T2V-Distilled-LoRA" \
|
||||
--lora-nickname "distilled" \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
@@ -26,6 +37,7 @@ fastvideo generate \
|
||||
--num-inference-steps 16 \
|
||||
--fps 15 \
|
||||
--guidance-scale 1.0 \
|
||||
--prompt "In a realistic photography style, an asian boy around seven or eight years old sits on a park bench, wearing a light yellow T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene." \
|
||||
--prompt "In a realistic photography style, a white boy around seven or eight years old sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene." \
|
||||
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 42 \
|
||||
--output-path outputs_video/longcat_distill
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
#!/bin/bash
|
||||
|
||||
# LongCat Image-to-Video (I2V) Inference Script
|
||||
#
|
||||
# This script runs LongCat I2V inference using the fastvideo CLI.
|
||||
# LongCat I2V takes an input image and generates a video from it.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_i2v.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
# - Or use local weights if you have them
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# Model path options:
|
||||
# Option 1: HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-I2V-Diffusers
|
||||
|
||||
# Option 2: Local weights (uncomment if you have local weights)
|
||||
# export MODEL_BASE=weights/longcat-for-i2v
|
||||
|
||||
# Input image path (must be square for LongCat I2V)
|
||||
IMAGE_PATH="assets/girl.png"
|
||||
|
||||
# Check if image exists
|
||||
if [ ! -f "$IMAGE_PATH" ]; then
|
||||
echo "Error: Image not found at $IMAGE_PATH"
|
||||
echo "Please provide a valid image path"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa False \
|
||||
--image-path "$IMAGE_PATH" \
|
||||
--height 480 \
|
||||
--width 480 \
|
||||
--num-frames 93 \
|
||||
--num-inference-steps 50 \
|
||||
--fps 15 \
|
||||
--guidance-scale 4.0 \
|
||||
--prompt "A woman sits at a wooden table by the window in a cozy café. She reaches out with her right hand, picks up the white coffee cup from the saucer, and gently brings it to her lips to take a sip. After drinking, she places the cup back on the table and looks out the window, enjoying the peaceful atmosphere." \
|
||||
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 42 \
|
||||
--output-path outputs_video/longcat_i2v
|
||||
|
||||
|
||||
|
||||
@@ -1,18 +1,30 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
# For longcat, we must first convert the official weights to FastVideo native format
|
||||
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
|
||||
# --source /path/to/LongCat-Video/weights/LongCat-Video
|
||||
# --output weights/longcat-native
|
||||
export MODEL_BASE=weights/longcat-native
|
||||
# LongCat T2V Refinement Script (480p -> 720p)
|
||||
#
|
||||
# This script refines a 480p distilled video to 720p using the refinement LoRA.
|
||||
# Run v1_inference_longcat_distill.sh first to generate the 480p video.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_refine_fromvideo.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
# - Run v1_inference_longcat_distill.sh first to generate input video
|
||||
|
||||
INPUT_VIDEO="outputs_video/longcat_distill/In a realistic photography style, an asian boy around seven or eight years old sits on a park bench,.mp4"
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# Model path - HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-T2V-Diffusers
|
||||
|
||||
INPUT_VIDEO="outputs_video/longcat_distill/In a realistic photography style, a white boy around seven or eight years old sits on a park bench,.mp4"
|
||||
REFINE_OUTPUT="outputs_video/longcat_refine_720p"
|
||||
|
||||
# Prompt used for base generation
|
||||
PROMPT="In a realistic photography style, an asian boy around seven or eight years old sits on a park bench, wearing a light yellow T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene."
|
||||
# Prompt used for base generation (must match distill script)
|
||||
PROMPT="In a realistic photography style, a white boy around seven or eight years old sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene."
|
||||
|
||||
echo "=========================================="
|
||||
echo "LongCat 480p -> 720p Refinement"
|
||||
@@ -25,14 +37,14 @@ echo ""
|
||||
# Check if input video exists
|
||||
if [ ! -f "$INPUT_VIDEO" ]; then
|
||||
echo "Error: Input video not found: $INPUT_VIDEO"
|
||||
echo "Please set INPUT_VIDEO to your 480p video path"
|
||||
echo "Please run v1_inference_longcat_distill.sh first to generate the 480p video"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "🔧 Configuring refinement (BSA enabled, refinement LoRA)..."
|
||||
echo "✅ Input video: $INPUT_VIDEO"
|
||||
echo "✅ BSA enabled with sparsity=0.875"
|
||||
echo "✅ Refinement LoRA loaded"
|
||||
echo "Configuring refinement (BSA enabled, refinement LoRA)..."
|
||||
echo "Input video: $INPUT_VIDEO"
|
||||
echo "BSA enabled with sparsity=0.875"
|
||||
echo "Refinement LoRA loaded"
|
||||
echo ""
|
||||
|
||||
fastvideo generate \
|
||||
@@ -41,14 +53,14 @@ fastvideo generate \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload True \
|
||||
--vae-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa True \
|
||||
--bsa-sparsity 0.875 \
|
||||
--bsa-chunk-q 4 4 8 \
|
||||
--bsa-chunk-k 4 4 8 \
|
||||
--lora-path "$MODEL_BASE/lora/refinement" \
|
||||
--lora-path "FastVideo/LongCat-Video-T2V-Refinement-LoRA" \
|
||||
--lora-nickname "refinement" \
|
||||
--refine-from "$INPUT_VIDEO" \
|
||||
--t-thresh 0.5 \
|
||||
@@ -60,12 +72,13 @@ fastvideo generate \
|
||||
--fps 30 \
|
||||
--guidance-scale 1.0 \
|
||||
--prompt "$PROMPT" \
|
||||
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 42 \
|
||||
--output-path "$REFINE_OUTPUT"
|
||||
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "✓ Refinement Complete!"
|
||||
echo "Refinement Complete!"
|
||||
echo "=========================================="
|
||||
echo ""
|
||||
echo "Output directory: $REFINE_OUTPUT"
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
#!/bin/bash
|
||||
|
||||
# LongCat Video Continuation (VC) Inference Script
|
||||
#
|
||||
# This script runs LongCat VC inference using the fastvideo CLI.
|
||||
# LongCat VC takes an input video and generates a continuation of it.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/inference/v1_inference_longcat_vc.sh
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Install fastvideo: pip install -e .
|
||||
# - The model weights will be auto-downloaded from HuggingFace
|
||||
# - Or use local weights if you have them
|
||||
|
||||
num_gpus=1
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
|
||||
# Model path options:
|
||||
# Option 1: HuggingFace model (auto-downloaded)
|
||||
export MODEL_BASE=FastVideo/LongCat-Video-VC-Diffusers
|
||||
|
||||
# Option 2: Local weights (uncomment if you have local weights)
|
||||
# export MODEL_BASE=weights/longcat-vc-upload
|
||||
|
||||
# Input video path
|
||||
VIDEO_PATH="assets/motorcycle.mp4"
|
||||
|
||||
# Check if video exists
|
||||
if [ ! -f "$VIDEO_PATH" ]; then
|
||||
echo "Error: Video not found at $VIDEO_PATH"
|
||||
echo "Please provide a valid video path"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--dit-cpu-offload False \
|
||||
--vae-cpu-offload True \
|
||||
--text-encoder-cpu-offload True \
|
||||
--pin-cpu-memory False \
|
||||
--enable-bsa False \
|
||||
--video-path "$VIDEO_PATH" \
|
||||
--num-cond-frames 13 \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
--num-frames 93 \
|
||||
--num-inference-steps 50 \
|
||||
--fps 15 \
|
||||
--guidance-scale 4.0 \
|
||||
--prompt "A person rides a motorcycle along a long, straight road that stretches between a body of water and a forested hillside. The rider steadily accelerates, keeping the motorcycle centered between the guardrails, while the scenery passes by on both sides. The video captures the journey from the rider's perspective, emphasizing the sense of motion and adventure." \
|
||||
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 42 \
|
||||
--output-path outputs_video/longcat_vc
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user