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
60 changed files with 5178 additions and 706 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
+12 -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">
@@ -112,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

@@ -26,7 +26,7 @@ def main():
)
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
sampling_param.num_frames = 73
sampling_param.num_frames = 81
sampling_param.width = 832
sampling_param.height = 480
sampling_param.seed = 1000
@@ -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"
+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
@@ -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")))
+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
@@ -400,6 +400,10 @@ class CausalDMDDenosingStage(DenoisingStage):
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
@@ -490,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
+19 -19
View File
@@ -5,6 +5,7 @@ Input validation stage for diffusion pipelines.
import torch
import torchvision.transforms.functional as TF
from PIL import Image
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
@@ -13,6 +14,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import (StageValidators,
VerificationResult)
from fastvideo.utils import best_output_size
logger = init_logger(__name__)
@@ -108,27 +110,25 @@ class InputValidationStage(PipelineStage):
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 = 720 * 1280
# ow, oh = best_output_size(iw, ih, dw, dh, max_area)
ih, iw = img.height, img.width
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
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)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
# 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)
# # 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("img height: %s, img width: %s", img.height, img.width)
oh = img.height
ow = img.width
# to tensor
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
self.device).unsqueeze(1)
@@ -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)
@@ -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)
+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,
+15 -4
View File
@@ -47,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:
@@ -76,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
@@ -104,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:
@@ -152,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
@@ -167,4 +178,4 @@ extra:
# Custom CSS
extra_css:
- assets/custom.css
- assets/custom.css
+2 -1
View File
@@ -7,7 +7,8 @@ mkdocstrings-python>=1.8.0
mkdocs-mermaid2-plugin>=1.1.0
mkdocs-git-revision-date-localized-plugin>=1.2.0
mkdocs-git-committers-plugin-2>=1.1.0
mkdocs-macros-plugin>=0.8.0
mkdocs-macros-plugin>=0.8.0
pymdown-extensions>=10.0
mkdocs-api-autonav
mkdocs-autorefs
mdx-truly-sane-lists
+68
View File
@@ -0,0 +1,68 @@
# LoRA Extraction and Merging
Tools for extracting and merging LoRA adapters for FastVideo models.
## Extract LoRA Adapter
```bash
python extract_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--out adapter_r32.safetensors \
--rank 32
```
**Options:**
- `--base`: Base model (HuggingFace ID or local path)
- `--ft`: Fine-tuned model (HuggingFace ID or local path)
- `--out`: Output adapter file
- `--rank`: LoRA rank (16, 32, 64, 128)
- `--full-rank`: Extract full-rank adapter (optional)
## Merge Adapter
```bash
python merge_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--adapter adapter_r32.safetensors \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--output merged_model
```
**Options:**
- `--base`: Base model (HuggingFace ID or local path)
- `--adapter`: LoRA adapter file (.safetensors)
- `--ft`: Fine-tuned model (for configuration)
- `--output`: Output directory
## Validate Quality (Optional)
```bash
python lora_inference_comparison.py \
--base merged_model \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--adapter NONE \
--output-dir results \
--prompt "A cat sitting on a windowsill" \
--seed 42 \
--height 480 \
--width 480 \
--num-frames 49 \
--num-inference-steps 32 \
--compute-ssim \
--compute-lpips
```
**Options:**
- `--base`: Merged model or base model path
- `--ft`: Fine-tuned model (reference)
- `--adapter`: Path to adapter or NONE
- `--output-dir`: Output directory
- `--prompt`: Text prompt (default: "A cat sitting on a windowsill")
- `--seed`: Random seed (default: 42)
- `--height`: Video height (default: 480)
- `--width`: Video width (default: 832)
- `--num-frames`: Number of frames (default: 49)
- `--num-inference-steps`: Inference steps (default: 32)
- `--compute-ssim`: Compute SSIM metric
- `--compute-lpips`: Compute LPIPS metric
+370
View File
@@ -0,0 +1,370 @@
"""Extract FastVideo-style LoRA adapters from a fine-tuned model by SVDing (FT - base).
Usage:
python scripts/lora_extraction/extract_lora.py \
--base <base_model> --ft <fine_tuned_model> --out adapter.safetensors --rank 16
example: python extract_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--out fastvideo_adapter.safetensors \
--rank 16
"""
from __future__ import annotations
import os
# Set distributed env BEFORE any fastvideo imports
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29500")
os.environ.setdefault("WORLD_SIZE", "1")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("LOCAL_RANK", "0")
import argparse
import sys
import logging
from pathlib import Path
from typing import Dict, Optional
import torch
from tqdm import tqdm
# Optional safetensors support
_HAVE_SAFETENSORS = True
try:
from safetensors.torch import save_file as safetensors_save # type: ignore
except Exception:
_HAVE_SAFETENSORS = False
# Configure minimal logging
LOG = logging.getLogger("extract_lora")
def configure_logging(level: str = "INFO") -> None:
handler = logging.StreamHandler()
fmt = "%(asctime)s %(levelname)s %(message)s"
handler.setFormatter(logging.Formatter(fmt, datefmt="%Y-%m-%d %H:%M:%S"))
LOG.addHandler(handler)
LOG.setLevel(level)
def get_pipeline_class_for_model(model_path: str):
"""Return appropriate FastVideo Pipeline class for the model."""
from fastvideo.utils import maybe_download_model_index # local import
from fastvideo.pipelines.pipeline_registry import get_pipeline_registry, PipelineType
from fastvideo.fastvideo_args import WorkloadType
config = maybe_download_model_index(model_path)
pipeline_name = config.get("_class_name")
if pipeline_name is None:
raise ValueError(f"Model config for {model_path} missing _class_name (diffusers format expected).")
pipeline_registry = get_pipeline_registry(PipelineType.BASIC)
pipeline_cls = pipeline_registry.resolve_pipeline_cls(pipeline_name, PipelineType.BASIC, WorkloadType.T2V)
return pipeline_cls
def load_transformer_state_dict_from_model(
model_path: str,
num_gpus: int = 1,
dit_cpu_offload: bool = True,
vae_cpu_offload: bool = True,
text_encoder_cpu_offload: bool = True,
pin_cpu_memory: bool = True,
) -> Dict[str, torch.Tensor]:
"""Load pipeline and extract transformer.state_dict as CPU tensors."""
pipeline_cls = get_pipeline_class_for_model(model_path)
pipeline = pipeline_cls.from_pretrained(
model_path,
num_gpus=num_gpus,
inference_mode=True,
dit_cpu_offload=dit_cpu_offload,
vae_cpu_offload=vae_cpu_offload,
text_encoder_cpu_offload=text_encoder_cpu_offload,
pin_cpu_memory=pin_cpu_memory,
)
# Try to locate transformer in several typical attributes
transformer = getattr(pipeline, "transformer", None)
if transformer is None:
modules = getattr(pipeline, "modules", None)
if isinstance(modules, dict):
transformer = modules.get("transformer")
if transformer is None:
pipeline_attr = getattr(pipeline, "pipeline", None)
transformer = getattr(pipeline_attr, "transformer", None) if pipeline_attr else None
if transformer is None:
raise RuntimeError("Transformer not found in pipeline. Expected pipeline.transformer or pipeline.modules['transformer'].")
state_dict = transformer.state_dict()
# DTensor safe handling
try:
from torch.distributed.tensor import DTensor # type: ignore
_HAS_DTENSOR = True
except Exception:
DTensor = None # type: ignore
_HAS_DTENSOR = False
state_dict_cpu: Dict[str, torch.Tensor] = {}
for k, v in state_dict.items():
if _HAS_DTENSOR and isinstance(v, DTensor): # type: ignore
state_dict_cpu[k] = v.to_local().detach().cpu().contiguous()
else:
state_dict_cpu[k] = v.detach().cpu().contiguous()
# cleanup
try:
del pipeline, transformer
except Exception:
pass
torch.cuda.empty_cache()
return state_dict_cpu
def is_extractable_weight(key: str) -> bool:
"""Return True if key represents a weight suitable for LoRA extraction."""
if not key.endswith("weight"):
return False
low = key.lower()
for skip in ("norm", "bias", "embedding"):
if skip in low:
return False
return True
def save_adapter_state(adapter_state: Dict[str, torch.Tensor], out_path: Path) -> None:
"""Save adapter state dict to safetensors (if available) or torch.save."""
cleaned = {k: v.detach().cpu().contiguous() for k, v in adapter_state.items()}
out_str = str(out_path)
if out_path.suffix == ".safetensors" and _HAVE_SAFETENSORS:
safetensors_save(cleaned, out_str)
else:
torch.save(cleaned, out_str)
def build_adapter_from_states(
base_sd: Dict[str, torch.Tensor],
ft_sd: Dict[str, torch.Tensor],
rank: int,
full_rank: bool,
min_delta: float,
checkpoint_interval: int,
checkpoint_path: Optional[Path],
resume_from: int = 0,
) -> Dict[str, torch.Tensor]:
"""Compute low-rank LoRA adapters by SVD on (ft - base) for extractable weights."""
# DTensor detection
try:
from torch.distributed.tensor import DTensor # type: ignore
_HAS_DTENSOR = True
except Exception:
DTensor = None # type: ignore
_HAS_DTENSOR = False
keys = sorted(ft_sd.keys())
adapter_state: Dict[str, torch.Tensor] = {}
processed = 0
mean_deltas = []
for idx, key in enumerate(tqdm(keys, desc="scanning keys", unit="keys")):
if idx < resume_from:
continue
if not is_extractable_weight(key):
continue
if key not in base_sd:
continue
Wb_raw = base_sd[key]
Wf_raw = ft_sd[key]
# Convert DTensor if present
if _HAS_DTENSOR and isinstance(Wb_raw, DTensor): # type: ignore
Wb = Wb_raw.to_local().detach().cpu().to(torch.float32).contiguous()
else:
Wb = Wb_raw.detach().cpu().to(torch.float32).contiguous()
if _HAS_DTENSOR and isinstance(Wf_raw, DTensor): # type: ignore
Wf = Wf_raw.to_local().detach().cpu().to(torch.float32).contiguous()
else:
Wf = Wf_raw.detach().cpu().to(torch.float32).contiguous()
if Wb.shape != Wf.shape:
continue
delta = (Wf - Wb).contiguous()
mean_abs = float(delta.abs().mean().item())
mean_deltas.append(mean_abs)
if mean_abs < min_delta:
continue
# SVD (CPU)
try:
U, S, Vh = torch.linalg.svd(delta, full_matrices=False)
except RuntimeError:
# skip layers that fail SVD
continue
max_rank = S.numel()
chosen_rank = max_rank if full_rank or rank <= 0 else min(rank, max_rank)
if chosen_rank == 0:
continue
S_sqrt = torch.sqrt(S[:chosen_rank].to(torch.float32))
U_r = U[:, :chosen_rank].to(torch.float32) # (out, r)
Vh_r = Vh[:chosen_rank, :].to(torch.float32) # (r, in)
lora_B = (U_r * S_sqrt.unsqueeze(0)).contiguous() # (out, r)
tmp = (Vh_r.T * S_sqrt.unsqueeze(0)).contiguous() # (in, r)
lora_A = tmp.T.contiguous() # (r, in)
base_name = key[:-len(".weight")]
a_key = f"{base_name}.lora_A.weight"
b_key = f"{base_name}.lora_B.weight"
rank_key = f"{base_name}.lora_rank"
alpha_key = f"{base_name}.lora_alpha"
adapter_state[a_key] = lora_A.cpu()
adapter_state[b_key] = lora_B.cpu()
adapter_state[rank_key] = torch.tensor([chosen_rank], dtype=torch.int32)
adapter_state[alpha_key] = torch.tensor([float(chosen_rank)], dtype=torch.float32)
processed += 1
# checkpoint periodically
if checkpoint_path and checkpoint_interval > 0 and (idx + 1) % checkpoint_interval == 0:
try:
torch.save({"index": idx + 1, "adapter": adapter_state}, str(checkpoint_path))
except Exception:
# non-fatal; continue
pass
# free local large tensors
del delta, U, S, Vh, U_r, Vh_r, tmp, lora_A, lora_B
# final checkpoint
if checkpoint_path:
try:
torch.save({"index": len(keys), "adapter": adapter_state}, str(checkpoint_path))
except Exception:
pass
avg_delta = (sum(mean_deltas) / len(mean_deltas)) if mean_deltas else 0.0
LOG.info("Extraction complete: processed_keys=%d, extracted_layers=%d, avg_abs_delta=%.6e",
len(keys), processed, avg_delta)
return adapter_state
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Extract FastVideo-style LoRA adapter (CPU SVD).")
p.add_argument("--base", required=True, help="Base model id or local path")
p.add_argument("--ft", required=True, help="Fine-tuned model id or local path")
p.add_argument("--out", default="fastvideo_adapter.safetensors", help="Output adapter file (.safetensors or .pt)")
p.add_argument("--rank", type=int, default=16, help="Truncated SVD rank; <=0 for full rank")
p.add_argument("--full-rank", action="store_true", help="Use full SVD rank for every layer")
p.add_argument("--min-delta", type=float, default=1e-8, help="Minimum mean abs delta to consider a layer changed")
p.add_argument("--checkpoint", default="extract_lora_checkpoint.pt", help="Checkpoint path to resume/save progress")
p.add_argument("--resume", action="store_true", help="Resume from checkpoint if available")
p.add_argument("--log-level", default="INFO", choices=["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"])
return p.parse_args()
def extract_lora_adapter(
base: str,
ft: str,
out: str,
rank: int = 32,
full_rank: bool = False,
min_delta: float = 1e-6,
checkpoint: Optional[str] = None,
resume: bool = False,
log_level: str = "INFO",
) -> None:
"""Extract LoRA adapter from fine-tuned model.
Args:
base: Base model path or HuggingFace ID
ft: Fine-tuned model path or HuggingFace ID
out: Output adapter file path
rank: LoRA rank (default: 32)
full_rank: Extract full-rank adapter
min_delta: Minimum delta for extraction
checkpoint: Checkpoint file path
resume: Resume from checkpoint
log_level: Logging level
"""
configure_logging(log_level)
# ensure fastvideo import
try:
import fastvideo # noqa: F401
except Exception as exc:
LOG.error("Failed to import fastvideo: %s", exc)
sys.exit(2)
out_path = Path(out)
checkpoint_path = Path(checkpoint) if checkpoint else None
# load state_dicts
LOG.info("Loading base model: %s", base)
base_sd = load_transformer_state_dict_from_model(base)
LOG.info("Loading fine-tuned model: %s", ft)
ft_sd = load_transformer_state_dict_from_model(ft)
resume_idx = 0
adapter_existing: Dict[str, torch.Tensor] = {}
if resume and checkpoint_path and checkpoint_path.exists():
try:
ck = torch.load(str(checkpoint_path), map_location="cpu")
adapter_existing = ck.get("adapter", {}) or {}
resume_idx = int(ck.get("index", 0) or 0)
LOG.info("Resuming from checkpoint index=%d with %d existing entries", resume_idx, len(adapter_existing))
except Exception:
adapter_existing = {}
adapter_state = dict(adapter_existing) if adapter_existing else {}
new_adapter = build_adapter_from_states(
base_sd=base_sd,
ft_sd=ft_sd,
rank=rank,
full_rank=full_rank,
min_delta=min_delta,
checkpoint_interval=50,
checkpoint_path=checkpoint_path,
resume_from=resume_idx,
)
adapter_state.update(new_adapter)
# final save
save_adapter_state(adapter_state, out_path)
# cleanup checkpoint if present
if checkpoint_path and checkpoint_path.exists():
try:
checkpoint_path.unlink()
except Exception:
pass
LOG.info("Saved adapter to %s (entries=%d)", str(out_path), len(adapter_state) // 4)
def main() -> None:
"""CLI wrapper for extract_lora_adapter."""
args = parse_args()
extract_lora_adapter(
base=args.base,
ft=args.ft,
out=args.out,
rank=args.rank,
full_rank=args.full_rank,
min_delta=args.min_delta,
checkpoint=args.checkpoint,
resume=args.resume,
log_level=args.log_level,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,358 @@
#!/usr/bin/env python3
"""Compare Fine-Tuned vs Merged/Base+LoRA inference outputs.
Generates two videos with the same seed and computes SSIM.
Usage examples:
python lora_inference_comparison.py \
--base ./merged_model \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--adapter NONE \
--output-dir ./inference_comparison \
--compute-ssim \
--seed 41
or
python lora_inference_comparison.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--adapter adapter.safetensors \
--output-dir ./inference_comparison \
--compute-ssim \
--seed 41
"""
from __future__ import annotations
import argparse
import os
import sys
from pathlib import Path
from typing import Optional, Dict, Any
import logging
# minimal distributed env defaults (kept for compatibility)
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29500")
os.environ.setdefault("WORLD_SIZE", "1")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("LOCAL_RANK", "0")
# allow running from repo root where fastvideo is located
_FASTVIDEO_PATH = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "fastvideo_pr", "FastVideo"))
if _FASTVIDEO_PATH not in sys.path:
sys.path.insert(0, _FASTVIDEO_PATH)
logger = logging.getLogger("inference_comparison")
def configure_logging(level: str = "INFO") -> None:
handler = logging.StreamHandler()
fmt = "%(asctime)s %(levelname)s %(message)s"
handler.setFormatter(logging.Formatter(fmt, datefmt="%Y-%m-%d %H:%M:%S"))
logger.addHandler(handler)
logger.setLevel(level)
def _validate_adapter(adapter: Optional[str]) -> Optional[str]:
if not adapter:
return None
if adapter.upper() == "NONE":
return None
p = Path(adapter).expanduser()
if not p.exists():
raise FileNotFoundError(f"Adapter not found: {p}")
# Accept both files and directories (FastVideo expects directories for HF-style adapters)
if p.is_file():
if p.suffix != ".safetensors":
raise ValueError(f"Adapter file must be .safetensors, got: {p.suffix}")
if p.stat().st_size == 0:
raise ValueError(f"Adapter file is empty: {p}")
elif p.is_dir():
# Check if directory contains at least one .safetensors file
safetensors_files = list(p.glob("*.safetensors"))
if not safetensors_files:
raise ValueError(f"Adapter directory contains no .safetensors files: {p}")
else:
raise ValueError(f"Adapter must be a file or directory: {p}")
return str(p.resolve())
def generate_with_model(
model_path: str,
output_dir: str,
output_name: str,
prompt: str,
seed: int,
lora_path: Optional[str],
height: int,
width: int,
num_frames: int,
num_inference_steps: int,
guidance_scale: float,
flow_shift: Optional[float] = None,
embedded_guidance_scale: Optional[float] = None,
) -> str:
"""Produce a video with VideoGenerator.from_pretrained; returns video path."""
try:
from fastvideo import VideoGenerator # lazy import
except Exception as exc:
raise RuntimeError(f"Failed to import fastvideo.VideoGenerator: {exc}") from exc
init_kwargs: Dict[str, Any] = {
"num_gpus": 1,
"dit_cpu_offload": True,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": True,
}
if lora_path:
init_kwargs["lora_path"] = lora_path
init_kwargs["lora_nickname"] = "extracted"
generator = VideoGenerator.from_pretrained(model_path, **init_kwargs)
gen_kwargs = {
"height": height,
"width": width,
"num_frames": num_frames,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
"seed": seed,
"output_path": output_dir,
"output_video_name": output_name,
"save_video": True,
}
if flow_shift is not None:
gen_kwargs["flow_shift"] = flow_shift
if embedded_guidance_scale is not None:
gen_kwargs["embedded_guidance_scale"] = embedded_guidance_scale
result = generator.generate_video(prompt, **gen_kwargs)
# best-effort cleanup of internal executors
try:
if hasattr(generator, "executor") and hasattr(generator.executor, "shutdown"):
generator.executor.shutdown()
except Exception:
pass
# determine saved video path
expected = Path(output_dir) / f"{output_name}.mp4"
if expected.exists():
return str(expected)
# fallback: check result dict
if isinstance(result, dict) and "video_path" in result:
return str(result["video_path"])
raise FileNotFoundError(f"Video not found at expected path: {expected}")
def compute_metrics(output_dir: str, ft_video: str, other_video: str, num_inference_steps: int, prompt: str, compute_ssim: bool, compute_lpips: bool) -> dict:
results = {}
if compute_ssim:
try:
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results # type: ignore
ssim_values = compute_video_ssim_torchvision(ft_video, other_video, use_ms_ssim=True)
results["mean_ssim"] = float(ssim_values[0])
write_ssim_results(output_dir, ssim_values, ft_video, other_video, num_inference_steps, prompt)
except Exception as e:
logger.warning(f"SSIM computation failed: {e}")
if compute_lpips:
try:
import torch
import lpips
import torchvision.io as tv_io
loss_fn = lpips.LPIPS(net='alex')
# Load videos
vid1, _, _ = tv_io.read_video(ft_video, pts_unit='sec')
vid2, _, _ = tv_io.read_video(other_video, pts_unit='sec')
# Normalize to [-1, 1]
vid1 = (vid1.float() / 127.5 - 1.0).permute(0, 3, 1, 2) # (T, C, H, W)
vid2 = (vid2.float() / 127.5 - 1.0).permute(0, 3, 1, 2)
lpips_scores = []
with torch.no_grad():
for frame1, frame2 in zip(vid1, vid2):
score = loss_fn(frame1.unsqueeze(0), frame2.unsqueeze(0))
lpips_scores.append(float(score.item()))
results["mean_lpips"] = sum(lpips_scores) / len(lpips_scores)
# Write LPIPS results
import json
lpips_file = Path(output_dir) / f"steps{num_inference_steps}_{prompt.replace(' ', '_')[:30]}_lpips.json"
with open(lpips_file, 'w') as f:
json.dump({"mean_lpips": results["mean_lpips"], "lpips_per_frame": lpips_scores}, f, indent=2)
except Exception as e:
logger.warning(f"LPIPS computation failed: {e}")
return results
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Compare Fine-Tuned vs Merged/Base+LoRA inference outputs")
p.add_argument("--base", required=True, help="Base model ID or merged model path")
p.add_argument("--ft", required=True, help="Fine-tuned model ID or path (reference)")
p.add_argument("--adapter", default="NONE", help="Path to .safetensors adapter, or NONE to use merged model")
p.add_argument("--output-dir", default="./inference_comparison")
p.add_argument("--prompt", default="A cat sitting on a windowsill")
p.add_argument("--seed", type=int, default=42)
p.add_argument("--height", type=int, default=480)
p.add_argument("--width", type=int, default=832)
p.add_argument("--num-frames", type=int, default=49)
p.add_argument("--num-inference-steps", type=int, default=32)
p.add_argument("--guidance-scale", type=float, default=6.0)
p.add_argument("--compute-ssim", action="store_true")
p.add_argument("--compute-lpips", action="store_true")
p.add_argument("--flow-shift", type=float, default=None)
p.add_argument("--embedded-guidance-scale", type=float, default=None)
p.add_argument("--log-level", default="INFO")
return p.parse_args()
def compare_inference(
base: str,
ft: str,
adapter: Optional[str],
output_dir: str,
prompt: str = "A cat sitting on a windowsill",
seed: int = 42,
height: int = 480,
width: int = 832,
num_frames: int = 49,
num_inference_steps: int = 32,
guidance_scale: float = 5.0,
flow_shift: Optional[float] = None,
embedded_guidance_scale: Optional[float] = None,
compute_ssim: bool = False,
compute_lpips: bool = False,
log_level: str = "INFO",
) -> dict:
"""Compare inference between fine-tuned model and merged/base+adapter model.
Args:
base: Base or merged model ID/path
ft: Fine-tuned model ID/path
adapter: LoRA adapter path (or NONE for merged model)
output_dir: Output directory for videos
prompt: Generation prompt
seed: Random seed
height: Video height
width: Video width
num_frames: Number of frames
num_inference_steps: Inference steps
guidance_scale: CFG scale
flow_shift: Flow shift
embedded_guidance_scale: Embedded guidance scale
compute_ssim: Compute SSIM metric
compute_lpips: Compute LPIPS metric
log_level: Logging level
Returns:
Dictionary with metric results
"""
configure_logging(log_level)
os.makedirs(output_dir, exist_ok=True)
try:
adapter_path = _validate_adapter(adapter)
except Exception as exc:
logger.error("Adapter validation failed: %s", exc)
sys.exit(2)
# 1) generate with fine-tuned model (reference)
logger.info("Generating reference (fine-tuned): %s", ft)
try:
ft_video = generate_with_model(
model_path=ft,
output_dir=output_dir,
output_name="fine_tuned",
prompt=prompt,
seed=seed,
lora_path=None,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
flow_shift=flow_shift,
embedded_guidance_scale=embedded_guidance_scale,
)
logger.info("Reference video saved: %s", ft_video)
except Exception as exc:
logger.error("Reference generation failed: %s", exc)
sys.exit(3)
# 2) generate with merged model OR base + adapter
use_merged = adapter_path is None
mode = "merged model" if use_merged else "base+adapter"
logger.info("Generating target (%s): %s", mode, base)
try:
target_video = generate_with_model(
model_path=base,
output_dir=output_dir,
output_name="merged_model" if use_merged else "base_plus_lora",
prompt=prompt,
seed=seed,
lora_path=None if use_merged else adapter_path,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
flow_shift=flow_shift,
embedded_guidance_scale=embedded_guidance_scale,
)
logger.info("Target video saved: %s", target_video)
except Exception as exc:
logger.error("Target generation failed: %s", exc)
sys.exit(4)
# 3) compute metrics
results = {}
if compute_ssim or compute_lpips:
results = compute_metrics(output_dir, ft_video, target_video, num_inference_steps, prompt, compute_ssim, compute_lpips)
if results.get("mean_ssim") is not None:
logger.info("Mean SSIM: %.4f", results["mean_ssim"])
if results.get("mean_lpips") is not None:
logger.info("Mean LPIPS: %.4f", results["mean_lpips"])
logger.info("Comparison complete. Videos in: %s", output_dir)
return results
def main() -> None:
"""CLI wrapper for compare_inference."""
args = parse_args()
compare_inference(
base=args.base,
ft=args.ft,
adapter=args.adapter,
output_dir=args.output_dir,
prompt=args.prompt,
seed=args.seed,
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
flow_shift=args.flow_shift,
embedded_guidance_scale=args.embedded_guidance_scale,
compute_ssim=args.compute_ssim,
compute_lpips=args.compute_lpips,
log_level=args.log_level,
)
if __name__ == "__main__":
main()
+339
View File
@@ -0,0 +1,339 @@
"""Merge LoRA adapter into base model weights.
Usage:
python merge_lora_updated.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--adapter adapter.safetensors \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--output ./merged_model
"""
from __future__ import annotations
import argparse
import os
import sys
import shutil
import logging
from pathlib import Path
from collections import defaultdict
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29500")
os.environ.setdefault("WORLD_SIZE", "1")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("LOCAL_RANK", "0")
_FASTVIDEO_PATH = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "fastvideo_pr", "FastVideo"))
if _FASTVIDEO_PATH not in sys.path:
sys.path.insert(0, _FASTVIDEO_PATH)
import torch
from safetensors.torch import load_file, save_file
from extract_lora import load_transformer_state_dict_from_model, get_pipeline_class_for_model
from fastvideo.training.training_utils import custom_to_hf_state_dict
from fastvideo.models.loader.utils import get_param_names_mapping, hf_to_custom_state_dict
LOG = logging.getLogger(__name__)
def configure_logging(level: str = "INFO"):
handler = logging.StreamHandler()
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
LOG.addHandler(handler)
LOG.setLevel(level)
def fix_adapter_naming(adapter: dict) -> dict:
fixed = {}
renamed = 0
for key, tensor in adapter.items():
new_key = key
if ".lora_A.weight" in key:
base = key.replace(".lora_A.weight", "")
if not base.endswith(".weight"):
new_key = base + ".weight.lora_A.weight"
renamed += 1
elif ".lora_B.weight" in key:
base = key.replace(".lora_B.weight", "")
if not base.endswith(".weight"):
new_key = base + ".weight.lora_B.weight"
renamed += 1
elif ".lora_rank" in key:
base = key.replace(".lora_rank", "")
if not base.endswith(".weight"):
new_key = base + ".weight.lora_rank"
renamed += 1
elif ".lora_alpha" in key:
base = key.replace(".lora_alpha", "")
if not base.endswith(".weight"):
new_key = base + ".weight.lora_alpha"
renamed += 1
fixed[new_key] = tensor
if renamed > 0:
LOG.info(f"Fixed {renamed} adapter key names")
return fixed
def load_adapter(adapter_path: str) -> dict:
abs_path = os.path.abspath(adapter_path)
if not os.path.exists(abs_path):
raise FileNotFoundError(f"Adapter file not found: {abs_path}")
if not abs_path.endswith('.safetensors'):
raise ValueError(f"Adapter must be .safetensors: {abs_path}")
LOG.info(f"Loading adapter: {abs_path}")
adapter = load_file(abs_path)
file_size_mb = os.path.getsize(abs_path) / (1024 * 1024)
LOG.info(f"Loaded {len(adapter)} tensors ({file_size_mb:.1f} MB)")
return fix_adapter_naming(adapter)
def group_adapter_keys(adapter: dict) -> dict:
grouped = defaultdict(dict)
for key, tensor in adapter.items():
if key.endswith(".lora_A.weight"):
grouped[key.replace(".lora_A.weight", "")]["A"] = tensor
elif key.endswith(".lora_B.weight"):
grouped[key.replace(".lora_B.weight", "")]["B"] = tensor
elif key.endswith(".lora_rank"):
grouped[key.replace(".lora_rank", "")]["rank"] = tensor
elif key.endswith(".lora_alpha"):
grouped[key.replace(".lora_alpha", "")]["alpha"] = tensor
LOG.info(f"Grouped {len(grouped)} LoRA layers")
return grouped
def get_reverse_param_mapping(base_model_path: str):
LOG.info("Loading base model for parameter mapping")
pipeline_cls = get_pipeline_class_for_model(base_model_path)
pipeline = pipeline_cls.from_pretrained(
base_model_path,
num_gpus=1,
inference_mode=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
transformer = None
if hasattr(pipeline, "transformer"):
transformer = pipeline.transformer
elif hasattr(pipeline, "modules") and isinstance(pipeline.modules, dict):
if "transformer" in pipeline.modules:
transformer = pipeline.modules["transformer"]
if transformer is None:
raise RuntimeError("Could not find transformer in pipeline")
if hasattr(transformer, "reverse_param_names_mapping"):
reverse_mapping = transformer.reverse_param_names_mapping
elif hasattr(transformer, "config") and hasattr(transformer.config, "arch_config"):
arch_config = transformer.config.arch_config
if hasattr(arch_config, "reverse_param_names_mapping"):
reverse_mapping = arch_config.reverse_param_names_mapping
else:
param_mapping = arch_config.param_names_mapping
param_names_mapping_fn = get_param_names_mapping(param_mapping)
from diffusers import DiffusionPipeline
from huggingface_hub import snapshot_download
if not os.path.exists(base_model_path) or not os.path.isdir(base_model_path):
model_path = snapshot_download(
repo_id=base_model_path,
ignore_patterns=["*.onnx", "*.msgpack"]
)
else:
model_path = base_model_path
hf_pipeline = DiffusionPipeline.from_pretrained(model_path, torch_dtype=torch.float32)
hf_transformer = hf_pipeline.transformer
hf_sd = hf_transformer.state_dict()
_, reverse_mapping = hf_to_custom_state_dict(hf_sd, param_names_mapping_fn)
del hf_pipeline
del hf_transformer
torch.cuda.empty_cache()
else:
raise RuntimeError("Could not find reverse_param_names_mapping in transformer or config")
del pipeline
del transformer
torch.cuda.empty_cache()
return reverse_mapping
def merge_lora_into_base(base_sd: dict, adapter: dict) -> dict:
LOG.info("Merging LoRA into base weights")
adapter_layers = group_adapter_keys(adapter)
merged_sd = dict(base_sd)
merged_count = 0
skipped_count = 0
for base_name, parts in adapter_layers.items():
weight_key = base_name if base_name.endswith(".weight") else base_name + ".weight"
if weight_key not in base_sd:
skipped_count += 1
continue
if "A" not in parts or "B" not in parts:
skipped_count += 1
continue
lora_A = parts["A"].to(torch.float32)
lora_B = parts["B"].to(torch.float32)
base_weight = base_sd[weight_key].to(torch.float32)
out_dim, in_dim = base_weight.shape
if lora_B.shape[0] != out_dim or lora_A.shape[1] != in_dim or lora_B.shape[1] != lora_A.shape[0]:
skipped_count += 1
continue
delta = lora_B @ lora_A
rank = int(parts.get("rank", torch.tensor([lora_A.shape[0]])).item()) if "rank" in parts else lora_A.shape[0]
alpha = float(parts.get("alpha", torch.tensor([rank])).item()) if "alpha" in parts else float(rank)
if rank != 0 and alpha != rank:
delta = delta * (alpha / float(rank))
merged_weight = base_weight + delta
merged_sd[weight_key] = merged_weight.to(base_sd[weight_key].dtype)
merged_count += 1
LOG.info(f"Merged {merged_count} layers, skipped {skipped_count}")
return merged_sd
def save_merged_model(merged_sd: dict, base_model_path: str, ft_model_path: str, output_dir: str, reverse_mapping: dict):
output_path = Path(output_dir)
output_path.mkdir(parents=True, exist_ok=True)
transformer_dir = output_path / "transformer"
transformer_dir.mkdir(parents=True, exist_ok=True)
LOG.info("Converting to HuggingFace format")
hf_merged_sd = custom_to_hf_state_dict(merged_sd, reverse_mapping)
LOG.info(f"Converted {len(hf_merged_sd)} parameters")
base_path = Path(base_model_path)
if not base_path.exists() or not base_path.is_dir():
from huggingface_hub import snapshot_download
base_path = Path(snapshot_download(
repo_id=base_model_path,
ignore_patterns=["*.onnx", "*.msgpack"]
))
LOG.info("Copying model components")
for component in ["scheduler", "text_encoder", "tokenizer", "vae"]:
src = base_path / component
if src.exists():
dst = output_path / component
if dst.exists():
shutil.rmtree(dst)
shutil.copytree(src, dst)
LOG.info("Copying finetuned model config")
ft_path = Path(ft_model_path)
if not ft_path.exists() or not ft_path.is_dir():
from huggingface_hub import snapshot_download
ft_path = Path(snapshot_download(
repo_id=ft_model_path,
allow_patterns=["model_index.json"],
ignore_patterns=["*.onnx", "*.msgpack"]
))
ft_index = ft_path / "model_index.json"
if ft_index.exists():
shutil.copy2(ft_index, output_path / "model_index.json")
else:
LOG.warning("Finetuned model_index.json not found, using base")
src_index = base_path / "model_index.json"
if src_index.exists():
shutil.copy2(src_index, output_path / "model_index.json")
weight_path = transformer_dir / "diffusion_pytorch_model.safetensors"
LOG.info(f"Saving merged weights to {weight_path}")
to_save_hf = {k: v.detach().cpu() for k, v in hf_merged_sd.items()}
save_file(to_save_hf, str(weight_path))
config_src = base_path / "transformer" / "config.json"
if config_src.exists():
shutil.copy2(config_src, transformer_dir / "config.json")
file_size_mb = weight_path.stat().st_size / (1024 * 1024)
LOG.info(f"Saved to {output_dir} ({file_size_mb:.0f} MB, {len(hf_merged_sd)} params)")
def merge_lora(
base: str,
adapter: str,
ft: str,
output: str,
log_level: str = "INFO",
) -> None:
"""Merge LoRA adapter into base model.
Args:
base: Base model ID or path
adapter: LoRA adapter .safetensors file
ft: Finetuned model ID (for config)
output: Output directory
log_level: Logging level
"""
configure_logging(log_level)
LOG.info(f"Base: {base}")
LOG.info(f"Adapter: {adapter}")
LOG.info(f"Output: {output}")
reverse_mapping = get_reverse_param_mapping(base)
LOG.info(f"Loading base model: {base}")
base_sd = load_transformer_state_dict_from_model(base)
LOG.info(f"Loaded { len(base_sd)} parameters")
adapter_sd = load_adapter(adapter)
merged_sd = merge_lora_into_base(base_sd, adapter_sd)
save_merged_model(merged_sd, base, ft, output, reverse_mapping)
LOG.info("Merge complete")
def main():
"""CLI wrapper for merge_lora."""
parser = argparse.ArgumentParser(description="Merge LoRA adapter into base model")
parser.add_argument("--base", required=True, help="Base model ID or path")
parser.add_argument("--adapter", required=True, help="LoRA adapter .safetensors file")
parser.add_argument("--ft", required=True, help="Finetuned model ID (for config)")
parser.add_argument("--output", required=True, help="Output directory")
parser.add_argument("--log-level", default="INFO", help="Logging level")
args = parser.parse_args()
merge_lora(
base=args.base,
adapter=args.adapter,
ft=args.ft,
output=args.output,
log_level=args.log_level,
)
if __name__ == "__main__":
main()
+141
View File
@@ -0,0 +1,141 @@
"""Verify merged LoRA model matches finetuned model numerically.
Usage:
python verify_lora.py \
--merged merged_model \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers
"""
from __future__ import annotations
import argparse
import logging
from pathlib import Path
import torch
from safetensors.torch import load_file
LOG = logging.getLogger(__name__)
def configure_logging(level: str = "INFO"):
handler = logging.StreamHandler()
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
LOG.addHandler(handler)
LOG.setLevel(level)
def load_transformer_weights(model_path: str | Path) -> dict:
"""Load transformer weights from model directory."""
model_path = Path(model_path)
if not model_path.exists() or not model_path.is_dir():
from huggingface_hub import snapshot_download
LOG.info(f"Downloading {model_path} from HuggingFace Hub...")
model_path = Path(snapshot_download(
repo_id=str(model_path),
ignore_patterns=["*.onnx", "*.msgpack"]
))
transformer_dir = model_path / "transformer"
if not transformer_dir.exists():
raise FileNotFoundError(f"Transformer directory not found: {transformer_dir}")
weight_files = sorted(transformer_dir.glob("*.safetensors"))
if not weight_files:
raise FileNotFoundError(f"No safetensors files in {transformer_dir}")
LOG.info(f"Loading {len(weight_files)} file(s) from {transformer_dir}")
state_dict = {}
for f in weight_files:
if "custom" in f.name:
continue
state_dict.update(load_file(str(f)))
return state_dict
def compare_models(merged_sd: dict, ft_sd: dict) -> dict:
"""Compare merged and finetuned model weights."""
common_keys = set(merged_sd.keys()) & set(ft_sd.keys())
merged_only = set(merged_sd.keys()) - set(ft_sd.keys())
ft_only = set(ft_sd.keys()) - set(merged_sd.keys())
LOG.info(f"Common keys: {len(common_keys)}")
if merged_only:
LOG.warning(f"Keys only in merged: {len(merged_only)}")
if ft_only:
LOG.warning(f"Keys only in finetuned: {len(ft_only)}")
results = []
for key in sorted(common_keys):
merged_param = merged_sd[key]
ft_param = ft_sd[key]
if merged_param.shape != ft_param.shape:
LOG.error(f"{key}: shape mismatch {merged_param.shape} vs {ft_param.shape}")
continue
diff = (merged_param.float() - ft_param.float()).abs()
max_abs = diff.max().item()
mean_abs = diff.mean().item()
merged_norm = merged_param.float().norm().item()
rel_mean = (mean_abs / merged_norm * 100) if merged_norm > 0 else 0
results.append({
"key": key,
"shape": tuple(merged_param.shape),
"max_abs": max_abs,
"mean_abs": mean_abs,
"rel_mean": rel_mean
})
return results
def main():
parser = argparse.ArgumentParser(description="Verify merged LoRA model matches finetuned model")
parser.add_argument("--merged", required=True, help="Merged model directory")
parser.add_argument("--ft", required=True, help="Finetuned model ID or path")
parser.add_argument("--log-level", default="INFO", help="Logging level")
args = parser.parse_args()
configure_logging(args.log_level)
LOG.info(f"Loading merged model: {args.merged}")
merged_sd = load_transformer_weights(args.merged)
LOG.info(f"Loaded {len(merged_sd)} parameters")
LOG.info(f"Loading finetuned model: {args.ft}")
ft_sd = load_transformer_weights(args.ft)
LOG.info(f"Loaded {len(ft_sd)} parameters")
LOG.info("Comparing models...")
results = compare_models(merged_sd, ft_sd)
results.sort(key=lambda x: x["max_abs"], reverse=True)
LOG.info(f"\nTop 10 mismatches by max_abs_error:")
for i, r in enumerate(results[:10], 1):
LOG.info(f"{i:2d}. {r['key']}")
LOG.info(f" shape={r['shape']}, max_abs={r['max_abs']:.3e}, mean_abs={r['mean_abs']:.3e}, rel_mean={r['rel_mean']:.4f}%")
overall_mean = sum(r["mean_abs"] for r in results) / len(results)
overall_max = max(r["max_abs"] for r in results)
LOG.info(f"\nOverall metrics:")
LOG.info(f" Layers compared: {len(results)}")
LOG.info(f" Mean(mean_abs): {overall_mean:.3e}")
LOG.info(f" Max(max_abs): {overall_max:.3e}")
if overall_mean < 1e-4:
LOG.info("\nVerification PASSED: Merge is numerically accurate")
else:
LOG.warning(f"\nVerification WARNING: Mean error {overall_mean:.3e} > 1e-4")
if __name__ == "__main__":
main()