Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9d95b1c10 | ||
|
|
8b55e9706c | ||
|
|
3ff640b2e6 | ||
|
|
c722429ab5 | ||
|
|
e04a192de6 | ||
|
|
c9ca6d1298 | ||
|
|
754292c419 | ||
|
|
0082bc66fc | ||
|
|
8b1937422e | ||
|
|
fb6cbf23e6 | ||
|
|
c8fdd5ed7b | ||
|
|
1c19a6a00c | ||
|
|
d44409c704 | ||
|
|
5d1c7852b7 | ||
|
|
77a211d006 | ||
|
|
bef8169bb1 | ||
|
|
681f1583f9 | ||
|
|
e3b4564d5a | ||
|
|
c0d03fc43d | ||
|
|
404ee8538e | ||
|
|
e57ac59462 | ||
|
|
8c55fdaf7e |
@@ -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..."
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -10,6 +10,7 @@ exclude: |
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
|
||||
@@ -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. [](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. [](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. [](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. [](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. [](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. [](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. [](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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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__"]
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
After Width: | Height: | Size: 113 KiB |
|
After Width: | Height: | Size: 229 KiB |
|
After Width: | Height: | Size: 168 KiB |
|
After Width: | Height: | Size: 103 KiB |
|
After Width: | Height: | Size: 148 KiB |
|
After Width: | Height: | Size: 155 KiB |
|
After Width: | Height: | Size: 723 KiB |
|
After Width: | Height: | Size: 723 KiB |
|
After Width: | Height: | Size: 875 KiB |
|
After Width: | Height: | Size: 664 KiB |
|
After Width: | Height: | Size: 62 KiB |
|
After Width: | Height: | Size: 686 KiB |
|
After Width: | Height: | Size: 957 KiB |
|
After Width: | Height: | Size: 585 KiB |
|
After Width: | Height: | Size: 558 KiB |
|
After Width: | Height: | Size: 942 KiB |
|
After Width: | Height: | Size: 890 KiB |
|
After Width: | Height: | Size: 433 KiB |
|
After Width: | Height: | Size: 595 KiB |
|
After Width: | Height: | Size: 781 KiB |
|
After Width: | Height: | Size: 783 KiB |
|
After Width: | Height: | Size: 762 KiB |
|
After Width: | Height: | Size: 68 KiB |
|
After Width: | Height: | Size: 147 KiB |
|
After Width: | Height: | Size: 89 KiB |
|
After Width: | Height: | Size: 133 KiB |
|
After Width: | Height: | Size: 213 KiB |
@@ -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
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
@@ -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()
|
||||
|
||||