Compare commits

..
Author SHA1 Message Date
Y-aang 5931a65da8 update 2025-11-26 02:50:52 +00:00
103 changed files with 1269 additions and 3659 deletions
+1 -1
View File
@@ -75,7 +75,7 @@ case "$TEST_TYPE" in
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
-10
View File
@@ -3,18 +3,8 @@ name: Deploy Documentation
on:
push:
branches: [ main ]
paths:
- 'docs/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.txt'
- '.github/workflows/docs.yml'
pull_request:
branches: [ main ]
paths:
- 'docs/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.txt'
- '.github/workflows/docs.yml'
permissions:
contents: read
-1
View File
@@ -10,7 +10,6 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
+23 -20
View File
@@ -1,12 +1,13 @@
<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/TM8JyJCd" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><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-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/tMwknPLY" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -14,7 +15,6 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
</div>
## NEWS
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
- ```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!
@@ -49,10 +49,10 @@ conda activate fastvideo
pip install fastvideo
```
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
## Sparse Distillation
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
See below for recipes and datasets:
@@ -64,7 +64,7 @@ See below for recipes and datasets:
## Inference
### Generating Your First Video
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation/). Create a file called `example.py` with the following code:
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Create a file called `example.py` with the following code:
```python
import os
@@ -100,32 +100,35 @@ Run the script with:
python example.py
```
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html).
### Other docs:
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## Awesome work using FastVideo or our research projects
## 📑 Development Plan
<!-- - More distillation methods -->
<!-- - [ ] Add Distribution Matching Distillation -->
More FastWan Models Coming Soon!
- [ ] Add FastWan2.1-T2V-14B
- [ ] Add FastWan2.2-T2V-14B
- [ ] Add FastWan2.2-I2V-14B
<!-- - Optimization features
- Code updates -->
<!-- - [ ] fp8 support -->
<!-- - [ ] faster load model and save model support -->
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025. [![Star](https://img.shields.io/github/stars/sgl-project/sglang.svg?style=social&label=Star)](https://github.com/sgl-project/sglang)
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/XueZeyue/DanceGRPO.svg?style=social&label=Star)](https://github.com/XueZeyue/DanceGRPO)
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/SRPO.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/SRPO)
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Vchitect/DCM.svg?style=social&label=Star)](https://github.com/Vchitect/DCM)
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/HunyuanVideo-1.5.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch. [![Star](https://img.shields.io/github/stars/kandinskylab/kandinsky-5.svg?style=social&label=Star)](https://github.com/kandinskylab/kandinsky-5)
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention. [![Star](https://img.shields.io/github/stars/meituan-longcat/LongCat-Video.svg?style=social&label=Star)](https://github.com/meituan-longcat/LongCat-Video)
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
## Acknowledgement
We learned and reused code from the following projects:
- [Wan-Video](https://github.com/Wan-Video)
@@ -247,11 +247,11 @@ def _attn_bwd_dq(dq, q, K, V, #
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
block_size = tl.load(variable_block_sizes + q_blk)
for blk_idx in range(kv_blocks*2):
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
block_size = tl.load(variable_block_sizes + blk_idx//2) - (blk_idx%2) * step_n
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
+80
View File
@@ -0,0 +1,80 @@
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Install Miniconda and create conda environment
COPY Miniconda3-latest-Linux-x86_64.sh /tmp/miniconda.sh
RUN bash /tmp/miniconda.sh -b -p /opt/conda && rm /tmp/miniconda.sh
ENV PATH=/opt/conda/bin:${PATH}
RUN /opt/conda/bin/conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/main
RUN /opt/conda/bin/conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/r
RUN /opt/conda/bin/conda update -y -n base conda && \
/opt/conda/bin/conda create -y -n fastvideo python=3.12 && \
/opt/conda/bin/conda clean -afy
ENV PATH=/opt/conda/envs/fastvideo/bin:/opt/conda/bin:${PATH}
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
# Install project dependencies inside the conda environment
RUN source /opt/conda/etc/profile.d/conda.sh && \
conda activate fastvideo && \
pip install --no-cache-dir --upgrade pip && \
pip install --no-cache-dir .[dev] && \
pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
# Install dependencies using pip inside conda env and set up shell configuration
RUN source /opt/conda/etc/profile.d/conda.sh && \
conda activate fastvideo && \
pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/conda/etc/profile.d/conda.sh && conda activate fastvideo' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install STA (Sliding Tile Attention)
RUN source /opt/conda/etc/profile.d/conda.sh && \
conda activate fastvideo && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# Install VSA
RUN source /opt/conda/etc/profile.d/conda.sh && \
conda activate fastvideo && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
EXPOSE 22
ENTRYPOINT ["/bin/bash", "-lc", "source /opt/conda/etc/profile.d/conda.sh && conda activate fastvideo && exec /FastVideo/examples/inference/gradio/start.sh"]
Binary file not shown.

Before

Width:  |  Height:  |  Size: 122 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 378 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 575 KiB

-4
View File
@@ -70,7 +70,3 @@ uv pip install ninja
python setup.py install
```
## Testing
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
-129
View File
@@ -1,129 +0,0 @@
# Testing in FastVideo
This guide explains how to add and run tests in FastVideo. The testing suite is divided into several categories to ensure correctness across components, training workflows, and inference quality.
## Test Types
* **Unit Tests**: Located in `fastvideo/tests/dataset`, `fastvideo/tests/entrypoints`, and `fastvideo/tests/workflow`. These test individual functions and classes.
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
For now, we will focus on **SSIM Tests**.
## SSIM Tests
SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos using specific models and parameters, and compare them against reference videos to ensure that changes in the codebase do not degrade generation quality or alter the output unexpectedly.
!!! note
If you are adding an SSIM test, this serves as a safeguard. Any future code changes that break or cause errors with the specific arguments and configurations you defined will trigger a failure. Therefore, it is important to include multiple settings and arguments that cover the core features of your new pipeline to ensure robust regression testing.
### Directory Structure
```
fastvideo/tests/ssim/
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
│ ├── <Model_Name>/
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
│ │ │ └── <Video_File>
├── test_causal_similarity.py
├── test_inference_similarity.py
├── update_reference_videos.sh
└── ...
```
### Adding a New SSIM Test
To add a new SSIM test, follow these steps:
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
```python
MY_MODEL_PARAMS = {
"num_gpus": 1,
"model_path": "organization/model-name",
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 20,
# ... other parameters
}
```
3. **Implement the Test Function**:
* Use `pytest.mark.parametrize` to run the test with different prompts, backends, and models.
* Set the attention backend environment variable.
* Initialize the `VideoGenerator`.
* Generate the video.
* Compare the generated video with the reference video using `compute_video_ssim_torchvision`.
Example structure:
```python
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_my_model_similarity(prompt, ATTENTION_BACKEND):
# Setup output directories
# ...
# Initialize Generator
generator = VideoGenerator.from_pretrained(...)
generator.generate_video(prompt, ...)
# Compare with Reference
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
assert ssim_values[0] >= 0.98 # Threshold
```
4. **Reference Videos**:
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
* Inspect the generated video to ensure it meets quality expectations.
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
### Running Tests Locally
To run the SSIM tests locally:
```bash
pytest fastvideo/tests/ssim/ -vs
```
Ensure you have the necessary GPUs available as defined in your test parameters.
## Modal Workflow
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
### `pr_test.py`
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
### Updating Modal Configuration
If you add a new test that requires:
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
* **Longer Execution Time**: Increase the `timeout` parameter.
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
```python
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
def run_ssim_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
```
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
### Workflow Scripts
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
!!! note
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
-8
View File
@@ -106,8 +106,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
return batch
```
![Pipeline execution and data flow](../assets/images/pipeline.png)
### ForwardBatch
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
@@ -210,10 +208,6 @@ def step(
return prev_sample
```
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
![Model loading flow](../assets/images/load_models.png)
## Optimized Attention
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
@@ -237,8 +231,6 @@ self.attn = LocalAttention(
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
```
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
+2 -3
View File
@@ -24,13 +24,12 @@ FastVideo is an inference and post-training framework for diffusion models. It f
## Key Features
FastVideo has the following features:
- State-of-the-art performance optimizations for inference
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- E2E post-training support
- Data preprocessing pipeline for video data
- Data preprocessing pipeline for video data.
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
- 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.
@@ -43,7 +42,7 @@ Use the navigation menu on the left to explore different sections:
- **Getting Started**: Installation and quick start guides
- **Inference**: Learn how to use FastVideo for video generation
- **Training**: Data preprocessing and fine-tuning workflows
- **Training**: Data preprocessing and fine-tuning workflows
- **Distillation**: Post-training optimization techniques
- **Sliding Tile Attention**: Advanced attention mechanisms
- **Video Sparse Attention**: Efficient attention for video models
+51
View File
@@ -0,0 +1,51 @@
# Seed Parameter Behavior in vLLM
## Overview
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
## Default Behavior
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
## Specifying a Seed
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
## Example Usage
### Without Specifying a Seed
```python
import random
from vllm import LLM
# Initialize a vLLM model without specifying a seed
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
# Try generating random numbers
print(random.randint(0, 100)) # Outputs different numbers across runs
```
### Specifying a Seed
```python
import random
from vllm import LLM
# Initialize a vLLM model with a specific seed
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
# Try generating random numbers
print(random.randint(0, 100)) # Outputs the same number across runs
```
## Important Notes
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
## Conclusion
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

@@ -1,44 +0,0 @@
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
from fastvideo import VideoGenerator, SamplingParam
import json
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
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.
generator = VideoGenerator.from_pretrained(
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
dit_precision="fp32",
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
# 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,
)
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
sampling_param.num_frames = 81
sampling_param.width = 832
sampling_param.height = 480
sampling_param.seed = 1000
with open("prompts/mixkit_i2v.jsonl", "r") as f:
prompt_image_pairs = json.load(f)
for prompt_image_pair in prompt_image_pairs:
prompt = prompt_image_pair["prompt"]
image_path = prompt_image_pair["image_path"]
_ = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
main()
@@ -1,43 +0,0 @@
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
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.
generator = VideoGenerator.from_pretrained(
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
# 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,
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
num_frame_per_block=7,
# image_encoder_cpu_offload=False,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
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."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
if __name__ == "__main__":
main()
@@ -3,9 +3,6 @@ import os
import requests
import base64
import time
import json
from pathlib import Path
import tempfile
import gradio as gr
@@ -15,7 +12,6 @@ from fastvideo.configs.sample.base import SamplingParam
MODEL_PATH_MAPPING = {
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"CausalWan2.2-I2V-A14B-Preview": "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
}
@@ -41,7 +37,7 @@ class RayServeClient:
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=900 # 15 minutes timeout for longer video generation
timeout=300
)
round_trip_time = time.time() - start_time
@@ -85,78 +81,49 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
return None
def encode_image_to_base64(image_path: str) -> str:
"""Encode an image file to base64 string."""
if not image_path or not os.path.exists(image_path):
return None
def create_timing_display(inference_time, encoding_time, network_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"
try:
with open(image_path, 'rb') as f:
image_bytes = f.read()
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
# Determine image type from extension
ext = os.path.splitext(image_path)[1].lower()
mime_types = {
'.jpg': 'image/jpeg',
'.jpeg': 'image/jpeg',
'.png': 'image/png',
'.gif': 'image/gif',
'.webp': 'image/webp',
}
mime_type = mime_types.get(ext, 'image/jpeg')
return f"data:{mime_type};base64,{image_base64}"
except Exception as e:
print(f"Failed to encode image: {e}")
return None
# def create_timing_display(inference_time, encoding_time, network_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;">{encoding_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;">Network Transfer</div>
# <div style="font-size: 16px; color: #059669;">{network_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;">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>"
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;">{encoding_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;">Network Transfer</div>
<div style="font-size: 16px; color: #059669;">{network_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;">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 load_example_prompts():
@@ -177,83 +144,26 @@ def load_example_prompts():
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
# Load prompts from prompts.txt
examples, example_labels = load_from_file("examples/inference/gradio/serving/prompts.txt")
examples, example_labels = load_from_file("prompts/prompts_final.txt")
if not examples:
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
example_labels = ["Crowded rooftop bar at night"]
# Load image mappings from JSON file
prompt_to_image = {}
# Try to find the JSON file relative to project root
possible_json_paths = [
Path("prompts/mixkit_i2v.jsonl"),
Path(__file__).parent.parent.parent.parent / "prompts" / "mixkit_i2v.jsonl",
]
json_path = None
for path in possible_json_paths:
if path.exists():
json_path = path
break
if json_path and json_path.exists():
try:
with open(json_path, "r", encoding='utf-8') as f:
data = json.load(f)
# Get the project root directory (parent of prompts directory)
project_root = json_path.parent.parent
for item in data:
prompt_text = item.get("prompt", "").strip()
image_path = item.get("image_path", "")
if prompt_text and image_path:
# Resolve image path relative to project root
full_image_path = project_root / image_path
if full_image_path.exists():
prompt_to_image[prompt_text] = str(full_image_path.absolute())
except Exception as e:
print(f"Warning: Could not load image mappings from {json_path}: {e}")
# Create image paths list matching the prompts
example_images = []
for prompt in examples:
# Try exact match first
image_path = prompt_to_image.get(prompt)
if not image_path:
# Try fuzzy match (case-insensitive, whitespace normalized)
normalized_prompt = " ".join(prompt.split())
for json_prompt, img_path in prompt_to_image.items():
normalized_json = " ".join(json_prompt.split())
if normalized_prompt.lower() == normalized_json.lower():
image_path = img_path
break
example_images.append(image_path if image_path and os.path.exists(image_path) else None)
return examples, example_labels, example_images
return examples, example_labels
def create_gradio_interface(backend_url: str, default_params: dict[str, SamplingParam]):
client = RayServeClient(backend_url)
def is_i2v_model(model_name: str) -> bool:
"""Check if the model is an I2V model."""
return "I2V" in model_name
def generate_video(
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, progress
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
):
# Use default seed value (randomize_seed disabled)
seed = 1000
randomize_seed = False
if not client.check_health():
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
# Check if I2V model requires an image
if is_i2v_model(model_selection) and not input_image:
return None, "I2V models require an input image. Please upload an image.", ""
# Validate dimensions
max_pixels = 720 * 1280
if height * width > max_pixels:
@@ -262,15 +172,6 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
if progress:
progress(0.1, desc="Checking backend health...")
# Encode image if provided
image_data = None
if input_image:
if progress:
progress(0.2, desc="Encoding input image...")
image_data = encode_image_to_base64(input_image)
if not image_data:
return None, "Failed to encode input image", ""
request_data = {
"prompt": prompt,
"negative_prompt": negative_prompt,
@@ -282,7 +183,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
"width": width,
"randomize_seed": randomize_seed,
"return_frames": False,
"image_data": image_data,
"image_path": None,
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
}
@@ -297,16 +198,16 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
if response.get("success", False):
video_data = response.get("video_data", "")
used_seed = response.get("seed", seed)
# inference_time = response.get("inference_time", 0.0)
# encoding_time = response.get("encoding_time", 0.0)
# total_time = response.get("total_time", 0.0)
# network_time = response.get("network_time", 0.0)
# stage_execution_times = response.get("stage_execution_times", [])
inference_time = response.get("inference_time", 0.0)
encoding_time = response.get("encoding_time", 0.0)
total_time = response.get("total_time", 0.0)
network_time = response.get("network_time", 0.0)
stage_execution_times = response.get("stage_execution_times", [])
# timing_details = create_timing_display(
# inference_time, encoding_time, network_time, total_time,
# stage_execution_times, num_frames
# )
timing_details = create_timing_display(
inference_time, encoding_time, network_time, total_time,
stage_execution_times, num_frames
)
if video_data:
if progress:
@@ -318,7 +219,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
progress(1.0, desc="Generation complete!")
if video_path and os.path.exists(video_path):
return video_path, used_seed, ""
return video_path, used_seed, timing_details
else:
return None, "Failed to save video", ""
else:
@@ -327,7 +228,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
error_msg = response.get("error_message", "Unknown error occurred")
return None, f"Generation failed: {error_msg}", ""
examples, example_labels, example_images = load_example_prompts()
examples, example_labels = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
@@ -338,39 +239,33 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
)
def get_default_values(model_name):
# model_path = MODEL_PATH_MAPPING.get(model_name)
# if model_path and model_path in default_params:
# params = default_params[model_path]
# return {
# 'height': params.height,
# 'width': params.width,
# 'num_frames': params.num_frames,
# 'guidance_scale': params.guidance_scale,
# }
model_path = MODEL_PATH_MAPPING.get(model_name)
if model_path and model_path in default_params:
params = default_params[model_path]
return {
'height': params.height,
'width': params.width,
'num_frames': params.num_frames,
'guidance_scale': params.guidance_scale,
'seed': params.seed,
}
return {
'height': 480,
'height': 448,
'width': 832,
'num_frames': 73,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
}
# Get available models based on what's loaded
available_models = []
for model_name, model_path in MODEL_PATH_MAPPING.items():
if model_path in default_params:
available_models.append(model_name)
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
# Select first available model as default
default_model = available_models[0] if available_models else "FastWan2.1-T2V-1.3B"
initial_values = get_default_values(default_model)
initial_show_image = is_i2v_model(default_model)
with gr.Blocks(title="CausalWan", theme=theme) as demo:
with gr.Blocks(title="FastWan", theme=theme) as demo:
gr.Image("assets/logos/logo.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_causalwan_preview/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </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>
""")
@@ -385,8 +280,8 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
with gr.Row():
model_selection = gr.Dropdown(
choices=available_models,
value=default_model,
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
label="Select Model",
interactive=True
)
@@ -417,70 +312,69 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", visible=False)
# timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
with gr.Row(equal_height=False):
with gr.Column(scale=1):
with gr.Tabs():
with gr.Tab("Input Image", visible=initial_show_image) as image_tab:
gr.Markdown("**Please make sure you upload a 480x832 image**")
input_image = gr.Image(
label="",
type="filepath",
height=400,
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;'>Advanced Options</div>")
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
)
with gr.Tab("Advanced Options"):
with gr.Group():
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Number(
label="Guidance Scale",
value=1.0,
interactive=False,
container=True
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=initial_values['guidance_scale'],
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
# randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed", value=1000)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
with gr.Column(scale=1):
with gr.Column(scale=1, elem_classes="video-column"):
result = gr.Video(
label="Generated Video",
show_label=True,
height=500,
height=466,
width=600,
container=True,
autoplay=True,
elem_classes="video-component"
)
gr.HTML("""
@@ -493,10 +387,116 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
}
.gradio-container {
max-width: 1400px !important;
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;
@@ -511,20 +511,18 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
def on_example_select(example_label):
if example_label and example_label in example_labels:
index = example_labels.index(example_label)
selected_prompt = examples[index]
selected_image = example_images[index] if index < len(example_images) else None
return selected_prompt, selected_image
return "", None
return examples[index]
return ""
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=[prompt, input_image],
outputs=prompt,
)
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant as a preview of our distilled I2V model. Outside of few-step distillation, we have not yet fully optimized it for speed. Stay tuned for updates!</p>
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected. We are also rate-limiting users to 3 requests per minute.</p>
</div>
""")
@@ -539,7 +537,6 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
selected_model = "FastWan2.1-T2V-1.3B"
model_path = MODEL_PATH_MAPPING.get(selected_model)
show_image_input = is_i2v_model(selected_model)
if model_path and model_path in default_params:
params = default_params[model_path]
@@ -548,29 +545,29 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
gr.update(value=params.width),
gr.update(value=params.num_frames),
gr.update(value=params.guidance_scale),
gr.update(visible=show_image_input),
gr.update(value=params.seed),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=20),
gr.update(value=61),
gr.update(value=3.0),
gr.update(visible=show_image_input),
gr.update(value=1024),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, image_tab],
outputs=[height, width, num_frames, guidance_scale, seed],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, guidance_scale, num_frames, height, width, input_image = args
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
result_path, seed_or_error, _ = generate_video(
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, progress
result_path, seed_or_error, timing_details = generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
)
if result_path and os.path.exists(result_path):
@@ -578,12 +575,14 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
result_path,
seed_or_error,
gr.update(visible=False),
gr.update(visible=True, value=timing_details),
)
else:
return (
None,
seed_or_error,
gr.update(visible=True, value=seed_or_error),
gr.update(visible=False),
)
run_button.click(
@@ -593,14 +592,14 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
# randomize_seed,
input_image,
randomize_seed,
],
outputs=[result, seed_output, error_output], # timing_display removed
outputs=[result, seed_output, error_output, timing_display],
concurrency_limit=20,
)
@@ -612,11 +611,8 @@ def main():
parser.add_argument("--backend_url", type=str, default="http://localhost:8000",
help="URL of the Ray Serve backend")
parser.add_argument("--t2v_model_paths", type=str,
default="",
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--i2v_model_paths", type=str,
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
help="Comma separated list of paths to the I2V model(s)")
parser.add_argument("--host", type=str, default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port", type=int, default=7860,
@@ -625,15 +621,8 @@ def main():
args = parser.parse_args()
default_params = {}
# Load T2V model params
t2v_paths = [p.strip() for p in args.t2v_model_paths.split(",") if p.strip()]
for model_path in t2v_paths:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
# Load I2V model params
i2v_paths = [p.strip() for p in args.i2v_model_paths.split(",") if p.strip()]
for model_path in i2v_paths:
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(args.backend_url, default_params)
@@ -641,8 +630,6 @@ def main():
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
print(f"Backend URL: {args.backend_url}")
print(f"T2V Models: {args.t2v_model_paths}")
if args.i2v_model_paths:
print(f"I2V Models: {args.i2v_model_paths}")
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
@@ -687,23 +674,23 @@ def main():
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>CausalWan</title>
<meta name="title" content="CausalWan">
<title>FastWan</title>
<meta name="title" content="FastWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, CausalWan">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="CausalWan">
<meta property="og:title" content="FastWan">
<meta property="og:description" content="Make video generation go blurrrrrrr">
<meta property="og:image" content="{base_url}/logo.svg">
<meta property="og:image:width" content="1200">
<meta property="og:image:height" content="630">
<meta property="og:site_name" content="CausalWan">
<meta property="og:site_name" content="FastWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="CausalWan">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
<meta property="twitter:image" content="{base_url}/logo.svg">
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
@@ -733,14 +720,7 @@ def main():
app,
demo,
path="/gradio",
allowed_paths=[
os.path.abspath("outputs"),
os.path.abspath("fastvideo-logos"),
os.path.abspath("prompts"),
os.path.abspath("images"),
os.path.abspath(tempfile.gettempdir()),
os.path.abspath(os.path.join(tempfile.gettempdir(), "gradio")),
]
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
)
uvicorn.run(app, host=args.host, port=args.port)
@@ -1,15 +0,0 @@
Some friends dancing and having fun together in circles, at a party surrounded by colored lights at a party, in a fancy old place, in a view from below them.
A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.
Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.
Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.
In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.
A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.
A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.
A saxophonist wearing a blazer dances while playing a song in a park.
Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.
Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.
Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.
Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.
A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.
Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.
Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.
@@ -26,7 +26,6 @@ SEED_RANGE_MAX = 1_000_000
SUPPORTED_MODELS = [
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
]
MODEL_CONFIGS = {
@@ -43,13 +42,6 @@ MODEL_CONFIGS = {
"dit_cpu_offload": True,
"vae_cpu_offload": False,
"VSA_sparsity": 0.9,
},
"I2V-A14B": {
"num_cpus": 15,
"text_encoder_cpu_offload": True,
"dit_cpu_offload": True,
"vae_cpu_offload": False,
"VSA_sparsity": 0.0,
}
}
@@ -66,7 +58,6 @@ class VideoGenerationRequest(BaseModel):
randomize_seed: bool = False
return_frames: bool = False
model_path: Optional[str] = None
image_data: Optional[str] = None # Base64 encoded image for I2V
class VideoGenerationResponse(BaseModel):
@@ -100,38 +91,11 @@ def encode_video_to_base64(frames: List[np.ndarray], fps: int = DEFAULT_FPS) ->
return ""
def save_image_from_base64(image_data: str, output_dir: str) -> Optional[str]:
"""Save base64 image data to a temporary file and return the path."""
if not image_data:
return None
try:
# Remove data URL prefix if present
if image_data.startswith('data:image/'):
image_data = image_data.split(',')[1]
image_bytes = base64.b64decode(image_data)
# Save to temporary file
os.makedirs(output_dir, exist_ok=True)
temp_image_path = os.path.join(output_dir, f"temp_input_{int(time.time() * 1000)}.png")
with open(temp_image_path, 'wb') as f:
f.write(image_bytes)
return temp_image_path
except Exception as e:
print(f"Warning: Failed to save image: {e}")
return None
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"
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_STAGE_LOGGING"] = "1"
@@ -193,41 +157,22 @@ class BaseModelDeployment:
num_gpus=1,
use_fsdp_inference=True,
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125], # TODO: hardocde for I2V
dit_precision="fp32", # TODO: hardocde for I2V
dit_cpu_offload=config["dit_cpu_offload"],
vae_cpu_offload=config["vae_cpu_offload"],
VSA_sparsity=config["VSA_sparsity"],
enable_stage_verification=False,
)
self.default_params = SamplingParam.from_pretrained(self.model_path)
self.default_params.seed = 1000
self.default_params.num_frames = 73
self.default_params.width = 832
self.default_params.height = 480
def generate_video(self, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
total_start_time = time.time()
params = prepare_sampling_params(video_request, self.default_params)
# Save image if provided (for I2V)
image_path = None
if video_request.image_data:
image_path = save_image_from_base64(video_request.image_data, self.output_path)
if image_path is None:
return VideoGenerationResponse(
video_data=None,
seed=params.seed,
success=False,
error_message="Failed to save input image",
)
inference_start_time = time.time()
result = self.generator.generate_video(
prompt=video_request.prompt,
sampling_param=params,
image_path=image_path,
save_video=False,
return_frames=False,
)
@@ -240,13 +185,6 @@ class BaseModelDeployment:
encoding_time = time.time() - encoding_start_time
total_time = time.time() - total_start_time
# Clean up temporary image file
if image_path and os.path.exists(image_path):
try:
os.remove(image_path)
except Exception as e:
print(f"Warning: Failed to remove temporary image file {image_path}: {e}")
return VideoGenerationResponse(
video_data=video_data,
@@ -262,7 +200,7 @@ class BaseModelDeployment:
@serve.deployment(
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "fastvideo"}},
)
class T2VModelDeployment(BaseModelDeployment):
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
@@ -272,7 +210,7 @@ class T2VModelDeployment(BaseModelDeployment):
@serve.deployment(
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fastvideo"}},
)
class T2V14BModelDeployment(BaseModelDeployment):
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
@@ -283,32 +221,18 @@ class T2V14BModelDeployment(BaseModelDeployment):
print("✅ T2V 14B model initialized successfully")
@serve.deployment(
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class I2VModelDeployment(BaseModelDeployment):
def __init__(self, i2v_model_path: str, output_path: str = "outputs"):
super().__init__(i2v_model_path, output_path)
# Override environment for I2V model
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
self._initialize_generator(MODEL_CONFIGS["I2V-A14B"])
print("✅ I2V model initialized successfully")
app = FastAPI()
limiter = Limiter(key_func=get_remote_address)
app.state.limiter = limiter
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 1})
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 2})
@serve.ingress(app)
class FastVideoAPI:
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle], i2v_deployments: Dict[str, DeploymentHandle] = None):
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
self.t2v_deployments = t2v_deployments
self.i2v_deployments = i2v_deployments or {}
self.all_deployments = {**self.t2v_deployments, **self.i2v_deployments}
# Initialize Prometheus metrics
self.request_count = Counter('fastvideo_requests_total', 'Total FastVideo requests', ['model_type', 'status'])
@@ -333,10 +257,10 @@ class FastVideoAPI:
model_name = self._get_model_name(video_request.model_path)
try:
if video_request.model_path not in self.all_deployments:
if video_request.model_path not in self.t2v_deployments:
raise ValueError(f"Model {video_request.model_path} not found")
response_ref = self.all_deployments[video_request.model_path].generate_video.remote(video_request)
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
response = await response_ref
self._record_metrics(model_name, "success", time.time() - start_time, response)
@@ -367,21 +291,18 @@ class FastVideoAPI:
def validate_configuration(model_paths: List[str], replicas: List[int]) -> None:
assert len(model_paths) > 0, "At least one model must be specified"
assert len(model_paths) == len(replicas), "Number of models and replicas must match"
assert sum(replicas) <= NUM_GPUS, f"Total replicas ({sum(replicas)}) must be <= {NUM_GPUS}"
for model, replica_count in zip(model_paths, replicas):
assert model in SUPPORTED_MODELS, f"Model {model} not supported. Supported models: {SUPPORTED_MODELS}"
assert model in SUPPORTED_MODELS, f"Model {model} not supported"
assert replica_count > 0, f"Replicas must be greater than 0"
def start_ray_serve(
*,
t2v_model_paths: str = "",
t2v_model_replicas: str = "",
i2v_model_paths: str = "",
i2v_model_replicas: str = "",
t2v_model_paths: str,
t2v_model_replicas: str,
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
@@ -389,39 +310,21 @@ def start_ray_serve(
if not ray.is_initialized():
ray.init()
# Parse T2V models
t2v_paths = [p.strip() for p in t2v_model_paths.split(",") if p.strip()]
t2v_reps = [int(r.strip()) for r in t2v_model_replicas.split(",") if r.strip()] if t2v_model_replicas else []
# Parse I2V models
i2v_paths = [p.strip() for p in i2v_model_paths.split(",") if p.strip()]
i2v_reps = [int(r.strip()) for r in i2v_model_replicas.split(",") if r.strip()] if i2v_model_replicas else []
# Validate configurations
all_paths = t2v_paths + i2v_paths
all_replicas = t2v_reps + i2v_reps
validate_configuration(all_paths, all_replicas)
model_paths = t2v_model_paths.split(",")
replicas = [int(r) for r in t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# Create T2V deployments
t2v_deps = {}
for model_path, replica_count in zip(t2v_paths, t2v_reps):
for model_path, replica_count in zip(model_paths, replicas):
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
# Create I2V deployments
i2v_deps = {}
for model_path, replica_count in zip(i2v_paths, i2v_reps):
i2v_dep = I2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
i2v_deps[model_path] = i2v_dep
api = FastVideoAPI.bind(t2v_deps, i2v_deps)
api = FastVideoAPI.bind(t2v_deps)
serve.run(api, route_prefix="/", name="fast_video")
print(f"Ray Serve backend started at http://{host}:{port}")
for model_path, replica_count in zip(t2v_paths, t2v_reps):
for model_path, replica_count in zip(model_paths, replicas):
print(f"T2V Model: {model_path} | Replicas: {replica_count}")
for model_path, replica_count in zip(i2v_paths, i2v_reps):
print(f"I2V Model: {model_path} | Replicas: {replica_count}")
print(f"Health check: http://{host}:{port}/health")
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
@@ -437,20 +340,12 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
parser.add_argument("--t2v_model_paths",
type=str,
default="",
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="",
default="4,4",
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--i2v_model_paths",
type=str,
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
help="Comma separated list of paths to the I2V model(s)")
parser.add_argument("--i2v_model_replicas",
type=str,
default="1",
help="Comma separated list of number of replicas for the I2V model(s)")
parser.add_argument("--output_path",
type=str,
default="outputs",
@@ -466,21 +361,13 @@ if __name__ == "__main__":
args = parser.parse_args()
# Parse and validate all models
t2v_paths = [p.strip() for p in args.t2v_model_paths.split(",") if p.strip()]
t2v_reps = [int(r.strip()) for r in args.t2v_model_replicas.split(",") if r.strip()] if args.t2v_model_replicas else []
i2v_paths = [p.strip() for p in args.i2v_model_paths.split(",") if p.strip()]
i2v_reps = [int(r.strip()) for r in args.i2v_model_replicas.split(",") if r.strip()] if args.i2v_model_replicas else []
all_paths = t2v_paths + i2v_paths
all_replicas = t2v_reps + i2v_reps
validate_configuration(all_paths, all_replicas)
model_paths = args.t2v_model_paths.split(",")
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
start_ray_serve(
t2v_model_paths=args.t2v_model_paths,
t2v_model_replicas=args.t2v_model_replicas,
i2v_model_paths=args.i2v_model_paths,
i2v_model_replicas=args.i2v_model_replicas,
output_path=args.output_path,
host=args.host,
port=args.port,
@@ -489,4 +376,4 @@ if __name__ == "__main__":
setup_signal_handlers()
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
while True:
time.sleep(3600)
time.sleep(3600)
+2 -4
View File
@@ -1,5 +1,3 @@
python examples/inference/gradio/serving/start_ray_serve_app.py \
--t2v_model_paths "" \
--t2v_model_replicas "" \
--i2v_model_paths "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers" \
--i2v_model_replicas "1"
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers" \
--t2v_model_replicas "1"
@@ -20,10 +20,8 @@ DEFAULT_BACKEND_PORT = 8000
DEFAULT_FRONTEND_HOST = "0.0.0.0"
DEFAULT_FRONTEND_PORT = 7860
DEFAULT_OUTPUT_PATH = "outputs"
DEFAULT_T2V_MODELS = ""
DEFAULT_T2V_REPLICAS = ""
DEFAULT_I2V_MODELS = "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers"
DEFAULT_I2V_REPLICAS = "1"
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_T2V_REPLICAS = "4,4"
HEALTH_CHECK_TIMEOUT = 5
HEALTH_CHECK_MAX_RETRIES = 100
@@ -102,12 +100,6 @@ class ServiceManager:
"port": self.args.backend_port
}
# Add I2V parameters if provided
if self.args.i2v_model_paths:
backend_args["i2v_model_paths"] = self.args.i2v_model_paths
if self.args.i2v_model_replicas:
backend_args["i2v_model_replicas"] = self.args.i2v_model_replicas
self.backend_process = self._start_service("ray_serve_backend.py", backend_args, "backend")
return self.backend_process
@@ -119,10 +111,6 @@ class ServiceManager:
"port": self.args.frontend_port
}
# Add I2V parameters if provided
if self.args.i2v_model_paths:
frontend_args["i2v_model_paths"] = self.args.i2v_model_paths
self.frontend_process = self._start_service("gradio_frontend.py", frontend_args, "frontend")
return self.frontend_process
@@ -185,9 +173,6 @@ def print_startup_info(args: argparse.Namespace) -> None:
print("=" * 50)
print(f"T2V Models: {args.t2v_model_paths}")
print(f"T2V Model Replicas: {args.t2v_model_replicas}")
if args.i2v_model_paths:
print(f"I2V Models: {args.i2v_model_paths}")
print(f"I2V Model Replicas: {args.i2v_model_replicas}")
print(f"Output: {args.output_path}")
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
@@ -205,14 +190,6 @@ def parse_arguments() -> argparse.Namespace:
type=str,
default=DEFAULT_T2V_REPLICAS,
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--i2v_model_paths",
type=str,
default=DEFAULT_I2V_MODELS,
help="Comma separated list of paths to the I2V model(s)")
parser.add_argument("--i2v_model_replicas",
type=str,
default=DEFAULT_I2V_REPLICAS,
help="Comma separated list of number of replicas for the I2V model(s)")
parser.add_argument("--output_path",
type=str,
default=DEFAULT_OUTPUT_PATH,
@@ -5,12 +5,12 @@ These are e2e example scripts for finetuning Wan2.1 T2V 1.3B on the crush-smol d
### Download crush-smol dataset:
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/download_dataset.sh`
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
### Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v.sh`
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh`
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
@@ -54,7 +54,7 @@ validation_args=(
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "3.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
+1 -1
View File
@@ -3,4 +3,4 @@ from fastvideo.configs.sample import SamplingParam
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.version import __version__
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
-20
View File
@@ -83,9 +83,6 @@ class PreprocessConfig:
speed_factor: float = 1.0
drop_short_ratio: float = 1.0
do_temporal_sample: bool = False
enable_smart_resize: bool = False
smart_resize_max_area: int | None = None
hw_aspect_threshold: float = 1.5
# Model configuration
training_cfg_rate: float = 0.0
@@ -187,23 +184,6 @@ class PreprocessConfig:
action=StoreBoolean,
default=PreprocessConfig.do_temporal_sample,
help="Whether to do temporal sampling")
preprocess_args.add_argument(
f"--{prefix_with_dot}enable-smart-resize",
action=StoreBoolean,
default=PreprocessConfig.enable_smart_resize,
help="Whether to enable smart resizing")
preprocess_args.add_argument(
f"--{prefix_with_dot}smart-resize-max-area",
type=int,
default=PreprocessConfig.smart_resize_max_area,
help="Maximum area for smart resizing")
preprocess_args.add_argument(
f"--{prefix_with_dot}hw-aspect-threshold",
type=float,
default=PreprocessConfig.hw_aspect_threshold,
help=
"Height/Width aspect ratio threshold. Allowed range is [1/threshold * target_aspect, threshold * target_aspect]."
)
# Model Training configuration
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
+1 -2
View File
@@ -1,10 +1,9 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
__all__ = [
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
"CosmosVideoConfig", "Cosmos25VideoConfig"
"CosmosVideoConfig"
]
-181
View File
@@ -1,181 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_transformer_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class Cosmos25ArchConfig(DiTArchConfig):
"""Configuration for Cosmos 2.5 architecture (MiniTrainDIT)."""
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_transformer_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
# Remove "net." prefix and map official structure to FastVideo
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
r"^net\.x_embedder\.proj\.1\.(.*)$":
r"patch_embed.proj.\1",
# Time embedding: net.t_embedder.1.linear_1.weight -> time_embed.t_embedder.linear_1.weight
r"^net\.t_embedder\.1\.linear_1\.(.*)$":
r"time_embed.t_embedder.linear_1.\1",
r"^net\.t_embedder\.1\.linear_2\.(.*)$":
r"time_embed.t_embedder.linear_2.\1",
# Time embedding norm: net.t_embedding_norm.weight -> time_embed.norm.weight
# Note: This also handles _extra_state if present
r"^net\.t_embedding_norm\.(.*)$":
r"time_embed.norm.\1",
# Cross-attention projection (optional): net.crossattn_proj.0.weight -> crossattn_proj.0.weight
r"^net\.crossattn_proj\.0\.weight$":
r"crossattn_proj.0.weight",
r"^net\.crossattn_proj\.0\.bias$":
r"crossattn_proj.0.bias",
# Transformer blocks: net.blocks.N -> transformer_blocks.N
# Self-attention (self_attn -> attn1)
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$":
r"transformer_blocks.\1.attn1.to_q.\2",
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$":
r"transformer_blocks.\1.attn1.to_k.\2",
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$":
r"transformer_blocks.\1.attn1.to_v.\2",
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$":
r"transformer_blocks.\1.attn1.to_out.\2",
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$":
r"transformer_blocks.\1.attn1.norm_q.weight",
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$":
r"transformer_blocks.\1.attn1.norm_k.weight",
# RMSNorm _extra_state keys (internal PyTorch state, will be recomputed automatically)
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\._extra_state$":
r"transformer_blocks.\1.attn1.norm_q._extra_state",
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\._extra_state$":
r"transformer_blocks.\1.attn1.norm_k._extra_state",
# Cross-attention (cross_attn -> attn2)
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$":
r"transformer_blocks.\1.attn2.to_q.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$":
r"transformer_blocks.\1.attn2.to_k.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$":
r"transformer_blocks.\1.attn2.to_v.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$":
r"transformer_blocks.\1.attn2.to_out.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$":
r"transformer_blocks.\1.attn2.norm_q.weight",
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$":
r"transformer_blocks.\1.attn2.norm_k.weight",
# RMSNorm _extra_state keys for cross-attention
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\._extra_state$":
r"transformer_blocks.\1.attn2.norm_q._extra_state",
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\._extra_state$":
r"transformer_blocks.\1.attn2.norm_k._extra_state",
# MLP: net.blocks.N.mlp.layer1 -> transformer_blocks.N.mlp.fc_in
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$":
r"transformer_blocks.\1.mlp.fc_in.\2",
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$":
r"transformer_blocks.\1.mlp.fc_out.\2",
# AdaLN-LoRA modulations: net.blocks.N.adaln_modulation_* -> transformer_blocks.N.adaln_modulation_*
# These are now at the block level, not inside norm layers
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.1\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_self_attn.1.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.2\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_self_attn.2.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.1\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_cross_attn.1.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.2\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_cross_attn.2.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
# Layer norms: net.blocks.N.layer_norm_* -> transformer_blocks.N.norm*.norm
r"^net\.blocks\.(\d+)\.layer_norm_self_attn\._extra_state$":
r"transformer_blocks.\1.norm1.norm._extra_state",
r"^net\.blocks\.(\d+)\.layer_norm_cross_attn\._extra_state$":
r"transformer_blocks.\1.norm2.norm._extra_state",
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$":
r"transformer_blocks.\1.norm3.norm._extra_state",
# Final layer: net.final_layer.linear -> final_layer.proj_out
r"^net\.final_layer\.linear\.(.*)$":
r"final_layer.proj_out.\1",
# Final layer AdaLN-LoRA: net.final_layer.adaln_modulation -> final_layer.linear_*
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$":
r"final_layer.linear_1.\1",
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$":
r"final_layer.linear_2.\1",
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
# - net.pos_embedder.* (seq, dim_spatial_range, dim_temporal_range) - These are computed dynamically
# in FastVideo's Cosmos25RotaryPosEmbed forward() method, so they don't need to be loaded.
# - net.accum_* keys (training metadata) - These are skipped during checkpoint loading.
})
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$":
r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$":
r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$":
r"transformer_blocks.\1.mlp.\2",
})
# Cosmos 2.5 specific config parameters
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 16
attention_head_dim: int = 128 # 2048 / 16
num_layers: int = 28
mlp_ratio: float = 4.0
text_embed_dim: int = 1024
adaln_lora_dim: int = 256
use_adaln_lora: bool = True
max_size: tuple[int, int, int] = (128, 240, 240)
patch_size: tuple[int, int, int] = (1, 2, 2)
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0) # T, H, W scaling
concat_padding_mask: bool = True
extra_pos_embed_type: str | None = None # "learnable" or None
# Note: Official checkpoint has use_crossattn_projection=True with 100K-dim input from Qwen 7B.
# When enabled, must provide 100,352-dim embeddings to match the projection layer in checkpoint.
use_crossattn_projection: bool = False
crossattn_proj_in_channels: int = 100352 # Qwen 7B embedding dimension
rope_enable_fps_modulation: bool = True
qk_norm: str = "rms_norm"
eps: float = 1e-6
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels
@dataclass
class Cosmos25VideoConfig(DiTConfig):
"""Configuration for Cosmos 2.5 video generation model."""
arch_config: DiTArchConfig = field(default_factory=Cosmos25ArchConfig)
prefix: str = "Cosmos25"
-1
View File
@@ -45,7 +45,6 @@ class PipelineConfig:
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
disable_autocast: bool = False
is_causal: bool = False
# Model configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
+1 -4
View File
@@ -14,7 +14,7 @@ from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
SelfForcingWanT2V480PConfig, WANV2VConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -38,9 +38,6 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers":
SelfForcingWan2_2_T2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
-14
View File
@@ -176,17 +176,3 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
@dataclass
class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
is_causal: bool = True
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
warp_denoising_step: bool = True
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
+5 -17
View File
@@ -11,7 +11,7 @@ from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_Sampl
# isort: off
from fastvideo.configs.sample.wan import (
FastWanT2V480P_SamplingParam,
FastWanT2V480PConfig,
Wan2_1_Fun_1_3B_InP_SamplingParam,
Wan2_2_I2V_A14B_SamplingParam,
Wan2_2_T2V_A14B_SamplingParam,
@@ -21,8 +21,7 @@ from fastvideo.configs.sample.wan import (
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
Wan2_1_Fun_1_3B_Control_SamplingParam,
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
SelfForcingWanT2V480PConfig,
)
# isort: on
from fastvideo.logger import init_logger
@@ -65,7 +64,7 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
FastWanT2V480P_SamplingParam,
FastWanT2V480PConfig,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
@@ -73,18 +72,11 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
# Causal Self-Forcing Wan2.2
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
SelfForcingWanT2V480PConfig,
# Cosmos2
"nvidia/Cosmos-Predict2-2B-Video2World":
Cosmos_Predict2_2B_Video2World_SamplingParam,
# Add other specific weight variants
}
@@ -94,8 +86,6 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -106,9 +96,7 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
"wandmdpipeline": FastWanT2V480P_SamplingParam,
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam,
"stepvideo": StepVideoT2VSamplingParam
# Other fallbacks by architecture
}
+2 -13
View File
@@ -97,7 +97,7 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
@dataclass
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
# DMD parameters
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
num_inference_steps: int = 3
@@ -183,16 +183,5 @@ class Wan2_2_Fun_A14B_Control_SamplingParam(
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
Wan2_1_Fun_1_3B_InP_SamplingParam):
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
pass
@dataclass
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
Wan2_2_T2V_A14B_SamplingParam):
num_inference_steps: int = 8
num_frames: int = 81
height: int = 448
width: int = 832
fps: int = 16
-41
View File
@@ -152,44 +152,3 @@ class TemporalRandomCrop:
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
def best_output_size(
width: int,
height: int,
width_stride: int,
height_stride: int,
max_area: int,
) -> tuple[int, int]:
"""
Calculate the best output size (width, height) given the original dimensions, strides and max area.
The aspect ratio is preserved as much as possible.
Args:
width (int): Original width
height (int): Original height
width_stride (int): Width stride requirement
height_stride (int): Height stride requirement
max_area (int): Maximum allowed area (width * height)
Returns:
tuple[int, int]: (new_width, new_height)
"""
aspect_ratio = width / height
# Scale dimensions if they exceed max_area
current_area = width * height
if current_area > max_area:
scale = (max_area / current_area)**0.5
width = int(width * scale)
height = int(height * scale)
# Round to the nearest multiple of stride
width = round(width / width_stride) * width_stride
height = round(height / height_stride) * height_stride
# Ensure dimensions are at least one stride
width = max(width, width_stride)
height = max(height, height_stride)
return width, height
+4 -21
View File
@@ -101,19 +101,12 @@ class BaseLayerWithLoRA(nn.Module):
def set_lora_weights(self,
A: torch.Tensor,
B: torch.Tensor,
lora_alpha: float | None = None,
training_mode: bool = False,
lora_path: str | None = None) -> None:
self.lora_A = torch.nn.Parameter(
A) # share storage with weights in the pipeline
self.lora_B = torch.nn.Parameter(B)
self.disable_lora = False
# Store rank and alpha directly
rank = A.shape[0] # rank is the first dimension of A
self.lora_rank = rank
self.lora_alpha = int(lora_alpha) if lora_alpha is not None else rank
if not training_mode:
self.merge_lora_weights()
self.lora_path = lora_path
@@ -141,13 +134,8 @@ class BaseLayerWithLoRA(nn.Module):
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
get_local_torch_device()).full_tensor()
# Apply LoRA with alpha scaling
lora_delta = (self.slice_lora_b_weights(self.lora_B).to(data)
@ self.slice_lora_a_weights(self.lora_A).to(data))
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
lora_delta *= (self.lora_alpha / self.lora_rank)
data += lora_delta
data += (self.slice_lora_b_weights(self.lora_B).to(data)
@ self.slice_lora_a_weights(self.lora_A).to(data))
unsharded_base_layer.weight = nn.Parameter(data.to(current_device))
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
unsharded_base_layer.bias = nn.Parameter(
@@ -166,13 +154,8 @@ class BaseLayerWithLoRA(nn.Module):
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(get_local_torch_device())
# Apply LoRA with alpha scaling
lora_delta = (self.slice_lora_b_weights(self.lora_B.to(data))
@ self.slice_lora_a_weights(self.lora_A.to(data)))
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
lora_delta *= (self.lora_alpha / self.lora_rank)
data += lora_delta
data += \
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
self.base_layer.weight.data = data.to(current_device,
non_blocking=True)
-961
View File
@@ -1,961 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import numpy as np
import torch
import torch.nn as nn
from torchvision import transforms
from fastvideo.attention import DistributedAttention, LocalAttention
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import apply_rotary_emb
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
class Cosmos25PatchEmbed(nn.Module):
"""
COSMOS 2.5 patch embedding - converts video (B, C, T, H, W) to patches (B, T', H', W', D).
Uses linear projection after rearranging patches.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
patch_size: tuple[int, int, int],
) -> None:
super().__init__()
self.patch_size = patch_size
self.dim = in_channels * patch_size[0] * patch_size[1] * patch_size[2]
self.proj = nn.Linear(self.dim, out_channels, bias=False)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""
Args:
hidden_states: (B, C, T, H, W)
Returns:
(B, T', H', W', D) where T'=T//pt, H'=H//ph, W'=W//pw
"""
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
# Rearrange: b c (t pt) (h ph) (w pw) -> b t h w (c pt ph pw)
hidden_states = hidden_states.reshape(
batch_size, num_channels,
num_frames // p_t, p_t,
height // p_h, p_h,
width // p_w, p_w
)
hidden_states = hidden_states.permute(0, 2, 4, 6, 1, 3, 5, 7)
hidden_states = hidden_states.flatten(4, 7) # Flatten patch dimensions
# Project to model dimension
hidden_states = self.proj(hidden_states)
return hidden_states
class Cosmos25TimestepEmbedding(nn.Module):
"""
COSMOS 2.5 timestep embedding with AdaLN-LoRA support.
Generates both standard embedding and AdaLN-LoRA parameters.
"""
def __init__(
self,
in_features: int,
out_features: int,
use_adaln_lora: bool = True,
adaln_lora_dim: int = 256,
) -> None:
super().__init__()
self.use_adaln_lora = use_adaln_lora
self.linear_1 = nn.Linear(in_features, out_features, bias=False)
self.activation = nn.SiLU()
if use_adaln_lora:
self.linear_2 = nn.Linear(out_features, 3 * out_features, bias=False)
else:
self.linear_2 = nn.Linear(out_features, out_features, bias=False)
def forward(self, sample: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]:
"""
Returns:
emb: Standard embedding (B, T, D)
adaln_lora: AdaLN-LoRA parameters (B, T, 3D) or None
"""
emb = self.linear_1(sample)
emb = self.activation(emb)
emb = self.linear_2(emb)
if self.use_adaln_lora:
adaln_lora = emb # (B, T, 3D)
emb_standard = sample # Use input as standard embedding
else:
emb_standard = emb
adaln_lora = None
return emb_standard, adaln_lora
class Cosmos25Embedding(nn.Module):
"""
COSMOS 2.5 timestep conditioning embedding.
Generates sinusoidal embeddings and processes them through MLP.
"""
def __init__(
self,
embedding_dim: int,
condition_dim: int,
use_adaln_lora: bool = True,
adaln_lora_dim: int = 256,
) -> None:
super().__init__()
self.time_proj = Timesteps(embedding_dim, flip_sin_to_cos=True, downscale_freq_shift=0.0)
self.t_embedder = Cosmos25TimestepEmbedding(
embedding_dim,
condition_dim,
use_adaln_lora=use_adaln_lora,
adaln_lora_dim=adaln_lora_dim,
)
self.norm = RMSNorm(embedding_dim, eps=1e-6)
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""
Args:
timestep: (B, T) tensor of timesteps
Returns:
embedded_timestep: Normalized timestep embedding (B, T, D)
adaln_lora: AdaLN-LoRA parameters (B, T, 3D) or None
"""
# Handle 2D timestep input (B, T) like the official model
assert timestep.ndim == 2, f"Expected 2D timestep, got {timestep.ndim}D with shape {timestep.shape}"
B, T = timestep.shape
# Flatten for Timesteps layer which expects 1D, then reshape back
timestep_flat = timestep.flatten() # (B*T,)
timesteps_proj = self.time_proj(timestep_flat).type_as(hidden_states) # (B*T, D)
timesteps_proj = timesteps_proj.reshape(B, T, -1) # (B, T, D)
embedded_timestep, adaln_lora = self.t_embedder(timesteps_proj)
embedded_timestep = self.norm(embedded_timestep)
return embedded_timestep, adaln_lora
class Cosmos25AdaLayerNormZero(nn.Module):
"""
COSMOS 2.5 Adaptive Layer Normalization with zero initialization and gate.
This is a simplified version that expects pre-computed shift/scale/gate parameters.
"""
def __init__(
self,
in_features: int,
) -> None:
super().__init__()
self.norm = nn.LayerNorm(in_features, elementwise_affine=False, eps=1e-6)
def forward(
self,
hidden_states: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor,
) -> torch.Tensor:
"""
Args:
hidden_states: Input tensor
shift: Shift parameter for modulation
scale: Scale parameter for modulation
Returns:
normalized_hidden_states: Modulated normalized hidden states
"""
# Apply layer norm and modulation
hidden_states = self.norm(hidden_states)
hidden_states = hidden_states * (1 + scale) + shift
return hidden_states
class Cosmos25SelfAttention(nn.Module):
"""
COSMOS 2.5 self-attention with QK normalization and RoPE.
"""
def __init__(
self,
dim: int,
num_heads: int,
qk_norm: bool = True,
eps: float = 1e-6,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.to_q = nn.Linear(dim, dim, bias=False)
self.to_k = nn.Linear(dim, dim, bias=False)
self.to_v = nn.Linear(dim, dim, bias=False)
self.to_out = nn.Linear(dim, dim, bias=False)
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
# Use DistributedAttention for flexible backend support (torch SDPA / FlashAttention)
# For single-GPU (non-distributed), use LocalAttention to avoid distributed requirements
if supported_attention_backends is None:
supported_attention_backends = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
# Always use DistributedAttention (requires distributed environment to be initialized)
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix="self_attn"
)
def forward(
self,
hidden_states: torch.Tensor,
rope_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""
Args:
hidden_states: (B, S, D) where S = T*H*W
rope_emb: Tuple of (cos, sin) for RoPE
"""
# Get QKV
query = self.to_q(hidden_states)
key = self.to_k(hidden_states)
value = self.to_v(hidden_states)
# Reshape for multi-head attention: (B, S, D) -> (B, S, H, D_h) -> (B, H, S, D_h)
query = query.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
key = key.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
value = value.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
# Apply QK normalization
query = self.norm_q(query)
key = self.norm_k(key)
# Apply RoPE if provided (query/key are now in (B, H, S, D_h) format)
if rope_emb is not None:
cos, sin = rope_emb
query = apply_rotary_emb(query, (cos, sin), use_real=True, use_real_unbind_dim=-2)
key = apply_rotary_emb(key, (cos, sin), use_real=True, use_real_unbind_dim=-2)
# Attention computation using DistributedAttention or LocalAttention
# Both expect (B, S, H, D_h), so transpose first
query = query.transpose(1, 2) # (B, H, S, D_h) -> (B, S, H, D_h)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attn_output, _ = self.attn(query, key, value)
# Reshape back: (B, S, H, D_h) -> (B, S, H*D_h)
attn_output = attn_output.flatten(-2, -1)
# Output projection
attn_output = self.to_out(attn_output)
return attn_output
class Cosmos25CrossAttention(nn.Module):
"""
COSMOS 2.5 cross-attention for text conditioning.
"""
def __init__(
self,
dim: int,
cross_attention_dim: int,
num_heads: int,
qk_norm: bool = True,
eps: float = 1e-6,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.cross_attention_dim = cross_attention_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.to_q = nn.Linear(dim, dim, bias=False)
self.to_k = nn.Linear(cross_attention_dim, dim, bias=False)
self.to_v = nn.Linear(cross_attention_dim, dim, bias=False)
self.to_out = nn.Linear(dim, dim, bias=False)
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
if supported_attention_backends is None:
supported_attention_backends = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
# Use LocalAttention for cross-attention since text embeddings are not sharded
# in sequence parallelism (replicated across ranks)
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Args:
hidden_states: (B, S, D)
encoder_hidden_states: (B, N, D_text)
"""
# Get QKV
query = self.to_q(hidden_states)
key = self.to_k(encoder_hidden_states)
value = self.to_v(encoder_hidden_states)
# Reshape for multi-head attention
query = query.unflatten(-1, (self.num_heads, self.head_dim))
key = key.unflatten(-1, (self.num_heads, self.head_dim))
value = value.unflatten(-1, (self.num_heads, self.head_dim))
# Apply QK normalization
query = self.norm_q(query)
key = self.norm_k(key)
# LocalAttention expects (B, S, H, D_h), which is what we already have
attn_output = self.attn(query, key, value)
# Reshape back: (B, S, H, D_h) -> (B, S, H*D_h)
attn_output = attn_output.flatten(-2, -1)
# Output projection
attn_output = self.to_out(attn_output)
return attn_output
class Cosmos25TransformerBlock(nn.Module):
"""
COSMOS 2.5 transformer block with self-attention, cross-attention, and MLP.
Uses AdaLN-LoRA for conditioning.
Matches the official architecture where modulation parameters are computed once per block.
"""
def __init__(
self,
num_attention_heads: int,
attention_head_dim: int,
cross_attention_dim: int,
mlp_ratio: float = 4.0,
adaln_lora_dim: int = 256,
use_adaln_lora: bool = True,
qk_norm: bool = True,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
super().__init__()
hidden_size = num_attention_heads * attention_head_dim
self.use_adaln_lora = use_adaln_lora
# Layer norms (no modulation logic inside)
self.norm1 = Cosmos25AdaLayerNormZero(hidden_size)
self.norm2 = Cosmos25AdaLayerNormZero(hidden_size)
self.norm3 = Cosmos25AdaLayerNormZero(hidden_size)
# Attention and MLP layers
self.attn1 = Cosmos25SelfAttention(
dim=hidden_size,
num_heads=num_attention_heads,
qk_norm=qk_norm,
supported_attention_backends=supported_attention_backends,
)
self.attn2 = Cosmos25CrossAttention(
dim=hidden_size,
cross_attention_dim=cross_attention_dim,
num_heads=num_attention_heads,
qk_norm=qk_norm,
supported_attention_backends=supported_attention_backends,
)
self.mlp = MLP(hidden_size, int(hidden_size * mlp_ratio), act_type="gelu", bias=False)
# AdaLN modulation layers (compute shift/scale/gate for each sub-layer)
# These match the official model's adaln_modulation_* layers
if use_adaln_lora:
self.adaln_modulation_self_attn = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
)
self.adaln_modulation_cross_attn = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
)
self.adaln_modulation_mlp = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
)
else:
self.adaln_modulation_self_attn = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
)
self.adaln_modulation_cross_attn = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
)
self.adaln_modulation_mlp = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
embedded_timestep: torch.Tensor,
adaln_lora: torch.Tensor | None = None,
rope_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
extra_pos_emb: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Args:
hidden_states: (B, T, H, W, D)
encoder_hidden_states: (B, N, D_text)
embedded_timestep: (B, T, D)
adaln_lora: (B, T, 3D) AdaLN-LoRA parameters
rope_emb: Tuple of (cos, sin) for RoPE
extra_pos_emb: Optional learnable positional embeddings
"""
# Add extra positional embeddings if provided
if extra_pos_emb is not None:
hidden_states = hidden_states + extra_pos_emb
B, T, H, W, D = hidden_states.shape
# Step 1: Compute ALL modulation parameters once (matches official model)
if self.use_adaln_lora and adaln_lora is not None:
shift_self_attn, scale_self_attn, gate_self_attn = (
self.adaln_modulation_self_attn(embedded_timestep) + adaln_lora
).chunk(3, dim=-1)
shift_cross_attn, scale_cross_attn, gate_cross_attn = (
self.adaln_modulation_cross_attn(embedded_timestep) + adaln_lora
).chunk(3, dim=-1)
shift_mlp, scale_mlp, gate_mlp = (
self.adaln_modulation_mlp(embedded_timestep) + adaln_lora
).chunk(3, dim=-1)
else:
shift_self_attn, scale_self_attn, gate_self_attn = self.adaln_modulation_self_attn(
embedded_timestep
).chunk(3, dim=-1)
shift_cross_attn, scale_cross_attn, gate_cross_attn = self.adaln_modulation_cross_attn(
embedded_timestep
).chunk(3, dim=-1)
shift_mlp, scale_mlp, gate_mlp = self.adaln_modulation_mlp(embedded_timestep).chunk(3, dim=-1)
# Reshape modulation parameters from (B, T, D) to (B, T, 1, 1, D) for broadcasting
shift_self_attn = shift_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
scale_self_attn = scale_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
gate_self_attn = gate_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
shift_cross_attn = shift_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
scale_cross_attn = scale_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
gate_cross_attn = gate_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
shift_mlp = shift_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
scale_mlp = scale_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
gate_mlp = gate_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
# Step 2: Self-attention block
norm_hidden_states = self.norm1(hidden_states, shift_self_attn, scale_self_attn)
# Flatten for attention: (B, T, H, W, D) -> (B, THW, D)
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
attn_output = self.attn1(norm_hidden_states_flat, rope_emb=rope_emb)
# Reshape back and apply residual
attn_output = attn_output.unflatten(1, (T, H, W)) # (B, T, H, W, D)
hidden_states = hidden_states + gate_self_attn * attn_output
# Step 3: Cross-attention block
norm_hidden_states = self.norm2(hidden_states, shift_cross_attn, scale_cross_attn)
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
attn_output = self.attn2(
norm_hidden_states_flat,
encoder_hidden_states=encoder_hidden_states,
attention_mask=attention_mask,
)
attn_output = attn_output.unflatten(1, (T, H, W))
hidden_states = hidden_states + gate_cross_attn * attn_output
# Step 4: MLP block
norm_hidden_states = self.norm3(hidden_states, shift_mlp, scale_mlp)
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
mlp_output = self.mlp(norm_hidden_states_flat)
mlp_output = mlp_output.unflatten(1, (T, H, W))
hidden_states = hidden_states + gate_mlp * mlp_output
return hidden_states
class Cosmos25RotaryPosEmbed(nn.Module):
"""
COSMOS 2.5 3D Rotary Position Embedding with NTK-aware extrapolation.
"""
def __init__(
self,
hidden_size: int,
max_size: tuple[int, int, int] = (128, 240, 240),
patch_size: tuple[int, int, int] = (1, 2, 2),
base_fps: int = 24,
rope_scale: tuple[float, float, float] = (1.0, 1.0, 1.0),
enable_fps_modulation: bool = True,
) -> None:
super().__init__()
self.max_size = [size // patch for size, patch in zip(max_size, patch_size, strict=True)]
self.patch_size = patch_size
self.base_fps = base_fps
self.enable_fps_modulation = enable_fps_modulation
# Split dimensions: 1/3 for T, 1/3 for H, 1/3 for W
self.dim_h = hidden_size // 6 * 2
self.dim_w = hidden_size // 6 * 2
self.dim_t = hidden_size - self.dim_h - self.dim_w
# NTK-aware extrapolation factors
self.h_ntk_factor = rope_scale[1] ** (self.dim_h / (self.dim_h - 2))
self.w_ntk_factor = rope_scale[2] ** (self.dim_w / (self.dim_w - 2))
self.t_ntk_factor = rope_scale[0] ** (self.dim_t / (self.dim_t - 2))
def forward(
self, hidden_states: torch.Tensor, fps: int | None = None
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate 3D RoPE embeddings.
Args:
hidden_states: (B, T, H, W, D) - patch-embedded features
fps: Frames per second for temporal scaling
Returns:
cos, sin: RoPE embeddings (THW, D)
"""
batch_size, T, H, W, input_dim = hidden_states.shape
device = hidden_states.device
# T, H, W are already patch dimensions after patch_embed
# No need to divide by patch_size
# Generate frequency scales with NTK
h_theta = 10000.0 * self.h_ntk_factor
w_theta = 10000.0 * self.w_ntk_factor
t_theta = 10000.0 * self.t_ntk_factor
seq = torch.arange(max(self.max_size), device=device, dtype=torch.float32)
# Use self.dim_h/w/t which were set during initialization
dim_h_range = torch.arange(0, self.dim_h, 2, device=device, dtype=torch.float32)[: (self.dim_h // 2)] / self.dim_h
dim_w_range = torch.arange(0, self.dim_w, 2, device=device, dtype=torch.float32)[: (self.dim_w // 2)] / self.dim_w
dim_t_range = torch.arange(0, self.dim_t, 2, device=device, dtype=torch.float32)[: (self.dim_t // 2)] / self.dim_t
h_spatial_freqs = 1.0 / (h_theta ** dim_h_range)
w_spatial_freqs = 1.0 / (w_theta ** dim_w_range)
temporal_freqs = 1.0 / (t_theta ** dim_t_range)
# Generate positional embeddings
half_emb_h = torch.outer(seq[:H], h_spatial_freqs)
half_emb_w = torch.outer(seq[:W], w_spatial_freqs)
if self.enable_fps_modulation and fps is not None:
# Apply FPS scaling
half_emb_t = torch.outer(seq[:T] / fps * self.base_fps, temporal_freqs)
else:
half_emb_t = torch.outer(seq[:T], temporal_freqs)
# Broadcast and concatenate embeddings
emb_t = half_emb_t[:, None, None, :].repeat(1, H, W, 1)
emb_h = half_emb_h[None, :, None, :].repeat(T, 1, W, 1)
emb_w = half_emb_w[None, None, :, :].repeat(T, H, 1, 1)
# Concatenate [t, h, w, t, h, w] for sin/cos pairs
freqs = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1)
freqs = freqs.flatten(0, 2).float() # (THW, D)
cos = torch.cos(freqs) # (THW, D)
sin = torch.sin(freqs) # (THW, D)
return cos, sin
class Cosmos25LearnablePositionalEmbed(nn.Module):
"""
COSMOS 2.5 learnable absolute positional embeddings (optional).
"""
def __init__(
self,
hidden_size: int,
max_size: tuple[int, int, int],
patch_size: tuple[int, int, int],
eps: float = 1e-6,
) -> None:
super().__init__()
self.max_size = [size // patch for size, patch in zip(max_size, patch_size, strict=True)]
self.patch_size = patch_size
self.eps = eps
self.pos_emb_t = nn.Parameter(torch.zeros(self.max_size[0], hidden_size))
self.pos_emb_h = nn.Parameter(torch.zeros(self.max_size[1], hidden_size))
self.pos_emb_w = nn.Parameter(torch.zeros(self.max_size[2], hidden_size))
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""
Args:
hidden_states: (B, T, H, W, D)
Returns:
pos_emb: (B, T, H, W, D)
"""
B, T, H, W, D = hidden_states.shape
emb_t = self.pos_emb_t[:T][None, :, None, None, :].repeat(B, 1, H, W, 1)
emb_h = self.pos_emb_h[:H][None, None, :, None, :].repeat(B, T, 1, W, 1)
emb_w = self.pos_emb_w[:W][None, None, None, :, :].repeat(B, T, H, 1, 1)
emb = emb_t + emb_h + emb_w
# Normalize
norm = torch.linalg.vector_norm(emb, dim=-1, keepdim=True, dtype=torch.float32)
norm = torch.add(self.eps, norm, alpha=np.sqrt(norm.numel() / emb.numel()))
return (emb / norm).type_as(hidden_states)
class Cosmos25FinalLayer(nn.Module):
"""
COSMOS 2.5 final layer with AdaLN modulation and unpatchification.
"""
def __init__(
self,
hidden_size: int,
out_channels: int,
patch_size: tuple[int, int, int],
adaln_lora_dim: int = 256,
use_adaln_lora: bool = True,
) -> None:
super().__init__()
self.hidden_size = hidden_size
self.use_adaln_lora = use_adaln_lora
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.activation = nn.SiLU()
if use_adaln_lora:
self.linear_1 = nn.Linear(hidden_size, adaln_lora_dim, bias=False)
self.linear_2 = nn.Linear(adaln_lora_dim, 2 * hidden_size, bias=False)
else:
self.linear_1 = nn.Identity()
self.linear_2 = nn.Linear(hidden_size, 2 * hidden_size, bias=False)
# Output projection
output_dim = out_channels * patch_size[0] * patch_size[1] * patch_size[2]
self.proj_out = nn.Linear(hidden_size, output_dim, bias=False)
def forward(
self,
hidden_states: torch.Tensor,
embedded_timestep: torch.Tensor,
adaln_lora: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Args:
hidden_states: (B, T, H, W, D)
embedded_timestep: (B, T, D) or (B, D)
adaln_lora: (B, T, 3D) or None
"""
# Generate modulation parameters
embedded_timestep = self.activation(embedded_timestep)
embedded_timestep = self.linear_1(embedded_timestep)
embedded_timestep = self.linear_2(embedded_timestep)
if self.use_adaln_lora and adaln_lora is not None:
# Use first 2*hidden_size elements for shift/scale
embedded_timestep = embedded_timestep + adaln_lora[..., : 2 * self.hidden_size]
shift, scale = embedded_timestep.chunk(2, dim=-1)
# Apply normalization and modulation
hidden_states = self.norm(hidden_states)
# Reshape for broadcasting if needed
if embedded_timestep.ndim == 2:
shift, scale = (x.unsqueeze(1) for x in (shift, scale))
elif embedded_timestep.ndim == 3 and hidden_states.ndim == 5:
shift, scale = (x.unsqueeze(2).unsqueeze(2) for x in (shift, scale))
hidden_states = hidden_states * (1 + scale) + shift
# Project to output
hidden_states = self.proj_out(hidden_states)
return hidden_states
class Cosmos25Transformer3DModel(BaseDiT):
"""
COSMOS 2.5 DiT - MiniTrainDIT architecture adapted for FastVideo.
Key features:
- AdaLN-LoRA conditioning
- 3D RoPE with NTK-aware extrapolation
- Optional learnable positional embeddings
- QK normalization
- Cross-attention projection (optional)
"""
_fsdp_shard_conditions = Cosmos25VideoConfig()._fsdp_shard_conditions
_compile_conditions = Cosmos25VideoConfig()._compile_conditions
param_names_mapping = Cosmos25VideoConfig().param_names_mapping
lora_param_names_mapping = Cosmos25VideoConfig().lora_param_names_mapping
def __init__(self, config: Cosmos25VideoConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = inner_dim
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.max_size = config.max_size
self.rope_scale = config.rope_scale
self.concat_padding_mask = config.concat_padding_mask
self.use_adaln_lora = getattr(config, "use_adaln_lora", True)
self.adaln_lora_dim = getattr(config, "adaln_lora_dim", 256)
self.extra_pos_embed_type = getattr(config, "extra_pos_embed_type", None)
self.use_crossattn_projection = getattr(config, "use_crossattn_projection", False)
# 1. Patch Embedding
# Account for: VAE channels + condition_mask (1) + padding_mask (1 if concat_padding_mask)
patch_embed_in_channels = config.in_channels # Base VAE channels (16)
patch_embed_in_channels += 1 # Always add 1 for condition_mask
if config.concat_padding_mask:
patch_embed_in_channels += 1 # Add 1 for padding_mask
# Total: 16 + 1 + 1 = 18 (with concat_padding_mask=True)
self.patch_embed = Cosmos25PatchEmbed(
patch_embed_in_channels, inner_dim, config.patch_size
)
# 2. Positional Embeddings
self.rope = Cosmos25RotaryPosEmbed(
hidden_size=config.attention_head_dim,
max_size=config.max_size,
patch_size=config.patch_size,
rope_scale=config.rope_scale,
enable_fps_modulation=getattr(config, "rope_enable_fps_modulation", True),
)
self.learnable_pos_embed = None
if self.extra_pos_embed_type == "learnable":
self.learnable_pos_embed = Cosmos25LearnablePositionalEmbed(
hidden_size=inner_dim,
max_size=config.max_size,
patch_size=config.patch_size,
)
# 3. Time Embedding
self.time_embed = Cosmos25Embedding(
inner_dim,
inner_dim,
use_adaln_lora=self.use_adaln_lora,
adaln_lora_dim=self.adaln_lora_dim,
)
# 4. Cross-attention projection (optional)
if self.use_crossattn_projection:
crossattn_proj_in_channels = getattr(config, "crossattn_proj_in_channels", config.text_embed_dim)
self.crossattn_proj = nn.Sequential(
nn.Linear(crossattn_proj_in_channels, config.text_embed_dim, bias=True),
nn.GELU(),
)
# 5. Transformer Blocks
self.transformer_blocks = nn.ModuleList([
Cosmos25TransformerBlock(
num_attention_heads=config.num_attention_heads,
attention_head_dim=config.attention_head_dim,
cross_attention_dim=config.text_embed_dim,
mlp_ratio=config.mlp_ratio,
adaln_lora_dim=self.adaln_lora_dim,
use_adaln_lora=self.use_adaln_lora,
qk_norm=(config.qk_norm == "rms_norm"),
supported_attention_backends=config._supported_attention_backends,
)
for i in range(config.num_layers)
])
# 6. Final Layer
self.final_layer = Cosmos25FinalLayer(
hidden_size=inner_dim,
out_channels=config.out_channels,
patch_size=config.patch_size,
adaln_lora_dim=self.adaln_lora_dim,
use_adaln_lora=self.use_adaln_lora,
)
self.gradient_checkpointing = False
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
attention_mask: torch.Tensor | None = None,
fps: int | None = None,
condition_mask: torch.Tensor | None = None,
padding_mask: torch.Tensor | None = None,
**kwargs,
) -> torch.Tensor:
"""
Args:
hidden_states: (B, C, T, H, W) latent video
timestep: (B,) or (B, T) diffusion timesteps
encoder_hidden_states: (B, N, D_text) text embeddings
attention_mask: Optional attention mask
fps: Frames per second
condition_mask: (B, 1, T, H, W) conditioning mask
padding_mask: (B, 1, H, W) padding mask
"""
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
batch_size, num_channels, num_frames, height, width = hidden_states.shape
# 1. Concatenate condition mask if provided
if condition_mask is not None:
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
# 2. Concatenate padding mask if needed
if self.concat_padding_mask and padding_mask is not None:
padding_mask = transforms.functional.resize(
padding_mask,
list(hidden_states.shape[-2:]),
interpolation=transforms.InterpolationMode.NEAREST,
)
hidden_states = torch.cat(
[hidden_states, padding_mask.unsqueeze(2).repeat(1, 1, num_frames, 1, 1)],
dim=1,
)
# 3. Patchify input
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
hidden_states = self.patch_embed(hidden_states) # (B, T', H', W', D)
# 4. Generate RoPE embeddings (after patchify, using patch dimensions)
rope_emb = self.rope(hidden_states, fps=fps)
# 5. Generate learnable positional embeddings (if used)
extra_pos_emb = None
if self.learnable_pos_embed is not None:
extra_pos_emb = self.learnable_pos_embed(hidden_states)
# 6. Timestep embeddings
# Official model expects timestep in (B, T) format, so ensure it has 2D shape
if timestep.ndim == 1:
# Scalar timestep per sample: (B,) -> (B, 1)
timestep = timestep.unsqueeze(1)
elif timestep.ndim == 2:
# Already in (B, T) format
pass
else:
raise ValueError(f"Unsupported timestep shape: {timestep.shape}")
# Now timestep is always (B, T), pass directly to time_embed
embedded_timestep, adaln_lora = self.time_embed(hidden_states, timestep)
# 7. Apply cross-attention projection (if used)
if self.use_crossattn_projection:
encoder_hidden_states = self.crossattn_proj(encoder_hidden_states)
# Prepare attention mask
if attention_mask is not None:
attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) # (B, 1, 1, N)
# 8. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.transformer_blocks:
hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
embedded_timestep,
adaln_lora,
rope_emb,
extra_pos_emb,
attention_mask,
)
else:
for i, block in enumerate(self.transformer_blocks):
hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
embedded_timestep=embedded_timestep,
adaln_lora=adaln_lora,
rope_emb=rope_emb,
extra_pos_emb=extra_pos_emb,
attention_mask=attention_mask,
)
# 9. Final layer - output norm & projection
hidden_states = self.final_layer(hidden_states, embedded_timestep, adaln_lora)
# 10. Unpatchify: (B, T', H', W', P) -> (B, C, T, H, W)
# After unflatten: (B, T', H', W', p_t, p_h, p_w, C) with dims [0,1,2,3,4,5,6,7]
hidden_states = hidden_states.unflatten(-1, (p_t, p_h, p_w, self.out_channels))
# Permute to: (B, C, T', p_t, H', p_h, W', p_w)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
# Flatten pairs to get (B, C, T, H, W)
hidden_states = hidden_states.flatten(2, 3).flatten(3, 4).flatten(4, 5)
return hidden_states
+18 -39
View File
@@ -4,8 +4,6 @@
# Copyright 2024 The TorchTune Authors.
# Copyright 2025 The FastVideo Authors.
from __future__ import annotations
import os
import contextlib
from collections.abc import Callable, Generator
from itertools import chain
@@ -199,11 +197,10 @@ def shard_model(
Raises:
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
"""
# Check if we should use size-based filtering
use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1"
if not fsdp_shard_conditions:
logger.warning("No FSDP shard conditions provided; nothing will be sharded.")
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
logger.warning(
"The FSDP shard condition list is empty or None. No modules will be sharded in %s",
type(model).__name__)
return
fsdp_kwargs = {
@@ -218,38 +215,20 @@ def shard_model(
# iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
if use_size_filtering:
# Size-based filtering mode
min_params = int(os.environ.get("FASTVIDEO_FSDP2_MIN_PARAMS", "10000000"))
logger.info("Using size-based filtering with threshold: %.2fM", min_params / 1e6)
for n, m in reversed(list(model.named_modules())):
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
# Count all parameters
param_count = sum(p.numel() for p in m.parameters(recurse=True))
# Skip small modules
if param_count < min_params:
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)",
n, param_count / 1e6, min_params / 1e6)
continue
# Shard this module
logger.info("Sharding module %s (%.2fM params)", n, param_count / 1e6)
fully_shard(m, **fsdp_kwargs)
num_layers_sharded += 1
else:
# Shard all modules matching conditions
for n, m in reversed(list(model.named_modules())):
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
fully_shard(m, **fsdp_kwargs)
num_layers_sharded += 1
if num_layers_sharded == 0:
raise ValueError(
"No layer modules were sharded. Please check if shard conditions are working as expected."
)
# TODO(will): don't reshard after forward for the last layer to save on the
# all-gather that will immediately happen Shard the model with FSDP,
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
for shard_condition in fsdp_shard_conditions
]):
fully_shard(m, **fsdp_kwargs)
num_layers_sharded += 1
if num_layers_sharded == 0:
raise ValueError(
"No layer modules were sharded. Please check if shard conditions are working as expected."
)
# Finally shard the entire model to account for any stragglers
fully_shard(model, **fsdp_kwargs)
@@ -86,12 +86,6 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
return (prev_sample, )
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
@staticmethod
def calculate_alpha_beta_high(sigma, sigma_bound):
alpha = (1 - sigma) / (1 - sigma_bound)
beta = torch.sqrt(sigma ** 2 - (alpha * sigma_bound) ** 2)
return alpha, beta
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
@@ -111,32 +105,6 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def add_noise_high(self, original_samples, noise, timestep, boundary_timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B*T, C, H, W]
- noise: the noise with shape [B*T, C, H, W]
- timestep: the timestep with shape [B*T]
Output: the corrupted latent with shape [B*T, C, H, W]
"""
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
if boundary_timestep.ndim == 2:
boundary_timestep = boundary_timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_boundary = self.sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
alpha, beta = self.calculate_alpha_beta_high(sigma, sigma_boundary)
sample = alpha * original_samples + beta * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
-48
View File
@@ -180,51 +180,3 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def pred_noise_to_x_bound(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
boundary_timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
boundary_timestep: the boundary timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_t_boundary = sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - (sigma_t - sigma_t_boundary) * pred_noise
return pred_video.to(dtype)
@@ -50,8 +50,7 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@@ -59,8 +59,7 @@ class WanDMDPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None),
use_btchw_layout=True))
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
@@ -62,8 +62,7 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
use_btchw_layout=True))
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
+3 -28
View File
@@ -28,8 +28,7 @@ class LoRAPipeline(ComposedPipelineBase):
TODO: support training.
"""
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
dict
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
dict) # state dicts of loaded lora adapters
cur_adapter_name: str = ""
cur_adapter_path: str = ""
lora_layers: dict[str, BaseLayerWithLoRA] = {}
@@ -184,26 +183,11 @@ class LoRAPipeline(ComposedPipelineBase):
lora_param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"].lora_param_names_mapping)
# Extract alpha values and weights in a single pass
to_merge_params: defaultdict[Hashable,
dict[Any, Any]] = defaultdict(dict)
for name, weight in lora_state_dict.items():
# Extract weights (lora_A, lora_B, and lora_alpha)
name = name.replace("diffusion_model.", "")
name = name.replace(".weight", "")
if "lora_alpha" in name:
# Store alpha with minimal mapping - same processing as lora_A/lora_B
# but store in lora_adapters with ".lora_alpha" suffix
layer_name = name.replace(".lora_alpha", "")
layer_name, _, _ = lora_param_names_mapping_fn(layer_name)
target_name, _, _ = param_names_mapping_fn(layer_name)
# Store alpha alongside weights with same target_name base
alpha_key = target_name + ".lora_alpha"
self.lora_adapters[lora_nickname][alpha_key] = weight.item(
) if weight.numel() == 1 else float(weight.mean())
continue
name, _, _ = lora_param_names_mapping_fn(name)
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
name)
@@ -241,20 +225,11 @@ class LoRAPipeline(ComposedPipelineBase):
for name, layer in self.lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = self.lora_adapters[lora_nickname].get(
lora_alpha_name) if adapter_updated else None
layer.set_lora_weights(
lora_A,
lora_B,
lora_alpha=alpha,
self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name],
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
@@ -5,19 +5,13 @@ from typing import cast
import numpy as np
import torch
import torchvision
import torchvision.transforms.functional as TF
from PIL import Image
from einops import rearrange
from torchvision import transforms
from fastvideo.configs.configs import VideoLoaderType
from fastvideo.dataset.transform import (CenterCropResizeVideo,
TemporalRandomCrop, best_output_size)
TemporalRandomCrop)
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.logger import init_logger
logger = init_logger(__name__)
from fastvideo.pipelines.pipeline_batch_info import (ForwardBatch,
PreprocessBatch)
from fastvideo.pipelines.stages.base import PipelineStage
@@ -47,7 +41,6 @@ class VideoTransformStage(PipelineStage):
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.fps, list)
assert isinstance(batch.num_frames, list)
assert fastvideo_args.preprocess_config is not None
if batch.data_type != "video":
return batch
@@ -56,17 +49,8 @@ class VideoTransformStage(PipelineStage):
raise ValueError("Video loader is not set")
video_pixel_batch = []
pil_image_batch = []
enable_smart_resize = fastvideo_args.preprocess_config.enable_smart_resize
smart_resize_max_area = fastvideo_args.preprocess_config.smart_resize_max_area
if smart_resize_max_area is None:
smart_resize_max_area = 480 * 832
calculated_size = None
for i in range(len(batch.video_loader)):
# logger.info(f"Processing video {i+1}/{len(batch.video_loader)}")
frame_interval = batch.fps[i] / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch.num_frames[i],
@@ -79,28 +63,8 @@ class VideoTransformStage(PipelineStage):
else:
frame_indices = frame_indices[:self.num_frames]
logger.info(
f"Frame indices selected (count={len(frame_indices)}): [{frame_indices[0]}, ..., {frame_indices[-1]}]"
)
if fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
try:
video = batch.video_loader[i].get_frames_at(
frame_indices).data
except Exception as e:
# Try to get filename if available in PreprocessBatch
video_path = "unknown"
print(f"batch: {batch}")
if isinstance(batch, PreprocessBatch) and hasattr(
batch, 'video_file_name') and i < len(
batch.video_file_name):
video_path = batch.video_file_name[i]
logger.error(
f"Failed to load frames for video {video_path}: {e}")
logger.error(
f"Attempting to load frame indices: {frame_indices}")
raise e
video = batch.video_loader[i].get_frames_at(frame_indices).data
elif fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHVISION:
video, _, _ = torchvision.io.read_video(batch.video_loader[i],
output_format="TCHW")
@@ -109,75 +73,16 @@ class VideoTransformStage(PipelineStage):
raise ValueError(
f"Invalid video loader type: {fastvideo_args.preprocess_config.video_loader_type}"
)
logger.info(f"Video tensor shape after loading: {video.shape}")
if enable_smart_resize:
if calculated_size is None:
_, _, h_in, w_in = video.shape
# Get config values
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[
2] * vae_stride
ow, oh = best_output_size(w_in, h_in, dw, dh,
smart_resize_max_area)
calculated_size = (oh, ow)
logger.info(
f"Smart resize: input=({h_in}, {w_in}), output=({oh}, {ow})"
)
# Resize video frames using CenterCropResizeVideo (efficient)
processed_video = CenterCropResizeVideo(calculated_size)(video)
logger.info(
f"Processed video shape after resize: {processed_video.shape}"
)
video_pixel_batch.append(processed_video)
# Process pil_image (condition) with high quality Lanczos if I2V
if fastvideo_args.workload_type == WorkloadType.I2V:
# Extract first frame
img_tensor = video[0] # C, H, W
img = TF.to_pil_image(img_tensor)
iw, ih = img.width, img.height
ow, oh = calculated_size[1], calculated_size[0]
# Smart Resize logic for PIL image
scale = max(ow / iw, oh / ih)
resampling = Image.Resampling.LANCZOS if hasattr(
Image, 'Resampling') else Image.LANCZOS
img = img.resize((round(iw * scale), round(ih * scale)),
resampling)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
# to tensor [0, 255] uint8
img_t = torch.from_numpy(np.array(img)).permute(
2, 0, 1).unsqueeze(0)
pil_image_batch.append(img_t)
else:
video = self.video_transform(video)
video_pixel_batch.append(video)
video = self.video_transform(video)
video_pixel_batch.append(video)
video_pixel_values = torch.stack(video_pixel_batch)
logger.info(
f"Final stacked video batch shape: {video_pixel_values.shape}")
video_pixel_values = rearrange(video_pixel_values,
"b t c h w -> b c t h w")
video_pixel_values = video_pixel_values.to(torch.uint8)
if fastvideo_args.workload_type == WorkloadType.I2V:
if enable_smart_resize and len(pil_image_batch) > 0:
batch.pil_image = torch.cat(
pil_image_batch, dim=0).to(self.device if hasattr(
self, 'device') else video_pixel_values.device)
else:
batch.pil_image = video_pixel_values[:, :, 0, :, :]
batch.pil_image = video_pixel_values[:, :, 0, :, :]
video_pixel_values = video_pixel_values.float() / 255.0
batch.latents = video_pixel_values
+123 -172
View File
@@ -4,7 +4,7 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
@@ -34,16 +34,13 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self,
transformer,
scheduler,
transformer_2=None,
vae=None) -> None:
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
super().__init__(transformer, scheduler, transformer_2)
# KV and cross-attention cache state (initialized on first forward)
self.transformer = transformer
self.transformer_2 = transformer_2
self.vae = vae
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
@@ -83,13 +80,6 @@ class CausalDMDDenosingStage(DenoisingStage):
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -113,110 +103,113 @@ class CausalDMDDenosingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset 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:
# Initialize the low noise kv cache
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
if self.kv_cache1 is None:
self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device)
else:
assert self.crossattn_cache is not None
# reset cross-attention cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index][
"is_init"] = False # type: ignore
# reset kv cache pointers
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index][
"global_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
self.kv_cache1[block_index][
"local_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
def _get_kv_cache(timestep: float) -> list[dict]:
if boundary_timestep is not None:
if timestep >= 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=fastvideo_args.pipeline_config.text_encoder_configs[0].
arch_config.text_len,
dtype=target_dtype,
device=latents.device)
pos_start_base = 0
# Determine block sizes
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
# For now hardcode the first block to be 1 frame assuming the model is Wan2.2-MoE
if boundary_timestep is not None:
block_sizes[0] = 1
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert self.vae is not None, "VAE is not provided for causal video gen task"
self.vae = self.vae.to(get_local_torch_device())
first_frame_latent = self.vae.encode(batch.pil_image).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent = first_frame_latent * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
t_zero = torch.zeros([latents.shape[0], 1],
# Optional: cache context features from provided image latents prior to generation
current_start_frame = 0
if getattr(batch, "image_latent", None) is not None:
image_latent = batch.image_latent
assert image_latent is not None
input_frames = image_latent.shape[2]
# timestep zero (or configured context noise) for cache warm-up
t_zero = torch.zeros([latents.shape[0]],
device=latents.device,
dtype=torch.long)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.transformer(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
if boundary_timestep is not None:
self.transformer_2(
first_frame_latent.to(target_dtype),
if independent_first_frame and input_frames >= 1:
# warm-up with the very first frame independently
image_first_btchw = image_latent[:, :, :1, :, :].to(
target_dtype).permute(0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
image_first_btchw,
prompt_embeds,
t_zero,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += 1
remaining_frames = input_frames - 1
else:
remaining_frames = input_frames
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# process remaining input frames in blocks of num_frame_per_block
while remaining_frames > 0:
block = min(self.num_frames_per_block, remaining_frames)
ref_btchw = image_latent[:, :, current_start_frame:
current_start_frame +
block, :, :].to(target_dtype).permute(
0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
ref_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += block
remaining_frames -= block
# Base position offset from any cache warm-up
pos_start_base = current_start_frame
# Determine block sizes
if not independent_first_frame or (independent_first_frame
and batch.image_latent is not None):
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
else:
if (t - 1) % self.num_frames_per_block != 0:
raise ValueError(
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
)
num_blocks = (t - 1) // self.num_frames_per_block
block_sizes = [1] + [self.num_frames_per_block] * num_blocks
start_index = 0
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
@@ -229,7 +222,7 @@ class CausalDMDDenosingStage(DenoisingStage):
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
if boundary_timestep is not None and t_cur < boundary_timestep:
if self.transformer_2 is not None and fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None and t_cur < fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps:
current_model = self.transformer_2
else:
current_model = self.transformer
@@ -287,8 +280,8 @@ class CausalDMDDenosingStage(DenoisingStage):
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -297,22 +290,12 @@ class CausalDMDDenosingStage(DenoisingStage):
).permute(0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
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])
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(
@@ -326,23 +309,11 @@ class CausalDMDDenosingStage(DenoisingStage):
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
if boundary_timestep 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 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])
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:
@@ -370,44 +341,24 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
if boundary_timestep is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_expanded_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(
_ = current_model(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frames_per_block - 1
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
def _initialize_kv_cache(self, batch_size, dtype, device) -> None:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
@@ -441,10 +392,10 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0], dtype=torch.long, device=device),
})
return kv_cache1
self.kv_cache1 = kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> list[dict]:
device) -> None:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
@@ -470,7 +421,7 @@ class CausalDMDDenosingStage(DenoisingStage):
"is_init":
False,
})
return crossattn_cache
self.crossattn_cache = crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
@@ -494,4 +445,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result
+1
View File
@@ -1085,6 +1085,7 @@ class DmdDenoisingStage(DenoisingStage):
# Get latents and embeddings
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
latents = latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = latents.shape
prompt_embeds = batch.prompt_embeds
@@ -106,28 +106,26 @@ class InputValidationStage(PipelineStage):
batch.pil_image = image
# further processing for ti2v task
if (fastvideo_args.pipeline_config.ti2v_task
or fastvideo_args.pipeline_config.is_causal
) and batch.pil_image is not None:
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
img = batch.pil_image
ih, iw = img.height, img.width
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
max_area = 704 * 1280
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
logger.info("resized img height: %s, img width: %s", img.height,
img.width)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
assert img.width == ow and img.height == oh
logger.info("final processed img height: %s, img width: %s",
img.height, img.width)
# to tensor
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
@@ -28,14 +28,10 @@ class LatentPreparationStage(PipelineStage):
denoised during the diffusion process.
"""
def __init__(self,
scheduler,
transformer,
use_btchw_layout: bool = False) -> None:
def __init__(self, scheduler, transformer) -> None:
super().__init__()
self.scheduler = scheduler
self.transformer = transformer
self.use_btchw_layout = use_btchw_layout
def forward(
self,
@@ -82,26 +78,15 @@ class LatentPreparationStage(PipelineStage):
raise ValueError("Height and width must be provided")
# Calculate latent shape
if self.use_btchw_layout:
shape = (
batch_size,
num_frames,
self.transformer.num_channels_latents,
height // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
else:
shape = (
batch_size,
self.transformer.num_channels_latents,
num_frames,
height // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
shape = (
batch_size,
self.transformer.num_channels_latents,
num_frames,
height // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
# Validate generator if it's a list
if isinstance(generator, list) and len(generator) != batch_size:
+32 -4
View File
@@ -22,8 +22,8 @@ os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
local_dir=os.path.join(
"data", BASE_MODEL_PATH))
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder_2")
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer_2")
@@ -130,6 +130,17 @@ def test_clip_encoder():
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info("Maximum difference in last hidden states: %f",
max_diff_hidden.item())
logger.info("Mean difference in last hidden states: %f",
mean_diff_hidden.item())
# Compare pooler outputs
pooler_output1 = outputs1.pooler_output
pooler_output2 = outputs2.pooler_output
@@ -137,5 +148,22 @@ def test_clip_encoder():
assert pooler_output1.shape == pooler_output2.shape, \
f"Pooler output shapes don't match: {pooler_output1.shape} vs {pooler_output2.shape}"
assert_close(pooler_output1, pooler_output2, atol=1e-2, rtol=1e-3)
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-2, rtol=1e-3)
max_diff_pooler = torch.max(
torch.abs(pooler_output1 - pooler_output2))
mean_diff_pooler = torch.mean(
torch.abs(pooler_output1 - pooler_output2))
logger.info("Maximum difference in pooler outputs: %f",
max_diff_pooler.item())
logger.info("Mean difference in pooler outputs: %f",
mean_diff_pooler.item())
# Check if outputs are similar (allowing for small numerical differences)
assert mean_diff_hidden < 1e-2, \
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
assert mean_diff_pooler < 1e-2, \
f"Pooler outputs differ significantly: mean diff = {mean_diff_pooler.item()}"
assert max_diff_hidden < 1e-1, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
assert max_diff_pooler < 2e-2, \
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
+20 -16
View File
@@ -22,8 +22,8 @@ os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder")
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer")
@@ -68,7 +68,8 @@ def test_llama_encoder():
logger.info("Model1 has %d parameters", len(params1))
logger.info("Model2 has %d parameters", len(params2))
# Compare a few key parameters
weight_diffs = []
# check if embed_tokens are the same
device = model1.embed_tokens.weight.device
assert torch.allclose(model1.embed_tokens.weight,
@@ -77,18 +78,6 @@ def test_llama_encoder():
"layers.{}.input_layernorm.weight",
"layers.{}.post_attention_layernorm.weight"
]
for layer_idx in range(hf_config.num_hidden_layers):
for w in weights:
name1 = w.format(layer_idx)
name2 = w.format(layer_idx)
p1 = params1[name1]
p2 = params2[name2]
if "gate_up" in name2:
# print("skipping gate_up")
continue
p1 = p1.to_local().to(device) if isinstance(p1, DTensor) else p1.to(device)
p2 = p2.to_local().to(device) if isinstance(p2, DTensor) else p2.to(device)
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
for name1, param1 in sorted(params1.items()):
name2 = name1
@@ -150,4 +139,19 @@ def test_llama_encoder():
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-1, rtol=1e-4)
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info("Maximum difference in last hidden states: %f",
max_diff_hidden.item())
logger.info("Mean difference in last hidden states: %f",
mean_diff_hidden.item())
# Check if outputs are similar (allowing for small numerical differences)
assert mean_diff_hidden < 1e-2, \
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
assert max_diff_hidden < 1e-1, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
+33 -2
View File
@@ -133,7 +133,24 @@ def test_t5_encoder(t5_model_paths):
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-4, rtol=1e-4)
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info("Maximum difference in last hidden states: %s",
max_diff_hidden.item())
logger.info("Mean difference in last hidden states: %s",
mean_diff_hidden.item())
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
# Check if outputs are similar (allowing for small numerical differences)
assert mean_diff_hidden < 1e-4, \
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
assert max_diff_hidden < 1e-4, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
@pytest.mark.usefixtures("distributed_setup")
@@ -235,4 +252,18 @@ def test_t5_large_encoder(t5_large_model_paths):
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-4, rtol=1e-4)
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info("Maximum difference in last hidden states: %s",
max_diff_hidden.item())
logger.info("Mean difference in last hidden states: %s",
mean_diff_hidden.item())
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
# Check if outputs are similar (allowing for small numerical differences)
assert mean_diff_hidden < 1e-4, \
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
assert max_diff_hidden < 1e-4, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import re
import pytest
@@ -53,38 +52,12 @@ LORA_CONFIGS = [
"negative_prompt": "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
"ssim_threshold": 0.79
}
# TODO: Add a LoRA with lora_alpha values to test alpha scaling
#
# Context: This change is mainly for an in-progress ticket porting over LongCat-Video,
# where they used an alpha value that is two times smaller than their rank. This fix
# ensures that LoRA weights are correctly scaled by the alpha/rank ratio when merged.
#
# Issue: Currently, we cannot add a test for LoRA adapters with alpha values because:
# - The existing public LoRAs for Wan-AI/Wan2.1-T2V-1.3B-Diffusers don't store lora_alpha
# - No publicly available LoRA for this model includes lora_alpha tensors in their weights
# - This is why the alpha/rank scaling bug wasn't caught by existing tests
#
# The fix has been validated with:
# - LongCat-Video distilled LoRA (which includes alpha values)
# - Manual testing shows correct alpha/rank scaling behavior
# - Backward compatibility confirmed with LoRAs without alpha values
#
# Future work:
# - Add a synthetic LoRA test fixture with alpha values when feasible
# - Or wait for public Wan LoRAs with alpha to become available
]
MODEL_TO_PARAMS = {
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WAN_LORA_PARAMS,
}
def _sanitize_filename_component(name: str) -> str:
"""Sanitize filename to remove invalid characters (same logic as VideoGenerator)"""
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
sanitized = sanitized.strip().strip('.')
sanitized = re.sub(r'\s+', ' ', sanitized)
return sanitized or "video"
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_merge_lora_weights(model_id):
lora_config = LORA_CONFIGS[0] # test only one
@@ -164,16 +137,14 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
generation_kwargs["negative_prompt"] = lora_config["negative_prompt"]
generator.set_lora_adapter(lora_nickname=lora_nickname, lora_path=lora_path)
# Sanitize the filename before adding .mp4 extension to match VideoGenerator's behavior
output_video_name = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
output_video_name = _sanitize_filename_component(output_video_name)
generated_video_path = os.path.join(output_dir, f"{output_video_name}.mp4")
generation_kwargs["output_path"] = generated_video_path
generation_kwargs["output_path"] = output_dir
generation_kwargs["output_video_name"] = output_video_name
generator.generate_video(prompt, **generation_kwargs)
assert os.path.exists(
generated_video_path), f"Output video was not generated at {generated_video_path}"
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, 'L40S_reference_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
@@ -182,25 +153,13 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
# Find the matching reference video - try exact match first, then fuzzy match
# The reference might have different sanitization (e.g., trailing spaces)
# Find the matching reference video for the switched LoRA
reference_video_name = None
unsanitized_prefix = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
for filename in os.listdir(reference_folder):
if not filename.endswith('.mp4'):
continue
# Try exact match with sanitized name
if filename.startswith(output_video_name):
reference_video_name = filename
break
# Try match with unsanitized prefix (for legacy reference videos)
# Remove .mp4 and compare the base names after sanitization
base_filename = filename[:-4] # Remove .mp4
if _sanitize_filename_component(base_filename) == output_video_name:
reference_video_name = filename
# Check if the filename starts with the expected output_video_name and ends with .mp4
if filename.startswith(output_video_name) and filename.endswith('.mp4'):
reference_video_name = filename # Remove .mp4 extension to match the logic below
break
if not reference_video_name:
@@ -208,6 +167,7 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
raise FileNotFoundError(f"Reference video missing for adapter {lora_path}")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name + ".mp4")
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
+2 -2
View File
@@ -74,9 +74,9 @@ def run_vae_tests():
def run_transformer_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(gpu="L40S:2", image=image, timeout=2700)
def run_ssim_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
run_test("pytest ./fastvideo/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
@@ -0,0 +1,19 @@
#!/bin/bash
num_gpus=8
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
--master_port 29503 \
tp_example.py
num_gpus=2
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
--master_port 29503 \
fastvideo/tests/test_hunyuanvideo_load.py --sequence_model_parallel_size $num_gpus
torchrun --nnodes=1 --nproc_per_node=1 --master_port 29503 fastvideo/tests/test_llama_encoder.py
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
torchrun --nnodes=1 --nproc_per_node=1 --master_port 29503 fastvideo/tests/test_clip_encoder.py
@@ -0,0 +1,164 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import numpy as np
import torch
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
from fastvideo.distributed import (maybe_init_distributed_environment_and_model_parallel)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def setup_args():
parser = argparse.ArgumentParser(description='T5 Encoder Test')
parser.add_argument('--model_path', type=str, default="google/umt5-xxl")
parser.add_argument(
'--dit-precision',
type=str,
default="float32",
help='Precision to use for the model (float32, float16, bfloat16)')
return parser.parse_args()
def test_t5_encoder():
maybe_init_distributed_environment_and_model_parallel(1, 1)
# Set fixed random seed for reproducibility
torch.manual_seed(42)
np.random.seed(42)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
# Initialize the two model implementations
model_path = "/workspace/data/Wan2.1-T2V-1.3B-Diffusers/text_encoder"
tokenizer_path = "/workspace/data/Wan2.1-T2V-1.3B-Diffusers/tokenizer"
hf_config = AutoConfig.from_pretrained(model_path)
print(hf_config)
precision = torch.float16 # It must be float16 because the weight loader is in float16
# Load our implementation using the loader from text_encoder/__init__.py
model1 = UMT5EncoderModel.from_pretrained(model_path).to(precision).to(
device).eval()
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
from fastvideo.models.loader.component_loader import TextEncoderLoader
loader = TextEncoderLoader()
model2 = loader.load_model(model_path, hf_config, device)
# Convert to float16 and move to device
model2 = model2.to(precision)
model2 = model2.to(device)
model2.eval()
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Check number of parameters
logger.info(f"Model1 has {len(params1)} parameters")
logger.info(f"Model2 has {len(params2)} parameters")
weight_diffs = []
# check if embed_tokens are the same
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.weight"]
# for (name1, param1), (name2, param2) in zip(
# sorted(params1.items()), sorted(params2.items())
# ):
for l in range(hf_config.num_hidden_layers):
for w in weights:
name1 = w.format(l)
name2 = w.format(l)
p1 = params1[name1]
p2 = params2[name2]
assert p1.dtype == p2.dtype
try:
logger.info(f"Parameter: {name1} vs {name2}")
max_diff = torch.max(torch.abs(p1 - p2)).item()
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
weight_diffs.append((name1, name2, max_diff, mean_diff))
logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
except Exception as e:
logger.info(f"Error comparing {name1} and {name2}: {e}")
total_params = sum(p.numel() for p in model1.parameters())
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
weight_mean_model1 = weight_sum_model1 / total_params
print("Model 1 Weight Sum: ", weight_sum_model1)
print("Model 1 Weight Mean: ", weight_mean_model1)
total_params = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params
print("Model 2 Weight Sum: ", weight_sum_model2)
print("Model 2 Weight Mean: ", weight_mean_model2)
# Test with some sample prompts
prompts = [
"Once upon a time", "The quick brown fox jumps over",
"In a galaxy far, far away"
]
logger.info("Testing T5 encoder with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info(f"Testing prompt: '{prompt}'")
# Tokenize the prompt
tokens = tokenizer(prompt,
padding="max_length",
max_length=512,
truncation=True,
return_tensors="pt").to(device)
# Get outputs from our implementation
# filter out padding input_ids
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
outputs1 = model1(input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True).last_hidden_state
print("--------------------------------")
logger.info("Testing model2")
# Get outputs from HuggingFace implementation
outputs2 = model2(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
)
# Compare last hidden states
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info(
f"Maximum difference in last hidden states: {max_diff_hidden.item()}"
)
logger.info(
f"Mean difference in last hidden states: {mean_diff_hidden.item()}"
)
logger.info(
"Test passed! Both T5 encoder implementations produce similar outputs.")
logger.info("Test completed successfully")
if __name__ == "__main__":
test_t5_encoder()
+123
View File
@@ -0,0 +1,123 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import numpy as np
import torch
from diffusers import AutoencoderKLWan
from safetensors.torch import load_file
from fastvideo.logger import init_logger
from fastvideo.models.vaes.wanvae import AutoencoderKLWan as MyWanVAE
logger = init_logger(__name__)
def test_wan_vae():
# Set fixed random seed for reproducibility
torch.manual_seed(42)
np.random.seed(42)
device = torch.device("cuda:0")
# Initialize the two model implementations
path = "/workspace/data/Wan2.1-T2V-1.3B-Diffusers/vae"
config_path = os.path.join(path, "config.json")
config = json.load(open(config_path))
config.pop("_class_name")
config.pop("_diffusers_version")
model1 = MyWanVAE(**config).to(torch.bfloat16)
model2 = AutoencoderKLWan(**config).to(torch.bfloat16)
loaded = load_file(os.path.join(path,
"diffusion_pytorch_model.safetensors"))
model1.load_state_dict(loaded)
model2.load_state_dict(loaded)
# Set both models to eval mode
model1.eval()
model2.eval()
# Move to GPU
model1 = model1.to(device)
model2 = model2.to(device)
# model1.enable_tiling(
# tile_sample_min_height=32,
# tile_sample_min_width=32,
# tile_sample_min_num_frames=8,
# tile_sample_stride_height=16,
# tile_sample_stride_width=16,
# tile_sample_stride_num_frames=4
# )
# Create identical inputs for both models
batch_size = 1
# Video input [B, C, T, H, W]
input_tensor = torch.randn(batch_size,
3,
81,
32,
32,
device=device,
dtype=torch.bfloat16)
latent_tensor = torch.randn(batch_size,
16,
21,
32,
32,
device=device,
dtype=torch.bfloat16)
# Disable gradients for inference
with torch.no_grad():
# Test encoding
logger.info("Testing encoding...")
latent2 = model2.encode(input_tensor).latent_dist.mean
print("--------------------------------")
latent1 = model1.encode(input_tensor).mean
# Check if latents have the same shape
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
# Check if latents are similar
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
logger.info(
f"Maximum difference between encoded latents: {max_diff_encode.item()}"
)
logger.info(
f"Mean difference between encoded latents: {mean_diff_encode.item()}"
)
assert mean_diff_encode < 5e-1, f"Encoded latents differ significantly: mean diff = {mean_diff_encode.item()}"
# Test decoding
logger.info("Testing decoding...")
latent1 = latent2 = latent_tensor
latents_mean = (torch.tensor(model2.config.latents_mean).view(
1, model2.config.z_dim, 1, 1, 1).to(latent2.device, latent2.dtype))
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
1, model2.config.z_dim, 1, 1, 1).to(latent2.device, latent2.dtype)
latent2 = latent2 / latents_std + latents_mean
output1 = model1.decode(latent1)
output2 = model2.decode(latent2).sample
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
# Check if outputs are similar
max_diff_decode = torch.max(torch.abs(output1 - output2))
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
logger.info(
f"Maximum difference between decoded outputs: {max_diff_decode.item()}"
)
logger.info(
f"Mean difference between decoded outputs: {mean_diff_decode.item()}"
)
assert mean_diff_decode < 1e-1, f"Decoded outputs differ significantly: mean diff = {mean_diff_decode.item()}"
logger.info(
"Test passed! Both VAE implementations produce similar outputs.")
logger.info("Test completed successfully")
if __name__ == "__main__":
test_wan_vae()
+152
View File
@@ -0,0 +1,152 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import torch
import torch.nn as nn
from fastvideo.distributed.parallel_state import (
cleanup_dist_env_and_memory, destroy_distributed_environment,
destroy_model_parallel, get_tp_rank,
get_tp_world_size, maybe_init_distributed_environment_and_model_parallel, get_world_group)
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class SimpleTPModel(nn.Module):
"""A simple model that uses tensor parallelism."""
def __init__(self, hidden_size=1024, intermediate_size=4096):
super().__init__()
# Column parallel linear layer (splits output dimension)
self.fc1 = ColumnParallelLinear(
input_size=hidden_size,
output_size=intermediate_size,
bias=True,
gather_output=
False, # Don't gather output since we're passing to row parallel
skip_bias_add=False)
# Row parallel linear layer (splits input dimension)
self.fc2 = RowParallelLinear(
input_size=intermediate_size,
output_size=hidden_size,
bias=True,
input_is_parallel=True, # Input is already split from previous layer
skip_bias_add=False)
self.activation = nn.GELU()
def forward(self, x):
# Forward through column parallel layer
hidden_states, _ = self.fc1(x)
# Apply activation
hidden_states = self.activation(hidden_states)
# Forward through row parallel layer
output, _ = self.fc2(hidden_states)
return output
def initialize_random_weights(model, seed=42):
"""Initialize the model with random weights using a fixed seed for reproducibility."""
# Set seed for reproducibility
torch.manual_seed(seed)
# Initialize weights for each layer
with torch.no_grad():
# For ColumnParallelLinear layers
if hasattr(model, 'fc1'):
nn.init.normal_(model.fc1.weight, mean=0.0, std=0.02)
if model.fc1.bias is not None:
nn.init.zeros_(model.fc1.bias)
# For RowParallelLinear layers
if hasattr(model, 'fc2'):
nn.init.normal_(model.fc2.weight, mean=0.0, std=0.02)
if model.fc2.bias is not None:
nn.init.zeros_(model.fc2.bias)
logger.info("Model initialized with random weights")
return model
def setup_args():
parser = argparse.ArgumentParser(
description='Simple Tensor Parallelism Example')
parser.add_argument('--tensor-model-parallel-size',
type=int,
default=8,
help='Degree of tensor model parallelism')
parser.add_argument('--batch-size',
type=int,
default=8,
help='Batch size for the example')
parser.add_argument('--hidden-size',
type=int,
default=1024,
help='Hidden size for the model')
parser.add_argument('--intermediate-size',
type=int,
default=4096,
help='Intermediate size for the model')
return parser.parse_args()
def main():
args = setup_args()
maybe_init_distributed_environment_and_model_parallel(args.tensor_model_parallel_size, args.tensor_model_parallel_size)
rank = get_world_group().rank
local_rank = get_world_group().local_rank
# Get tensor parallel info
tp_rank = get_tp_rank()
tp_world_size = get_tp_world_size()
logger.info(
f"Process rank {rank} initialized with TP rank {tp_rank} in TP world size {tp_world_size}"
)
# Create a simple model
model = SimpleTPModel(hidden_size=args.hidden_size,
intermediate_size=args.intermediate_size)
# Initialize with random weights
model = initialize_random_weights(model)
# Create a random input tensor
batch_size = args.batch_size
hidden_size = args.hidden_size
x = torch.randn(batch_size, hidden_size, dtype=torch.float)
# Move to GPU if available
device = torch.device(
f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
model = model.to(device)
x = x.to(device)
# Forward pass
logger.info(f"Running forward pass on TP rank {tp_rank}")
with torch.no_grad():
output = model(x)
# Print output shape and statistics
logger.info(f"Output shape: {output.shape}")
logger.info(
f"Output mean: {output.mean().item()}, std: {output.std().item()}")
# Clean up
logger.info("Cleaning up distributed environment")
destroy_model_parallel()
destroy_distributed_environment()
cleanup_dist_env_and_memory()
logger.info("Example completed successfully")
if __name__ == "__main__":
main()
@@ -1,10 +1 @@
{
"step_time": 0.6983645600266755,
"grad_norm": 0.245593056678772,
"avg_step_time": 1.002151239803061,
"_timestamp": 1751181952.70901,
"vsa_sparsity": 0.05,
"learning_rate": 1e-05,
"train_loss": 0.2530866410434246,
"_runtime": 107.325113071
}
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":1.260593056678772,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.2620866410434246,"_runtime":107.325113071}
@@ -1,211 +0,0 @@
import os
import sys
from pathlib import Path
# Set Python path to current folder
current_dir = str(Path(__file__).parent.parent.parent.parent.parent)
if current_dir not in sys.path:
sys.path.insert(0, current_dir)
os.environ["PYTHONPATH"] = current_dir + ":" + os.environ.get("PYTHONPATH", "")
import subprocess
import torch
import json
from huggingface_hub import snapshot_download
from fastvideo.utils import logger
# Import the training pipeline
from fastvideo.training.wan_training_pipeline import main
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
from fastvideo.training.wan_training_pipeline import WanTrainingPipeline
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_PATH = "data/crush-smol_processed_t2v/training_dataset/worker_1/worker_0/"
VALIDATION_DATASET_FILE = "examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json"
OUTPUT_DIR = Path("checkpoints/wan_t2v_finetune")
PROFILER_TRACE_ROOT = Path("/mnt/fast-disks/hao_lab/ohm/profiler_traces/wan_t2v_finetune")
WANDB_SUMMARY_FILE = OUTPUT_DIR / "tracker/wandb/latest-run/files/wandb-summary.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
GRAD_ACCUM = "1"
MASTER_PORT = "29504"
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = MASTER_PORT
def run_worker():
"""Worker function that will be run on each GPU"""
# Create and populate args
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
# Set the arguments as they are in finetune_t2v.sh
args = parser.parse_args([
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", DATA_PATH,
"--dataloader_num_workers", "1",
"--train_batch_size", "4",
"--train_sp_batch_size", "1",
"--gradient_accumulation_steps", GRAD_ACCUM,
"--num_latent_t", "20",
"--num_height", "720",
"--num_width", "1280",
"--num_frames", "77",
"--enable_gradient_checkpointing_type", "full",
"--max_train_steps", "20",
"--learning_rate", "5e-5",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "250",
"--training_state_checkpointing_steps", "250",
"--weight_decay", "1e-4",
"--max_grad_norm", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--not_apply_cfg_solver",
"--training_cfg_rate", "0.1",
"--ema_start_step", "0",
"--dit_precision", "fp32",
"--output_dir", str(OUTPUT_DIR),
"--tracker_project_name", "wan_t2v_finetune",
"--checkpoints_total_limit", "3",
"--validation_dataset_file", VALIDATION_DATASET_FILE,
"--validation_steps", "200",
"--validation_sampling_steps", "50",
"--validation_guidance_scale", "6.0",
#"--enable_torch_compile",
#"--log_validation",
"--num_gpus", NUM_GPUS_PER_NODE,
"--sp_size", NUM_GPUS_PER_NODE,
"--tp_size", "1",
"--hsdp_replicate_dim", NUM_GPUS_PER_NODE,
"--hsdp_shard_dim", "1"
])
# Call the main training function
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "online"
data_dir = Path("data/crush-smol_processed_t2v")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="wlsaidhi/crush-smol_processed_t2v",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
)
# Get the current file path
current_file = Path(__file__).resolve()
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
"--master_port", MASTER_PORT,
str(current_file)
]
process = subprocess.run(cmd, capture_output=True, text=True)
# Print stdout and stderr for debugging
if process.stdout:
print("STDOUT:", process.stdout)
if process.stderr:
print("STDERR:", process.stderr)
# Check if the process failed
if process.returncode != 0:
print(f"Process failed with return code: {process.returncode}")
raise subprocess.CalledProcessError(process.returncode, cmd, process.stdout, process.stderr)
summary_file = WANDB_SUMMARY_FILE
with summary_file.open() as f:
wandb_summary = json.load(f)
# Calculate and print MFU metrics
device_name = torch.cuda.get_device_name()
try:
# Get actual values from training run (logged from training_batch.raw_latent_shape)
batch_size = wandb_summary.get("batch_size")
seq_len = wandb_summary.get("dit_seq_len")
context_len = wandb_summary.get("context_len")
avg_step_time = wandb_summary.get("avg_step_time")
hidden_dim = wandb_summary.get("hidden_dim")
num_layers = wandb_summary.get("num_layers")
ffn_dim = wandb_summary.get("ffn_dim")
# FLOPs per layer (forward pass)
# - QKV + out proj: 8 * hidden_dim^2 * seq_len
# - Cross-attn proj: 4 * hidden_dim^2 * seq_len + 4 * hidden_dim^2 * context_len
# - MLP: 4 * hidden_dim * ffn_dim * seq_len
# - Self-attn matmuls: 4 * seq_len^2 * hidden_dim
# - Cross-attn matmuls: 4 * seq_len * context_len * hidden_dim
qkv_out_flops = 8 * hidden_dim * hidden_dim * seq_len
cross_attn_proj_flops = (
(4 * hidden_dim * hidden_dim * seq_len) +
(4 * hidden_dim * hidden_dim * context_len)
)
mlp_flops = 4 * hidden_dim * ffn_dim * seq_len
self_attn_flops = 4 * seq_len * seq_len * hidden_dim
cross_attn_flops = 4 * seq_len * context_len * hidden_dim
flops_per_layer = (
qkv_out_flops + cross_attn_proj_flops + mlp_flops + self_attn_flops + cross_attn_flops
)
# With full activation checkpointing: 1 forward + 3 backward (1 recompute + 2 gradient)
achieved_flops = batch_size * flops_per_layer * num_layers * 4
# Account for gradient accumulation (from config)
grad_accum = int(GRAD_ACCUM)
achieved_flops *= grad_accum
# Peak FLOPs based on device
if "H100" in device_name:
peak_flops_per_gpu = 989e12
elif "A100" in device_name:
peak_flops_per_gpu = 312e12
elif "A40" in device_name:
peak_flops_per_gpu = 312e12
elif "L40S" in device_name:
peak_flops_per_gpu = 362e12
else:
raise ValueError(f"Device {device_name} not supported")
# Total peak (2 GPUs)
world_size = int(NUM_GPUS_PER_NODE)
total_peak_flops = peak_flops_per_gpu * world_size
# Calculate MFU
achieved_flops_per_sec = achieved_flops / avg_step_time if avg_step_time > 0 else 0
mfu = (achieved_flops_per_sec / total_peak_flops * 100) if total_peak_flops > 0 else 0
print(f"Per-Step MFU: {mfu:.4f}%")
except Exception as e:
print(f"Could not calculate MFU: {e}")
if __name__ == "__main__":
if os.environ.get("LOCAL_RANK") is not None:
# We're being run by torchrun
run_worker()
else:
# We're being run directly
test_distributed_training()
@@ -1,598 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Test COSMOS 2.5 DiT implementation against reference.
Compares FastVideo's Cosmos25Transformer3DModel with the official MinimalV1LVGDiT from cosmos-predict2.5.
"""
import os
import sys
import pytest
import torch
# Add cosmos-predict2.5 to Python path for loading reference model
TEST_DIR = os.path.dirname(os.path.abspath(__file__))
COSMOS_PREDICT2_5_PATH = os.path.join(TEST_DIR, '..', '..', '..', '..', 'cosmos-predict2.5')
COSMOS_PREDICT2_5_PATH = os.path.normpath(COSMOS_PREDICT2_5_PATH)
if os.path.exists(COSMOS_PREDICT2_5_PATH) and COSMOS_PREDICT2_5_PATH not in sys.path:
sys.path.insert(0, COSMOS_PREDICT2_5_PATH)
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.utils import maybe_download_model
# Use Cosmos 2.5 specific config
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
# Log the cosmos-predict2.5 path after logger is initialized
if os.path.exists(COSMOS_PREDICT2_5_PATH):
logger.info(f"cosmos-predict2.5 found at: {COSMOS_PREDICT2_5_PATH}")
else:
logger.warning(f"cosmos-predict2.5 not found at: {COSMOS_PREDICT2_5_PATH}")
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29505"
# COSMOS 2.5 model path - update this based on the actual HuggingFace model ID
# The model has subdirectories: base/pre-trained, base/post-trained, auto/multiview, robot/action-cond
BASE_MODEL_PATH = "nvidia/Cosmos-Predict2.5-2B"
CHECKPOINT_SUBDIR = "base/post-trained"
CHECKPOINT_FILENAME = "81edfebe-bd6a-4039-8c1d-737df1a790bf_ema_bf16.pt"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH, local_dir=None)
TRANSFORMER_PATH = os.path.join(MODEL_PATH, CHECKPOINT_SUBDIR, "transformer")
if not os.path.exists(TRANSFORMER_PATH):
# Try without subdirectory
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
def load_reference_cosmos25_model(checkpoint_path: str, device, dtype):
"""
Load the reference COSMOS 2.5 model from cosmos-predict2.5 repo.
This assumes the cosmos-predict2.5 repo is available in the Python path.
"""
try:
# Try to import from cosmos-predict2.5 repo
from cosmos_predict2._src.predict2.networks.minimal_v1_lvg_dit import MinimalV1LVGDiT
# COSMOS 2.5 2B model configuration
model_config = {
'max_img_h': 240,
'max_img_w': 240,
'max_frames': 128,
'in_channels': 16,
'out_channels': 16,
'patch_spatial': 2,
'patch_temporal': 1,
'model_channels': 2048, # 2B model
'num_blocks': 28,
'num_heads': 16,
'mlp_ratio': 4.0,
'crossattn_emb_channels': 1024,
'pos_emb_cls': 'rope3d',
'pos_emb_learnable': True,
'pos_emb_interpolation': 'crop',
'use_adaln_lora': True,
'adaln_lora_dim': 256,
'rope_h_extrapolation_ratio': 3.0,
'rope_w_extrapolation_ratio': 3.0,
'rope_t_extrapolation_ratio': 1.0,
'extra_per_block_abs_pos_emb': False,
'rope_enable_fps_modulation': False,
'use_crossattn_projection': True,
'crossattn_proj_in_channels': 100352,
'concat_padding_mask': True,
'atten_backend': 'torch',
}
model = MinimalV1LVGDiT(**model_config)
# Load checkpoint if path exists
if os.path.exists(checkpoint_path):
logger.info(f"Loading reference model from {checkpoint_path}")
checkpoint = torch.load(checkpoint_path, map_location='cpu')
# Extract state dict
if 'state_dict' in checkpoint:
checkpoint_state = checkpoint['state_dict']
elif 'model' in checkpoint:
checkpoint_state = checkpoint['model']
else:
checkpoint_state = checkpoint
# Filter to only model parameters (remove training metadata)
model_state = {k: v for k, v in checkpoint_state.items()
if k.startswith('net.') and 'accum_' not in k}
# Transform checkpoint keys to match model's expected format
# 1. Strip 'net.' prefix (e.g., 'net.blocks.0.self_attn.*' -> 'blocks.0.self_attn.*')
# 2. Add '_checkpoint_wrapped_module' after 'blocks.N.' if model expects it
transformed_state = {}
# First, check what the model expects
model_state_dict = model.state_dict()
needs_checkpoint_wrapper = any('_checkpoint_wrapped_module' in k for k in model_state_dict.keys())
for key, value in model_state.items():
# Strip 'net.' prefix
if key.startswith('net.'):
new_key = key[4:] # Remove 'net.' prefix
else:
new_key = key
# Add '_checkpoint_wrapped_module' if needed
if needs_checkpoint_wrapper and new_key.startswith('blocks.'):
# Pattern: 'blocks.N.something' -> 'blocks.N._checkpoint_wrapped_module.something'
parts = new_key.split('.', 2)
if len(parts) >= 3 and parts[0] == 'blocks' and parts[1].isdigit():
new_key = f"{parts[0]}.{parts[1]}._checkpoint_wrapped_module.{parts[2]}"
transformed_state[new_key] = value
# Load with strict=False to handle any remaining mismatches
missing_keys, unexpected_keys = model.load_state_dict(transformed_state, strict=False)
if missing_keys:
logger.warning(f"Missing keys when loading reference model: {len(missing_keys)} keys")
# Show all missing keys for debugging
logger.warning("All missing keys:")
for k in missing_keys:
logger.warning(f" - {k}")
# Filter out _extra_state and pos_embedder keys as they're optional
missing_important = [k for k in missing_keys
if '_extra_state' not in k and 'pos_embedder' not in k and 'accum_' not in k]
if missing_important:
logger.warning(f"Missing important keys ({len(missing_important)} total):")
for k in missing_important[:10]: # Show first 10
logger.warning(f" - {k}")
if len(missing_important) > 10:
logger.warning(f" ... and {len(missing_important) - 10} more")
if unexpected_keys:
logger.warning(f"Unexpected keys when loading reference model: {len(unexpected_keys)} keys")
logger.warning("All unexpected keys:")
for k in unexpected_keys:
logger.warning(f" - {k}")
logger.info(f"Successfully loaded {len(transformed_state)} parameters into reference model")
else:
logger.warning(f"Checkpoint path {checkpoint_path} not found, using random weights")
model = model.to(device, dtype=dtype)
model.eval()
return model
except ImportError as e:
logger.error(f"Failed to import cosmos-predict2.5: {e}")
logger.info("Make sure cosmos-predict2.5 is in your Python path")
return None
@pytest.mark.usefixtures("distributed_setup")
def test_cosmos25_transformer():
"""Test COSMOS 2.5 transformer against reference implementation."""
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
# Create COSMOS 2.5 specific config
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25ArchConfig
arch_config = Cosmos25ArchConfig(
num_attention_heads=16,
attention_head_dim=128, # 2048 / 16
in_channels=16,
out_channels=16,
num_layers=28,
patch_size=(1, 2, 2),
max_size=(128, 240, 240),
rope_scale=(1.0, 3.0, 3.0), # T, H, W
text_embed_dim=1024,
mlp_ratio=4.0,
adaln_lora_dim=256,
use_adaln_lora=True,
concat_padding_mask=True,
extra_pos_embed_type=None,
use_crossattn_projection=True,
rope_enable_fps_modulation=False,
qk_norm="rms_norm",
)
cosmos25_config = Cosmos25VideoConfig(arch_config=arch_config)
# Create FastVideo model directly (Cosmos 2.5 is not in diffusers format)
logger.info("Creating FastVideo COSMOS 2.5 model...")
from fastvideo.models.dits.cosmos2_5 import Cosmos25Transformer3DModel
# Get hf_config from the arch_config for model initialization
hf_config = {
'in_channels': arch_config.in_channels,
'out_channels': arch_config.out_channels,
'num_attention_heads': arch_config.num_attention_heads,
'attention_head_dim': arch_config.attention_head_dim,
'num_layers': arch_config.num_layers,
'patch_size': arch_config.patch_size,
'max_size': arch_config.max_size,
'rope_scale': arch_config.rope_scale,
'text_embed_dim': arch_config.text_embed_dim,
'mlp_ratio': arch_config.mlp_ratio,
'adaln_lora_dim': arch_config.adaln_lora_dim,
'use_adaln_lora': arch_config.use_adaln_lora,
'concat_padding_mask': arch_config.concat_padding_mask,
'extra_pos_embed_type': arch_config.extra_pos_embed_type,
'use_crossattn_projection': arch_config.use_crossattn_projection,
'rope_enable_fps_modulation': arch_config.rope_enable_fps_modulation,
'qk_norm': arch_config.qk_norm,
}
fastvideo_model = Cosmos25Transformer3DModel(config=cosmos25_config, hf_config=hf_config)
fastvideo_model = fastvideo_model.to(device, dtype=precision)
fastvideo_model.eval()
# Construct checkpoint path using relative paths
checkpoint_file = os.path.join(MODEL_PATH, CHECKPOINT_SUBDIR, CHECKPOINT_FILENAME)
if not os.path.exists(checkpoint_file):
logger.warning(f"Checkpoint file not found at {checkpoint_file}")
logger.info("Will test architecture without loading checkpoint weights")
checkpoint_file = None
# Load checkpoint into FastVideo model using param_names_mapping
if checkpoint_file:
logger.info(f"Loading checkpoint into FastVideo model from {checkpoint_file}")
from fastvideo.models.loader.utils import hf_to_custom_state_dict, get_param_names_mapping
checkpoint = torch.load(checkpoint_file, map_location='cpu')
# Extract state dict (checkpoint might have 'state_dict', 'model', or be the dict itself)
if 'state_dict' in checkpoint:
checkpoint_state = checkpoint['state_dict']
elif 'model' in checkpoint:
checkpoint_state = checkpoint['model']
else:
checkpoint_state = checkpoint
# Filter to only model parameters (remove training metadata like accum_*)
model_state = {k: v for k, v in checkpoint_state.items()
if k.startswith('net.') and 'accum_' not in k}
# Convert checkpoint keys to FastVideo format using param_names_mapping
param_names_mapping_fn = get_param_names_mapping(
cosmos25_config.arch_config.param_names_mapping
)
custom_state_dict, reverse_mapping = hf_to_custom_state_dict(
model_state, param_names_mapping_fn
)
# Only load keys that exist in the model
model_param_names = set(fastvideo_model.state_dict().keys())
filtered_state_dict = {
k: v.to(device=device, dtype=precision)
for k, v in custom_state_dict.items()
if k in model_param_names
}
# Load into FastVideo model
missing_keys, unexpected_keys = fastvideo_model.load_state_dict(
filtered_state_dict, strict=False
)
if missing_keys:
logger.warning(f"Missing keys when loading checkpoint: {len(missing_keys)} keys")
# Filter out _extra_state keys as they're optional
missing_non_extra = [k for k in missing_keys if '_extra_state' not in k]
if missing_non_extra:
logger.warning(f"Missing non-extra keys (first 10): {missing_non_extra[:10]}")
if unexpected_keys:
logger.warning(f"Unexpected keys when loading checkpoint: {len(unexpected_keys)} keys")
logger.info(f"Successfully loaded {len(filtered_state_dict)} parameters into FastVideo model")
# Try to load reference model from the raw checkpoint
logger.info("Loading reference COSMOS 2.5 model...")
reference_model = load_reference_cosmos25_model(checkpoint_file, device, precision) if checkpoint_file else None
# Set models to eval mode
fastvideo_model = fastvideo_model.eval()
if reference_model is not None:
reference_model = reference_model.eval()
# Create test inputs
batch_size = 1
seq_len = 77 # Typical T5 sequence length
# Video latents [B, C, T, H, W]
# COSMOS 2.5: 16 channels (VAE latent), no condition mask in input
hidden_states = torch.randn(
batch_size,
16, # VAE channels only (condition mask added internally)
1, # Single frame for image generation (or 16 for video)
64, # Height (720p / 8 / 2 patch = 45, use 64 for testing)
64, # Width
device=device,
dtype=precision
)
# Condition mask [B, 1, T, H, W] - for video2world conditioning
condition_mask = torch.zeros(
batch_size,
1,
1,
64,
64,
device=device,
dtype=precision
)
# Text embeddings [B, L, D] - Qwen 7B embeddings (100,352 dims)
# Using 100,352 dimensions to match the crossattn_projection layer
encoder_hidden_states = torch.randn(
batch_size,
seq_len,
100352,
device=device,
dtype=precision
)
# Timestep [B, T] - official model expects [B, T] shape with dtype matching model precision
# For single frame, use [B, 1]
timestep = torch.full((batch_size, 1), 500.0, device=device, dtype=precision)
# Padding mask [B, H, W] - official model expects NO channel dimension
# It's added internally via unsqueeze(1) if needed
padding_mask = torch.ones(
batch_size,
64,
64,
device=device,
dtype=precision
)
# FPS for temporal scaling
fps = 16
forward_batch = ForwardBatch(
data_type="dummy",
)
logger.info("Running inference...")
with torch.no_grad():
with torch.autocast('cuda', dtype=precision):
# FastVideo model
with set_forward_context(
current_timestep=500,
attn_metadata=None,
forward_batch=forward_batch,
):
# FastVideo expects padding_mask in [B, 1, H, W] format
padding_mask_fv = padding_mask.unsqueeze(1) # Add channel dimension for FastVideo
# FastVideo supports both [B] and [B, T] formats - use [B, T] to match official model
# This ensures each frame gets its own timestep embedding (even if values are the same)
timestep_fv = timestep # Already in [B, T] format
output_fv = fastvideo_model(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep_fv,
condition_mask=condition_mask,
padding_mask=padding_mask_fv,
fps=fps,
)
# Reference model (if available)
if reference_model is not None:
# Prepare input for reference model
# MinimalV1LVGDiT adds condition mask internally, so pass them separately
from cosmos_predict2._src.predict2.conditioner import DataType
# Determine data_type based on temporal dimension
num_frames = hidden_states.shape[2]
ref_data_type = DataType.VIDEO if num_frames > 1 else DataType.IMAGE
# Reference model expects different input format
# Pass hidden_states without condition_mask (model concatenates it internally)
# timestep is already in [B, T] format with correct dtype
# padding_mask is already in [B, H, W] format (no channel dimension)
# FPS should be a tensor [B] or scalar
fps_tensor = torch.tensor([fps], device=device, dtype=precision)
output_ref = reference_model(
x_B_C_T_H_W=hidden_states, # [B, 16, T, H, W] - model will add condition mask
timesteps_B_T=timestep, # Already in [B, T] format
crossattn_emb=encoder_hidden_states,
condition_video_input_mask_B_C_T_H_W=condition_mask if ref_data_type == DataType.VIDEO else None,
fps=fps_tensor,
padding_mask=padding_mask, # [B, H, W] format
data_type=ref_data_type,
)
# Check FastVideo output shape and dtype
logger.info(f"FastVideo output shape: {output_fv.shape}")
logger.info(f"FastVideo output dtype: {output_fv.dtype}")
assert output_fv.shape[0] == batch_size, "Batch size mismatch"
assert output_fv.shape[1] == 16, "Output channels should be 16"
assert output_fv.dtype == precision, f"Output dtype mismatch: {output_fv.dtype} vs {precision}"
# Compare with reference if available
if reference_model is not None:
logger.info(f"Reference output shape: {output_ref.shape}")
# Check if outputs have the same shape
assert output_fv.shape == output_ref.shape, \
f"Output shapes don't match: {output_fv.shape} vs {output_ref.shape}"
assert output_fv.dtype == output_ref.dtype, \
f"Output dtype don't match: {output_fv.dtype} vs {output_ref.dtype}"
# Check if outputs are similar
max_diff = torch.max(torch.abs(output_fv - output_ref))
mean_diff = torch.mean(torch.abs(output_fv - output_ref))
relative_diff = mean_diff / (torch.mean(torch.abs(output_ref)) + 1e-8)
logger.info(f"Max difference: {max_diff.item():.6f}")
logger.info(f"Mean difference: {mean_diff.item():.6f}")
logger.info(f"Relative difference: {relative_diff.item():.6f}")
# Allow for some numerical differences due to implementation details
assert max_diff < 1e-1, f"Maximum difference too large: {max_diff.item()}"
assert mean_diff < 1e-2, f"Mean difference too large: {mean_diff.item()}"
logger.info("✓ COSMOS 2.5 FastVideo implementation matches reference!")
else:
logger.warning("Reference model not available, skipping comparison")
logger.info("✓ COSMOS 2.5 FastVideo model runs successfully!")
@pytest.mark.usefixtures("distributed_setup")
def test_cosmos25_transformer_video():
"""Test COSMOS 2.5 transformer with video input (multiple frames)."""
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
# Create COSMOS 2.5 specific config
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25ArchConfig
arch_config = Cosmos25ArchConfig(
num_attention_heads=16,
attention_head_dim=128,
in_channels=16,
out_channels=16,
num_layers=28,
patch_size=(1, 2, 2),
max_size=(128, 240, 240),
rope_scale=(1.0, 3.0, 3.0),
text_embed_dim=1024,
mlp_ratio=4.0,
adaln_lora_dim=256,
use_adaln_lora=True,
concat_padding_mask=True,
extra_pos_embed_type=None,
use_crossattn_projection=True, # Enable to match official model
rope_enable_fps_modulation=False,
qk_norm="rms_norm",
)
cosmos25_config = Cosmos25VideoConfig(arch_config=arch_config)
# Create FastVideo model directly (Cosmos 2.5 is not in diffusers format)
logger.info("Creating FastVideo COSMOS 2.5 model for video test...")
from fastvideo.models.dits.cosmos2_5 import Cosmos25Transformer3DModel
# Get hf_config from the arch_config for model initialization
hf_config = {
'in_channels': arch_config.in_channels,
'out_channels': arch_config.out_channels,
'num_attention_heads': arch_config.num_attention_heads,
'attention_head_dim': arch_config.attention_head_dim,
'num_layers': arch_config.num_layers,
'patch_size': arch_config.patch_size,
'max_size': arch_config.max_size,
'rope_scale': arch_config.rope_scale,
'text_embed_dim': arch_config.text_embed_dim,
'mlp_ratio': arch_config.mlp_ratio,
'adaln_lora_dim': arch_config.adaln_lora_dim,
'use_adaln_lora': arch_config.use_adaln_lora,
'concat_padding_mask': arch_config.concat_padding_mask,
'extra_pos_embed_type': arch_config.extra_pos_embed_type,
'use_crossattn_projection': arch_config.use_crossattn_projection,
'rope_enable_fps_modulation': arch_config.rope_enable_fps_modulation,
'qk_norm': arch_config.qk_norm,
}
model = Cosmos25Transformer3DModel(config=cosmos25_config, hf_config=hf_config)
model = model.to(device, dtype=precision)
model.eval()
# Create video input with multiple frames
batch_size = 1
num_frames = 16 # Video with 16 frames
seq_len = 77
hidden_states = torch.randn(
batch_size,
16,
num_frames, # Multiple frames
64,
64,
device=device,
dtype=precision
)
condition_mask = torch.zeros(
batch_size,
1,
num_frames,
64,
64,
device=device,
dtype=precision
)
# Set first 2 frames as conditioning
condition_mask[:, :, :2, :, :] = 1.0
encoder_hidden_states = torch.randn(
batch_size,
seq_len,
100352, # Qwen 7B embedding dimension (matches crossattn_proj input)
device=device,
dtype=precision
)
timestep = torch.tensor([500], device=device, dtype=torch.long)
padding_mask = torch.ones(
batch_size,
1,
64,
64,
device=device,
dtype=precision
)
fps = 16
forward_batch = ForwardBatch(
data_type="dummy",
)
logger.info("Running video inference...")
with torch.no_grad():
with torch.autocast('cuda', dtype=precision):
with set_forward_context(
current_timestep=500,
attn_metadata=None,
forward_batch=forward_batch,
):
output = model(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
condition_mask=condition_mask,
padding_mask=padding_mask,
fps=fps,
)
logger.info(f"Video output shape: {output.shape}")
logger.info(f"Video output dtype: {output.dtype}")
# Check output shape
assert output.shape[0] == batch_size, "Batch size mismatch"
assert output.shape[1] == 16, "Output channels should be 16"
assert output.shape[2] == num_frames, "Number of frames mismatch"
assert output.dtype == precision, f"Output dtype mismatch"
logger.info("✓ COSMOS 2.5 video inference successful!")
if __name__ == "__main__":
# Run tests directly
test_cosmos25_transformer()
test_cosmos25_transformer_video()
@@ -25,8 +25,8 @@ os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
local_dir=os.path.join(
"data", BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
CONFIG_PATH = os.path.join(TRANSFORMER_PATH, "config.json")
@@ -5,7 +5,6 @@ import numpy as np
import pytest
import torch
from diffusers import WanTransformer3DModel
from torch.testing import assert_close
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
@@ -24,8 +23,8 @@ os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@@ -121,4 +120,10 @@ 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)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
+2 -2
View File
@@ -23,8 +23,8 @@ os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
local_dir=os.path.join(
"data", BASE_MODEL_PATH))
VAE_PATH = os.path.join(MODEL_PATH, "vae")
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
+18 -7
View File
@@ -12,7 +12,6 @@ from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import VAELoader
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.utils import maybe_download_model
from torch.testing import assert_close
logger = init_logger(__name__)
@@ -21,16 +20,16 @@ os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
VAE_PATH = os.path.join(MODEL_PATH, "vae")
@pytest.mark.usefixtures("distributed_setup")
def test_wan_vae():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.float32
precision_str = "fp32"
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=WanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_cpu_offload = False
@@ -71,7 +70,13 @@ def test_wan_vae():
# Check if latents have the same shape
assert latent1.mean.shape == latent2.mean.shape, f"Latent shapes don't match: {latent1.mean.shape} vs {latent2.mean.shape}"
# Check if latents are similar
assert_close(latent1.mean, latent2.mean, atol=1e-4, rtol=1e-4)
max_diff_encode = torch.max(torch.abs(latent1.mean - latent2.mean))
mean_diff_encode = torch.mean(torch.abs(latent1.mean - latent2.mean))
logger.info("Maximum difference between encoded latents: %s",
max_diff_encode.item())
logger.info("Mean difference between encoded latents: %s",
mean_diff_encode.item())
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {mean_diff_encode.item()}"
# Test decoding
logger.info("Testing decoding...")
latent1_tensor = latent1.mode()
@@ -93,4 +98,10 @@ def test_wan_vae():
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
# Check if outputs are similar
assert_close(output1, output2, atol=1e-5, rtol=1e-3)
max_diff_decode = torch.max(torch.abs(output1 - output2))
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
logger.info("Maximum difference between decoded outputs: %s",
max_diff_decode.item())
logger.info("Mean difference between decoded outputs: %s",
mean_diff_decode.item())
assert max_diff_decode < 1e-5, f"Decoded outputs differ significantly: max diff = {mean_diff_decode.item()}"
@@ -729,6 +729,9 @@ class DistillationPipeline(TrainingPipeline):
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
world_group = get_world_group()
if world_group.world_size > 1:
world_group.broadcast(timestep, src=0)
timestep = shift_timestep(
timestep,
@@ -841,6 +844,9 @@ class DistillationPipeline(TrainingPipeline):
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
world_group = get_world_group()
if world_group.world_size > 1:
world_group.broadcast(fake_score_timestep, src=0)
fake_score_timestep = shift_timestep(
fake_score_timestep,
@@ -1,6 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import os
import time
from collections import deque
from typing import Any
@@ -46,13 +45,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
def initialize_training_pipeline(self, training_args: TrainingArgs):
"""Initialize the self-forcing training pipeline."""
# Check if FSDP2 auto wrap is enabled - not supported for self-forcing distillation
if os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1":
raise NotImplementedError(
"FASTVIDEO_FSDP2_AUTOWRAP is not implemented for self-forcing distillation. "
"Please set FASTVIDEO_FSDP2_AUTOWRAP=0 or unset the environment variable."
)
logger.info("Initializing self-forcing distillation pipeline...")
self.generator_ema: EMA_FSDP | None = None
+1 -21
View File
@@ -470,7 +470,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# local_main_process_only=False)
with self.tracker.timed("timing/reduce_loss"):
world_group = get_world_group()
avg_loss = world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
training_batch.total_loss += avg_loss.item()
return training_batch
@@ -656,23 +656,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
context_len = int(training_batch.encoder_hidden_states.shape[1])
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
arch_config = self.training_args.pipeline_config.dit_config.arch_config
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
with self.profiler_controller.region(
@@ -758,9 +741,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
sampling_param.width = training_args.num_width
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
if training_args.validation_guidance_scale:
sampling_param.guidance_scale = float(
training_args.validation_guidance_scale)
assert self.seed is not None
sampling_param.seed = self.seed
+1 -2
View File
@@ -510,8 +510,7 @@ def load_checkpoint(transformer,
return 0
# Extract step number from checkpoint path
step = int(
os.path.basename(os.path.normpath(checkpoint_path)).split('-')[-1])
step = int(os.path.basename(checkpoint_path).split('-')[-1])
if rank == 0:
logger.info("Loading checkpoint from step %s", step)
+7 -10
View File
@@ -48,26 +48,23 @@ class Worker:
# This env var set by Ray causes exceptions with graph building.
os.environ.pop("NCCL_ASYNC_ERROR_HANDLING", None)
# Set environment variables BEFORE calling get_local_torch_device()
# so that each worker uses the correct device
if self.fastvideo_args.distributed_executor_backend == "mp":
os.environ["LOCAL_RANK"] = str(self.local_rank)
os.environ["RANK"] = str(self.rank)
os.environ["WORLD_SIZE"] = str(self.fastvideo_args.num_gpus)
# Platform-agnostic device initialization
self.device = get_local_torch_device()
from fastvideo.platforms import current_platform
# Set the CUDA device BEFORE any CUDA calls
# _check_if_gpu_supports_dtype(self.model_config.dtype)
if current_platform.is_cuda_alike():
torch.cuda.set_device(self.device)
self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0]
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
else:
# For MPS, we can't get memory info the same way
self.init_gpu_memory = 0
if self.fastvideo_args.distributed_executor_backend == "mp":
os.environ["LOCAL_RANK"] = str(self.local_rank)
os.environ["RANK"] = str(self.rank)
os.environ["WORLD_SIZE"] = str(self.fastvideo_args.num_gpus)
# Initialize the distributed environment.
maybe_init_distributed_environment_and_model_parallel(
self.fastvideo_args.tp_size, self.fastvideo_args.sp_size,
-4
View File
@@ -467,10 +467,6 @@ class WorkerMultiprocProc:
"output_batch": output_batch.output.cpu(),
"logging_info": logging_info
})
else:
result = self.worker.execute_method(
method, *args, **kwargs)
self.pipe.send(result)
else:
result = self.worker.execute_method(method, *args, **kwargs)
self.pipe.send(result)
+8 -66
View File
@@ -71,7 +71,6 @@ class PreprocessingDataValidator:
for name, validator in self.validators.items():
if not validator(batch):
logger.info(f"Failed validation for {name}")
self.filter_counts[name] += 1
return False
@@ -88,8 +87,6 @@ class PreprocessingDataValidator:
"""Validate resolution constraints"""
aspect = self.max_height / self.max_width
height = None
width = None
if batch["resolution"] is not None:
height = batch["resolution"].get("height", None)
width = batch["resolution"].get("width", None)
@@ -97,15 +94,12 @@ class PreprocessingDataValidator:
if height is None or width is None:
return False
ret = self._filter_resolution(
return self._filter_resolution(
height,
width,
max_h_div_w_ratio=self.hw_aspect_threshold * aspect,
min_h_div_w_ratio=1 / self.hw_aspect_threshold * aspect,
)
if not ret:
logger.info(f"failed in resolution: {batch['caption']}")
return ret
def _filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
min_h_div_w_ratio: float) -> bool:
@@ -119,19 +113,14 @@ class PreprocessingDataValidator:
if (batch["num_frames"] / batch["fps"]
> self.video_length_tolerance_range *
(self.num_frames / self.train_fps * self.speed_factor)):
logger.info("Failed in 1")
return False
frame_interval = (batch["fps"] / self.train_fps) * self.speed_factor
frame_interval = batch["fps"] / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch["num_frames"],
frame_interval).astype(int)
# logger.info("Failed in 2")
result = not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
if not result:
logger.info(f"failed in frame_sampling: {batch['caption']}")
return result
return not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
def log_validation_stats(self):
info = ""
@@ -291,46 +280,9 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
dataset = dataset.shard(num_shards=get_world_size(),
index=get_world_rank())
elif preprocess_config.dataset_type == DatasetType.MERGED:
merge_txt_path = os.path.join(preprocess_config.dataset_path,
"merge.txt")
if os.path.exists(merge_txt_path):
logger.info(f"Found merge.txt at {merge_txt_path}")
with open(merge_txt_path) as f:
line = f.read().strip()
if "," not in line:
raise ValueError(
f"Invalid format in {merge_txt_path}: expected 'video_folder,metadata_json_path'"
)
video_folder, metadata_json_path = line.split(",", 1)
video_folder = video_folder.strip()
metadata_json_path = metadata_json_path.strip()
if not os.path.isabs(video_folder):
video_folder = os.path.join(preprocess_config.dataset_path,
video_folder)
if not os.path.isabs(metadata_json_path):
metadata_json_path = os.path.join(
preprocess_config.dataset_path, metadata_json_path)
else:
logger.info(
f"merge.txt not found at {merge_txt_path}, using default paths")
metadata_json_path = os.path.join(preprocess_config.dataset_path,
"videos2caption.json")
video_folder = os.path.join(preprocess_config.dataset_path,
"videos")
if not os.path.exists(metadata_json_path):
logger.error(f"Metadata file not found: {metadata_json_path}")
raise FileNotFoundError(
f"Metadata file not found: {metadata_json_path}")
if not os.path.exists(video_folder):
logger.error(f"Video folder not found: {video_folder}")
raise FileNotFoundError(f"Video folder not found: {video_folder}")
logger.info(f"Using metadata file: {metadata_json_path}")
logger.info(f"Using video folder: {video_folder}")
metadata_json_path = os.path.join(preprocess_config.dataset_path,
"videos2caption.json")
video_folder = os.path.join(preprocess_config.dataset_path, "videos")
dataset = load_dataset("json",
data_files=metadata_json_path,
split=split)
@@ -341,18 +293,9 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
if "path" in column_names:
dataset = dataset.rename_column("path", "name")
print(f"Length of dataset before filtering: {len(dataset)}")
if len(dataset) > 0:
print(f"DEBUG: First item in dataset: {dataset[0]}")
# Disable caching to ensure our print statements run
dataset = dataset.filter(validator, load_from_cache_file=False)
validator.log_validation_stats()
print(f"Length of dataset after filtering: {len(dataset)}")
dataset = dataset.filter(validator)
dataset = dataset.shard(num_shards=get_world_size(),
index=get_world_rank())
print(f"Length of dataset after sharding: {len(dataset)}")
# add video column
def add_video_column(item: dict[str, Any]) -> dict[str, Any]:
@@ -360,7 +303,6 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
return item
dataset = dataset.map(add_video_column)
print(f"Length of dataset after mapping: {len(dataset)}")
if preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
dataset = dataset.cast_column("video", Video())
else:
@@ -40,7 +40,6 @@ class PreprocessWorkflow(WorkflowBase):
video_length_tolerance_range=preprocess_config.
video_length_tolerance_range,
drop_short_ratio=preprocess_config.drop_short_ratio,
hw_aspect_threshold=preprocess_config.hw_aspect_threshold,
)
self.add_component("raw_data_validator", raw_data_validator)
Binary file not shown.

Before

Width:  |  Height:  |  Size: 113 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 229 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 168 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 148 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 155 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 875 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 664 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 62 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 686 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 957 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 585 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 558 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 942 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 890 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 433 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 595 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 781 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 783 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 762 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 68 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 147 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 89 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 133 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 213 KiB

+4 -16
View File
@@ -14,7 +14,6 @@ edit_uri: edit/main/docs/
# Configuration
theme:
name: material
favicon: assets/logos/icon_simple.svg
palette:
- scheme: default
toggle:
@@ -47,18 +46,11 @@ plugins:
hooks:
on_pre_build: "docs.generate_examples:on_pre_build_hook"
- autorefs
# - awesome-nav
# - glightbox
- git-revision-date-localized:
# exclude autogenerated files
exclude:
- examples/*
- api-autonav:
modules: ["fastvideo"]
modules: ["fastvideo"]
api_root_uri: "api"
exclude:
- "re:fastvideo\\._.*"
- "fastvideo.third_party"
- "re:fastvideo\\._.*"
- mkdocstrings:
handlers:
python:
@@ -83,10 +75,9 @@ plugins:
inventories:
- https://docs.python.org/3/objects.inv
# Markdown extensions
markdown_extensions:
- admonition
- pymdownx.highlight:
anchor_linenums: true
line_spans: __span
@@ -112,10 +103,8 @@ markdown_extensions:
- pymdownx.tasklist:
custom_checkbox: true
- pymdownx.tilde
# For in page [TOC] (not sidebar)
- toc:
permalink: true
- mdx_truly_sane_lists
# Page tree
nav:
@@ -162,7 +151,6 @@ nav:
- Index: contributing/developer_env/index.md
- Docker: contributing/developer_env/docker.md
- RunPod: contributing/developer_env/runpod.md
- Testing: contributing/testing.md
- Profiling: contributing/profiling.md
- API Reference:
- FastVideo: api/fastvideo.md
@@ -178,4 +166,4 @@ extra:
# Custom CSS
extra_css:
- assets/custom.css
- assets/custom.css

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