Compare commits

...
Author SHA1 Message Date
SolitaryThinker 19c1d164c3 fix 2025-12-12 22:51:21 +00:00
William Lin b6fa3d24d8 [misc] update wechat image (#931) 2025-12-11 21:22:21 -08:00
Ketaki Tank 55c2e7cd76 [feat] Add fvd implementation (#923) 2025-12-11 19:06:19 -08:00
Tuyabei 5a549af823 [bugfix] [VSA] Fix block_size computation in backward kernel (#925) 2025-12-10 14:36:40 -08:00
Shreejith SG 92fb660c2e Add LoRA extraction, verification, and comparison scripts (#865) 2025-12-08 16:07:58 -08:00
William Lin 3ff640b2e6 [bigfix] [distillation] Fix DMD inference pipeline noise initialization shape (#921) 2025-12-08 13:00:48 -08:00
William Lin c722429ab5 [docs] fix testing.md visibility (#920) 2025-12-08 00:44:53 -08:00
KyleShaoandKyleS1016 e04a192de6 [feat]: add COSMOS 2.5 DiT implementation (#897)
Co-authored-by: KyleS1016 <kyle.s@gmicloud.ai>
2025-12-07 21:48:32 -08:00
William Lin c9ca6d1298 [docs] add docs for ssim testing (#918) 2025-12-06 18:20:04 -08:00
Wenxuan TanandSolitaryThinker 754292c419 Use assert_close in tests (#429)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-12-06 18:18:25 -08:00
Qi Jia 0082bc66fc fix: correct mp backend GPU assignment on multi-GPU systems (#912) 2025-11-30 23:00:22 -08:00
Ohm-Rishabh 8b1937422e [feat] training mfu calculation scripts (#871) 2025-11-27 16:54:17 -08:00
fb6cbf23e6 Fix the docs (#905)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-11-27 00:37:03 -08:00
Mihir Jagtap c8fdd5ed7b [docs] modified the .github/workflows/docs.yml file to include path filtering (#906) 2025-11-26 17:34:21 -08:00
Loay Rashid 1c19a6a00c [Bugfix] Minor bugfixes (#889) 2025-11-26 17:20:45 -08:00
William Lin d44409c704 [CI] fix VSA training CI (#900) 2025-11-24 17:47:59 -08:00
Zhang Peiyuan 5d1c7852b7 + Awesome work using FastVideo or our research projects (#898) 2025-11-23 22:22:27 -08:00
Wenxuan Tan 77a211d006 [misc] Update wechat link (#893) 2025-11-20 19:59:05 -08:00
Wei Zhou bef8169bb1 [Feat] [I2V] resize all image sizes to below 480*832 (#890) 2025-11-20 00:08:36 -08:00
William Lin 681f1583f9 [readme] update link to inference code (#887) 2025-11-19 13:24:13 -08:00
e3b4564d5a [feat] Add inference for MoE SF (#880)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-11-19 13:16:24 -08:00
Shao Duan c0d03fc43d [bugfix] [lora] [CI] Fix LoRA alpha scaling factor & Fix LoRA Inference CI (#870) 2025-11-19 01:02:01 -08:00
Wei Zhou 404ee8538e [Bugfix] [DMD Distillation] Each rank should have its own timestep sampled (#885) 2025-11-18 14:03:25 -08:00
Shao Duan e57ac59462 Fix mp worker busy loop to handle all string RPC methods (#881) 2025-11-16 13:26:44 -08:00
Mihir Jagtap 8c55fdaf7e [docs] add favicon (#878) 2025-11-15 13:44:16 -08:00
107 changed files with 6154 additions and 1137 deletions
+11
View File
@@ -222,3 +222,14 @@ steps:
- TEST_TYPE=unit_test
agents:
queue: "default"
- path:
- "scripts/lora_extraction/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
+5 -1
View File
@@ -75,7 +75,7 @@ case "$TEST_TYPE" in
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
@@ -126,6 +126,10 @@ case "$TEST_TYPE" in
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
;;
"lora_extraction")
log "Running LoRA extraction tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
+10
View File
@@ -3,8 +3,18 @@ name: Deploy Documentation
on:
push:
branches: [ main ]
paths:
- 'docs/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.txt'
- '.github/workflows/docs.yml'
pull_request:
branches: [ main ]
paths:
- 'docs/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.txt'
- '.github/workflows/docs.yml'
permissions:
contents: read
+1
View File
@@ -10,6 +10,7 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
+13 -16
View File
@@ -1,13 +1,12 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
</div>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/XcY0Cpv" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/c7g1qdD" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -15,6 +14,7 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
</div>
## NEWS
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
@@ -111,24 +111,21 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## 📑 Development Plan
<!-- - More distillation methods -->
<!-- - [ ] Add Distribution Matching Distillation -->
More FastWan Models Coming Soon!
- [ ] Add FastWan2.1-T2V-14B
- [ ] Add FastWan2.2-T2V-14B
- [ ] Add FastWan2.2-I2V-14B
<!-- - Optimization features
- Code updates -->
<!-- - [ ] fp8 support -->
<!-- - [ ] faster load model and save model support -->
## Awesome work using FastVideo or our research projects
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025. [![Star](https://img.shields.io/github/stars/sgl-project/sglang.svg?style=social&label=Star)](https://github.com/sgl-project/sglang)
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/XueZeyue/DanceGRPO.svg?style=social&label=Star)](https://github.com/XueZeyue/DanceGRPO)
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/SRPO.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/SRPO)
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Vchitect/DCM.svg?style=social&label=Star)](https://github.com/Vchitect/DCM)
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/HunyuanVideo-1.5.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch. [![Star](https://img.shields.io/github/stars/kandinskylab/kandinsky-5.svg?style=social&label=Star)](https://github.com/kandinskylab/kandinsky-5)
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention. [![Star](https://img.shields.io/github/stars/meituan-longcat/LongCat-Video.svg?style=social&label=Star)](https://github.com/meituan-longcat/LongCat-Video)
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
## Acknowledgement
We learned and reused code from the following projects:
- [Wan-Video](https://github.com/Wan-Video)
+103
View File
@@ -0,0 +1,103 @@
# FVD (Fréchet Video Distance) Benchmark
Evaluate generated video quality using FVD with the I3D feature extractor.
## Quick Start
**Run the benchmark:**
```bash
bash benchmarks/scripts/run.sh
```
That's it! The script auto-installs dependencies and runs the benchmark.
**To customize:** Edit `benchmarks/fvd/run_fvd.py` to change:
- Video paths (`real_dir`, `gen_dir`)
- Number of videos, frames, sampling strategy
- Device, batch size, caching, etc.
## Advanced Usage (CLI)
For more control without editing Python files, use the CLI.
**First-time setup** (one-time per pod/environment):
```bash
bash benchmarks/scripts/setup_fvd.sh
```
Then run any configuration you want:
```bash
# Custom configuration
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--num-videos 1024 \
--num-frames 32 \
--clip-strategy random \
--batch-size 32 \
--seed 42
```
**Standard protocols:**
```bash
# Use predefined protocols
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
```
**Feature caching** (speed up repeated evaluations):
```bash
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f \
--cache-real-features cache/real # Directory path (will save/load cache/real/real_features.pkl)
```
Run `python -m benchmarks.fvd.cli --help` for all options.
## Available Protocols
- `fvd2048_16f` - Standard (2048 videos, 16 frames)
- `fvd2048_128f` - Long videos (128 frames)
- `fvd2048_128f_subsample8` - Subsampled long videos
- `quick_test` - Fast testing (10 videos)
## Configuration Options
Key options in `FVDConfig`:
```python
num_videos=2048, # Videos to evaluate
num_frames_per_clip=16, # Frames per clip
clip_strategy='beginning', # beginning|random|uniform|middle|sliding
frame_stride=1, # Frame subsampling
batch_size=32, # GPU batch size
device='cuda', # cuda|cpu
cache_real_features=None, # Cache path for speed
seed=42, # Reproducibility
```
## Programmatic Usage
```python
from benchmarks.fvd import compute_fvd_with_config, FVDConfig
config = FVDConfig.fvd2048_16f() # or custom config
results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
print(f"FVD: {results['fvd']:.2f}")
```
## Notes
- I3D model auto-downloads from Hugging Face on first run
- Requires minimum 10 frames per clip
- Supports both video files (.mp4, .avi, etc.) and frame directories
- `--cache-real-features` expects a **directory path** (e.g., `cache/real`), it will automatically create/load `real_features.pkl` inside that directory
+35
View File
@@ -0,0 +1,35 @@
"""
FastVideo Frechet Video Distance (FVD) Benchmark Module.
>>> from fastvideo.benchmarks.fvd import compute_fvd_with_config, FVDConfig
>>> config = FVDConfig.fvd2048_16f() # Standard protocol
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
"""
from .fvd import (
compute_fvd,
compute_fvd_with_config,
compute_frechet_distance,
compute_statistics,
FVDConfig,
)
from .i3d_model import I3DFeatureExtractor
from .video_utils import (
load_video_auto,
sample_clips_from_video,
load_video_clips_streaming,
ClipSamplingStrategy,
)
__all__ = [
'compute_fvd',
'compute_fvd_with_config',
'compute_frechet_distance',
'compute_statistics',
'FVDConfig',
'I3DFeatureExtractor',
'load_video_auto',
'sample_clips_from_video',
'load_video_clips_streaming',
'ClipSamplingStrategy',
]
+185
View File
@@ -0,0 +1,185 @@
import argparse
import json
import sys
from pathlib import Path
from .fvd import compute_fvd_with_config, FVDConfig
def main() -> int:
parser = argparse.ArgumentParser(
description='Compute Fréchet Video Distance (FVD)',
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Standard FVD2048_16f protocol
python -m fastvideo.benchmarks.fvd.cli \\
--real-path data/real/ \\
--gen-path outputs/gen/ \\
--protocol fvd2048_16f
# Custom configuration
python -m fastvideo.benchmarks.fvd.cli \\
--real-path data/real/ \\
--gen-path outputs/gen/ \\
--num-videos 1024 \\
--num-frames 32 \\
--clip-strategy random \\
--frame-stride 2
""")
# Required arguments
parser.add_argument('--real-path',
type=str,
required=True,
help='Path to real videos directory')
parser.add_argument('--gen-path',
type=str,
required=True,
help='Path to generated videos directory')
# Reproducibility
parser.add_argument(
'--seed',
type=int,
default=None,
help='Random seed for reproducibility (np.random, random, torch)')
# Protocol presets
parser.add_argument('--protocol',
type=str,
default=None,
choices=[
'fvd2048_16f', 'fvd2048_128f',
'fvd2048_128f_subsample8', 'quick_test'
],
help='Use standard protocol (overrides other settings)')
# Video selection
parser.add_argument('--num-videos',
type=int,
default=2048,
help='Number of videos to use (default: 2048)')
# Clip sampling
parser.add_argument('--num-frames',
type=int,
default=16,
help='Number of frames per clip (default: 16)')
parser.add_argument('--num-clips',
type=int,
default=1,
help='Number of clips per video (default: 1)')
parser.add_argument(
'--clip-strategy',
type=str,
default='beginning',
choices=['beginning', 'random', 'uniform', 'middle', 'sliding', 'all'],
help='Clip sampling strategy (default: beginning)')
parser.add_argument(
'--frame-stride',
type=int,
default=1,
help='Frame stride for FPS subsampling (default: 1, no subsampling)')
parser.add_argument('--temporal-stride',
type=int,
default=1,
help='Temporal stride for sliding window (default: 1)')
# Data processing
parser.add_argument('--no-frame-dirs',
action='store_true',
help='Disable frame directory support')
# Computation
parser.add_argument('--batch-size',
type=int,
default=32,
help='Batch size for feature extraction (default: 32)')
parser.add_argument('--device',
type=str,
default='cuda',
choices=['cuda', 'cpu'],
help='Device to use (default: cuda)')
# Caching
parser.add_argument('--cache-real-features',
type=str,
default=None,
help='Path to cache real video features')
parser.add_argument('--i3d-model-path',
type=str,
default=None,
help='Custom cache path for I3D model')
# Output
parser.add_argument('--output',
type=str,
default='fvd_results.json',
help='Output JSON file (default: fvd_results.json)')
parser.add_argument('--quiet',
action='store_true',
help='Suppress progress output')
args = parser.parse_args()
# Create config
if args.protocol:
protocol_map = {
'fvd2048_16f': FVDConfig.fvd2048_16f,
'fvd2048_128f': FVDConfig.fvd2048_128f,
'fvd2048_128f_subsample8': FVDConfig.fvd2048_128f_subsample8,
'quick_test': FVDConfig.quick_test,
}
config = protocol_map[args.protocol]()
# Override device and caching from args
config.device = args.device
config.cache_real_features = args.cache_real_features
config.i3d_model_path = args.i3d_model_path
config.batch_size = args.batch_size
config.seed = args.seed
else:
# Custom config from args
config = FVDConfig(num_videos=args.num_videos,
num_frames_per_clip=args.num_frames,
num_clips_per_video=args.num_clips,
clip_strategy=args.clip_strategy,
frame_stride=args.frame_stride,
temporal_stride=args.temporal_stride,
support_frame_dirs=not args.no_frame_dirs,
batch_size=args.batch_size,
device=args.device,
cache_real_features=args.cache_real_features,
i3d_model_path=args.i3d_model_path,
seed=args.seed)
# Compute FVD
try:
results = compute_fvd_with_config(real_videos=args.real_path,
gen_videos=args.gen_path,
config=config,
verbose=not args.quiet)
# Save results
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w') as f:
json.dump(results, f, indent=2)
print(f"\nResults saved to {output_path}")
print(f"FVD: {results['fvd']:.2f}")
print(f"Protocol: {results['protocol']}")
return 0
except Exception as e:
print(f"Error: {e}", file=sys.stderr)
import traceback
traceback.print_exc()
return 1
if __name__ == '__main__':
sys.exit(main())
+447
View File
@@ -0,0 +1,447 @@
import numpy as np
import scipy.linalg
import torch
from pathlib import Path
from collections.abc import Iterator
import pickle
from dataclasses import dataclass, field
from .i3d_model import I3DFeatureExtractor
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
def compute_statistics(features: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
"""Compute mean and covariance."""
mu = np.mean(features, axis=0)
sigma = np.cov(features, rowvar=False)
return mu, sigma
def compute_frechet_distance(mu1: np.ndarray,
sigma1: np.ndarray,
mu2: np.ndarray,
sigma2: np.ndarray,
eps: float = 1e-6) -> float:
"""
Compute Fréchet distance between two Gaussians.
"""
sigma1 = sigma1 + eps * np.eye(sigma1.shape[0])
sigma2 = sigma2 + eps * np.eye(sigma2.shape[0])
diff = mu1 - mu2
mean_distance = np.sum(diff**2)
trace_sum = np.trace(sigma1 + sigma2)
covmean = scipy.linalg.sqrtm(sigma1 @ sigma2)
if np.iscomplexobj(covmean):
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
print(
f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
covmean = covmean.real
trace_product = np.trace(covmean)
fvd = mean_distance + trace_sum - 2 * trace_product
return float(fvd)
@dataclass
class FVDConfig:
# default configuration for FVD computation:
# Video selection
num_videos: int = 2048
# Clip sampling
num_frames_per_clip: int = 16
num_clips_per_video: int = 1
clip_strategy: str | ClipSamplingStrategy = 'beginning'
# Temporal subsampling
frame_stride: int = 1 # 1=no subsampling, 2=every 2nd, 8=every 8th
temporal_stride: int = 1 # For sliding window clips
# Data processing
video_extensions: list[str] = field(
default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
support_frame_dirs: bool = True
# Computation
batch_size: int = 32
device: str = 'cuda'
use_streaming: bool = True
resize_before_extraction: bool = True
# Caching
cache_real_features: str | None = None
i3d_model_path: str | None = None
# Reproducibility
seed: int | None = None
@classmethod
def fvd2048_16f(cls) -> 'FVDConfig':
"""
Standard FVD protocol: 2048 videos, 16 frames, beginning clip.
most common FVD configuration used in papers
"""
return cls(num_videos=2048,
num_frames_per_clip=16,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def fvd2048_128f(cls) -> 'FVDConfig':
"""Long video protocol: 2048 videos, 128 frames."""
return cls(num_videos=2048,
num_frames_per_clip=128,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def fvd2048_128f_subsample8(cls) -> 'FVDConfig':
"""
Long video with FPS subsampling: 2048 videos, 128 frames (every 8th).
Used for very long videos - samples every 8th frame
"""
return cls(num_videos=2048,
num_frames_per_clip=16,
frame_stride=8,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def quick_test(cls) -> 'FVDConfig':
"""Quick test config: 100 videos, 16 frames."""
return cls(num_videos=100,
num_frames_per_clip=16,
clip_strategy='beginning')
def to_dict(self) -> dict:
"""Export config to dict for logging"""
return {
'num_videos': self.num_videos,
'num_frames_per_clip': self.num_frames_per_clip,
'num_clips_per_video': self.num_clips_per_video,
'clip_strategy': str(self.clip_strategy),
'frame_stride': self.frame_stride,
'temporal_stride': self.temporal_stride,
'batch_size': self.batch_size,
'device': self.device,
'seed': self.seed,
'use_streaming': self.use_streaming,
}
def __str__(self) -> str:
"""Human-readable protocol name"""
desc = f"FVD{self.num_videos}_{self.num_frames_per_clip}f"
if self.frame_stride > 1:
desc += f"_subsample{self.frame_stride}"
if self.num_clips_per_video > 1:
desc += f"_{self.num_clips_per_video}clips"
if self.clip_strategy != 'beginning':
desc += f"_{self.clip_strategy}"
return desc
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
extractor: I3DFeatureExtractor,
batch_size: int = 32,
max_clips: int | None = None,
verbose: bool = True) -> np.ndarray:
"""
Extract features from a video clip generator using streaming.
Args:
video_generator: Iterator yielding clips [T, C, H, W]
extractor: I3D feature extractor
batch_size: Batch size for processing
max_clips: Maximum clips to process (for validation)
verbose: Show progress
Returns:
features: [N, 400] numpy array
"""
all_features = []
batch = []
clip_count = 0
if verbose:
print(f"Extracting features with batch_size={batch_size}...")
for clip_count, clip in enumerate(video_generator):
batch.append(clip)
# Process batch when full
if len(batch) == batch_size:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features(batch_tensor,
batch_size=batch_size,
verbose=False)
all_features.append(features.cpu().numpy())
batch = [] # Clear batch
if verbose and clip_count % (batch_size * 10) == 0:
print(f"Processed {clip_count} clips...")
# Stop if we've reached max_clips
if max_clips is not None and clip_count >= max_clips:
break
# Process remaining clips
if len(batch) > 0:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features(batch_tensor,
batch_size=len(batch),
verbose=False)
all_features.append(features.cpu().numpy())
if len(all_features) == 0:
raise RuntimeError("No features extracted - check video loading")
features = np.concatenate(all_features, axis=0)
if verbose:
print(f"Extracted {len(features)} feature vectors")
return features
def load_or_compute_features(videos: str | Path | torch.Tensor,
extractor: I3DFeatureExtractor,
config: FVDConfig,
cache_path: str | None = None,
cache_name: str = "real_features") -> np.ndarray:
"""Load features from cache or compute (with streaming support)"""
if cache_path is not None:
cache_file = Path(cache_path) / f"{cache_name}.pkl"
if cache_file.exists():
print(f"Loading cached features from {cache_file}")
with open(cache_file, 'rb') as f:
features = pickle.load(f)
# Validate and limit based on config
max_features = config.num_videos * config.num_clips_per_video
if len(features) < max_features:
print(
f"WARNING: Cache has {len(features)} features but need {max_features}"
)
print("Recomputing features...")
elif len(features) > max_features:
features = features[:max_features]
return features
else:
return features
# Compute features
if isinstance(videos, str | Path):
target_size = (224, 224) if config.resize_before_extraction else None
video_generator = load_video_clips_streaming(
videos,
num_frames=config.num_frames_per_clip,
max_videos=config.num_videos,
clip_strategy=config.clip_strategy,
frame_stride=config.frame_stride,
num_clips_per_video=config.num_clips_per_video,
video_extensions=config.video_extensions,
support_frame_dirs=config.support_frame_dirs,
target_size=target_size,
verbose=True)
max_clips = config.num_videos * config.num_clips_per_video
features = extract_features_streaming(video_generator,
extractor,
batch_size=config.batch_size,
max_clips=max_clips,
verbose=True)
else:
# Already a tensor
print(f"Extracting features from {len(videos)} video tensors...")
features = extractor.extract_features(videos,
batch_size=config.batch_size,
verbose=True)
features = features.numpy()
# Validate feature count
expected_count = config.num_videos * config.num_clips_per_video
if len(features) < expected_count:
raise ValueError(
f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
f"Found fewer videos than expected. Check your video directory.")
elif len(features) > expected_count:
print(f"Truncating {len(features)} features to {expected_count}")
features = features[:expected_count]
# Cache features if requested
if cache_path is not None:
cache_dir = Path(cache_path)
cache_dir.mkdir(parents=True, exist_ok=True)
cache_file = cache_dir / f"{cache_name}.pkl"
print(f"Caching features to {cache_file}")
with open(cache_file, 'wb') as f:
pickle.dump(features, f)
return features
def compute_fvd(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
num_frames: int = 16,
batch_size: int = 32,
device: str = 'cuda',
num_videos: int | None = 2048,
cache_real_features: str | None = None,
i3d_model_path: str | None = None,
seed: int | None = None,
verbose: bool = True) -> float:
"""
Compute Fréchet Video Distance (FVD)
For advanced control, use compute_fvd_with_config() instead.
Args:
real_videos: Path to real videos or tensor [N, T, C, H, W]
gen_videos: Path to generated videos or tensor [N, T, C, H, W]
num_frames: Frames per video (default: 16)
batch_size: Batch size (default: 32)
device: 'cuda' or 'cpu' (default: 'cuda')
num_videos: Max videos (default: 2048)
cache_real_features: Cache path for real features
i3d_model_path: Custom I3D model cache path
seed: Random seed for reproducibility
verbose: Print progress
Returns:
FVD score (float). Lower is better.
"""
num_videos = num_videos if num_videos is not None else 2048
config = FVDConfig(
num_videos=num_videos,
num_frames_per_clip=num_frames,
batch_size=batch_size,
device=device,
cache_real_features=cache_real_features,
i3d_model_path=i3d_model_path,
seed=seed,
)
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
return result['fvd']
def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
config: FVDConfig,
verbose: bool = True) -> dict:
"""
Compute FVD using a standardized configuration.
This is the recommended way to compute FVD for reproducibility.
Args:
real_videos: Path or tensors
gen_videos: Path or tensors
config: FVDConfig specifying protocol
verbose: Print progress
Returns:
results: Dictionary with:
- 'fvd': FVD score (float)
- 'protocol': Protocol name (str)
- 'config': Configuration dict
Example:
>>> config = FVDConfig.fvd2048_16f()
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
>>> print(f"Protocol: {results['protocol']}") # "FVD2048_16f"
"""
# Seed for reproducibility
if config.seed is not None:
import random as _rnd
_rnd.seed(config.seed)
np.random.seed(config.seed)
torch.manual_seed(config.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(config.seed)
if verbose:
print("=" * 70)
print(f"Computing FVD with protocol: {config}")
print("=" * 70)
print("\nConfiguration:")
for key, value in config.to_dict().items():
print(f" {key}: {value}")
print()
# Initialize I3D
if verbose:
print(f"\nInitializing I3D model on {config.device}...")
extractor = I3DFeatureExtractor(device=config.device,
cache_dir=config.i3d_model_path)
# Extract features
if verbose:
print(f"\n{'='*70}")
print("Extracting REAL video features...")
print(f"{'='*70}")
real_features = load_or_compute_features(
videos=real_videos,
extractor=extractor,
config=config,
cache_path=config.cache_real_features,
cache_name="real_features")
if verbose:
print(f"\n{'='*70}")
print("Extracting GENERATED video features...")
print(f"{'='*70}")
gen_features = load_or_compute_features(videos=gen_videos,
extractor=extractor,
config=config,
cache_path=None,
cache_name="gen_features")
if verbose:
print(f"\nReal videos/clips: {len(real_features)}")
print(f"Generated videos/clips: {len(gen_features)}")
print(f"\n{'='*70}")
print("Computing statistics...")
print(f"{'='*70}")
mu_real, sigma_real = compute_statistics(real_features)
mu_gen, sigma_gen = compute_statistics(gen_features)
if verbose:
print(f"\n{'='*70}")
print("Computing Fréchet distance...")
print(f"{'='*70}")
fvd = compute_frechet_distance(mu_real, sigma_real, mu_gen, sigma_gen)
if verbose:
print(f"\n{'='*70}")
print(f"FVD Score: {fvd:.4f}")
print(f"Protocol: {config}")
print(f"{'='*70}\n")
results = {
'fvd': fvd,
'protocol': str(config),
'config': config.to_dict(),
}
return results
+142
View File
@@ -0,0 +1,142 @@
"""I3D Feature Extractor for FVD Computation"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from pathlib import Path
from huggingface_hub import hf_hub_download
from tqdm import tqdm
from contextlib import suppress
class I3DFeatureExtractor(nn.Module):
"""
I3D feature extractor for FVD computation.
Extracts 400-dimensional features from videos using I3D model
trained on Kinetics-400.
"""
REPO_ID = 'flateon/FVD-I3D-torchscript'
MODEL_FILENAME = 'i3d_torchscript.pt'
def __init__(self,
device: str = 'cuda',
cache_dir: str | Path | None = None):
super().__init__()
self.device_str = device
if device == 'cuda' and not torch.cuda.is_available():
print(
"Warning: CUDA requested but not available – falling back to CPU"
)
self.device = torch.device('cpu')
else:
self.device = torch.device(device)
self.cache_dir: str | None
if cache_dir is not None:
self.cache_dir = str(Path(cache_dir).resolve())
else:
self.cache_dir = None # Use HF default cache
self.model = self._load_model()
self.model.eval()
with suppress(Exception):
self.model.to(self.device)
def _load_model(self) -> torch.nn.Module:
"""Download and load I3D TorchScript model from Hugging Face Hub."""
print(f"Loading I3D model from Hugging Face Hub ({self.REPO_ID})...")
try:
# Download model from Hugging Face Hub
model_path = hf_hub_download(repo_id=self.REPO_ID,
filename=self.MODEL_FILENAME,
cache_dir=self.cache_dir)
# Load directly to chosen device
model = torch.jit.load(model_path, map_location=self.device)
print("I3D model loaded successfully")
return model
except Exception as e:
raise RuntimeError(
f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
f"Ensure you have internet connection and huggingface_hub installed:\n"
f"pip install huggingface_hub") from e
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Preprocess videos for I3D.
Args:
videos: [B, T, C, H, W], values in [0, 255]
Returns:
Preprocessed videos [B, C, T, 224, 224] (normalized and resized)
"""
B, T, C, H, W = videos.shape
if T < 10:
raise ValueError(f"I3D requires at least 10 frames, got {T}")
# Normalize to [0, 1] if needed
if videos.max() > 1.0:
videos = videos / 255.0
# Resize to 224x224 if needed
if H != 224 or W != 224:
videos = videos.reshape(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.reshape(B, T, C, 224, 224)
# Convert to [B, C, T, H, W] format
videos = videos.permute(0, 2, 1, 3, 4).contiguous()
return videos
@torch.no_grad()
def extract_features(self,
videos: torch.Tensor,
batch_size: int = 32,
verbose: bool = True) -> torch.Tensor:
"""
Extract I3D features
Args:
videos: [N, T, C, H, W], values in [0, 255]
batch_size: Batch size for processing
verbose: Show progress bar
Returns:
Features [N, 400]
"""
N = len(videos)
all_features = []
iterator = range(0, N, batch_size)
if verbose:
iterator = tqdm(iterator, desc="Extracting I3D features")
for i in iterator:
batch = videos[i:i + batch_size].to(self.device)
batch = self.preprocess(batch) # Now returns [B, C, T, H, W]
# Use the HF model without rescale/resize (we handle it in preprocess)
features = self.model(batch,
rescale=False,
resize=False,
return_features=True)
all_features.append(features.cpu())
return torch.cat(all_features, dim=0)
def __call__(self,
videos: torch.Tensor,
batch_size: int = 32) -> torch.Tensor:
return self.extract_features(videos, batch_size=batch_size)
+34
View File
@@ -0,0 +1,34 @@
import sys
from pathlib import Path
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config
root_dir = Path(__file__).parent.parent.parent
sys.path.insert(0, str(root_dir))
def main() -> None:
# Get script directory
script_dir = Path(__file__).parent.resolve()
clip_strategy = 'beginning' # Options: 'uniform', 'random', 'beginning', 'end', 'all'
cfg = FVDConfig(
num_videos=650,
num_frames_per_clip=16,
num_clips_per_video=1,
clip_strategy=clip_strategy,
frame_stride=1,
batch_size=32,
device='cuda',
seed=42,
cache_real_features=str(script_dir / f'fvd-cache/{clip_strategy}'),
)
real_dir = "benchmarks/data/real_videos"
gen_dir = "benchmarks/data/generated_videos"
results = compute_fvd_with_config(real_dir, gen_dir, cfg, verbose=True)
print(f"FVD = {results['fvd']:.2f}")
if __name__ == '__main__':
main()
+97
View File
@@ -0,0 +1,97 @@
#!/usr/bin/env python3
import sys
from pathlib import Path
import shutil
import random
from fvd import compute_fvd_with_config, FVDConfig
script_path = Path(__file__).resolve()
fastvideo_root = script_path.parent.parent.parent
sys.path.insert(0, str(fastvideo_root))
def split_videos(video_dir: Path, n_per_subset: int = 128, seed: int = 42):
subset_a = video_dir.parent / 'bair_full_subset_A'
subset_b = video_dir.parent / 'bair_full_subset_B'
if subset_a.exists():
shutil.rmtree(subset_a)
if subset_b.exists():
shutil.rmtree(subset_b)
subset_a.mkdir(parents=True)
subset_b.mkdir(parents=True)
videos = sorted(video_dir.glob('*.mp4'))
random.seed(seed)
shuffled = list(videos)
random.shuffle(shuffled)
needed = n_per_subset * 2
if len(shuffled) > needed:
shuffled = shuffled[:needed]
mid = len(shuffled) // 2
print(f"\nSplitting {len(shuffled)} BAIR FULL videos:")
print(f" Subset A: {mid} videos")
print(f" Subset B: {len(shuffled) - mid} videos")
for v in shuffled[:mid]:
shutil.copy2(v, subset_a / v.name)
for v in shuffled[mid:]:
shutil.copy2(v, subset_b / v.name)
return subset_a, subset_b, mid
def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
config = FVDConfig(num_videos=num_videos,
num_frames_per_clip=16,
clip_strategy='beginning',
batch_size=8,
device='cuda',
seed=42)
print("\n" + "=" * 70)
print("TEST 1: Identity Test")
print("=" * 70)
result1 = compute_fvd_with_config(real_videos=str(subset_a),
gen_videos=str(subset_a),
config=config,
verbose=False)
fvd_identity = result1['fvd']
print(f"\nIdentity FVD: {fvd_identity:.2f}")
print("\n" + "=" * 70)
print("TEST 2: Real vs Real")
print("=" * 70)
result2 = compute_fvd_with_config(real_videos=str(subset_a),
gen_videos=str(subset_b),
config=config,
verbose=False)
fvd_real = result2['fvd']
print(f"\nReal vs Real FVD: {fvd_real:.2f}")
print("\n" + "=" * 70)
print("RESULTS")
print("=" * 70)
print(f"Identity: {fvd_identity:.2f}")
print(f"Real vs Real: {fvd_real:.2f}")
def main() -> None:
bair_dir = Path('benchmarks/data/bair_full_videos')
subset_a, subset_b, count = split_videos(bair_dir,
n_per_subset=128,
seed=42)
validate_fvd(subset_a, subset_b, count)
if __name__ == '__main__':
main()
+490
View File
@@ -0,0 +1,490 @@
import torch
import cv2
import numpy as np
from pathlib import Path
from collections.abc import Iterator
from tqdm import tqdm
from enum import Enum
class ClipSamplingStrategy(Enum):
"""Clip sampling strategies for FVD evaluation."""
BEGINNING = 'beginning' # Take first N frames (most common)
RANDOM = 'random' # Random N consecutive frames
UNIFORM = 'uniform' # Uniformly spaced frames across video
MIDDLE = 'middle' # Middle N frames
SLIDING = 'sliding' # Multiple sliding windows
ALL = 'all' # All possible clips
def _load_video_cv2(video_path: str | Path,
num_frames: int | None = 16,
sample_strategy: str = 'uniform') -> torch.Tensor:
"""
Load video from video file using OpenCV.
Args:
video_path: Path to video file (MP4, AVI, MOV, MKV)
num_frames: Number of frames to extract
sample_strategy: 'uniform' or 'random'
Returns:
video: [T, C, H, W]
"""
video_path = str(video_path)
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise RuntimeError(f"Cannot open video: {video_path}")
frames = []
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
if num_frames is None:
# Read all available frames
while True:
ret, frame = cap.read()
if not ret:
break
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
cap.release()
if len(frames) == 0:
raise RuntimeError(f"Video has 0 frames: {video_path}")
frames = np.stack(frames) # [T, H, W, C]
frames = torch.from_numpy(frames).permute(0, 3, 1,
2).float() # [T, C, H, W]
return frames
if total_frames == 0:
raise RuntimeError(f"Video has 0 frames: {video_path}")
# Determine frame indices for sampling
if total_frames < num_frames:
frame_indices = list(range(
total_frames)) + [total_frames - 1] * (num_frames - total_frames)
elif sample_strategy == 'uniform':
frame_indices = np.linspace(0, total_frames - 1, num_frames,
dtype=int).tolist()
elif sample_strategy == 'random':
frame_indices = sorted(
np.random.choice(total_frames, num_frames, replace=False))
else:
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
# Extract frames
for idx in frame_indices:
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
ret, frame = cap.read()
if not ret:
if len(frames) > 0:
frames.append(frames[-1].copy())
else:
h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
frames.append(np.zeros((h, w, 3), dtype=np.uint8))
continue
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
cap.release()
frames = np.stack(frames) # [T, H, W, C]
frames = torch.from_numpy(frames).permute(0, 3, 1,
2).float() # [T, C, H, W]
return frames
def _load_video_from_frames(
frame_dir: str | Path,
num_frames: int | None = 16,
sample_strategy: str = 'uniform',
frame_extensions: list[str] | None = None) -> torch.Tensor:
"""
Load video from directory of frame images.
Args:
frame_dir: Directory containing frames
num_frames: Number of frames to sample
sample_strategy: 'uniform' or 'random'
frame_extensions: Image file extensions to look for
Returns:
video: [T, C, H, W]
"""
if frame_extensions is None:
frame_extensions = ['.jpg', '.png', '.jpeg', '.bmp']
frame_dir = Path(frame_dir)
if not frame_dir.exists():
raise FileNotFoundError(f"Frame directory not found: {frame_dir}")
# Find all frames
frame_files: list[Path] = []
for ext in frame_extensions:
frame_files.extend(frame_dir.glob(f"*{ext}"))
if len(frame_files) == 0:
raise ValueError(
f"No frames found in {frame_dir} with extensions {frame_extensions}"
)
frame_files = sorted(frame_files, key=lambda x: x.name)
total_frames = len(frame_files)
# Determine frame indices
if num_frames is None:
frame_indices = list(range(total_frames))
else:
if total_frames < num_frames:
frame_indices = list(range(total_frames)) + [total_frames - 1] * (
num_frames - total_frames)
elif sample_strategy == 'uniform':
frame_indices = np.linspace(0,
total_frames - 1,
num_frames,
dtype=int).tolist()
elif sample_strategy == 'random':
frame_indices = sorted(
np.random.choice(total_frames, num_frames, replace=False))
else:
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
# Load frames
frames = []
for idx in frame_indices:
frame_path = frame_files[idx]
frame = cv2.imread(str(frame_path))
if frame is None:
raise RuntimeError(f"Failed to load frame: {frame_path}")
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
# Stack and convert to tensor
frames = np.stack(frames) # [T, H, W, C]
frames = torch.from_numpy(frames).permute(0, 3, 1,
2).float() # [T, C, H, W]
return frames
def _detect_video_format(path: str | Path) -> str:
"""
Detect if path is a video file or frame directory.
Returns:
'video_file', 'frame_directory', or 'unknown'
"""
path = Path(path)
if path.is_file():
return 'video_file'
elif path.is_dir():
# Check if contains image files
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
for ext in image_extensions:
if list(path.glob(f"*{ext}")):
return 'frame_directory'
return 'unknown'
else:
raise ValueError(f"Path does not exist: {path}")
def load_video_auto(video_path: str | Path,
num_frames: int | None = 16,
sample_strategy: str = 'uniform') -> torch.Tensor:
"""
Automatically detect format and load video.
Supports:
- Video files (MP4, AVI, MOV, MKV)
- Frame directories (JPG, PNG)
Args:
video_path: Path to video file or frame directory
num_frames: Number of frames to extract
sample_strategy: 'uniform' or 'random'
Returns:
video: [T, C, H, W]
"""
format_type = _detect_video_format(video_path)
if format_type == 'video_file':
return _load_video_cv2(video_path, num_frames, sample_strategy)
elif format_type == 'frame_directory':
return _load_video_from_frames(video_path, num_frames, sample_strategy)
else:
raise ValueError(f"Unknown video format at {video_path}")
def sample_clips_from_video(
video: torch.Tensor,
num_frames_per_clip: int = 16,
num_clips: int = 1,
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
frame_stride: int = 1,
temporal_stride: int = 1) -> list[torch.Tensor]:
"""
Sample clips from a video with various strategies.
Args:
video: [T, C, H, W] full video
num_frames_per_clip: Frames per clip
num_clips: Number of clips to extract
strategy: ClipSamplingStrategy or string ('beginning', 'random', etc.)
frame_stride: Skip frames (FPS control: 1=all, 2=every 2nd, 8=every 8th)
temporal_stride: Stride between clips for sliding window
Returns:
List of clips, each [num_frames_per_clip, C, H, W]
Examples:
>>> # Beginning clip (most common for FVD)
>>> clips = sample_clips_from_video(video, 16, strategy='beginning')
>>> # Multiple random clips
>>> clips = sample_clips_from_video(video, 16, num_clips=4, strategy='random')
>>> # Subsample FPS by 2x (every 2nd frame)
>>> clips = sample_clips_from_video(video, 16, frame_stride=2)
>>> # Sliding window with overlap
>>> clips = sample_clips_from_video(video, 16, strategy='sliding', temporal_stride=8)
"""
# Convert string to enum if needed
if isinstance(strategy, str):
strategy = ClipSamplingStrategy(strategy)
T, C, H, W = video.shape
# Apply frame stride (FPS subsampling)
if frame_stride > 1:
video = video[::frame_stride]
T = len(video)
effective_clip_length = num_frames_per_clip
# Handle videos shorter than clip length
if effective_clip_length > T:
pad_length = effective_clip_length - T
last_frame = video[-1:].repeat(pad_length, 1, 1, 1)
video = torch.cat([video, last_frame], dim=0)
T = len(video)
clips = []
if strategy == ClipSamplingStrategy.BEGINNING:
# Take first clip (most common for FVD evaluation)
clip = video[:effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.MIDDLE:
# Take middle clip
start = (T - effective_clip_length) // 2
clip = video[start:start + effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.RANDOM:
# Sample N random clips
for _ in range(num_clips):
if effective_clip_length == T:
start = 0
else:
start = np.random.randint(0, T - effective_clip_length + 1)
clip = video[start:start + effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.UNIFORM:
# Uniformly spaced clips
if num_clips == 1:
# Single clip from middle
start = (T - effective_clip_length) // 2
clip = video[start:start + effective_clip_length]
clips.append(clip)
else:
# Multiple uniformly spaced clips
step = (T - effective_clip_length) / (num_clips -
1) if num_clips > 1 else 0
for i in range(num_clips):
start = int(i * step)
start = min(start, T - effective_clip_length)
clip = video[start:start + effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.SLIDING:
# Sliding window with stride
for start in range(0, T - effective_clip_length + 1, temporal_stride):
clip = video[start:start + effective_clip_length]
clips.append(clip)
if len(clips) >= num_clips:
break
elif strategy == ClipSamplingStrategy.ALL:
# All possible clips (overlapping)
for start in range(T - effective_clip_length + 1):
clip = video[start:start + effective_clip_length]
clips.append(clip)
else:
raise ValueError(f"Unknown strategy: {strategy}")
return clips
def load_video_clips_streaming(directory: str | Path,
num_frames: int = 16,
max_videos: int | None = None,
clip_strategy: str
| ClipSamplingStrategy = 'beginning',
frame_stride: int = 1,
num_clips_per_video: int = 1,
video_extensions: list[str] | None = None,
support_frame_dirs: bool = True,
target_size: tuple[int, int] | None = (224, 224),
verbose: bool = True) -> Iterator[torch.Tensor]:
"""
This generator yields clips one-by-one instead of loading all videos into RAM.
Perfect for large datasets where memory is limited.
Args:
directory: Path to directory with videos
num_frames: Frames per clip
max_videos: Max videos to load
clip_strategy: 'beginning', 'random', 'uniform', etc.
frame_stride: Frame skip (1=all, 2=every 2nd, 8=every 8th)
num_clips_per_video: Number of clips per video
video_extensions: Video file extensions
support_frame_dirs: Also load frame directories
target_size: Resize clips to (H, W). If None, keep original size.
verbose: Show progress
Yields:
clip: [T, C, H, W] individual clips
Example:
>>> for clip in load_video_clips_streaming('data/videos/', num_frames=16):
>>> features = model.extract_features(clip.unsqueeze(0))
>>> # Process one clip at a time - low memory usage!
"""
if video_extensions is None:
video_extensions = ['.mp4', '.avi', '.mov', '.mkv']
directory = Path(directory)
if not directory.exists():
raise FileNotFoundError(f"Directory not found: {directory}")
# Find video paths
video_paths: list[Path] = []
# Find video files
for ext in video_extensions:
video_paths.extend(directory.glob(f"**/*{ext}"))
# Find frame directories if enabled
if support_frame_dirs:
for subdir in directory.iterdir():
if subdir.is_dir():
# Check if it contains frames
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
for ext in image_extensions:
if list(subdir.glob(f"*{ext}")):
video_paths.append(subdir)
break
if len(video_paths) == 0:
raise ValueError(f"No videos found in {directory}")
video_paths = sorted(video_paths)
if max_videos is not None:
video_paths = video_paths[:max_videos]
if verbose:
print(f"Found {len(video_paths)} videos in {directory}")
if num_clips_per_video > 1:
print(f"Extracting {num_clips_per_video} clips per video...")
if frame_stride > 1:
print(f"Subsampling frames with stride {frame_stride}...")
if target_size:
print(f"Resizing clips to {target_size}...")
# Track statistics
failed_count = 0
total_clips = 0
iterator = tqdm(video_paths,
desc="Loading videos") if verbose else video_paths
for video_path in iterator:
try:
# Load full video
video = load_video_auto(video_path,
num_frames=None,
sample_strategy='uniform')
# Sample clips from video
clips = sample_clips_from_video(video,
num_frames_per_clip=num_frames,
num_clips=num_clips_per_video,
strategy=clip_strategy,
frame_stride=frame_stride)
if target_size is not None:
resized_clips = []
for clip in clips:
T, C, H, W = clip.shape
if target_size != (H, W):
# Resize to target size
clip = clip.contiguous(
) # Fix non-contiguous tensors first
clip_flat = clip.view(T * C, H,
W).unsqueeze(0) # [1, T*C, H, W]
clip_resized = torch.nn.functional.interpolate(
clip_flat,
size=target_size,
mode='bilinear',
align_corners=False)
clip = clip_resized.squeeze(0).view(
T, C, target_size[0],
target_size[1]) # Back to [T, C, H, W]
resized_clips.append(clip)
clips = resized_clips
# Yield clips one by one
for clip in clips:
yield clip
total_clips += 1
# Free memory
del video, clips
except Exception as e:
failed_count += 1
if verbose:
print(f"\nWarning: Failed to load {video_path}: {e}")
continue
# Validate
if total_clips == 0:
raise RuntimeError(f"Failed to load any videos from {directory}")
failure_rate = failed_count / len(video_paths)
if failure_rate > 0.1: # More than 10% failed
print(
f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})"
)
if verbose:
print(
f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos"
)
+7
View File
@@ -0,0 +1,7 @@
#!/bin/bash
# 1. Install missing dependency
pip install -q opencv-python-headless
# 2. Run FVD script
python benchmarks/fvd/run_fvd.py
+4
View File
@@ -0,0 +1,4 @@
#!/bin/bash
# 1. Install missing dependency
pip install -q opencv-python-headless
@@ -247,11 +247,12 @@ def _attn_bwd_dq(dq, q, K, V, #
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
block_size = tl.load(variable_block_sizes + q_blk)
for blk_idx in range(kv_blocks*2):
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
kv_idx = tl.load(kv_ptr + blk_idx//2).to(tl.int32)
block_size = tl.load(variable_block_sizes + kv_idx) - (blk_idx % 2) * step_n
block_sparse_offset = (kv_idx*2 + blk_idx%2) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
+4
View File
@@ -70,3 +70,7 @@ uv pip install ninja
python setup.py install
```
## Testing
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
+129
View File
@@ -0,0 +1,129 @@
# Testing in FastVideo
This guide explains how to add and run tests in FastVideo. The testing suite is divided into several categories to ensure correctness across components, training workflows, and inference quality.
## Test Types
* **Unit Tests**: Located in `fastvideo/tests/dataset`, `fastvideo/tests/entrypoints`, and `fastvideo/tests/workflow`. These test individual functions and classes.
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
For now, we will focus on **SSIM Tests**.
## SSIM Tests
SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos using specific models and parameters, and compare them against reference videos to ensure that changes in the codebase do not degrade generation quality or alter the output unexpectedly.
!!! note
If you are adding an SSIM test, this serves as a safeguard. Any future code changes that break or cause errors with the specific arguments and configurations you defined will trigger a failure. Therefore, it is important to include multiple settings and arguments that cover the core features of your new pipeline to ensure robust regression testing.
### Directory Structure
```
fastvideo/tests/ssim/
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
│ ├── <Model_Name>/
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
│ │ │ └── <Video_File>
├── test_causal_similarity.py
├── test_inference_similarity.py
├── update_reference_videos.sh
└── ...
```
### Adding a New SSIM Test
To add a new SSIM test, follow these steps:
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
```python
MY_MODEL_PARAMS = {
"num_gpus": 1,
"model_path": "organization/model-name",
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 20,
# ... other parameters
}
```
3. **Implement the Test Function**:
* Use `pytest.mark.parametrize` to run the test with different prompts, backends, and models.
* Set the attention backend environment variable.
* Initialize the `VideoGenerator`.
* Generate the video.
* Compare the generated video with the reference video using `compute_video_ssim_torchvision`.
Example structure:
```python
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_my_model_similarity(prompt, ATTENTION_BACKEND):
# Setup output directories
# ...
# Initialize Generator
generator = VideoGenerator.from_pretrained(...)
generator.generate_video(prompt, ...)
# Compare with Reference
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
assert ssim_values[0] >= 0.98 # Threshold
```
4. **Reference Videos**:
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
* Inspect the generated video to ensure it meets quality expectations.
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
### Running Tests Locally
To run the SSIM tests locally:
```bash
pytest fastvideo/tests/ssim/ -vs
```
Ensure you have the necessary GPUs available as defined in your test parameters.
## Modal Workflow
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
### `pr_test.py`
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
### Updating Modal Configuration
If you add a new test that requires:
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
* **Longer Execution Time**: Increase the `timeout` parameter.
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
```python
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
def run_ssim_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
```
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
### Workflow Scripts
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
!!! note
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
+3 -2
View File
@@ -24,12 +24,13 @@ FastVideo is an inference and post-training framework for diffusion models. It f
## Key Features
FastVideo has the following features:
- State-of-the-art performance optimizations for inference
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- E2E post-training support
- Data preprocessing pipeline for video data.
- Data preprocessing pipeline for video data
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
@@ -42,7 +43,7 @@ Use the navigation menu on the left to explore different sections:
- **Getting Started**: Installation and quick start guides
- **Inference**: Learn how to use FastVideo for video generation
- **Training**: Data preprocessing and fine-tuning workflows
- **Training**: Data preprocessing and fine-tuning workflows
- **Distillation**: Post-training optimization techniques
- **Sliding Tile Attention**: Advanced attention mechanisms
- **Video Sparse Attention**: Efficient attention for video models
-51
View File
@@ -1,51 +0,0 @@
# Seed Parameter Behavior in vLLM
## Overview
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
## Default Behavior
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
## Specifying a Seed
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
## Example Usage
### Without Specifying a Seed
```python
import random
from vllm import LLM
# Initialize a vLLM model without specifying a seed
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
# Try generating random numbers
print(random.randint(0, 100)) # Outputs different numbers across runs
```
### Specifying a Seed
```python
import random
from vllm import LLM
# Initialize a vLLM model with a specific seed
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
# Try generating random numbers
print(random.randint(0, 100)) # Outputs the same number across runs
```
## Important Notes
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
## Conclusion
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 98 KiB

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

After

Width:  |  Height:  |  Size: 113 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 229 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 168 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 148 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 155 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 875 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 664 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 62 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 686 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 957 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 585 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 558 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 942 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 890 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 433 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 595 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 781 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 783 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 762 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 68 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 147 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 89 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 133 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 213 KiB

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

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