Compare commits

...
Author SHA1 Message Date
SolitaryThinker e9d95b1c10 wip 2025-12-12 06:30:08 +00:00
SolitaryThinker 8b55e9706c wip 2025-12-10 01:59:07 +00:00
William Lin 3ff640b2e6 [bigfix] [distillation] Fix DMD inference pipeline noise initialization shape (#921) 2025-12-08 13:00:48 -08:00
William Lin c722429ab5 [docs] fix testing.md visibility (#920) 2025-12-08 00:44:53 -08:00
KyleShaoandKyleS1016 e04a192de6 [feat]: add COSMOS 2.5 DiT implementation (#897)
Co-authored-by: KyleS1016 <kyle.s@gmicloud.ai>
2025-12-07 21:48:32 -08:00
William Lin c9ca6d1298 [docs] add docs for ssim testing (#918) 2025-12-06 18:20:04 -08:00
Wenxuan TanandSolitaryThinker 754292c419 Use assert_close in tests (#429)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-12-06 18:18:25 -08:00
Qi Jia 0082bc66fc fix: correct mp backend GPU assignment on multi-GPU systems (#912) 2025-11-30 23:00:22 -08:00
Ohm-Rishabh 8b1937422e [feat] training mfu calculation scripts (#871) 2025-11-27 16:54:17 -08:00
fb6cbf23e6 Fix the docs (#905)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-11-27 00:37:03 -08:00
Mihir Jagtap c8fdd5ed7b [docs] modified the .github/workflows/docs.yml file to include path filtering (#906) 2025-11-26 17:34:21 -08:00
Loay Rashid 1c19a6a00c [Bugfix] Minor bugfixes (#889) 2025-11-26 17:20:45 -08:00
William Lin d44409c704 [CI] fix VSA training CI (#900) 2025-11-24 17:47:59 -08:00
Zhang Peiyuan 5d1c7852b7 + Awesome work using FastVideo or our research projects (#898) 2025-11-23 22:22:27 -08:00
Wenxuan Tan 77a211d006 [misc] Update wechat link (#893) 2025-11-20 19:59:05 -08:00
Wei Zhou bef8169bb1 [Feat] [I2V] resize all image sizes to below 480*832 (#890) 2025-11-20 00:08:36 -08:00
William Lin 681f1583f9 [readme] update link to inference code (#887) 2025-11-19 13:24:13 -08:00
e3b4564d5a [feat] Add inference for MoE SF (#880)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-11-19 13:16:24 -08:00
Shao Duan c0d03fc43d [bugfix] [lora] [CI] Fix LoRA alpha scaling factor & Fix LoRA Inference CI (#870) 2025-11-19 01:02:01 -08:00
Wei Zhou 404ee8538e [Bugfix] [DMD Distillation] Each rank should have its own timestep sampled (#885) 2025-11-18 14:03:25 -08:00
Shao Duan e57ac59462 Fix mp worker busy loop to handle all string RPC methods (#881) 2025-11-16 13:26:44 -08:00
Mihir Jagtap 8c55fdaf7e [docs] add favicon (#878) 2025-11-15 13:44:16 -08:00
94 changed files with 3518 additions and 1162 deletions
+1 -1
View File
@@ -75,7 +75,7 @@ case "$TEST_TYPE" in
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
+10
View File
@@ -3,8 +3,18 @@ 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,6 +10,7 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
+13 -16
View File
@@ -1,13 +1,12 @@
<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/XcY0Cpv" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/TM8JyJCd" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -15,6 +14,7 @@ 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!
@@ -111,24 +111,21 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## 📑 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 -->
## Awesome work using FastVideo or our research projects
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
- [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)
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
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).
## Acknowledgement
We learned and reused code from the following projects:
- [Wan-Video](https://github.com/Wan-Video)
+4
View File
@@ -70,3 +70,7 @@ 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
@@ -0,0 +1,129 @@
# 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.
+3 -2
View File
@@ -24,12 +24,13 @@ 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.
@@ -42,7 +43,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
@@ -1,51 +0,0 @@
# 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.

Before

Width:  |  Height:  |  Size: 98 KiB

@@ -0,0 +1,44 @@
# 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()
@@ -3,6 +3,9 @@ import os
import requests
import base64
import time
import json
from pathlib import Path
import tempfile
import gradio as gr
@@ -12,6 +15,7 @@ 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",
}
@@ -37,7 +41,7 @@ class RayServeClient:
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=300
timeout=900 # 15 minutes timeout for longer video generation
)
round_trip_time = time.time() - start_time
@@ -81,49 +85,78 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
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"
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
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>"
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>"
def load_example_prompts():
@@ -144,26 +177,83 @@ def load_example_prompts():
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("prompts/prompts_final.txt")
# Load prompts from prompts.txt
examples, example_labels = load_from_file("examples/inference/gradio/serving/prompts.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"]
return examples, example_labels
# 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
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, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, 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:
@@ -172,6 +262,15 @@ 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,
@@ -183,7 +282,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
"width": width,
"randomize_seed": randomize_seed,
"return_frames": False,
"image_path": None,
"image_data": image_data,
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
}
@@ -198,16 +297,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:
@@ -219,7 +318,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, timing_details
return video_path, used_seed, ""
else:
return None, "Failed to save video", ""
else:
@@ -228,7 +327,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 = load_example_prompts()
examples, example_labels, example_images = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
@@ -239,33 +338,39 @@ 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,
'seed': params.seed,
}
# 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,
# }
return {
'height': 448,
'height': 480,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
'num_frames': 73,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
# 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)
with gr.Blocks(title="FastWan", theme=theme) as demo:
# 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:
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_post_training/" 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_causalwan_preview/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
@@ -280,8 +385,8 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
choices=available_models,
value=default_model,
label="Select Model",
interactive=True
)
@@ -312,69 +417,70 @@ 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=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.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():
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,
)
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,
)
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")
# randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed", value=1000)
with gr.Column(scale=1, elem_classes="video-column"):
with gr.Column(scale=1):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
height=500,
container=True,
elem_classes="video-component"
autoplay=True,
)
gr.HTML("""
@@ -387,116 +493,10 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
}
.gradio-container {
max-width: 1200px !important;
max-width: 1400px !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,18 +511,20 @@ 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)
return examples[index]
return ""
selected_prompt = examples[index]
selected_image = example_images[index] if index < len(example_images) else None
return selected_prompt, selected_image
return "", None
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
outputs=[prompt, input_image],
)
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 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>
<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>
</div>
""")
@@ -537,6 +539,7 @@ 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]
@@ -545,29 +548,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(value=params.seed),
gr.update(visible=show_image_input),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=20),
gr.update(value=3.0),
gr.update(value=1024),
gr.update(visible=show_image_input),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
outputs=[height, width, num_frames, guidance_scale, image_tab],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
model_selection, prompt, negative_prompt, use_negative_prompt, guidance_scale, num_frames, height, width, input_image = args
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
result_path, seed_or_error, _ = generate_video(
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, progress
)
if result_path and os.path.exists(result_path):
@@ -575,14 +578,12 @@ 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(
@@ -592,14 +593,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,
# randomize_seed,
input_image,
],
outputs=[result, seed_output, error_output, timing_display],
outputs=[result, seed_output, error_output], # timing_display removed
concurrency_limit=20,
)
@@ -611,8 +612,11 @@ 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="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
default="",
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,
@@ -621,8 +625,15 @@ def main():
args = parser.parse_args()
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
# 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:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(args.backend_url, default_params)
@@ -630,6 +641,8 @@ 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
@@ -674,23 +687,23 @@ def main():
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<title>CausalWan</title>
<meta name="title" content="CausalWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, CausalWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:title" content="CausalWan">
<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="FastWan">
<meta property="og:site_name" content="CausalWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:title" content="CausalWan">
<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">
@@ -720,7 +733,14 @@ def main():
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
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")),
]
)
uvicorn.run(app, host=args.host, port=args.port)
@@ -0,0 +1,15 @@
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,6 +26,7 @@ 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 = {
@@ -42,6 +43,13 @@ 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,
}
}
@@ -58,6 +66,7 @@ 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):
@@ -91,11 +100,38 @@ 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"
# if "fullattn" in model_path.lower():
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
# else:
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
@@ -157,22 +193,41 @@ 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,
)
@@ -185,6 +240,13 @@ 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,
@@ -200,7 +262,7 @@ class BaseModelDeployment:
@serve.deployment(
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class T2VModelDeployment(BaseModelDeployment):
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
@@ -210,7 +272,7 @@ class T2VModelDeployment(BaseModelDeployment):
@serve.deployment(
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class T2V14BModelDeployment(BaseModelDeployment):
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
@@ -221,18 +283,32 @@ 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=50, ray_actor_options={"num_cpus": 2})
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 1})
@serve.ingress(app)
class FastVideoAPI:
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle], i2v_deployments: Dict[str, DeploymentHandle] = None):
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'])
@@ -257,10 +333,10 @@ class FastVideoAPI:
model_name = self._get_model_name(video_request.model_path)
try:
if video_request.model_path not in self.t2v_deployments:
if video_request.model_path not in self.all_deployments:
raise ValueError(f"Model {video_request.model_path} not found")
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
response_ref = self.all_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)
@@ -291,18 +367,21 @@ 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"
assert model in SUPPORTED_MODELS, f"Model {model} not supported. Supported models: {SUPPORTED_MODELS}"
assert replica_count > 0, f"Replicas must be greater than 0"
def start_ray_serve(
*,
t2v_model_paths: str,
t2v_model_replicas: str,
t2v_model_paths: str = "",
t2v_model_replicas: str = "",
i2v_model_paths: str = "",
i2v_model_replicas: str = "",
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
@@ -310,21 +389,39 @@ def start_ray_serve(
if not ray.is_initialized():
ray.init()
model_paths = t2v_model_paths.split(",")
replicas = [int(r) for r in t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# 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)
# Create T2V deployments
t2v_deps = {}
for model_path, replica_count in zip(model_paths, replicas):
for model_path, replica_count in zip(t2v_paths, t2v_reps):
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
api = FastVideoAPI.bind(t2v_deps)
# 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)
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(model_paths, replicas):
for model_path, replica_count in zip(t2v_paths, t2v_reps):
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")
@@ -340,12 +437,20 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
parser.add_argument("--t2v_model_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
default="",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
default="",
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",
@@ -361,13 +466,21 @@ if __name__ == "__main__":
args = parser.parse_args()
model_paths = args.t2v_model_paths.split(",")
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# 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)
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,
@@ -376,4 +489,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)
+5 -3
View File
@@ -1,3 +1,5 @@
python examples/inference/gradio/start_ray_serve_app.py \
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" \
--t2v_model_replicas "4,4"
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"
@@ -20,8 +20,10 @@ DEFAULT_BACKEND_PORT = 8000
DEFAULT_FRONTEND_HOST = "0.0.0.0"
DEFAULT_FRONTEND_PORT = 7860
DEFAULT_OUTPUT_PATH = "outputs"
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_T2V_REPLICAS = "4,4"
DEFAULT_T2V_MODELS = ""
DEFAULT_T2V_REPLICAS = ""
DEFAULT_I2V_MODELS = "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers"
DEFAULT_I2V_REPLICAS = "1"
HEALTH_CHECK_TIMEOUT = 5
HEALTH_CHECK_MAX_RETRIES = 100
@@ -100,6 +102,12 @@ 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
@@ -111,6 +119,10 @@ 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
@@ -173,6 +185,9 @@ 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}")
@@ -190,6 +205,14 @@ 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 "6.0"
--validation_guidance_scale "3.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,6 +83,9 @@ 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
@@ -184,6 +187,23 @@ 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",
+2 -1
View File
@@ -1,9 +1,10 @@
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"
"CosmosVideoConfig", "Cosmos25VideoConfig"
]
+181
View File
@@ -0,0 +1,181 @@
# 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,6 +45,7 @@ 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)
+2
View File
@@ -39,6 +39,8 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"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,
+4
View File
@@ -186,3 +186,7 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
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
+2
View File
@@ -78,6 +78,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
# 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,
# Cosmos2
"nvidia/Cosmos-Predict2-2B-Video2World":
-2
View File
@@ -191,8 +191,6 @@ class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
@dataclass
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
Wan2_2_T2V_A14B_SamplingParam):
guidance_scale: float = 2.0
guidance_scale_2: float = 2.0
num_inference_steps: int = 8
num_frames: int = 81
height: int = 448
+41
View File
@@ -152,3 +152,44 @@ 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
+21 -4
View File
@@ -101,12 +101,19 @@ 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
@@ -134,8 +141,13 @@ 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()
data += (self.slice_lora_b_weights(self.lora_B).to(data)
@ self.slice_lora_a_weights(self.lora_A).to(data))
# 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
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(
@@ -154,8 +166,13 @@ class BaseLayerWithLoRA(nn.Module):
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(get_local_torch_device())
data += \
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
# 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
self.base_layer.weight.data = data.to(current_device,
non_blocking=True)
+961
View File
@@ -0,0 +1,961 @@
# 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
@@ -86,6 +86,12 @@ 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.
@@ -105,6 +111,32 @@ 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,3 +180,51 @@ 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,7 +50,8 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler")))
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@@ -59,7 +59,8 @@ 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)))
transformer=self.get_module("transformer", None),
use_btchw_layout=True))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
@@ -62,7 +62,8 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
transformer=self.get_module("transformer"),
use_btchw_layout=True))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
+28 -3
View File
@@ -28,7 +28,8 @@ class LoRAPipeline(ComposedPipelineBase):
TODO: support training.
"""
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
dict) # state dicts of loaded lora adapters
dict
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
cur_adapter_name: str = ""
cur_adapter_path: str = ""
lora_layers: dict[str, BaseLayerWithLoRA] = {}
@@ -183,11 +184,26 @@ 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)
@@ -225,11 +241,20 @@ 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(
self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name],
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
@@ -5,13 +5,19 @@ 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)
TemporalRandomCrop, best_output_size)
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
@@ -41,6 +47,7 @@ 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
@@ -49,8 +56,17 @@ 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],
@@ -63,8 +79,28 @@ 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:
video = batch.video_loader[i].get_frames_at(frame_indices).data
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
elif fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHVISION:
video, _, _ = torchvision.io.read_video(batch.video_loader[i],
output_format="TCHW")
@@ -73,16 +109,75 @@ class VideoTransformStage(PipelineStage):
raise ValueError(
f"Invalid video loader type: {fastvideo_args.preprocess_config.video_loader_type}"
)
video = self.video_transform(video)
video_pixel_batch.append(video)
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_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:
batch.pil_image = video_pixel_values[:, :, 0, :, :]
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, :, :]
video_pixel_values = video_pixel_values.float() / 255.0
batch.latents = video_pixel_values
+176 -127
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
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
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,13 +34,16 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
def __init__(self,
transformer,
scheduler,
transformer_2=None,
vae=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.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
self.vae = vae
# 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
@@ -80,6 +83,13 @@ 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 = {}
@@ -103,113 +113,110 @@ class CausalDMDDenosingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
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)
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)
# 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)
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=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 += 1
remaining_frames = input_frames - 1
else:
remaining_frames = input_frames
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
# 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
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)
# Base position offset from any cache warm-up
pos_start_base = current_start_frame
pos_start_base = 0
# 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"
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],
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,
)
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
if boundary_timestep is not None:
self.transformer_2(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
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,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
@@ -222,7 +229,7 @@ class CausalDMDDenosingStage(DenoisingStage):
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
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:
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2
else:
current_model = self.transformer
@@ -280,8 +287,8 @@ class CausalDMDDenosingStage(DenoisingStage):
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -290,12 +297,22 @@ class CausalDMDDenosingStage(DenoisingStage):
).permute(0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
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 boundary_timestep is not None and t_cur >= boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
@@ -309,11 +326,23 @@ class CausalDMDDenosingStage(DenoisingStage):
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
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])
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])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
@@ -341,24 +370,44 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = current_model(
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(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
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,
)
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) -> None:
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
@@ -392,10 +441,10 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0], dtype=torch.long, device=device),
})
self.kv_cache1 = kv_cache1
return kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> None:
device) -> list[dict]:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
@@ -421,7 +470,7 @@ class CausalDMDDenosingStage(DenoisingStage):
"is_init":
False,
})
self.crossattn_cache = crossattn_cache
return crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
@@ -445,4 +494,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,7 +1085,6 @@ 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,26 +106,28 @@ class InputValidationStage(PipelineStage):
batch.pil_image = image
# further processing for ti2v task
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
if (fastvideo_args.pipeline_config.ti2v_task
or fastvideo_args.pipeline_config.is_causal
) 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 = 704 * 1280
max_area = 480 * 832
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,10 +28,14 @@ class LatentPreparationStage(PipelineStage):
denoised during the diffusion process.
"""
def __init__(self, scheduler, transformer) -> None:
def __init__(self,
scheduler,
transformer,
use_btchw_layout: bool = False) -> None:
super().__init__()
self.scheduler = scheduler
self.transformer = transformer
self.use_btchw_layout = use_btchw_layout
def forward(
self,
@@ -78,15 +82,26 @@ class LatentPreparationStage(PipelineStage):
raise ValueError("Height and width must be provided")
# Calculate latent shape
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,
)
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,
)
# Validate generator if it's a list
if isinstance(generator, list) and len(generator) != batch_size:
+4 -32
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))
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder_2")
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer_2")
@@ -130,17 +130,6 @@ 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
@@ -148,22 +137,5 @@ 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}"
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()}"
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)
+16 -20
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))
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder")
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer")
@@ -68,8 +68,7 @@ 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,
@@ -78,6 +77,18 @@ 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
@@ -139,19 +150,4 @@ 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}"
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()}"
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-1, rtol=1e-4)
+2 -33
View File
@@ -133,24 +133,7 @@ def test_t5_encoder(t5_model_paths):
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("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()}"
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-4, rtol=1e-4)
@pytest.mark.usefixtures("distributed_setup")
@@ -252,18 +235,4 @@ 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}"
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()}"
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-4, rtol=1e-4)
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import re
import pytest
@@ -52,12 +53,38 @@ 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
@@ -137,14 +164,16 @@ 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]}"
generation_kwargs["output_path"] = output_dir
generation_kwargs["output_video_name"] = output_video_name
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
generator.generate_video(prompt, **generation_kwargs)
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
generated_video_path), f"Output video was not generated at {generated_video_path}"
reference_folder = os.path.join(script_dir, 'L40S_reference_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
@@ -153,13 +182,25 @@ 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 for the switched LoRA
# Find the matching reference video - try exact match first, then fuzzy match
# The reference might have different sanitization (e.g., trailing spaces)
reference_video_name = None
unsanitized_prefix = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
for filename in os.listdir(reference_folder):
# 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
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
break
if not reference_video_name:
@@ -167,7 +208,6 @@ 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)
@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("pytest ./fastvideo/tests/ssim -vs")
run_test("hf auth login --token $HF_API_KEY && 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():
@@ -1,19 +0,0 @@
#!/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
@@ -1,164 +0,0 @@
# 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
@@ -1,123 +0,0 @@
# 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
@@ -1,152 +0,0 @@
# 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 +1,10 @@
{"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}
{
"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
}
@@ -0,0 +1,211 @@
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()
@@ -0,0 +1,598 @@
# 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))
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
CONFIG_PATH = os.path.join(TRANSFORMER_PATH, "config.json")
@@ -5,6 +5,7 @@ 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
@@ -23,8 +24,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))
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@@ -120,10 +121,4 @@ def test_wan_transformer():
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
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()}"
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
+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))
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
VAE_PATH = os.path.join(MODEL_PATH, "vae")
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
+7 -18
View File
@@ -12,6 +12,7 @@ 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__)
@@ -20,16 +21,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))
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
)
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.bfloat16
precision_str = "bf16"
precision = torch.float32
precision_str = "fp32"
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=WanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_cpu_offload = False
@@ -70,13 +71,7 @@ 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
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()}"
assert_close(latent1.mean, latent2.mean, atol=1e-4, rtol=1e-4)
# Test decoding
logger.info("Testing decoding...")
latent1_tensor = latent1.mode()
@@ -98,10 +93,4 @@ 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
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()}"
assert_close(output1, output2, atol=1e-5, rtol=1e-3)
@@ -729,9 +729,6 @@ 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,
@@ -844,9 +841,6 @@ 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,
+21 -1
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()
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
avg_loss = world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
training_batch.total_loss += avg_loss.item()
return training_batch
@@ -656,6 +656,23 @@ 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(
@@ -741,6 +758,9 @@ 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
+2 -1
View File
@@ -510,7 +510,8 @@ def load_checkpoint(transformer,
return 0
# Extract step number from checkpoint path
step = int(os.path.basename(checkpoint_path).split('-')[-1])
step = int(
os.path.basename(os.path.normpath(checkpoint_path)).split('-')[-1])
if rank == 0:
logger.info("Loading checkpoint from step %s", step)
+10 -7
View File
@@ -48,23 +48,26 @@ 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
# _check_if_gpu_supports_dtype(self.model_config.dtype)
# Set the CUDA device BEFORE any CUDA calls
if current_platform.is_cuda_alike():
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
torch.cuda.set_device(self.device)
self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[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,6 +467,10 @@ 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)
+66 -8
View File
@@ -71,6 +71,7 @@ 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
@@ -87,6 +88,8 @@ 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)
@@ -94,12 +97,15 @@ class PreprocessingDataValidator:
if height is None or width is None:
return False
return self._filter_resolution(
ret = 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:
@@ -113,14 +119,19 @@ 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
frame_interval = (batch["fps"] / self.train_fps) * self.speed_factor
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch["num_frames"],
frame_interval).astype(int)
return not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
# 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
def log_validation_stats(self):
info = ""
@@ -280,9 +291,46 @@ 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:
metadata_json_path = os.path.join(preprocess_config.dataset_path,
"videos2caption.json")
video_folder = os.path.join(preprocess_config.dataset_path, "videos")
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}")
dataset = load_dataset("json",
data_files=metadata_json_path,
split=split)
@@ -293,9 +341,18 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
if "path" in column_names:
dataset = dataset.rename_column("path", "name")
dataset = dataset.filter(validator)
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.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]:
@@ -303,6 +360,7 @@ 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,6 +40,7 @@ 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.

After

Width:  |  Height:  |  Size: 113 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 229 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 168 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 148 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 155 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 875 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 664 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 62 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 686 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 957 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 585 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 558 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 942 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 890 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 433 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 595 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 781 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 783 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 762 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 68 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 147 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 89 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 133 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 213 KiB

+16 -4
View File
@@ -14,6 +14,7 @@ edit_uri: edit/main/docs/
# Configuration
theme:
name: material
favicon: assets/logos/icon_simple.svg
palette:
- scheme: default
toggle:
@@ -46,11 +47,18 @@ 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\\._.*"
- "re:fastvideo\\._.*"
- "fastvideo.third_party"
- mkdocstrings:
handlers:
python:
@@ -75,9 +83,10 @@ plugins:
inventories:
- https://docs.python.org/3/objects.inv
# Markdown extensions
markdown_extensions:
- admonition
- pymdownx.highlight:
anchor_linenums: true
line_spans: __span
@@ -103,8 +112,10 @@ 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:
@@ -151,6 +162,7 @@ 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
@@ -166,4 +178,4 @@ extra:
# Custom CSS
extra_css:
- assets/custom.css
- assets/custom.css
+110
View File
@@ -0,0 +1,110 @@
[
{
"prompt": "Young man skating with a skateboard on the ramps with graffiti of a park with trees, on a sunny day.",
"image_path": "images/mixkit-boy-skating-with-a-skateboard-in-a-park-with-ramps-34389.png"
},
{
"prompt": "In the midst of the joyous New Year's Eve celebration, the cheerful group of friends, their spirits lifted by the festivities, decides to immortalize the moment with a vibrant snapshot",
"image_path": "images/mixkit-a-cheerful-group-of-friends-celebrate-new-years-eve-and-51525.png"
},
{
"prompt": "A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.",
"image_path": "images/mixkit-a-cute-couple-playing-on-the-grass-4688.png"
},
{
"prompt": "Aerial view of a rocky mountain in the forest at a sunny day drone flight footage",
"image_path": "images/mixkit-aerial-view-of-a-rocky-mountain-in-the-forest-50589.png"
},
{
"prompt": "A little girl wearing a pink security helmet and denim overall discovers the art of cycling amidst the serene park, as the camera captures her graceful progress.",
"image_path": "images/mixkit-a-little-girl-cruises-through-the-forest-path-on-her-50088.png"
},
{
"prompt": "Aerial shot of a beach shore with sea waves. Big rocks on the sand at an alone beach.",
"image_path": "images/mixkit-aerial-shot-of-a-beach-with-sea-waves-1087.png"
},
{
"prompt": "Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.",
"image_path": "images/mixkit-woman-cleaning-her-house-dancing-happy-43379.png"
},
{
"prompt": "Aerial tour in a meadow surrounded by hills on the horizon, while some birds fly low over a lake.",
"image_path": "images/mixkit-birds-flying-low-over-a-lake-in-a-meadow-41417.png"
},
{
"prompt": "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.",
"image_path": "images/mixkit-a-rancher-riding-a-horse-at-sunset-1143.png"
},
{
"prompt": "In the video, a martial artist dressed in a traditional white uniform with a black belt demonstrates a series of precise movements against a stark black background. The individual gracefully transitions between stances, embodying a sense of focused discipline and control. Each motion is executed with a deliberate pace, showcasing the fluidity of martial arts techniques. The soft lighting creates subtle highlights on the uniform, adding depth to the figure as it moves. The practitioner begins with an open-hand pose, feet firmly grounded, gradually shifting to a powerful forward punch. The fluidity of the sequence displays a mastery of balance and poise. Every trajectory of the limbs is precise and deliberate, capturing the elegance and strength of martial arts. The serene, isolated setting enhances the intensity and concentration of the practitioner. This visual presentation is an elegant interplay of motion and stillness, displaying the art form's discipline and grace.",
"image_path": "images/mixkit-a-young-man-practicing-his-karate-moves-49635.png"
},
{
"prompt": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
"image_path": "images/mixkit-small-group-of-people-doing-yoga-together-43730.png"
},
{
"prompt": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
"image_path": "images/mixkit-dolphins-underwater-4133.png"
},
{
"prompt": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
"image_path": "images/mixkit-skiers-on-a-snowy-slope-3327.png"
},
{
"prompt": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
"image_path": "images/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306.png"
},
{
"prompt": "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.",
"image_path": "images/mixkit-curve-on-a-snowy-forest-road-3317.png"
},
{
"prompt": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
"image_path": "images/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996.png"
},
{
"prompt": "A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.",
"image_path": "images/gray_short_man.jpg"
},
{
"prompt": "Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.",
"image_path": "images/peninsula.jpg"
},
{
"prompt": "Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.",
"image_path": "images/cyclist.jpg"
},
{
"prompt": "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.",
"image_path": "images/friends.jpg"
},
{
"prompt": "A saxophonist wearing a blazer dances while playing a song in a park.",
"image_path": "images/saxophonist.jpg"
},
{
"prompt": "Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.",
"image_path": "images/romance.jpg"
},
{
"prompt": "Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.",
"image_path": "images/80s_dance.jpg"
},
{
"prompt": "Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.",
"image_path": "images/jazz.jpg"
},
{
"prompt": "A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.",
"image_path": "images/pink.jpg"
},
{
"prompt": "Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.",
"image_path": "images/natural.jpg"
},
{
"prompt": "Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.",
"image_path": "images/couple.jpg"
}
]
+2 -1
View File
@@ -7,7 +7,8 @@ mkdocstrings-python>=1.8.0
mkdocs-mermaid2-plugin>=1.1.0
mkdocs-git-revision-date-localized-plugin>=1.2.0
mkdocs-git-committers-plugin-2>=1.1.0
mkdocs-macros-plugin>=0.8.0
mkdocs-macros-plugin>=0.8.0
pymdown-extensions>=10.0
mkdocs-api-autonav
mkdocs-autorefs
mdx-truly-sane-lists
@@ -4,17 +4,37 @@ import os
import random
def generate_merged_validation_json(input_dir, output_file):
# read in video2caption.json
with open(os.path.join(input_dir, "video2caption_replace.json"), "r") as f:
def generate_merged_validation_json(args):
input_file = args.input_file
output_validation_file = args.output_validation_file
if args.output_train_file:
output_train_file = args.output_train_file
else:
base, ext = os.path.splitext(input_file)
output_train_file = f"{base}_train{ext}"
# read in input json
print(f"Reading from {input_file}")
with open(input_file, "r") as f:
video2caption = json.load(f)
# count how many elements are in the list
num_elements = len(video2caption)
print(f"Number of elements in video2caption.json: {num_elements}")
print(f"Number of elements in input file: {num_elements}")
# randomly sample 64 elements from the list
sampled_elements = random.sample(video2caption, 64)
# randomly sample elements from the list
num_sample = min(args.num_elements, num_elements)
indices = set(random.sample(range(num_elements), num_sample))
sampled_elements = []
remaining_elements = []
for i in range(num_elements):
if i in indices:
sampled_elements.append(video2caption[i])
else:
remaining_elements.append(video2caption[i])
# Transform sampled elements into validation.json format
validation_data = []
@@ -23,10 +43,10 @@ def generate_merged_validation_json(input_dir, output_file):
validation_entry = {
"caption": element["cap"],
"video_path": element.get("path", ""),
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
"num_inference_steps": args.num_inference_steps,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames
}
validation_data.append(validation_entry)
@@ -36,18 +56,25 @@ def generate_merged_validation_json(input_dir, output_file):
}
# Write the validation JSON to the output file
with open(output_file, "w") as f:
with open(output_validation_file, "w") as f:
json.dump(validation_json, f, indent=2)
print(f"Generated validation JSON with {len(validation_data)} entries and saved to {output_file}")
print(f"Generated validation JSON with {len(validation_data)} entries and saved to {output_validation_file}")
# Write the remaining JSON to the output train file
with open(output_train_file, "w") as f:
json.dump(remaining_elements, f, indent=2)
print(f"Saved remaining {len(remaining_elements)} entries to {output_train_file}")
def main():
parser = argparse.ArgumentParser()
# dataset_type: "mixkit"
# dataset_type: "merged"
parser.add_argument("--dataset_type", choices=["merged"], required=True)
parser.add_argument("--input_dir", type=str, required=True)
parser.add_argument("--output_file", type=str, required=True)
parser.add_argument("--input_file", type=str, required=True, help="Path to input json file")
parser.add_argument("--output_validation_file", type=str, required=True, help="Path to output validation json file")
parser.add_argument("--output_train_file", type=str, help="Path to output train json file (remaining data). Defaults to {input_filename}_train.json")
parser.add_argument("--num_elements", type=int, default=64)
parser.add_argument("--num_frames", type=int, default=77)
parser.add_argument("--height", type=int, default=480)
@@ -56,8 +83,8 @@ def main():
args = parser.parse_args()
if args.dataset_type == "merged":
generate_merged_validation_json(args.input_dir, args.output_file)
generate_merged_validation_json(args)
if __name__ == "__main__":
main()
main()