Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c1d7c51bc7 | ||
|
|
5de2d80ebd | ||
|
|
6cd227349c | ||
|
|
7e15e5dab1 | ||
|
|
c9ee4f8bf8 | ||
|
|
f764b43aaa | ||
|
|
592af8c954 | ||
|
|
4225a96a40 | ||
|
|
733c14a2cb | ||
|
|
81b302409e | ||
|
|
7fbd0cfec7 | ||
|
|
85b9934079 |
@@ -222,14 +222,3 @@ 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"
|
||||
|
||||
@@ -75,7 +75,7 @@ case "$TEST_TYPE" in
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
@@ -126,10 +126,6 @@ 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
|
||||
|
||||
@@ -3,18 +3,8 @@ name: Deploy Documentation
|
||||
on:
|
||||
push:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/docs.yml'
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/docs.yml'
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/c7g1qdD" 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/XcY0Cpv" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -14,7 +15,6 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
@@ -111,21 +111,24 @@ 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) -->
|
||||
|
||||
## Awesome work using FastVideo or our research projects
|
||||
## 📑 Development Plan
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
More FastWan Models Coming Soon!
|
||||
- [ ] Add FastWan2.1-T2V-14B
|
||||
- [ ] Add FastWan2.2-T2V-14B
|
||||
- [ ] Add FastWan2.2-I2V-14B
|
||||
<!-- - Optimization features
|
||||
- Code updates -->
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025. [](https://github.com/sgl-project/sglang)
|
||||
|
||||
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [](https://github.com/XueZeyue/DanceGRPO)
|
||||
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo. [](https://github.com/Tencent-Hunyuan/SRPO)
|
||||
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [](https://github.com/Vchitect/DCM)
|
||||
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention. [](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
|
||||
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch. [](https://github.com/kandinskylab/kandinsky-5)
|
||||
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention. [](https://github.com/meituan-longcat/LongCat-Video)
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
|
||||
@@ -1,103 +0,0 @@
|
||||
# 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
|
||||
@@ -1,35 +0,0 @@
|
||||
"""
|
||||
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',
|
||||
]
|
||||
@@ -1,185 +0,0 @@
|
||||
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())
|
||||
@@ -1,447 +0,0 @@
|
||||
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
|
||||
@@ -1,142 +0,0 @@
|
||||
"""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)
|
||||
@@ -1,34 +0,0 @@
|
||||
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()
|
||||
@@ -1,97 +0,0 @@
|
||||
#!/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()
|
||||
@@ -1,490 +0,0 @@
|
||||
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"
|
||||
)
|
||||
@@ -1,7 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
@@ -1,4 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
@@ -247,12 +247,11 @@ def _attn_bwd_dq(dq, q, K, V, #
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + q_blk)
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
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
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*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)
|
||||
|
||||
@@ -70,7 +70,3 @@ uv pip install ninja
|
||||
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Testing
|
||||
|
||||
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
|
||||
|
||||
@@ -1,129 +0,0 @@
|
||||
# Testing in FastVideo
|
||||
|
||||
This guide explains how to add and run tests in FastVideo. The testing suite is divided into several categories to ensure correctness across components, training workflows, and inference quality.
|
||||
|
||||
## Test Types
|
||||
|
||||
* **Unit Tests**: Located in `fastvideo/tests/dataset`, `fastvideo/tests/entrypoints`, and `fastvideo/tests/workflow`. These test individual functions and classes.
|
||||
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
|
||||
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
|
||||
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
|
||||
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
|
||||
|
||||
For now, we will focus on **SSIM Tests**.
|
||||
|
||||
## SSIM Tests
|
||||
|
||||
SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos using specific models and parameters, and compare them against reference videos to ensure that changes in the codebase do not degrade generation quality or alter the output unexpectedly.
|
||||
|
||||
!!! note
|
||||
If you are adding an SSIM test, this serves as a safeguard. Any future code changes that break or cause errors with the specific arguments and configurations you defined will trigger a failure. Therefore, it is important to include multiple settings and arguments that cover the core features of your new pipeline to ensure robust regression testing.
|
||||
|
||||
### Directory Structure
|
||||
|
||||
```
|
||||
fastvideo/tests/ssim/
|
||||
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
|
||||
│ ├── <Model_Name>/
|
||||
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
|
||||
│ │ │ └── <Video_File>
|
||||
├── test_causal_similarity.py
|
||||
├── test_inference_similarity.py
|
||||
├── update_reference_videos.sh
|
||||
└── ...
|
||||
```
|
||||
|
||||
### Adding a New SSIM Test
|
||||
|
||||
To add a new SSIM test, follow these steps:
|
||||
|
||||
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
|
||||
|
||||
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
|
||||
|
||||
```python
|
||||
MY_MODEL_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": "organization/model-name",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
# ... other parameters
|
||||
}
|
||||
```
|
||||
|
||||
3. **Implement the Test Function**:
|
||||
* Use `pytest.mark.parametrize` to run the test with different prompts, backends, and models.
|
||||
* Set the attention backend environment variable.
|
||||
* Initialize the `VideoGenerator`.
|
||||
* Generate the video.
|
||||
* Compare the generated video with the reference video using `compute_video_ssim_torchvision`.
|
||||
|
||||
Example structure:
|
||||
|
||||
```python
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
def test_my_model_similarity(prompt, ATTENTION_BACKEND):
|
||||
# Setup output directories
|
||||
# ...
|
||||
|
||||
# Initialize Generator
|
||||
generator = VideoGenerator.from_pretrained(...)
|
||||
generator.generate_video(prompt, ...)
|
||||
|
||||
# Compare with Reference
|
||||
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
|
||||
assert ssim_values[0] >= 0.98 # Threshold
|
||||
```
|
||||
|
||||
4. **Reference Videos**:
|
||||
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
|
||||
* Inspect the generated video to ensure it meets quality expectations.
|
||||
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
|
||||
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
|
||||
|
||||
### Running Tests Locally
|
||||
|
||||
To run the SSIM tests locally:
|
||||
|
||||
```bash
|
||||
pytest fastvideo/tests/ssim/ -vs
|
||||
```
|
||||
|
||||
Ensure you have the necessary GPUs available as defined in your test parameters.
|
||||
|
||||
## Modal Workflow
|
||||
|
||||
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
|
||||
|
||||
### `pr_test.py`
|
||||
|
||||
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
|
||||
|
||||
### Updating Modal Configuration
|
||||
|
||||
If you add a new test that requires:
|
||||
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
|
||||
* **Longer Execution Time**: Increase the `timeout` parameter.
|
||||
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
|
||||
|
||||
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
|
||||
|
||||
```python
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_ssim_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
```
|
||||
|
||||
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
|
||||
|
||||
### Workflow Scripts
|
||||
|
||||
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
|
||||
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
|
||||
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
|
||||
|
||||
!!! note
|
||||
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
|
||||
+2
-3
@@ -24,13 +24,12 @@ FastVideo is an inference and post-training framework for diffusion models. It f
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- E2E post-training support
|
||||
- Data preprocessing pipeline for video data
|
||||
- Data preprocessing pipeline for video data.
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
@@ -43,7 +42,7 @@ Use the navigation menu on the left to explore different sections:
|
||||
|
||||
- **Getting Started**: Installation and quick start guides
|
||||
- **Inference**: Learn how to use FastVideo for video generation
|
||||
- **Training**: Data preprocessing and fine-tuning workflows
|
||||
- **Training**: Data preprocessing and fine-tuning workflows
|
||||
- **Distillation**: Post-training optimization techniques
|
||||
- **Sliding Tile Attention**: Advanced attention mechanisms
|
||||
- **Video Sparse Attention**: Efficient attention for video models
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# Seed Parameter Behavior in vLLM
|
||||
|
||||
## Overview
|
||||
|
||||
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
|
||||
|
||||
## Default Behavior
|
||||
|
||||
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
|
||||
|
||||
## Specifying a Seed
|
||||
|
||||
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
|
||||
|
||||
## Example Usage
|
||||
|
||||
### Without Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model without specifying a seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs different numbers across runs
|
||||
```
|
||||
|
||||
### Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model with a specific seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs the same number across runs
|
||||
```
|
||||
|
||||
## Important Notes
|
||||
|
||||
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
|
||||
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
|
||||
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
|
||||
|
||||
## Conclusion
|
||||
|
||||
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 98 KiB |
@@ -26,7 +26,7 @@ def main():
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
|
||||
sampling_param.num_frames = 81
|
||||
sampling_param.num_frames = 73
|
||||
sampling_param.width = 832
|
||||
sampling_param.height = 480
|
||||
sampling_param.seed = 1000
|
||||
|
||||
@@ -6,8 +6,10 @@ import time
|
||||
import json
|
||||
from pathlib import Path
|
||||
import tempfile
|
||||
from io import BytesIO
|
||||
|
||||
import gradio as gr
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
@@ -85,28 +87,39 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
|
||||
return None
|
||||
|
||||
|
||||
def encode_image_to_base64(image_path: str) -> str:
|
||||
"""Encode an image file to base64 string."""
|
||||
if not image_path or not os.path.exists(image_path):
|
||||
def encode_image_to_base64(image_input) -> str:
|
||||
"""Encode an image file path or in-memory image to a base64 string."""
|
||||
if image_input is None:
|
||||
return None
|
||||
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
|
||||
try:
|
||||
with open(image_path, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
if isinstance(image_input, str):
|
||||
if not os.path.exists(image_input):
|
||||
return None
|
||||
|
||||
with open(image_input, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
|
||||
ext = os.path.splitext(image_input)[1].lower()
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
elif isinstance(image_input, Image.Image):
|
||||
buffer = BytesIO()
|
||||
image_to_save = image_input.convert("RGB")
|
||||
image_to_save.save(buffer, format="PNG")
|
||||
image_bytes = buffer.getvalue()
|
||||
mime_type = 'image/png'
|
||||
else:
|
||||
return None
|
||||
|
||||
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
|
||||
|
||||
# Determine image type from extension
|
||||
ext = os.path.splitext(image_path)[1].lower()
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
|
||||
return f"data:{mime_type};base64,{image_base64}"
|
||||
|
||||
except Exception as e:
|
||||
@@ -426,7 +439,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
gr.Markdown("**Please make sure you upload a 480x832 image**")
|
||||
input_image = gr.Image(
|
||||
label="",
|
||||
type="filepath",
|
||||
type="pil",
|
||||
height=400,
|
||||
)
|
||||
|
||||
@@ -512,7 +525,17 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
if example_label and example_label in example_labels:
|
||||
index = example_labels.index(example_label)
|
||||
selected_prompt = examples[index]
|
||||
selected_image = example_images[index] if index < len(example_images) else None
|
||||
selected_image_path = example_images[index] if index < len(example_images) else None
|
||||
|
||||
if selected_image_path and os.path.exists(selected_image_path):
|
||||
try:
|
||||
with Image.open(selected_image_path) as img:
|
||||
selected_image = img.convert("RGB")
|
||||
except Exception:
|
||||
selected_image = None
|
||||
else:
|
||||
selected_image = None
|
||||
|
||||
return selected_prompt, selected_image
|
||||
return "", None
|
||||
|
||||
@@ -733,6 +756,7 @@ def main():
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
root_path="/gradio",
|
||||
allowed_paths=[
|
||||
os.path.abspath("outputs"),
|
||||
os.path.abspath("fastvideo-logos"),
|
||||
|
||||
@@ -5,12 +5,12 @@ These are e2e example scripts for finetuning Wan2.1 T2V 1.3B on the crush-smol d
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/download_dataset.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
|
||||
|
||||
@@ -54,7 +54,7 @@ validation_args=(
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
|
||||
@@ -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__"]
|
||||
@@ -1,10 +1,9 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
|
||||
"CosmosVideoConfig", "Cosmos25VideoConfig"
|
||||
"CosmosVideoConfig"
|
||||
]
|
||||
|
||||
@@ -1,181 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_transformer_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25ArchConfig(DiTArchConfig):
|
||||
"""Configuration for Cosmos 2.5 architecture (MiniTrainDIT)."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Remove "net." prefix and map official structure to FastVideo
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$":
|
||||
r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding: net.t_embedder.1.linear_1.weight -> time_embed.t_embedder.linear_1.weight
|
||||
r"^net\.t_embedder\.1\.linear_1\.(.*)$":
|
||||
r"time_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.t_embedder\.1\.linear_2\.(.*)$":
|
||||
r"time_embed.t_embedder.linear_2.\1",
|
||||
# Time embedding norm: net.t_embedding_norm.weight -> time_embed.norm.weight
|
||||
# Note: This also handles _extra_state if present
|
||||
r"^net\.t_embedding_norm\.(.*)$":
|
||||
r"time_embed.norm.\1",
|
||||
|
||||
# Cross-attention projection (optional): net.crossattn_proj.0.weight -> crossattn_proj.0.weight
|
||||
r"^net\.crossattn_proj\.0\.weight$":
|
||||
r"crossattn_proj.0.weight",
|
||||
r"^net\.crossattn_proj\.0\.bias$":
|
||||
r"crossattn_proj.0.bias",
|
||||
|
||||
# Transformer blocks: net.blocks.N -> transformer_blocks.N
|
||||
# Self-attention (self_attn -> attn1)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.weight",
|
||||
# RMSNorm _extra_state keys (internal PyTorch state, will be recomputed automatically)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_q._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_k._extra_state",
|
||||
|
||||
# Cross-attention (cross_attn -> attn2)
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.weight",
|
||||
# RMSNorm _extra_state keys for cross-attention
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_q._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_k._extra_state",
|
||||
|
||||
# MLP: net.blocks.N.mlp.layer1 -> transformer_blocks.N.mlp.fc_in
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
|
||||
# AdaLN-LoRA modulations: net.blocks.N.adaln_modulation_* -> transformer_blocks.N.adaln_modulation_*
|
||||
# These are now at the block level, not inside norm layers
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
|
||||
|
||||
# Layer norms: net.blocks.N.layer_norm_* -> transformer_blocks.N.norm*.norm
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_self_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm1.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_cross_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm2.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$":
|
||||
r"transformer_blocks.\1.norm3.norm._extra_state",
|
||||
|
||||
# Final layer: net.final_layer.linear -> final_layer.proj_out
|
||||
r"^net\.final_layer\.linear\.(.*)$":
|
||||
r"final_layer.proj_out.\1",
|
||||
# Final layer AdaLN-LoRA: net.final_layer.adaln_modulation -> final_layer.linear_*
|
||||
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$":
|
||||
r"final_layer.linear_1.\1",
|
||||
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$":
|
||||
r"final_layer.linear_2.\1",
|
||||
|
||||
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
|
||||
# - net.pos_embedder.* (seq, dim_spatial_range, dim_temporal_range) - These are computed dynamically
|
||||
# in FastVideo's Cosmos25RotaryPosEmbed forward() method, so they don't need to be loaded.
|
||||
# - net.accum_* keys (training metadata) - These are skipped during checkpoint loading.
|
||||
})
|
||||
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.\2",
|
||||
})
|
||||
|
||||
# Cosmos 2.5 specific config parameters
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128 # 2048 / 16
|
||||
num_layers: int = 28
|
||||
mlp_ratio: float = 4.0
|
||||
text_embed_dim: int = 1024
|
||||
adaln_lora_dim: int = 256
|
||||
use_adaln_lora: bool = True
|
||||
max_size: tuple[int, int, int] = (128, 240, 240)
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0) # T, H, W scaling
|
||||
concat_padding_mask: bool = True
|
||||
extra_pos_embed_type: str | None = None # "learnable" or None
|
||||
# Note: Official checkpoint has use_crossattn_projection=True with 100K-dim input from Qwen 7B.
|
||||
# When enabled, must provide 100,352-dim embeddings to match the projection layer in checkpoint.
|
||||
use_crossattn_projection: bool = False
|
||||
crossattn_proj_in_channels: int = 100352 # Qwen 7B embedding dimension
|
||||
rope_enable_fps_modulation: bool = True
|
||||
qk_norm: str = "rms_norm"
|
||||
eps: float = 1e-6
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25VideoConfig(DiTConfig):
|
||||
"""Configuration for Cosmos 2.5 video generation model."""
|
||||
arch_config: DiTArchConfig = field(default_factory=Cosmos25ArchConfig)
|
||||
prefix: str = "Cosmos25"
|
||||
@@ -45,7 +45,6 @@ class PipelineConfig:
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
disable_autocast: bool = False
|
||||
is_causal: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
|
||||
@@ -101,19 +101,12 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
def set_lora_weights(self,
|
||||
A: torch.Tensor,
|
||||
B: torch.Tensor,
|
||||
lora_alpha: float | None = None,
|
||||
training_mode: bool = False,
|
||||
lora_path: str | None = None) -> None:
|
||||
self.lora_A = torch.nn.Parameter(
|
||||
A) # share storage with weights in the pipeline
|
||||
self.lora_B = torch.nn.Parameter(B)
|
||||
self.disable_lora = False
|
||||
|
||||
# Store rank and alpha directly
|
||||
rank = A.shape[0] # rank is the first dimension of A
|
||||
self.lora_rank = rank
|
||||
self.lora_alpha = int(lora_alpha) if lora_alpha is not None else rank
|
||||
|
||||
if not training_mode:
|
||||
self.merge_lora_weights()
|
||||
self.lora_path = lora_path
|
||||
@@ -141,13 +134,8 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(
|
||||
get_local_torch_device()).full_tensor()
|
||||
|
||||
# Apply LoRA with alpha scaling
|
||||
lora_delta = (self.slice_lora_b_weights(self.lora_B).to(data)
|
||||
@ self.slice_lora_a_weights(self.lora_A).to(data))
|
||||
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
|
||||
lora_delta *= (self.lora_alpha / self.lora_rank)
|
||||
data += lora_delta
|
||||
data += (self.slice_lora_b_weights(self.lora_B).to(data)
|
||||
@ self.slice_lora_a_weights(self.lora_A).to(data))
|
||||
unsharded_base_layer.weight = nn.Parameter(data.to(current_device))
|
||||
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
|
||||
unsharded_base_layer.bias = nn.Parameter(
|
||||
@@ -166,13 +154,8 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
else:
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(get_local_torch_device())
|
||||
|
||||
# Apply LoRA with alpha scaling
|
||||
lora_delta = (self.slice_lora_b_weights(self.lora_B.to(data))
|
||||
@ self.slice_lora_a_weights(self.lora_A.to(data)))
|
||||
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
|
||||
lora_delta *= (self.lora_alpha / self.lora_rank)
|
||||
data += lora_delta
|
||||
data += \
|
||||
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
|
||||
self.base_layer.weight.data = data.to(current_device,
|
||||
non_blocking=True)
|
||||
|
||||
|
||||
@@ -1,961 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision import transforms
|
||||
|
||||
from fastvideo.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.rotary_embedding import apply_rotary_emb
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
|
||||
class Cosmos25PatchEmbed(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 patch embedding - converts video (B, C, T, H, W) to patches (B, T', H', W', D).
|
||||
Uses linear projection after rearranging patches.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.dim = in_channels * patch_size[0] * patch_size[1] * patch_size[2]
|
||||
|
||||
self.proj = nn.Linear(self.dim, out_channels, bias=False)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, C, T, H, W)
|
||||
Returns:
|
||||
(B, T', H', W', D) where T'=T//pt, H'=H//ph, W'=W//pw
|
||||
"""
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
|
||||
# Rearrange: b c (t pt) (h ph) (w pw) -> b t h w (c pt ph pw)
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, num_channels,
|
||||
num_frames // p_t, p_t,
|
||||
height // p_h, p_h,
|
||||
width // p_w, p_w
|
||||
)
|
||||
hidden_states = hidden_states.permute(0, 2, 4, 6, 1, 3, 5, 7)
|
||||
hidden_states = hidden_states.flatten(4, 7) # Flatten patch dimensions
|
||||
|
||||
# Project to model dimension
|
||||
hidden_states = self.proj(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25TimestepEmbedding(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 timestep embedding with AdaLN-LoRA support.
|
||||
Generates both standard embedding and AdaLN-LoRA parameters.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
use_adaln_lora: bool = True,
|
||||
adaln_lora_dim: int = 256,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.use_adaln_lora = use_adaln_lora
|
||||
|
||||
self.linear_1 = nn.Linear(in_features, out_features, bias=False)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
if use_adaln_lora:
|
||||
self.linear_2 = nn.Linear(out_features, 3 * out_features, bias=False)
|
||||
else:
|
||||
self.linear_2 = nn.Linear(out_features, out_features, bias=False)
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""
|
||||
Returns:
|
||||
emb: Standard embedding (B, T, D)
|
||||
adaln_lora: AdaLN-LoRA parameters (B, T, 3D) or None
|
||||
"""
|
||||
emb = self.linear_1(sample)
|
||||
emb = self.activation(emb)
|
||||
emb = self.linear_2(emb)
|
||||
|
||||
if self.use_adaln_lora:
|
||||
adaln_lora = emb # (B, T, 3D)
|
||||
emb_standard = sample # Use input as standard embedding
|
||||
else:
|
||||
emb_standard = emb
|
||||
adaln_lora = None
|
||||
|
||||
return emb_standard, adaln_lora
|
||||
|
||||
|
||||
class Cosmos25Embedding(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 timestep conditioning embedding.
|
||||
Generates sinusoidal embeddings and processes them through MLP.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
condition_dim: int,
|
||||
use_adaln_lora: bool = True,
|
||||
adaln_lora_dim: int = 256,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(embedding_dim, flip_sin_to_cos=True, downscale_freq_shift=0.0)
|
||||
self.t_embedder = Cosmos25TimestepEmbedding(
|
||||
embedding_dim,
|
||||
condition_dim,
|
||||
use_adaln_lora=use_adaln_lora,
|
||||
adaln_lora_dim=adaln_lora_dim,
|
||||
)
|
||||
self.norm = RMSNorm(embedding_dim, eps=1e-6)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""
|
||||
Args:
|
||||
timestep: (B, T) tensor of timesteps
|
||||
|
||||
Returns:
|
||||
embedded_timestep: Normalized timestep embedding (B, T, D)
|
||||
adaln_lora: AdaLN-LoRA parameters (B, T, 3D) or None
|
||||
"""
|
||||
# Handle 2D timestep input (B, T) like the official model
|
||||
assert timestep.ndim == 2, f"Expected 2D timestep, got {timestep.ndim}D with shape {timestep.shape}"
|
||||
B, T = timestep.shape
|
||||
|
||||
# Flatten for Timesteps layer which expects 1D, then reshape back
|
||||
timestep_flat = timestep.flatten() # (B*T,)
|
||||
timesteps_proj = self.time_proj(timestep_flat).type_as(hidden_states) # (B*T, D)
|
||||
timesteps_proj = timesteps_proj.reshape(B, T, -1) # (B, T, D)
|
||||
|
||||
embedded_timestep, adaln_lora = self.t_embedder(timesteps_proj)
|
||||
embedded_timestep = self.norm(embedded_timestep)
|
||||
|
||||
return embedded_timestep, adaln_lora
|
||||
|
||||
|
||||
class Cosmos25AdaLayerNormZero(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 Adaptive Layer Normalization with zero initialization and gate.
|
||||
This is a simplified version that expects pre-computed shift/scale/gate parameters.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.norm = nn.LayerNorm(in_features, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: Input tensor
|
||||
shift: Shift parameter for modulation
|
||||
scale: Scale parameter for modulation
|
||||
|
||||
Returns:
|
||||
normalized_hidden_states: Modulated normalized hidden states
|
||||
"""
|
||||
# Apply layer norm and modulation
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25SelfAttention(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 self-attention with QK normalization and RoPE.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: bool = True,
|
||||
eps: float = 1e-6,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
self.to_q = nn.Linear(dim, dim, bias=False)
|
||||
self.to_k = nn.Linear(dim, dim, bias=False)
|
||||
self.to_v = nn.Linear(dim, dim, bias=False)
|
||||
self.to_out = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
# Use DistributedAttention for flexible backend support (torch SDPA / FlashAttention)
|
||||
# For single-GPU (non-distributed), use LocalAttention to avoid distributed requirements
|
||||
if supported_attention_backends is None:
|
||||
supported_attention_backends = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
|
||||
# Always use DistributedAttention (requires distributed environment to be initialized)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix="self_attn"
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rope_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, S, D) where S = T*H*W
|
||||
rope_emb: Tuple of (cos, sin) for RoPE
|
||||
"""
|
||||
# Get QKV
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(hidden_states)
|
||||
value = self.to_v(hidden_states)
|
||||
|
||||
# Reshape for multi-head attention: (B, S, D) -> (B, S, H, D_h) -> (B, H, S, D_h)
|
||||
query = query.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
|
||||
key = key.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
|
||||
value = value.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
|
||||
|
||||
# Apply QK normalization
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
# Apply RoPE if provided (query/key are now in (B, H, S, D_h) format)
|
||||
if rope_emb is not None:
|
||||
cos, sin = rope_emb
|
||||
query = apply_rotary_emb(query, (cos, sin), use_real=True, use_real_unbind_dim=-2)
|
||||
key = apply_rotary_emb(key, (cos, sin), use_real=True, use_real_unbind_dim=-2)
|
||||
|
||||
# Attention computation using DistributedAttention or LocalAttention
|
||||
# Both expect (B, S, H, D_h), so transpose first
|
||||
query = query.transpose(1, 2) # (B, H, S, D_h) -> (B, S, H, D_h)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
|
||||
attn_output, _ = self.attn(query, key, value)
|
||||
# Reshape back: (B, S, H, D_h) -> (B, S, H*D_h)
|
||||
attn_output = attn_output.flatten(-2, -1)
|
||||
|
||||
# Output projection
|
||||
attn_output = self.to_out(attn_output)
|
||||
return attn_output
|
||||
|
||||
|
||||
class Cosmos25CrossAttention(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 cross-attention for text conditioning.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
cross_attention_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: bool = True,
|
||||
eps: float = 1e-6,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
self.to_q = nn.Linear(dim, dim, bias=False)
|
||||
self.to_k = nn.Linear(cross_attention_dim, dim, bias=False)
|
||||
self.to_v = nn.Linear(cross_attention_dim, dim, bias=False)
|
||||
self.to_out = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
if supported_attention_backends is None:
|
||||
supported_attention_backends = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
|
||||
# Use LocalAttention for cross-attention since text embeddings are not sharded
|
||||
# in sequence parallelism (replicated across ranks)
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, S, D)
|
||||
encoder_hidden_states: (B, N, D_text)
|
||||
"""
|
||||
# Get QKV
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
# Reshape for multi-head attention
|
||||
query = query.unflatten(-1, (self.num_heads, self.head_dim))
|
||||
key = key.unflatten(-1, (self.num_heads, self.head_dim))
|
||||
value = value.unflatten(-1, (self.num_heads, self.head_dim))
|
||||
|
||||
# Apply QK normalization
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
# LocalAttention expects (B, S, H, D_h), which is what we already have
|
||||
attn_output = self.attn(query, key, value)
|
||||
|
||||
# Reshape back: (B, S, H, D_h) -> (B, S, H*D_h)
|
||||
attn_output = attn_output.flatten(-2, -1)
|
||||
|
||||
# Output projection
|
||||
attn_output = self.to_out(attn_output)
|
||||
return attn_output
|
||||
|
||||
|
||||
class Cosmos25TransformerBlock(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 transformer block with self-attention, cross-attention, and MLP.
|
||||
Uses AdaLN-LoRA for conditioning.
|
||||
Matches the official architecture where modulation parameters are computed once per block.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
cross_attention_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
adaln_lora_dim: int = 256,
|
||||
use_adaln_lora: bool = True,
|
||||
qk_norm: bool = True,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
self.use_adaln_lora = use_adaln_lora
|
||||
|
||||
# Layer norms (no modulation logic inside)
|
||||
self.norm1 = Cosmos25AdaLayerNormZero(hidden_size)
|
||||
self.norm2 = Cosmos25AdaLayerNormZero(hidden_size)
|
||||
self.norm3 = Cosmos25AdaLayerNormZero(hidden_size)
|
||||
|
||||
# Attention and MLP layers
|
||||
self.attn1 = Cosmos25SelfAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_attention_heads,
|
||||
qk_norm=qk_norm,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.attn2 = Cosmos25CrossAttention(
|
||||
dim=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_heads=num_attention_heads,
|
||||
qk_norm=qk_norm,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.mlp = MLP(hidden_size, int(hidden_size * mlp_ratio), act_type="gelu", bias=False)
|
||||
|
||||
# AdaLN modulation layers (compute shift/scale/gate for each sub-layer)
|
||||
# These match the official model's adaln_modulation_* layers
|
||||
if use_adaln_lora:
|
||||
self.adaln_modulation_self_attn = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
|
||||
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
|
||||
)
|
||||
self.adaln_modulation_cross_attn = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
|
||||
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
|
||||
)
|
||||
self.adaln_modulation_mlp = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
|
||||
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
|
||||
)
|
||||
else:
|
||||
self.adaln_modulation_self_attn = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
|
||||
)
|
||||
self.adaln_modulation_cross_attn = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
|
||||
)
|
||||
self.adaln_modulation_mlp = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
adaln_lora: torch.Tensor | None = None,
|
||||
rope_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
extra_pos_emb: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D)
|
||||
encoder_hidden_states: (B, N, D_text)
|
||||
embedded_timestep: (B, T, D)
|
||||
adaln_lora: (B, T, 3D) AdaLN-LoRA parameters
|
||||
rope_emb: Tuple of (cos, sin) for RoPE
|
||||
extra_pos_emb: Optional learnable positional embeddings
|
||||
"""
|
||||
# Add extra positional embeddings if provided
|
||||
if extra_pos_emb is not None:
|
||||
hidden_states = hidden_states + extra_pos_emb
|
||||
|
||||
B, T, H, W, D = hidden_states.shape
|
||||
|
||||
# Step 1: Compute ALL modulation parameters once (matches official model)
|
||||
if self.use_adaln_lora and adaln_lora is not None:
|
||||
shift_self_attn, scale_self_attn, gate_self_attn = (
|
||||
self.adaln_modulation_self_attn(embedded_timestep) + adaln_lora
|
||||
).chunk(3, dim=-1)
|
||||
shift_cross_attn, scale_cross_attn, gate_cross_attn = (
|
||||
self.adaln_modulation_cross_attn(embedded_timestep) + adaln_lora
|
||||
).chunk(3, dim=-1)
|
||||
shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.adaln_modulation_mlp(embedded_timestep) + adaln_lora
|
||||
).chunk(3, dim=-1)
|
||||
else:
|
||||
shift_self_attn, scale_self_attn, gate_self_attn = self.adaln_modulation_self_attn(
|
||||
embedded_timestep
|
||||
).chunk(3, dim=-1)
|
||||
shift_cross_attn, scale_cross_attn, gate_cross_attn = self.adaln_modulation_cross_attn(
|
||||
embedded_timestep
|
||||
).chunk(3, dim=-1)
|
||||
shift_mlp, scale_mlp, gate_mlp = self.adaln_modulation_mlp(embedded_timestep).chunk(3, dim=-1)
|
||||
|
||||
# Reshape modulation parameters from (B, T, D) to (B, T, 1, 1, D) for broadcasting
|
||||
shift_self_attn = shift_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
scale_self_attn = scale_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
gate_self_attn = gate_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
|
||||
shift_cross_attn = shift_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
scale_cross_attn = scale_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
gate_cross_attn = gate_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
|
||||
shift_mlp = shift_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
scale_mlp = scale_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
gate_mlp = gate_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
|
||||
# Step 2: Self-attention block
|
||||
norm_hidden_states = self.norm1(hidden_states, shift_self_attn, scale_self_attn)
|
||||
# Flatten for attention: (B, T, H, W, D) -> (B, THW, D)
|
||||
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
|
||||
|
||||
attn_output = self.attn1(norm_hidden_states_flat, rope_emb=rope_emb)
|
||||
|
||||
# Reshape back and apply residual
|
||||
attn_output = attn_output.unflatten(1, (T, H, W)) # (B, T, H, W, D)
|
||||
hidden_states = hidden_states + gate_self_attn * attn_output
|
||||
|
||||
# Step 3: Cross-attention block
|
||||
norm_hidden_states = self.norm2(hidden_states, shift_cross_attn, scale_cross_attn)
|
||||
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
|
||||
|
||||
attn_output = self.attn2(
|
||||
norm_hidden_states_flat,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
attn_output = attn_output.unflatten(1, (T, H, W))
|
||||
hidden_states = hidden_states + gate_cross_attn * attn_output
|
||||
|
||||
# Step 4: MLP block
|
||||
norm_hidden_states = self.norm3(hidden_states, shift_mlp, scale_mlp)
|
||||
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
|
||||
|
||||
mlp_output = self.mlp(norm_hidden_states_flat)
|
||||
|
||||
mlp_output = mlp_output.unflatten(1, (T, H, W))
|
||||
hidden_states = hidden_states + gate_mlp * mlp_output
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25RotaryPosEmbed(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 3D Rotary Position Embedding with NTK-aware extrapolation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
max_size: tuple[int, int, int] = (128, 240, 240),
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
base_fps: int = 24,
|
||||
rope_scale: tuple[float, float, float] = (1.0, 1.0, 1.0),
|
||||
enable_fps_modulation: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.max_size = [size // patch for size, patch in zip(max_size, patch_size, strict=True)]
|
||||
self.patch_size = patch_size
|
||||
self.base_fps = base_fps
|
||||
self.enable_fps_modulation = enable_fps_modulation
|
||||
|
||||
# Split dimensions: 1/3 for T, 1/3 for H, 1/3 for W
|
||||
self.dim_h = hidden_size // 6 * 2
|
||||
self.dim_w = hidden_size // 6 * 2
|
||||
self.dim_t = hidden_size - self.dim_h - self.dim_w
|
||||
|
||||
# NTK-aware extrapolation factors
|
||||
self.h_ntk_factor = rope_scale[1] ** (self.dim_h / (self.dim_h - 2))
|
||||
self.w_ntk_factor = rope_scale[2] ** (self.dim_w / (self.dim_w - 2))
|
||||
self.t_ntk_factor = rope_scale[0] ** (self.dim_t / (self.dim_t - 2))
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, fps: int | None = None
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Generate 3D RoPE embeddings.
|
||||
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D) - patch-embedded features
|
||||
fps: Frames per second for temporal scaling
|
||||
|
||||
Returns:
|
||||
cos, sin: RoPE embeddings (THW, D)
|
||||
"""
|
||||
batch_size, T, H, W, input_dim = hidden_states.shape
|
||||
device = hidden_states.device
|
||||
|
||||
# T, H, W are already patch dimensions after patch_embed
|
||||
# No need to divide by patch_size
|
||||
|
||||
# Generate frequency scales with NTK
|
||||
h_theta = 10000.0 * self.h_ntk_factor
|
||||
w_theta = 10000.0 * self.w_ntk_factor
|
||||
t_theta = 10000.0 * self.t_ntk_factor
|
||||
|
||||
seq = torch.arange(max(self.max_size), device=device, dtype=torch.float32)
|
||||
|
||||
# Use self.dim_h/w/t which were set during initialization
|
||||
dim_h_range = torch.arange(0, self.dim_h, 2, device=device, dtype=torch.float32)[: (self.dim_h // 2)] / self.dim_h
|
||||
dim_w_range = torch.arange(0, self.dim_w, 2, device=device, dtype=torch.float32)[: (self.dim_w // 2)] / self.dim_w
|
||||
dim_t_range = torch.arange(0, self.dim_t, 2, device=device, dtype=torch.float32)[: (self.dim_t // 2)] / self.dim_t
|
||||
|
||||
h_spatial_freqs = 1.0 / (h_theta ** dim_h_range)
|
||||
w_spatial_freqs = 1.0 / (w_theta ** dim_w_range)
|
||||
temporal_freqs = 1.0 / (t_theta ** dim_t_range)
|
||||
|
||||
# Generate positional embeddings
|
||||
half_emb_h = torch.outer(seq[:H], h_spatial_freqs)
|
||||
half_emb_w = torch.outer(seq[:W], w_spatial_freqs)
|
||||
|
||||
if self.enable_fps_modulation and fps is not None:
|
||||
# Apply FPS scaling
|
||||
half_emb_t = torch.outer(seq[:T] / fps * self.base_fps, temporal_freqs)
|
||||
else:
|
||||
half_emb_t = torch.outer(seq[:T], temporal_freqs)
|
||||
|
||||
# Broadcast and concatenate embeddings
|
||||
emb_t = half_emb_t[:, None, None, :].repeat(1, H, W, 1)
|
||||
emb_h = half_emb_h[None, :, None, :].repeat(T, 1, W, 1)
|
||||
emb_w = half_emb_w[None, None, :, :].repeat(T, H, 1, 1)
|
||||
|
||||
# Concatenate [t, h, w, t, h, w] for sin/cos pairs
|
||||
freqs = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1)
|
||||
freqs = freqs.flatten(0, 2).float() # (THW, D)
|
||||
|
||||
cos = torch.cos(freqs) # (THW, D)
|
||||
sin = torch.sin(freqs) # (THW, D)
|
||||
|
||||
return cos, sin
|
||||
|
||||
|
||||
class Cosmos25LearnablePositionalEmbed(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 learnable absolute positional embeddings (optional).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
max_size: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.max_size = [size // patch for size, patch in zip(max_size, patch_size, strict=True)]
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
self.pos_emb_t = nn.Parameter(torch.zeros(self.max_size[0], hidden_size))
|
||||
self.pos_emb_h = nn.Parameter(torch.zeros(self.max_size[1], hidden_size))
|
||||
self.pos_emb_w = nn.Parameter(torch.zeros(self.max_size[2], hidden_size))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D)
|
||||
Returns:
|
||||
pos_emb: (B, T, H, W, D)
|
||||
"""
|
||||
B, T, H, W, D = hidden_states.shape
|
||||
|
||||
emb_t = self.pos_emb_t[:T][None, :, None, None, :].repeat(B, 1, H, W, 1)
|
||||
emb_h = self.pos_emb_h[:H][None, None, :, None, :].repeat(B, T, 1, W, 1)
|
||||
emb_w = self.pos_emb_w[:W][None, None, None, :, :].repeat(B, T, H, 1, 1)
|
||||
|
||||
emb = emb_t + emb_h + emb_w
|
||||
|
||||
# Normalize
|
||||
norm = torch.linalg.vector_norm(emb, dim=-1, keepdim=True, dtype=torch.float32)
|
||||
norm = torch.add(self.eps, norm, alpha=np.sqrt(norm.numel() / emb.numel()))
|
||||
return (emb / norm).type_as(hidden_states)
|
||||
|
||||
|
||||
class Cosmos25FinalLayer(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 final layer with AdaLN modulation and unpatchification.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
out_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
adaln_lora_dim: int = 256,
|
||||
use_adaln_lora: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.use_adaln_lora = use_adaln_lora
|
||||
|
||||
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
if use_adaln_lora:
|
||||
self.linear_1 = nn.Linear(hidden_size, adaln_lora_dim, bias=False)
|
||||
self.linear_2 = nn.Linear(adaln_lora_dim, 2 * hidden_size, bias=False)
|
||||
else:
|
||||
self.linear_1 = nn.Identity()
|
||||
self.linear_2 = nn.Linear(hidden_size, 2 * hidden_size, bias=False)
|
||||
|
||||
# Output projection
|
||||
output_dim = out_channels * patch_size[0] * patch_size[1] * patch_size[2]
|
||||
self.proj_out = nn.Linear(hidden_size, output_dim, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
adaln_lora: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D)
|
||||
embedded_timestep: (B, T, D) or (B, D)
|
||||
adaln_lora: (B, T, 3D) or None
|
||||
"""
|
||||
# Generate modulation parameters
|
||||
embedded_timestep = self.activation(embedded_timestep)
|
||||
embedded_timestep = self.linear_1(embedded_timestep)
|
||||
embedded_timestep = self.linear_2(embedded_timestep)
|
||||
|
||||
if self.use_adaln_lora and adaln_lora is not None:
|
||||
# Use first 2*hidden_size elements for shift/scale
|
||||
embedded_timestep = embedded_timestep + adaln_lora[..., : 2 * self.hidden_size]
|
||||
|
||||
shift, scale = embedded_timestep.chunk(2, dim=-1)
|
||||
|
||||
# Apply normalization and modulation
|
||||
hidden_states = self.norm(hidden_states)
|
||||
|
||||
# Reshape for broadcasting if needed
|
||||
if embedded_timestep.ndim == 2:
|
||||
shift, scale = (x.unsqueeze(1) for x in (shift, scale))
|
||||
elif embedded_timestep.ndim == 3 and hidden_states.ndim == 5:
|
||||
shift, scale = (x.unsqueeze(2).unsqueeze(2) for x in (shift, scale))
|
||||
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
# Project to output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25Transformer3DModel(BaseDiT):
|
||||
"""
|
||||
COSMOS 2.5 DiT - MiniTrainDIT architecture adapted for FastVideo.
|
||||
|
||||
Key features:
|
||||
- AdaLN-LoRA conditioning
|
||||
- 3D RoPE with NTK-aware extrapolation
|
||||
- Optional learnable positional embeddings
|
||||
- QK normalization
|
||||
- Cross-attention projection (optional)
|
||||
"""
|
||||
|
||||
_fsdp_shard_conditions = Cosmos25VideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = Cosmos25VideoConfig()._compile_conditions
|
||||
param_names_mapping = Cosmos25VideoConfig().param_names_mapping
|
||||
lora_param_names_mapping = Cosmos25VideoConfig().lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: Cosmos25VideoConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = inner_dim
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.max_size = config.max_size
|
||||
self.rope_scale = config.rope_scale
|
||||
self.concat_padding_mask = config.concat_padding_mask
|
||||
self.use_adaln_lora = getattr(config, "use_adaln_lora", True)
|
||||
self.adaln_lora_dim = getattr(config, "adaln_lora_dim", 256)
|
||||
self.extra_pos_embed_type = getattr(config, "extra_pos_embed_type", None)
|
||||
self.use_crossattn_projection = getattr(config, "use_crossattn_projection", False)
|
||||
|
||||
# 1. Patch Embedding
|
||||
# Account for: VAE channels + condition_mask (1) + padding_mask (1 if concat_padding_mask)
|
||||
patch_embed_in_channels = config.in_channels # Base VAE channels (16)
|
||||
patch_embed_in_channels += 1 # Always add 1 for condition_mask
|
||||
if config.concat_padding_mask:
|
||||
patch_embed_in_channels += 1 # Add 1 for padding_mask
|
||||
# Total: 16 + 1 + 1 = 18 (with concat_padding_mask=True)
|
||||
|
||||
self.patch_embed = Cosmos25PatchEmbed(
|
||||
patch_embed_in_channels, inner_dim, config.patch_size
|
||||
)
|
||||
|
||||
# 2. Positional Embeddings
|
||||
self.rope = Cosmos25RotaryPosEmbed(
|
||||
hidden_size=config.attention_head_dim,
|
||||
max_size=config.max_size,
|
||||
patch_size=config.patch_size,
|
||||
rope_scale=config.rope_scale,
|
||||
enable_fps_modulation=getattr(config, "rope_enable_fps_modulation", True),
|
||||
)
|
||||
|
||||
self.learnable_pos_embed = None
|
||||
if self.extra_pos_embed_type == "learnable":
|
||||
self.learnable_pos_embed = Cosmos25LearnablePositionalEmbed(
|
||||
hidden_size=inner_dim,
|
||||
max_size=config.max_size,
|
||||
patch_size=config.patch_size,
|
||||
)
|
||||
|
||||
# 3. Time Embedding
|
||||
self.time_embed = Cosmos25Embedding(
|
||||
inner_dim,
|
||||
inner_dim,
|
||||
use_adaln_lora=self.use_adaln_lora,
|
||||
adaln_lora_dim=self.adaln_lora_dim,
|
||||
)
|
||||
|
||||
# 4. Cross-attention projection (optional)
|
||||
if self.use_crossattn_projection:
|
||||
crossattn_proj_in_channels = getattr(config, "crossattn_proj_in_channels", config.text_embed_dim)
|
||||
self.crossattn_proj = nn.Sequential(
|
||||
nn.Linear(crossattn_proj_in_channels, config.text_embed_dim, bias=True),
|
||||
nn.GELU(),
|
||||
)
|
||||
|
||||
# 5. Transformer Blocks
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
Cosmos25TransformerBlock(
|
||||
num_attention_heads=config.num_attention_heads,
|
||||
attention_head_dim=config.attention_head_dim,
|
||||
cross_attention_dim=config.text_embed_dim,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
adaln_lora_dim=self.adaln_lora_dim,
|
||||
use_adaln_lora=self.use_adaln_lora,
|
||||
qk_norm=(config.qk_norm == "rms_norm"),
|
||||
supported_attention_backends=config._supported_attention_backends,
|
||||
)
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
# 6. Final Layer
|
||||
self.final_layer = Cosmos25FinalLayer(
|
||||
hidden_size=inner_dim,
|
||||
out_channels=config.out_channels,
|
||||
patch_size=config.patch_size,
|
||||
adaln_lora_dim=self.adaln_lora_dim,
|
||||
use_adaln_lora=self.use_adaln_lora,
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
fps: int | None = None,
|
||||
condition_mask: torch.Tensor | None = None,
|
||||
padding_mask: torch.Tensor | None = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, C, T, H, W) latent video
|
||||
timestep: (B,) or (B, T) diffusion timesteps
|
||||
encoder_hidden_states: (B, N, D_text) text embeddings
|
||||
attention_mask: Optional attention mask
|
||||
fps: Frames per second
|
||||
condition_mask: (B, 1, T, H, W) conditioning mask
|
||||
padding_mask: (B, 1, H, W) padding mask
|
||||
"""
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
|
||||
# 1. Concatenate condition mask if provided
|
||||
if condition_mask is not None:
|
||||
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
|
||||
|
||||
# 2. Concatenate padding mask if needed
|
||||
if self.concat_padding_mask and padding_mask is not None:
|
||||
padding_mask = transforms.functional.resize(
|
||||
padding_mask,
|
||||
list(hidden_states.shape[-2:]),
|
||||
interpolation=transforms.InterpolationMode.NEAREST,
|
||||
)
|
||||
hidden_states = torch.cat(
|
||||
[hidden_states, padding_mask.unsqueeze(2).repeat(1, 1, num_frames, 1, 1)],
|
||||
dim=1,
|
||||
)
|
||||
|
||||
# 3. Patchify input
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
|
||||
hidden_states = self.patch_embed(hidden_states) # (B, T', H', W', D)
|
||||
|
||||
|
||||
# 4. Generate RoPE embeddings (after patchify, using patch dimensions)
|
||||
rope_emb = self.rope(hidden_states, fps=fps)
|
||||
|
||||
# 5. Generate learnable positional embeddings (if used)
|
||||
extra_pos_emb = None
|
||||
if self.learnable_pos_embed is not None:
|
||||
extra_pos_emb = self.learnable_pos_embed(hidden_states)
|
||||
|
||||
# 6. Timestep embeddings
|
||||
# Official model expects timestep in (B, T) format, so ensure it has 2D shape
|
||||
if timestep.ndim == 1:
|
||||
# Scalar timestep per sample: (B,) -> (B, 1)
|
||||
timestep = timestep.unsqueeze(1)
|
||||
elif timestep.ndim == 2:
|
||||
# Already in (B, T) format
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Unsupported timestep shape: {timestep.shape}")
|
||||
|
||||
# Now timestep is always (B, T), pass directly to time_embed
|
||||
embedded_timestep, adaln_lora = self.time_embed(hidden_states, timestep)
|
||||
|
||||
# 7. Apply cross-attention projection (if used)
|
||||
if self.use_crossattn_projection:
|
||||
encoder_hidden_states = self.crossattn_proj(encoder_hidden_states)
|
||||
|
||||
|
||||
|
||||
|
||||
# Prepare attention mask
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) # (B, 1, 1, N)
|
||||
|
||||
# 8. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block,
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
embedded_timestep,
|
||||
adaln_lora,
|
||||
rope_emb,
|
||||
extra_pos_emb,
|
||||
attention_mask,
|
||||
)
|
||||
else:
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
embedded_timestep=embedded_timestep,
|
||||
adaln_lora=adaln_lora,
|
||||
rope_emb=rope_emb,
|
||||
extra_pos_emb=extra_pos_emb,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
# 9. Final layer - output norm & projection
|
||||
hidden_states = self.final_layer(hidden_states, embedded_timestep, adaln_lora)
|
||||
|
||||
# 10. Unpatchify: (B, T', H', W', P) -> (B, C, T, H, W)
|
||||
# After unflatten: (B, T', H', W', p_t, p_h, p_w, C) with dims [0,1,2,3,4,5,6,7]
|
||||
hidden_states = hidden_states.unflatten(-1, (p_t, p_h, p_w, self.out_channels))
|
||||
# Permute to: (B, C, T', p_t, H', p_h, W', p_w)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
# Flatten pairs to get (B, C, T, H, W)
|
||||
hidden_states = hidden_states.flatten(2, 3).flatten(3, 4).flatten(4, 5)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -59,8 +59,7 @@ class WanDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None),
|
||||
use_btchw_layout=True))
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DmdDenoisingStage(
|
||||
|
||||
@@ -62,8 +62,7 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer"),
|
||||
use_btchw_layout=True))
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
@@ -28,8 +28,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
TODO: support training.
|
||||
"""
|
||||
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
|
||||
dict
|
||||
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
|
||||
dict) # state dicts of loaded lora adapters
|
||||
cur_adapter_name: str = ""
|
||||
cur_adapter_path: str = ""
|
||||
lora_layers: dict[str, BaseLayerWithLoRA] = {}
|
||||
@@ -184,26 +183,11 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
lora_param_names_mapping_fn = get_param_names_mapping(
|
||||
self.modules["transformer"].lora_param_names_mapping)
|
||||
|
||||
# Extract alpha values and weights in a single pass
|
||||
to_merge_params: defaultdict[Hashable,
|
||||
dict[Any, Any]] = defaultdict(dict)
|
||||
for name, weight in lora_state_dict.items():
|
||||
# Extract weights (lora_A, lora_B, and lora_alpha)
|
||||
name = name.replace("diffusion_model.", "")
|
||||
name = name.replace(".weight", "")
|
||||
|
||||
if "lora_alpha" in name:
|
||||
# Store alpha with minimal mapping - same processing as lora_A/lora_B
|
||||
# but store in lora_adapters with ".lora_alpha" suffix
|
||||
layer_name = name.replace(".lora_alpha", "")
|
||||
layer_name, _, _ = lora_param_names_mapping_fn(layer_name)
|
||||
target_name, _, _ = param_names_mapping_fn(layer_name)
|
||||
# Store alpha alongside weights with same target_name base
|
||||
alpha_key = target_name + ".lora_alpha"
|
||||
self.lora_adapters[lora_nickname][alpha_key] = weight.item(
|
||||
) if weight.numel() == 1 else float(weight.mean())
|
||||
continue
|
||||
|
||||
name, _, _ = lora_param_names_mapping_fn(name)
|
||||
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
|
||||
name)
|
||||
@@ -241,20 +225,11 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
for name, layer in self.lora_layers.items():
|
||||
lora_A_name = name + ".lora_A"
|
||||
lora_B_name = name + ".lora_B"
|
||||
lora_alpha_name = name + ".lora_alpha"
|
||||
if lora_A_name in self.lora_adapters[lora_nickname]\
|
||||
and lora_B_name in self.lora_adapters[lora_nickname]:
|
||||
# Get alpha value for this layer (defaults to None if not present)
|
||||
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
|
||||
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
|
||||
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
|
||||
alpha = self.lora_adapters[lora_nickname].get(
|
||||
lora_alpha_name) if adapter_updated else None
|
||||
|
||||
layer.set_lora_weights(
|
||||
lora_A,
|
||||
lora_B,
|
||||
lora_alpha=alpha,
|
||||
self.lora_adapters[lora_nickname][lora_A_name],
|
||||
self.lora_adapters[lora_nickname][lora_B_name],
|
||||
training_mode=self.fastvideo_args.training_mode,
|
||||
lora_path=lora_path)
|
||||
adapted_count += 1
|
||||
|
||||
@@ -115,7 +115,7 @@ class ForwardBatch:
|
||||
|
||||
# Latent tensors
|
||||
latents: torch.Tensor | None = None
|
||||
raw_latent_shape: tuple[int, ...] | None = None
|
||||
raw_latent_shape: torch.Tensor | 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: tuple[int, ...] | None = None
|
||||
raw_latent_shape: torch.Tensor | None = None
|
||||
noise_latents: torch.Tensor | None = None
|
||||
encoder_hidden_states: torch.Tensor | None = None
|
||||
encoder_attention_mask: torch.Tensor | None = None
|
||||
|
||||
@@ -85,10 +85,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
|
||||
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_kwargs: dict = {}
|
||||
@@ -142,18 +142,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
pos_start_base = 0
|
||||
|
||||
# Determine block sizes
|
||||
if t % self.num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
|
||||
)
|
||||
num_blocks = t // self.num_frames_per_block
|
||||
block_sizes = [self.num_frames_per_block] * num_blocks
|
||||
block_sizes = [self.num_frames_per_block] * 7
|
||||
block_sizes[0] = 1
|
||||
start_index = 0
|
||||
|
||||
# For now hardcode the first block to be 1 frame assuming the model is Wan2.2-MoE
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
first_frame_latent = None
|
||||
if batch.pil_image is not None:
|
||||
# Causal video gen directly replaces the first frame of the latent with
|
||||
@@ -400,10 +392,6 @@ 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
|
||||
|
||||
@@ -494,4 +482,4 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
return result
|
||||
|
||||
@@ -1085,6 +1085,7 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
# Get latents and embeddings
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
|
||||
video_raw_latent_shape = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
|
||||
@@ -5,7 +5,6 @@ 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
|
||||
@@ -14,7 +13,6 @@ 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__)
|
||||
|
||||
@@ -110,25 +108,27 @@ 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 = 480 * 832
|
||||
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 = 720 * 1280
|
||||
# ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||
|
||||
scale = max(ow / iw, oh / ih)
|
||||
img = img.resize((round(iw * scale), round(ih * scale)),
|
||||
Image.LANCZOS)
|
||||
|
||||
# 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)
|
||||
# scale = max(ow / iw, oh / ih)
|
||||
# img = img.resize((round(iw * scale), round(ih * scale)),
|
||||
# Image.LANCZOS)
|
||||
# logger.info("resized img height: %s, img width: %s", img.height,
|
||||
# img.width)
|
||||
|
||||
# # center-crop
|
||||
# x1 = (img.width - ow) // 2
|
||||
# y1 = (img.height - oh) // 2
|
||||
# img = img.crop((x1, y1, x1 + ow, y1 + oh))
|
||||
# assert img.width == ow and img.height == oh
|
||||
logger.info("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,14 +28,10 @@ class LatentPreparationStage(PipelineStage):
|
||||
denoised during the diffusion process.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
scheduler,
|
||||
transformer,
|
||||
use_btchw_layout: bool = False) -> None:
|
||||
def __init__(self, scheduler, transformer) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.transformer = transformer
|
||||
self.use_btchw_layout = use_btchw_layout
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -82,29 +78,15 @@ class LatentPreparationStage(PipelineStage):
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
# Calculate latent shape
|
||||
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
|
||||
shape = (
|
||||
batch_size,
|
||||
self.transformer.num_channels_latents,
|
||||
num_frames,
|
||||
height // fastvideo_args.pipeline_config.vae_config.arch_config.
|
||||
spatial_compression_ratio,
|
||||
width // fastvideo_args.pipeline_config.vae_config.arch_config.
|
||||
spatial_compression_ratio,
|
||||
)
|
||||
|
||||
# Validate generator if it's a list
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
@@ -126,7 +108,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
# Update batch with prepared latents
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = bcthw_shape
|
||||
batch.raw_latent_shape = latents.shape
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
@@ -22,8 +22,8 @@ os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
|
||||
)
|
||||
local_dir=os.path.join(
|
||||
"data", BASE_MODEL_PATH))
|
||||
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder_2")
|
||||
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer_2")
|
||||
|
||||
@@ -130,6 +130,17 @@ def test_clip_encoder():
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info("Maximum difference in last hidden states: %f",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %f",
|
||||
mean_diff_hidden.item())
|
||||
|
||||
# Compare pooler outputs
|
||||
pooler_output1 = outputs1.pooler_output
|
||||
pooler_output2 = outputs2.pooler_output
|
||||
@@ -137,5 +148,22 @@ def test_clip_encoder():
|
||||
assert pooler_output1.shape == pooler_output2.shape, \
|
||||
f"Pooler output shapes don't match: {pooler_output1.shape} vs {pooler_output2.shape}"
|
||||
|
||||
assert_close(pooler_output1, pooler_output2, atol=1e-2, rtol=1e-3)
|
||||
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-2, rtol=1e-3)
|
||||
max_diff_pooler = torch.max(
|
||||
torch.abs(pooler_output1 - pooler_output2))
|
||||
mean_diff_pooler = torch.mean(
|
||||
torch.abs(pooler_output1 - pooler_output2))
|
||||
|
||||
logger.info("Maximum difference in pooler outputs: %f",
|
||||
max_diff_pooler.item())
|
||||
logger.info("Mean difference in pooler outputs: %f",
|
||||
mean_diff_pooler.item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-2, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert mean_diff_pooler < 1e-2, \
|
||||
f"Pooler outputs differ significantly: mean diff = {mean_diff_pooler.item()}"
|
||||
assert max_diff_hidden < 1e-1, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
assert max_diff_pooler < 2e-2, \
|
||||
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
|
||||
|
||||
@@ -22,8 +22,8 @@ os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
|
||||
)
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder")
|
||||
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer")
|
||||
|
||||
@@ -68,7 +68,8 @@ def test_llama_encoder():
|
||||
logger.info("Model1 has %d parameters", len(params1))
|
||||
logger.info("Model2 has %d parameters", len(params2))
|
||||
|
||||
|
||||
# Compare a few key parameters
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
device = model1.embed_tokens.weight.device
|
||||
assert torch.allclose(model1.embed_tokens.weight,
|
||||
@@ -77,18 +78,6 @@ def test_llama_encoder():
|
||||
"layers.{}.input_layernorm.weight",
|
||||
"layers.{}.post_attention_layernorm.weight"
|
||||
]
|
||||
for layer_idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(layer_idx)
|
||||
name2 = w.format(layer_idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
if "gate_up" in name2:
|
||||
# print("skipping gate_up")
|
||||
continue
|
||||
p1 = p1.to_local().to(device) if isinstance(p1, DTensor) else p1.to(device)
|
||||
p2 = p2.to_local().to(device) if isinstance(p2, DTensor) else p2.to(device)
|
||||
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
|
||||
|
||||
for name1, param1 in sorted(params1.items()):
|
||||
name2 = name1
|
||||
@@ -150,4 +139,19 @@ def test_llama_encoder():
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-1, rtol=1e-4)
|
||||
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info("Maximum difference in last hidden states: %f",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %f",
|
||||
mean_diff_hidden.item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-2, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert max_diff_hidden < 1e-1, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
|
||||
@@ -133,7 +133,24 @@ def test_t5_encoder(t5_model_paths):
|
||||
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
|
||||
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
|
||||
|
||||
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-4, rtol=1e-4)
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info("Maximum difference in last hidden states: %s",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %s",
|
||||
mean_diff_hidden.item())
|
||||
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert max_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
@@ -235,4 +252,18 @@ def test_t5_large_encoder(t5_large_model_paths):
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
assert_close(last_hidden_state1, last_hidden_state2, atol=1e-4, rtol=1e-4)
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info("Maximum difference in last hidden states: %s",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %s",
|
||||
mean_diff_hidden.item())
|
||||
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert max_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -53,38 +52,12 @@ LORA_CONFIGS = [
|
||||
"negative_prompt": "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
"ssim_threshold": 0.79
|
||||
}
|
||||
# TODO: Add a LoRA with lora_alpha values to test alpha scaling
|
||||
#
|
||||
# Context: This change is mainly for an in-progress ticket porting over LongCat-Video,
|
||||
# where they used an alpha value that is two times smaller than their rank. This fix
|
||||
# ensures that LoRA weights are correctly scaled by the alpha/rank ratio when merged.
|
||||
#
|
||||
# Issue: Currently, we cannot add a test for LoRA adapters with alpha values because:
|
||||
# - The existing public LoRAs for Wan-AI/Wan2.1-T2V-1.3B-Diffusers don't store lora_alpha
|
||||
# - No publicly available LoRA for this model includes lora_alpha tensors in their weights
|
||||
# - This is why the alpha/rank scaling bug wasn't caught by existing tests
|
||||
#
|
||||
# The fix has been validated with:
|
||||
# - LongCat-Video distilled LoRA (which includes alpha values)
|
||||
# - Manual testing shows correct alpha/rank scaling behavior
|
||||
# - Backward compatibility confirmed with LoRAs without alpha values
|
||||
#
|
||||
# Future work:
|
||||
# - Add a synthetic LoRA test fixture with alpha values when feasible
|
||||
# - Or wait for public Wan LoRAs with alpha to become available
|
||||
]
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WAN_LORA_PARAMS,
|
||||
}
|
||||
|
||||
def _sanitize_filename_component(name: str) -> str:
|
||||
"""Sanitize filename to remove invalid characters (same logic as VideoGenerator)"""
|
||||
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
|
||||
sanitized = sanitized.strip().strip('.')
|
||||
sanitized = re.sub(r'\s+', ' ', sanitized)
|
||||
return sanitized or "video"
|
||||
|
||||
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
|
||||
def test_merge_lora_weights(model_id):
|
||||
lora_config = LORA_CONFIGS[0] # test only one
|
||||
@@ -164,16 +137,14 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
generation_kwargs["negative_prompt"] = lora_config["negative_prompt"]
|
||||
|
||||
generator.set_lora_adapter(lora_nickname=lora_nickname, lora_path=lora_path)
|
||||
# Sanitize the filename before adding .mp4 extension to match VideoGenerator's behavior
|
||||
output_video_name = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
|
||||
output_video_name = _sanitize_filename_component(output_video_name)
|
||||
generated_video_path = os.path.join(output_dir, f"{output_video_name}.mp4")
|
||||
generation_kwargs["output_path"] = generated_video_path
|
||||
generation_kwargs["output_path"] = output_dir
|
||||
generation_kwargs["output_video_name"] = output_video_name
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
assert os.path.exists(
|
||||
generated_video_path), f"Output video was not generated at {generated_video_path}"
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, 'L40S_reference_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
|
||||
|
||||
@@ -182,25 +153,13 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}")
|
||||
|
||||
# Find the matching reference video - try exact match first, then fuzzy match
|
||||
# The reference might have different sanitization (e.g., trailing spaces)
|
||||
# Find the matching reference video for the switched LoRA
|
||||
reference_video_name = None
|
||||
unsanitized_prefix = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
|
||||
|
||||
|
||||
for filename in os.listdir(reference_folder):
|
||||
if not filename.endswith('.mp4'):
|
||||
continue
|
||||
|
||||
# Try exact match with sanitized name
|
||||
if filename.startswith(output_video_name):
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
# Try match with unsanitized prefix (for legacy reference videos)
|
||||
# Remove .mp4 and compare the base names after sanitization
|
||||
base_filename = filename[:-4] # Remove .mp4
|
||||
if _sanitize_filename_component(base_filename) == output_video_name:
|
||||
reference_video_name = filename
|
||||
# Check if the filename starts with the expected output_video_name and ends with .mp4
|
||||
if filename.startswith(output_video_name) and filename.endswith('.mp4'):
|
||||
reference_video_name = filename # Remove .mp4 extension to match the logic below
|
||||
break
|
||||
|
||||
if not reference_video_name:
|
||||
@@ -208,6 +167,7 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
raise FileNotFoundError(f"Reference video missing for adapter {lora_path}")
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
generated_video_path = os.path.join(output_dir, output_video_name + ".mp4")
|
||||
|
||||
logger.info(
|
||||
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
|
||||
|
||||
@@ -1,60 +0,0 @@
|
||||
"""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")
|
||||
@@ -74,9 +74,9 @@ def run_vae_tests():
|
||||
def run_transformer_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700)
|
||||
def run_ssim_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
run_test("pytest ./fastvideo/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests():
|
||||
@@ -125,7 +125,3 @@ 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")
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=8
|
||||
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
|
||||
--master_port 29503 \
|
||||
tp_example.py
|
||||
|
||||
|
||||
|
||||
num_gpus=2
|
||||
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
|
||||
--master_port 29503 \
|
||||
fastvideo/tests/test_hunyuanvideo_load.py --sequence_model_parallel_size $num_gpus
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=1 --master_port 29503 fastvideo/tests/test_llama_encoder.py
|
||||
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
torchrun --nnodes=1 --nproc_per_node=1 --master_port 29503 fastvideo/tests/test_clip_encoder.py
|
||||
@@ -0,0 +1,164 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.distributed import (maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='T5 Encoder Test')
|
||||
parser.add_argument('--model_path', type=str, default="google/umt5-xxl")
|
||||
parser.add_argument(
|
||||
'--dit-precision',
|
||||
type=str,
|
||||
default="float32",
|
||||
help='Precision to use for the model (float32, float16, bfloat16)')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def test_t5_encoder():
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
# Initialize the two model implementations
|
||||
model_path = "/workspace/data/Wan2.1-T2V-1.3B-Diffusers/text_encoder"
|
||||
tokenizer_path = "/workspace/data/Wan2.1-T2V-1.3B-Diffusers/tokenizer"
|
||||
|
||||
hf_config = AutoConfig.from_pretrained(model_path)
|
||||
print(hf_config)
|
||||
precision = torch.float16 # It must be float16 because the weight loader is in float16
|
||||
# Load our implementation using the loader from text_encoder/__init__.py
|
||||
model1 = UMT5EncoderModel.from_pretrained(model_path).to(precision).to(
|
||||
device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load_model(model_path, hf_config, device)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
model2 = model2.to(precision)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Check number of parameters
|
||||
logger.info(f"Model1 has {len(params1)} parameters")
|
||||
logger.info(f"Model2 has {len(params2)} parameters")
|
||||
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.weight"]
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
for l in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(l)
|
||||
name2 = w.format(l)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
assert p1.dtype == p2.dtype
|
||||
try:
|
||||
logger.info(f"Parameter: {name1} vs {name2}")
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
except Exception as e:
|
||||
logger.info(f"Error comparing {name1} and {name2}: {e}")
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
print("Model 1 Weight Sum: ", weight_sum_model1)
|
||||
print("Model 1 Weight Mean: ", weight_mean_model1)
|
||||
|
||||
total_params = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params
|
||||
print("Model 2 Weight Sum: ", weight_sum_model2)
|
||||
print("Model 2 Weight Mean: ", weight_mean_model2)
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
"Once upon a time", "The quick brown fox jumps over",
|
||||
"In a galaxy far, far away"
|
||||
]
|
||||
|
||||
logger.info("Testing T5 encoder with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info(f"Testing prompt: '{prompt}'")
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens = tokenizer(prompt,
|
||||
padding="max_length",
|
||||
max_length=512,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(device)
|
||||
|
||||
# Get outputs from our implementation
|
||||
# filter out padding input_ids
|
||||
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
|
||||
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
|
||||
outputs1 = model1(input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True).last_hidden_state
|
||||
print("--------------------------------")
|
||||
logger.info("Testing model2")
|
||||
|
||||
# Get outputs from HuggingFace implementation
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
)
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
|
||||
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info(
|
||||
f"Maximum difference in last hidden states: {max_diff_hidden.item()}"
|
||||
)
|
||||
logger.info(
|
||||
f"Mean difference in last hidden states: {mean_diff_hidden.item()}"
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Test passed! Both T5 encoder implementations produce similar outputs.")
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_t5_encoder()
|
||||
@@ -0,0 +1,123 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.wanvae import AutoencoderKLWan as MyWanVAE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def test_wan_vae():
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
# Initialize the two model implementations
|
||||
path = "/workspace/data/Wan2.1-T2V-1.3B-Diffusers/vae"
|
||||
config_path = os.path.join(path, "config.json")
|
||||
config = json.load(open(config_path))
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
model1 = MyWanVAE(**config).to(torch.bfloat16)
|
||||
|
||||
model2 = AutoencoderKLWan(**config).to(torch.bfloat16)
|
||||
|
||||
loaded = load_file(os.path.join(path,
|
||||
"diffusion_pytorch_model.safetensors"))
|
||||
model1.load_state_dict(loaded)
|
||||
model2.load_state_dict(loaded)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Move to GPU
|
||||
model1 = model1.to(device)
|
||||
model2 = model2.to(device)
|
||||
|
||||
# model1.enable_tiling(
|
||||
# tile_sample_min_height=32,
|
||||
# tile_sample_min_width=32,
|
||||
# tile_sample_min_num_frames=8,
|
||||
# tile_sample_stride_height=16,
|
||||
# tile_sample_stride_width=16,
|
||||
# tile_sample_stride_num_frames=4
|
||||
# )
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
|
||||
# Video input [B, C, T, H, W]
|
||||
input_tensor = torch.randn(batch_size,
|
||||
3,
|
||||
81,
|
||||
32,
|
||||
32,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
latent_tensor = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
32,
|
||||
32,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Test encoding
|
||||
logger.info("Testing encoding...")
|
||||
latent2 = model2.encode(input_tensor).latent_dist.mean
|
||||
print("--------------------------------")
|
||||
latent1 = model1.encode(input_tensor).mean
|
||||
# Check if latents have the same shape
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
# Check if latents are similar
|
||||
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
|
||||
logger.info(
|
||||
f"Maximum difference between encoded latents: {max_diff_encode.item()}"
|
||||
)
|
||||
logger.info(
|
||||
f"Mean difference between encoded latents: {mean_diff_encode.item()}"
|
||||
)
|
||||
assert mean_diff_encode < 5e-1, f"Encoded latents differ significantly: mean diff = {mean_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
latent1 = latent2 = latent_tensor
|
||||
latents_mean = (torch.tensor(model2.config.latents_mean).view(
|
||||
1, model2.config.z_dim, 1, 1, 1).to(latent2.device, latent2.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
|
||||
1, model2.config.z_dim, 1, 1, 1).to(latent2.device, latent2.dtype)
|
||||
latent2 = latent2 / latents_std + latents_mean
|
||||
output1 = model1.decode(latent1)
|
||||
output2 = model2.decode(latent2).sample
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Check if outputs are similar
|
||||
max_diff_decode = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info(
|
||||
f"Maximum difference between decoded outputs: {max_diff_decode.item()}"
|
||||
)
|
||||
logger.info(
|
||||
f"Mean difference between decoded outputs: {mean_diff_decode.item()}"
|
||||
)
|
||||
assert mean_diff_decode < 1e-1, f"Decoded outputs differ significantly: mean diff = {mean_diff_decode.item()}"
|
||||
|
||||
logger.info(
|
||||
"Test passed! Both VAE implementations produce similar outputs.")
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_wan_vae()
|
||||
@@ -0,0 +1,152 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, destroy_distributed_environment,
|
||||
destroy_model_parallel, get_tp_rank,
|
||||
get_tp_world_size, maybe_init_distributed_environment_and_model_parallel, get_world_group)
|
||||
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SimpleTPModel(nn.Module):
|
||||
"""A simple model that uses tensor parallelism."""
|
||||
|
||||
def __init__(self, hidden_size=1024, intermediate_size=4096):
|
||||
super().__init__()
|
||||
# Column parallel linear layer (splits output dimension)
|
||||
self.fc1 = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=intermediate_size,
|
||||
bias=True,
|
||||
gather_output=
|
||||
False, # Don't gather output since we're passing to row parallel
|
||||
skip_bias_add=False)
|
||||
|
||||
# Row parallel linear layer (splits input dimension)
|
||||
self.fc2 = RowParallelLinear(
|
||||
input_size=intermediate_size,
|
||||
output_size=hidden_size,
|
||||
bias=True,
|
||||
input_is_parallel=True, # Input is already split from previous layer
|
||||
skip_bias_add=False)
|
||||
|
||||
self.activation = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
# Forward through column parallel layer
|
||||
hidden_states, _ = self.fc1(x)
|
||||
|
||||
# Apply activation
|
||||
hidden_states = self.activation(hidden_states)
|
||||
|
||||
# Forward through row parallel layer
|
||||
output, _ = self.fc2(hidden_states)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def initialize_random_weights(model, seed=42):
|
||||
"""Initialize the model with random weights using a fixed seed for reproducibility."""
|
||||
# Set seed for reproducibility
|
||||
torch.manual_seed(seed)
|
||||
|
||||
# Initialize weights for each layer
|
||||
with torch.no_grad():
|
||||
# For ColumnParallelLinear layers
|
||||
if hasattr(model, 'fc1'):
|
||||
nn.init.normal_(model.fc1.weight, mean=0.0, std=0.02)
|
||||
if model.fc1.bias is not None:
|
||||
nn.init.zeros_(model.fc1.bias)
|
||||
|
||||
# For RowParallelLinear layers
|
||||
if hasattr(model, 'fc2'):
|
||||
nn.init.normal_(model.fc2.weight, mean=0.0, std=0.02)
|
||||
if model.fc2.bias is not None:
|
||||
nn.init.zeros_(model.fc2.bias)
|
||||
|
||||
logger.info("Model initialized with random weights")
|
||||
return model
|
||||
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Simple Tensor Parallelism Example')
|
||||
parser.add_argument('--tensor-model-parallel-size',
|
||||
type=int,
|
||||
default=8,
|
||||
help='Degree of tensor model parallelism')
|
||||
parser.add_argument('--batch-size',
|
||||
type=int,
|
||||
default=8,
|
||||
help='Batch size for the example')
|
||||
parser.add_argument('--hidden-size',
|
||||
type=int,
|
||||
default=1024,
|
||||
help='Hidden size for the model')
|
||||
parser.add_argument('--intermediate-size',
|
||||
type=int,
|
||||
default=4096,
|
||||
help='Intermediate size for the model')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
args = setup_args()
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(args.tensor_model_parallel_size, args.tensor_model_parallel_size)
|
||||
rank = get_world_group().rank
|
||||
local_rank = get_world_group().local_rank
|
||||
# Get tensor parallel info
|
||||
tp_rank = get_tp_rank()
|
||||
tp_world_size = get_tp_world_size()
|
||||
|
||||
logger.info(
|
||||
f"Process rank {rank} initialized with TP rank {tp_rank} in TP world size {tp_world_size}"
|
||||
)
|
||||
|
||||
# Create a simple model
|
||||
model = SimpleTPModel(hidden_size=args.hidden_size,
|
||||
intermediate_size=args.intermediate_size)
|
||||
|
||||
# Initialize with random weights
|
||||
model = initialize_random_weights(model)
|
||||
|
||||
# Create a random input tensor
|
||||
batch_size = args.batch_size
|
||||
hidden_size = args.hidden_size
|
||||
x = torch.randn(batch_size, hidden_size, dtype=torch.float)
|
||||
|
||||
# Move to GPU if available
|
||||
device = torch.device(
|
||||
f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
|
||||
model = model.to(device)
|
||||
x = x.to(device)
|
||||
|
||||
# Forward pass
|
||||
logger.info(f"Running forward pass on TP rank {tp_rank}")
|
||||
with torch.no_grad():
|
||||
output = model(x)
|
||||
|
||||
# Print output shape and statistics
|
||||
logger.info(f"Output shape: {output.shape}")
|
||||
logger.info(
|
||||
f"Output mean: {output.mean().item()}, std: {output.std().item()}")
|
||||
|
||||
# Clean up
|
||||
logger.info("Cleaning up distributed environment")
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
logger.info("Example completed successfully")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,10 +1 @@
|
||||
{
|
||||
"step_time": 0.6983645600266755,
|
||||
"grad_norm": 0.245593056678772,
|
||||
"avg_step_time": 1.002151239803061,
|
||||
"_timestamp": 1751181952.70901,
|
||||
"vsa_sparsity": 0.05,
|
||||
"learning_rate": 1e-05,
|
||||
"train_loss": 0.2530866410434246,
|
||||
"_runtime": 107.325113071
|
||||
}
|
||||
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":1.260593056678772,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.2620866410434246,"_runtime":107.325113071}
|
||||
@@ -1,211 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# Set Python path to current folder
|
||||
current_dir = str(Path(__file__).parent.parent.parent.parent.parent)
|
||||
if current_dir not in sys.path:
|
||||
sys.path.insert(0, current_dir)
|
||||
os.environ["PYTHONPATH"] = current_dir + ":" + os.environ.get("PYTHONPATH", "")
|
||||
|
||||
import subprocess
|
||||
import torch
|
||||
import json
|
||||
from huggingface_hub import snapshot_download
|
||||
from fastvideo.utils import logger
|
||||
# Import the training pipeline
|
||||
from fastvideo.training.wan_training_pipeline import main
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
from fastvideo.training.wan_training_pipeline import WanTrainingPipeline
|
||||
|
||||
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_PATH = "data/crush-smol_processed_t2v/training_dataset/worker_1/worker_0/"
|
||||
VALIDATION_DATASET_FILE = "examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json"
|
||||
OUTPUT_DIR = Path("checkpoints/wan_t2v_finetune")
|
||||
PROFILER_TRACE_ROOT = Path("/mnt/fast-disks/hao_lab/ohm/profiler_traces/wan_t2v_finetune")
|
||||
WANDB_SUMMARY_FILE = OUTPUT_DIR / "tracker/wandb/latest-run/files/wandb-summary.json"
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "2"
|
||||
GRAD_ACCUM = "1"
|
||||
MASTER_PORT = "29504"
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = MASTER_PORT
|
||||
|
||||
|
||||
def run_worker():
|
||||
"""Worker function that will be run on each GPU"""
|
||||
# Create and populate args
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
|
||||
# Set the arguments as they are in finetune_t2v.sh
|
||||
args = parser.parse_args([
|
||||
"--model_path", MODEL_PATH,
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", DATA_PATH,
|
||||
"--dataloader_num_workers", "1",
|
||||
"--train_batch_size", "4",
|
||||
"--train_sp_batch_size", "1",
|
||||
"--gradient_accumulation_steps", GRAD_ACCUM,
|
||||
"--num_latent_t", "20",
|
||||
"--num_height", "720",
|
||||
"--num_width", "1280",
|
||||
"--num_frames", "77",
|
||||
"--enable_gradient_checkpointing_type", "full",
|
||||
"--max_train_steps", "20",
|
||||
"--learning_rate", "5e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--weight_only_checkpointing_steps", "250",
|
||||
"--training_state_checkpointing_steps", "250",
|
||||
"--weight_decay", "1e-4",
|
||||
"--max_grad_norm", "1.0",
|
||||
"--num_euler_timesteps", "50",
|
||||
"--multi_phased_distill_schedule", "4000-1",
|
||||
"--not_apply_cfg_solver",
|
||||
"--training_cfg_rate", "0.1",
|
||||
"--ema_start_step", "0",
|
||||
"--dit_precision", "fp32",
|
||||
"--output_dir", str(OUTPUT_DIR),
|
||||
"--tracker_project_name", "wan_t2v_finetune",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--validation_dataset_file", VALIDATION_DATASET_FILE,
|
||||
"--validation_steps", "200",
|
||||
"--validation_sampling_steps", "50",
|
||||
"--validation_guidance_scale", "6.0",
|
||||
#"--enable_torch_compile",
|
||||
#"--log_validation",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE,
|
||||
"--sp_size", NUM_GPUS_PER_NODE,
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", NUM_GPUS_PER_NODE,
|
||||
"--hsdp_shard_dim", "1"
|
||||
])
|
||||
# Call the main training function
|
||||
pipeline = WanTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
def test_distributed_training():
|
||||
"""Test the distributed training setup"""
|
||||
os.environ["WANDB_MODE"] = "online"
|
||||
|
||||
data_dir = Path("data/crush-smol_processed_t2v")
|
||||
|
||||
if not data_dir.exists():
|
||||
print(f"Downloading test dataset to {data_dir}...")
|
||||
snapshot_download(
|
||||
repo_id="wlsaidhi/crush-smol_processed_t2v",
|
||||
local_dir=str(data_dir),
|
||||
repo_type="dataset",
|
||||
local_dir_use_symlinks=False
|
||||
)
|
||||
|
||||
# Get the current file path
|
||||
current_file = Path(__file__).resolve()
|
||||
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE,
|
||||
"--master_port", MASTER_PORT,
|
||||
str(current_file)
|
||||
]
|
||||
process = subprocess.run(cmd, capture_output=True, text=True)
|
||||
|
||||
# Print stdout and stderr for debugging
|
||||
if process.stdout:
|
||||
print("STDOUT:", process.stdout)
|
||||
if process.stderr:
|
||||
print("STDERR:", process.stderr)
|
||||
|
||||
# Check if the process failed
|
||||
if process.returncode != 0:
|
||||
print(f"Process failed with return code: {process.returncode}")
|
||||
raise subprocess.CalledProcessError(process.returncode, cmd, process.stdout, process.stderr)
|
||||
|
||||
summary_file = WANDB_SUMMARY_FILE
|
||||
|
||||
with summary_file.open() as f:
|
||||
wandb_summary = json.load(f)
|
||||
|
||||
# Calculate and print MFU metrics
|
||||
device_name = torch.cuda.get_device_name()
|
||||
try:
|
||||
# Get actual values from training run (logged from training_batch.raw_latent_shape)
|
||||
batch_size = wandb_summary.get("batch_size")
|
||||
seq_len = wandb_summary.get("dit_seq_len")
|
||||
context_len = wandb_summary.get("context_len")
|
||||
avg_step_time = wandb_summary.get("avg_step_time")
|
||||
hidden_dim = wandb_summary.get("hidden_dim")
|
||||
num_layers = wandb_summary.get("num_layers")
|
||||
ffn_dim = wandb_summary.get("ffn_dim")
|
||||
|
||||
|
||||
|
||||
|
||||
# FLOPs per layer (forward pass)
|
||||
# - QKV + out proj: 8 * hidden_dim^2 * seq_len
|
||||
# - Cross-attn proj: 4 * hidden_dim^2 * seq_len + 4 * hidden_dim^2 * context_len
|
||||
# - MLP: 4 * hidden_dim * ffn_dim * seq_len
|
||||
# - Self-attn matmuls: 4 * seq_len^2 * hidden_dim
|
||||
# - Cross-attn matmuls: 4 * seq_len * context_len * hidden_dim
|
||||
qkv_out_flops = 8 * hidden_dim * hidden_dim * seq_len
|
||||
cross_attn_proj_flops = (
|
||||
(4 * hidden_dim * hidden_dim * seq_len) +
|
||||
(4 * hidden_dim * hidden_dim * context_len)
|
||||
)
|
||||
mlp_flops = 4 * hidden_dim * ffn_dim * seq_len
|
||||
self_attn_flops = 4 * seq_len * seq_len * hidden_dim
|
||||
cross_attn_flops = 4 * seq_len * context_len * hidden_dim
|
||||
flops_per_layer = (
|
||||
qkv_out_flops + cross_attn_proj_flops + mlp_flops + self_attn_flops + cross_attn_flops
|
||||
)
|
||||
|
||||
# With full activation checkpointing: 1 forward + 3 backward (1 recompute + 2 gradient)
|
||||
achieved_flops = batch_size * flops_per_layer * num_layers * 4
|
||||
|
||||
|
||||
# Account for gradient accumulation (from config)
|
||||
grad_accum = int(GRAD_ACCUM)
|
||||
achieved_flops *= grad_accum
|
||||
|
||||
# Peak FLOPs based on device
|
||||
if "H100" in device_name:
|
||||
peak_flops_per_gpu = 989e12
|
||||
elif "A100" in device_name:
|
||||
peak_flops_per_gpu = 312e12
|
||||
elif "A40" in device_name:
|
||||
peak_flops_per_gpu = 312e12
|
||||
elif "L40S" in device_name:
|
||||
peak_flops_per_gpu = 362e12
|
||||
else:
|
||||
raise ValueError(f"Device {device_name} not supported")
|
||||
|
||||
# Total peak (2 GPUs)
|
||||
world_size = int(NUM_GPUS_PER_NODE)
|
||||
total_peak_flops = peak_flops_per_gpu * world_size
|
||||
|
||||
# Calculate MFU
|
||||
achieved_flops_per_sec = achieved_flops / avg_step_time if avg_step_time > 0 else 0
|
||||
mfu = (achieved_flops_per_sec / total_peak_flops * 100) if total_peak_flops > 0 else 0
|
||||
|
||||
print(f"Per-Step MFU: {mfu:.4f}%")
|
||||
except Exception as e:
|
||||
print(f"Could not calculate MFU: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if os.environ.get("LOCAL_RANK") is not None:
|
||||
# We're being run by torchrun
|
||||
run_worker()
|
||||
else:
|
||||
# We're being run directly
|
||||
test_distributed_training()
|
||||
@@ -1,598 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Test COSMOS 2.5 DiT implementation against reference.
|
||||
Compares FastVideo's Cosmos25Transformer3DModel with the official MinimalV1LVGDiT from cosmos-predict2.5.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# Add cosmos-predict2.5 to Python path for loading reference model
|
||||
TEST_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
COSMOS_PREDICT2_5_PATH = os.path.join(TEST_DIR, '..', '..', '..', '..', 'cosmos-predict2.5')
|
||||
COSMOS_PREDICT2_5_PATH = os.path.normpath(COSMOS_PREDICT2_5_PATH)
|
||||
if os.path.exists(COSMOS_PREDICT2_5_PATH) and COSMOS_PREDICT2_5_PATH not in sys.path:
|
||||
sys.path.insert(0, COSMOS_PREDICT2_5_PATH)
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import maybe_download_model
|
||||
# Use Cosmos 2.5 specific config
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Log the cosmos-predict2.5 path after logger is initialized
|
||||
if os.path.exists(COSMOS_PREDICT2_5_PATH):
|
||||
logger.info(f"cosmos-predict2.5 found at: {COSMOS_PREDICT2_5_PATH}")
|
||||
else:
|
||||
logger.warning(f"cosmos-predict2.5 not found at: {COSMOS_PREDICT2_5_PATH}")
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29505"
|
||||
|
||||
# COSMOS 2.5 model path - update this based on the actual HuggingFace model ID
|
||||
# The model has subdirectories: base/pre-trained, base/post-trained, auto/multiview, robot/action-cond
|
||||
BASE_MODEL_PATH = "nvidia/Cosmos-Predict2.5-2B"
|
||||
CHECKPOINT_SUBDIR = "base/post-trained"
|
||||
CHECKPOINT_FILENAME = "81edfebe-bd6a-4039-8c1d-737df1a790bf_ema_bf16.pt"
|
||||
|
||||
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH, local_dir=None)
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, CHECKPOINT_SUBDIR, "transformer")
|
||||
if not os.path.exists(TRANSFORMER_PATH):
|
||||
# Try without subdirectory
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
def load_reference_cosmos25_model(checkpoint_path: str, device, dtype):
|
||||
"""
|
||||
Load the reference COSMOS 2.5 model from cosmos-predict2.5 repo.
|
||||
This assumes the cosmos-predict2.5 repo is available in the Python path.
|
||||
"""
|
||||
try:
|
||||
# Try to import from cosmos-predict2.5 repo
|
||||
from cosmos_predict2._src.predict2.networks.minimal_v1_lvg_dit import MinimalV1LVGDiT
|
||||
|
||||
# COSMOS 2.5 2B model configuration
|
||||
model_config = {
|
||||
'max_img_h': 240,
|
||||
'max_img_w': 240,
|
||||
'max_frames': 128,
|
||||
'in_channels': 16,
|
||||
'out_channels': 16,
|
||||
'patch_spatial': 2,
|
||||
'patch_temporal': 1,
|
||||
'model_channels': 2048, # 2B model
|
||||
'num_blocks': 28,
|
||||
'num_heads': 16,
|
||||
'mlp_ratio': 4.0,
|
||||
'crossattn_emb_channels': 1024,
|
||||
'pos_emb_cls': 'rope3d',
|
||||
'pos_emb_learnable': True,
|
||||
'pos_emb_interpolation': 'crop',
|
||||
'use_adaln_lora': True,
|
||||
'adaln_lora_dim': 256,
|
||||
'rope_h_extrapolation_ratio': 3.0,
|
||||
'rope_w_extrapolation_ratio': 3.0,
|
||||
'rope_t_extrapolation_ratio': 1.0,
|
||||
'extra_per_block_abs_pos_emb': False,
|
||||
'rope_enable_fps_modulation': False,
|
||||
'use_crossattn_projection': True,
|
||||
'crossattn_proj_in_channels': 100352,
|
||||
'concat_padding_mask': True,
|
||||
'atten_backend': 'torch',
|
||||
}
|
||||
|
||||
model = MinimalV1LVGDiT(**model_config)
|
||||
|
||||
# Load checkpoint if path exists
|
||||
if os.path.exists(checkpoint_path):
|
||||
logger.info(f"Loading reference model from {checkpoint_path}")
|
||||
checkpoint = torch.load(checkpoint_path, map_location='cpu')
|
||||
|
||||
# Extract state dict
|
||||
if 'state_dict' in checkpoint:
|
||||
checkpoint_state = checkpoint['state_dict']
|
||||
elif 'model' in checkpoint:
|
||||
checkpoint_state = checkpoint['model']
|
||||
else:
|
||||
checkpoint_state = checkpoint
|
||||
|
||||
# Filter to only model parameters (remove training metadata)
|
||||
model_state = {k: v for k, v in checkpoint_state.items()
|
||||
if k.startswith('net.') and 'accum_' not in k}
|
||||
|
||||
# Transform checkpoint keys to match model's expected format
|
||||
# 1. Strip 'net.' prefix (e.g., 'net.blocks.0.self_attn.*' -> 'blocks.0.self_attn.*')
|
||||
# 2. Add '_checkpoint_wrapped_module' after 'blocks.N.' if model expects it
|
||||
transformed_state = {}
|
||||
|
||||
# First, check what the model expects
|
||||
model_state_dict = model.state_dict()
|
||||
needs_checkpoint_wrapper = any('_checkpoint_wrapped_module' in k for k in model_state_dict.keys())
|
||||
|
||||
for key, value in model_state.items():
|
||||
# Strip 'net.' prefix
|
||||
if key.startswith('net.'):
|
||||
new_key = key[4:] # Remove 'net.' prefix
|
||||
else:
|
||||
new_key = key
|
||||
|
||||
# Add '_checkpoint_wrapped_module' if needed
|
||||
if needs_checkpoint_wrapper and new_key.startswith('blocks.'):
|
||||
# Pattern: 'blocks.N.something' -> 'blocks.N._checkpoint_wrapped_module.something'
|
||||
parts = new_key.split('.', 2)
|
||||
if len(parts) >= 3 and parts[0] == 'blocks' and parts[1].isdigit():
|
||||
new_key = f"{parts[0]}.{parts[1]}._checkpoint_wrapped_module.{parts[2]}"
|
||||
|
||||
transformed_state[new_key] = value
|
||||
|
||||
# Load with strict=False to handle any remaining mismatches
|
||||
missing_keys, unexpected_keys = model.load_state_dict(transformed_state, strict=False)
|
||||
|
||||
if missing_keys:
|
||||
logger.warning(f"Missing keys when loading reference model: {len(missing_keys)} keys")
|
||||
# Show all missing keys for debugging
|
||||
logger.warning("All missing keys:")
|
||||
for k in missing_keys:
|
||||
logger.warning(f" - {k}")
|
||||
# Filter out _extra_state and pos_embedder keys as they're optional
|
||||
missing_important = [k for k in missing_keys
|
||||
if '_extra_state' not in k and 'pos_embedder' not in k and 'accum_' not in k]
|
||||
if missing_important:
|
||||
logger.warning(f"Missing important keys ({len(missing_important)} total):")
|
||||
for k in missing_important[:10]: # Show first 10
|
||||
logger.warning(f" - {k}")
|
||||
if len(missing_important) > 10:
|
||||
logger.warning(f" ... and {len(missing_important) - 10} more")
|
||||
|
||||
if unexpected_keys:
|
||||
logger.warning(f"Unexpected keys when loading reference model: {len(unexpected_keys)} keys")
|
||||
logger.warning("All unexpected keys:")
|
||||
for k in unexpected_keys:
|
||||
logger.warning(f" - {k}")
|
||||
|
||||
logger.info(f"Successfully loaded {len(transformed_state)} parameters into reference model")
|
||||
else:
|
||||
logger.warning(f"Checkpoint path {checkpoint_path} not found, using random weights")
|
||||
|
||||
model = model.to(device, dtype=dtype)
|
||||
model.eval()
|
||||
|
||||
return model
|
||||
|
||||
except ImportError as e:
|
||||
logger.error(f"Failed to import cosmos-predict2.5: {e}")
|
||||
logger.info("Make sure cosmos-predict2.5 is in your Python path")
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_cosmos25_transformer():
|
||||
"""Test COSMOS 2.5 transformer against reference implementation."""
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
|
||||
# Create COSMOS 2.5 specific config
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25ArchConfig
|
||||
|
||||
arch_config = Cosmos25ArchConfig(
|
||||
num_attention_heads=16,
|
||||
attention_head_dim=128, # 2048 / 16
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
num_layers=28,
|
||||
patch_size=(1, 2, 2),
|
||||
max_size=(128, 240, 240),
|
||||
rope_scale=(1.0, 3.0, 3.0), # T, H, W
|
||||
text_embed_dim=1024,
|
||||
mlp_ratio=4.0,
|
||||
adaln_lora_dim=256,
|
||||
use_adaln_lora=True,
|
||||
concat_padding_mask=True,
|
||||
extra_pos_embed_type=None,
|
||||
use_crossattn_projection=True,
|
||||
rope_enable_fps_modulation=False,
|
||||
qk_norm="rms_norm",
|
||||
)
|
||||
|
||||
cosmos25_config = Cosmos25VideoConfig(arch_config=arch_config)
|
||||
|
||||
# Create FastVideo model directly (Cosmos 2.5 is not in diffusers format)
|
||||
logger.info("Creating FastVideo COSMOS 2.5 model...")
|
||||
from fastvideo.models.dits.cosmos2_5 import Cosmos25Transformer3DModel
|
||||
|
||||
# Get hf_config from the arch_config for model initialization
|
||||
hf_config = {
|
||||
'in_channels': arch_config.in_channels,
|
||||
'out_channels': arch_config.out_channels,
|
||||
'num_attention_heads': arch_config.num_attention_heads,
|
||||
'attention_head_dim': arch_config.attention_head_dim,
|
||||
'num_layers': arch_config.num_layers,
|
||||
'patch_size': arch_config.patch_size,
|
||||
'max_size': arch_config.max_size,
|
||||
'rope_scale': arch_config.rope_scale,
|
||||
'text_embed_dim': arch_config.text_embed_dim,
|
||||
'mlp_ratio': arch_config.mlp_ratio,
|
||||
'adaln_lora_dim': arch_config.adaln_lora_dim,
|
||||
'use_adaln_lora': arch_config.use_adaln_lora,
|
||||
'concat_padding_mask': arch_config.concat_padding_mask,
|
||||
'extra_pos_embed_type': arch_config.extra_pos_embed_type,
|
||||
'use_crossattn_projection': arch_config.use_crossattn_projection,
|
||||
'rope_enable_fps_modulation': arch_config.rope_enable_fps_modulation,
|
||||
'qk_norm': arch_config.qk_norm,
|
||||
}
|
||||
|
||||
fastvideo_model = Cosmos25Transformer3DModel(config=cosmos25_config, hf_config=hf_config)
|
||||
fastvideo_model = fastvideo_model.to(device, dtype=precision)
|
||||
fastvideo_model.eval()
|
||||
|
||||
# Construct checkpoint path using relative paths
|
||||
checkpoint_file = os.path.join(MODEL_PATH, CHECKPOINT_SUBDIR, CHECKPOINT_FILENAME)
|
||||
|
||||
if not os.path.exists(checkpoint_file):
|
||||
logger.warning(f"Checkpoint file not found at {checkpoint_file}")
|
||||
logger.info("Will test architecture without loading checkpoint weights")
|
||||
checkpoint_file = None
|
||||
|
||||
# Load checkpoint into FastVideo model using param_names_mapping
|
||||
if checkpoint_file:
|
||||
logger.info(f"Loading checkpoint into FastVideo model from {checkpoint_file}")
|
||||
from fastvideo.models.loader.utils import hf_to_custom_state_dict, get_param_names_mapping
|
||||
|
||||
checkpoint = torch.load(checkpoint_file, map_location='cpu')
|
||||
|
||||
# Extract state dict (checkpoint might have 'state_dict', 'model', or be the dict itself)
|
||||
if 'state_dict' in checkpoint:
|
||||
checkpoint_state = checkpoint['state_dict']
|
||||
elif 'model' in checkpoint:
|
||||
checkpoint_state = checkpoint['model']
|
||||
else:
|
||||
checkpoint_state = checkpoint
|
||||
|
||||
# Filter to only model parameters (remove training metadata like accum_*)
|
||||
model_state = {k: v for k, v in checkpoint_state.items()
|
||||
if k.startswith('net.') and 'accum_' not in k}
|
||||
|
||||
# Convert checkpoint keys to FastVideo format using param_names_mapping
|
||||
param_names_mapping_fn = get_param_names_mapping(
|
||||
cosmos25_config.arch_config.param_names_mapping
|
||||
)
|
||||
custom_state_dict, reverse_mapping = hf_to_custom_state_dict(
|
||||
model_state, param_names_mapping_fn
|
||||
)
|
||||
|
||||
# Only load keys that exist in the model
|
||||
model_param_names = set(fastvideo_model.state_dict().keys())
|
||||
filtered_state_dict = {
|
||||
k: v.to(device=device, dtype=precision)
|
||||
for k, v in custom_state_dict.items()
|
||||
if k in model_param_names
|
||||
}
|
||||
|
||||
# Load into FastVideo model
|
||||
missing_keys, unexpected_keys = fastvideo_model.load_state_dict(
|
||||
filtered_state_dict, strict=False
|
||||
)
|
||||
|
||||
if missing_keys:
|
||||
logger.warning(f"Missing keys when loading checkpoint: {len(missing_keys)} keys")
|
||||
# Filter out _extra_state keys as they're optional
|
||||
missing_non_extra = [k for k in missing_keys if '_extra_state' not in k]
|
||||
if missing_non_extra:
|
||||
logger.warning(f"Missing non-extra keys (first 10): {missing_non_extra[:10]}")
|
||||
|
||||
if unexpected_keys:
|
||||
logger.warning(f"Unexpected keys when loading checkpoint: {len(unexpected_keys)} keys")
|
||||
|
||||
logger.info(f"Successfully loaded {len(filtered_state_dict)} parameters into FastVideo model")
|
||||
|
||||
# Try to load reference model from the raw checkpoint
|
||||
logger.info("Loading reference COSMOS 2.5 model...")
|
||||
reference_model = load_reference_cosmos25_model(checkpoint_file, device, precision) if checkpoint_file else None
|
||||
|
||||
# Set models to eval mode
|
||||
fastvideo_model = fastvideo_model.eval()
|
||||
if reference_model is not None:
|
||||
reference_model = reference_model.eval()
|
||||
|
||||
# Create test inputs
|
||||
batch_size = 1
|
||||
seq_len = 77 # Typical T5 sequence length
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
# COSMOS 2.5: 16 channels (VAE latent), no condition mask in input
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
16, # VAE channels only (condition mask added internally)
|
||||
1, # Single frame for image generation (or 16 for video)
|
||||
64, # Height (720p / 8 / 2 patch = 45, use 64 for testing)
|
||||
64, # Width
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
# Condition mask [B, 1, T, H, W] - for video2world conditioning
|
||||
condition_mask = torch.zeros(
|
||||
batch_size,
|
||||
1,
|
||||
1,
|
||||
64,
|
||||
64,
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
# Text embeddings [B, L, D] - Qwen 7B embeddings (100,352 dims)
|
||||
# Using 100,352 dimensions to match the crossattn_projection layer
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
seq_len,
|
||||
100352,
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
# Timestep [B, T] - official model expects [B, T] shape with dtype matching model precision
|
||||
# For single frame, use [B, 1]
|
||||
timestep = torch.full((batch_size, 1), 500.0, device=device, dtype=precision)
|
||||
|
||||
# Padding mask [B, H, W] - official model expects NO channel dimension
|
||||
# It's added internally via unsqueeze(1) if needed
|
||||
padding_mask = torch.ones(
|
||||
batch_size,
|
||||
64,
|
||||
64,
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
# FPS for temporal scaling
|
||||
fps = 16
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
logger.info("Running inference...")
|
||||
|
||||
with torch.no_grad():
|
||||
with torch.autocast('cuda', dtype=precision):
|
||||
# FastVideo model
|
||||
with set_forward_context(
|
||||
current_timestep=500,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
# FastVideo expects padding_mask in [B, 1, H, W] format
|
||||
padding_mask_fv = padding_mask.unsqueeze(1) # Add channel dimension for FastVideo
|
||||
# FastVideo supports both [B] and [B, T] formats - use [B, T] to match official model
|
||||
# This ensures each frame gets its own timestep embedding (even if values are the same)
|
||||
timestep_fv = timestep # Already in [B, T] format
|
||||
|
||||
output_fv = fastvideo_model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep_fv,
|
||||
condition_mask=condition_mask,
|
||||
padding_mask=padding_mask_fv,
|
||||
fps=fps,
|
||||
)
|
||||
|
||||
# Reference model (if available)
|
||||
if reference_model is not None:
|
||||
# Prepare input for reference model
|
||||
# MinimalV1LVGDiT adds condition mask internally, so pass them separately
|
||||
from cosmos_predict2._src.predict2.conditioner import DataType
|
||||
|
||||
# Determine data_type based on temporal dimension
|
||||
num_frames = hidden_states.shape[2]
|
||||
ref_data_type = DataType.VIDEO if num_frames > 1 else DataType.IMAGE
|
||||
|
||||
# Reference model expects different input format
|
||||
# Pass hidden_states without condition_mask (model concatenates it internally)
|
||||
# timestep is already in [B, T] format with correct dtype
|
||||
# padding_mask is already in [B, H, W] format (no channel dimension)
|
||||
# FPS should be a tensor [B] or scalar
|
||||
fps_tensor = torch.tensor([fps], device=device, dtype=precision)
|
||||
|
||||
output_ref = reference_model(
|
||||
x_B_C_T_H_W=hidden_states, # [B, 16, T, H, W] - model will add condition mask
|
||||
timesteps_B_T=timestep, # Already in [B, T] format
|
||||
crossattn_emb=encoder_hidden_states,
|
||||
condition_video_input_mask_B_C_T_H_W=condition_mask if ref_data_type == DataType.VIDEO else None,
|
||||
fps=fps_tensor,
|
||||
padding_mask=padding_mask, # [B, H, W] format
|
||||
data_type=ref_data_type,
|
||||
)
|
||||
|
||||
# Check FastVideo output shape and dtype
|
||||
logger.info(f"FastVideo output shape: {output_fv.shape}")
|
||||
logger.info(f"FastVideo output dtype: {output_fv.dtype}")
|
||||
assert output_fv.shape[0] == batch_size, "Batch size mismatch"
|
||||
assert output_fv.shape[1] == 16, "Output channels should be 16"
|
||||
assert output_fv.dtype == precision, f"Output dtype mismatch: {output_fv.dtype} vs {precision}"
|
||||
|
||||
# Compare with reference if available
|
||||
if reference_model is not None:
|
||||
logger.info(f"Reference output shape: {output_ref.shape}")
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output_fv.shape == output_ref.shape, \
|
||||
f"Output shapes don't match: {output_fv.shape} vs {output_ref.shape}"
|
||||
assert output_fv.dtype == output_ref.dtype, \
|
||||
f"Output dtype don't match: {output_fv.dtype} vs {output_ref.dtype}"
|
||||
|
||||
# Check if outputs are similar
|
||||
max_diff = torch.max(torch.abs(output_fv - output_ref))
|
||||
mean_diff = torch.mean(torch.abs(output_fv - output_ref))
|
||||
relative_diff = mean_diff / (torch.mean(torch.abs(output_ref)) + 1e-8)
|
||||
|
||||
logger.info(f"Max difference: {max_diff.item():.6f}")
|
||||
logger.info(f"Mean difference: {mean_diff.item():.6f}")
|
||||
logger.info(f"Relative difference: {relative_diff.item():.6f}")
|
||||
|
||||
|
||||
# Allow for some numerical differences due to implementation details
|
||||
assert max_diff < 1e-1, f"Maximum difference too large: {max_diff.item()}"
|
||||
assert mean_diff < 1e-2, f"Mean difference too large: {mean_diff.item()}"
|
||||
|
||||
logger.info("✓ COSMOS 2.5 FastVideo implementation matches reference!")
|
||||
else:
|
||||
logger.warning("Reference model not available, skipping comparison")
|
||||
logger.info("✓ COSMOS 2.5 FastVideo model runs successfully!")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_cosmos25_transformer_video():
|
||||
"""Test COSMOS 2.5 transformer with video input (multiple frames)."""
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
|
||||
# Create COSMOS 2.5 specific config
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25ArchConfig
|
||||
|
||||
arch_config = Cosmos25ArchConfig(
|
||||
num_attention_heads=16,
|
||||
attention_head_dim=128,
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
num_layers=28,
|
||||
patch_size=(1, 2, 2),
|
||||
max_size=(128, 240, 240),
|
||||
rope_scale=(1.0, 3.0, 3.0),
|
||||
text_embed_dim=1024,
|
||||
mlp_ratio=4.0,
|
||||
adaln_lora_dim=256,
|
||||
use_adaln_lora=True,
|
||||
concat_padding_mask=True,
|
||||
extra_pos_embed_type=None,
|
||||
use_crossattn_projection=True, # Enable to match official model
|
||||
rope_enable_fps_modulation=False,
|
||||
qk_norm="rms_norm",
|
||||
)
|
||||
|
||||
cosmos25_config = Cosmos25VideoConfig(arch_config=arch_config)
|
||||
|
||||
# Create FastVideo model directly (Cosmos 2.5 is not in diffusers format)
|
||||
logger.info("Creating FastVideo COSMOS 2.5 model for video test...")
|
||||
from fastvideo.models.dits.cosmos2_5 import Cosmos25Transformer3DModel
|
||||
|
||||
# Get hf_config from the arch_config for model initialization
|
||||
hf_config = {
|
||||
'in_channels': arch_config.in_channels,
|
||||
'out_channels': arch_config.out_channels,
|
||||
'num_attention_heads': arch_config.num_attention_heads,
|
||||
'attention_head_dim': arch_config.attention_head_dim,
|
||||
'num_layers': arch_config.num_layers,
|
||||
'patch_size': arch_config.patch_size,
|
||||
'max_size': arch_config.max_size,
|
||||
'rope_scale': arch_config.rope_scale,
|
||||
'text_embed_dim': arch_config.text_embed_dim,
|
||||
'mlp_ratio': arch_config.mlp_ratio,
|
||||
'adaln_lora_dim': arch_config.adaln_lora_dim,
|
||||
'use_adaln_lora': arch_config.use_adaln_lora,
|
||||
'concat_padding_mask': arch_config.concat_padding_mask,
|
||||
'extra_pos_embed_type': arch_config.extra_pos_embed_type,
|
||||
'use_crossattn_projection': arch_config.use_crossattn_projection,
|
||||
'rope_enable_fps_modulation': arch_config.rope_enable_fps_modulation,
|
||||
'qk_norm': arch_config.qk_norm,
|
||||
}
|
||||
|
||||
model = Cosmos25Transformer3DModel(config=cosmos25_config, hf_config=hf_config)
|
||||
model = model.to(device, dtype=precision)
|
||||
model.eval()
|
||||
|
||||
# Create video input with multiple frames
|
||||
batch_size = 1
|
||||
num_frames = 16 # Video with 16 frames
|
||||
seq_len = 77
|
||||
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
16,
|
||||
num_frames, # Multiple frames
|
||||
64,
|
||||
64,
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
condition_mask = torch.zeros(
|
||||
batch_size,
|
||||
1,
|
||||
num_frames,
|
||||
64,
|
||||
64,
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
# Set first 2 frames as conditioning
|
||||
condition_mask[:, :, :2, :, :] = 1.0
|
||||
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
seq_len,
|
||||
100352, # Qwen 7B embedding dimension (matches crossattn_proj input)
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
timestep = torch.tensor([500], device=device, dtype=torch.long)
|
||||
|
||||
padding_mask = torch.ones(
|
||||
batch_size,
|
||||
1,
|
||||
64,
|
||||
64,
|
||||
device=device,
|
||||
dtype=precision
|
||||
)
|
||||
|
||||
fps = 16
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
logger.info("Running video inference...")
|
||||
|
||||
with torch.no_grad():
|
||||
with torch.autocast('cuda', dtype=precision):
|
||||
with set_forward_context(
|
||||
current_timestep=500,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output = model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
condition_mask=condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
fps=fps,
|
||||
)
|
||||
|
||||
logger.info(f"Video output shape: {output.shape}")
|
||||
logger.info(f"Video output dtype: {output.dtype}")
|
||||
|
||||
# Check output shape
|
||||
assert output.shape[0] == batch_size, "Batch size mismatch"
|
||||
assert output.shape[1] == 16, "Output channels should be 16"
|
||||
assert output.shape[2] == num_frames, "Number of frames mismatch"
|
||||
assert output.dtype == precision, f"Output dtype mismatch"
|
||||
|
||||
logger.info("✓ COSMOS 2.5 video inference successful!")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests directly
|
||||
test_cosmos25_transformer()
|
||||
test_cosmos25_transformer_video()
|
||||
|
||||
@@ -25,8 +25,8 @@ os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
|
||||
)
|
||||
local_dir=os.path.join(
|
||||
"data", BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
CONFIG_PATH = os.path.join(TRANSFORMER_PATH, "config.json")
|
||||
|
||||
|
||||
@@ -5,7 +5,6 @@ import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers import WanTransformer3DModel
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
@@ -24,8 +23,8 @@ os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
|
||||
)
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@@ -121,4 +120,10 @@ def test_wan_transformer():
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -23,8 +23,8 @@ os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
|
||||
)
|
||||
local_dir=os.path.join(
|
||||
"data", BASE_MODEL_PATH))
|
||||
VAE_PATH = os.path.join(MODEL_PATH, "vae")
|
||||
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
|
||||
|
||||
|
||||
@@ -12,7 +12,6 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import VAELoader
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from torch.testing import assert_close
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -21,16 +20,16 @@ os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join("data", BASE_MODEL_PATH) # store in the large /workspace disk on Runpod
|
||||
)
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
VAE_PATH = os.path.join(MODEL_PATH, "vae")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_vae():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.float32
|
||||
precision_str = "fp32"
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=WanVAEConfig(), vae_precision=precision_str))
|
||||
args.device = device
|
||||
args.vae_cpu_offload = False
|
||||
@@ -71,7 +70,13 @@ def test_wan_vae():
|
||||
# Check if latents have the same shape
|
||||
assert latent1.mean.shape == latent2.mean.shape, f"Latent shapes don't match: {latent1.mean.shape} vs {latent2.mean.shape}"
|
||||
# Check if latents are similar
|
||||
assert_close(latent1.mean, latent2.mean, atol=1e-4, rtol=1e-4)
|
||||
max_diff_encode = torch.max(torch.abs(latent1.mean - latent2.mean))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1.mean - latent2.mean))
|
||||
logger.info("Maximum difference between encoded latents: %s",
|
||||
max_diff_encode.item())
|
||||
logger.info("Mean difference between encoded latents: %s",
|
||||
mean_diff_encode.item())
|
||||
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {mean_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
latent1_tensor = latent1.mode()
|
||||
@@ -93,4 +98,10 @@ def test_wan_vae():
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Check if outputs are similar
|
||||
assert_close(output1, output2, atol=1e-5, rtol=1e-3)
|
||||
max_diff_decode = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Maximum difference between decoded outputs: %s",
|
||||
max_diff_decode.item())
|
||||
logger.info("Mean difference between decoded outputs: %s",
|
||||
mean_diff_decode.item())
|
||||
assert max_diff_decode < 1e-5, f"Decoded outputs differ significantly: max diff = {mean_diff_decode.item()}"
|
||||
|
||||
@@ -729,6 +729,9 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(timestep, src=0)
|
||||
|
||||
timestep = shift_timestep(
|
||||
timestep,
|
||||
@@ -841,6 +844,9 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(fake_score_timestep, src=0)
|
||||
|
||||
fake_score_timestep = shift_timestep(
|
||||
fake_score_timestep,
|
||||
|
||||
@@ -470,7 +470,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# local_main_process_only=False)
|
||||
with self.tracker.timed("timing/reduce_loss"):
|
||||
world_group = get_world_group()
|
||||
avg_loss = world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
training_batch.total_loss += avg_loss.item()
|
||||
|
||||
return training_batch
|
||||
@@ -656,23 +656,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
"grad_norm": grad_norm,
|
||||
"vsa_sparsity": current_vsa_sparsity,
|
||||
}
|
||||
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
|
||||
|
||||
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
|
||||
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
|
||||
training_batch.raw_latent_shape[3] //
|
||||
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
|
||||
context_len = int(training_batch.encoder_hidden_states.shape[1])
|
||||
|
||||
metrics["dit_seq_len"] = int(seq_len)
|
||||
metrics["context_len"] = context_len
|
||||
|
||||
arch_config = self.training_args.pipeline_config.dit_config.arch_config
|
||||
|
||||
metrics["hidden_dim"] = arch_config.hidden_size
|
||||
metrics["num_layers"] = arch_config.num_layers
|
||||
metrics["ffn_dim"] = arch_config.ffn_dim
|
||||
|
||||
self.tracker.log(metrics, step)
|
||||
if step % self.training_args.training_state_checkpointing_steps == 0:
|
||||
with self.profiler_controller.region(
|
||||
@@ -758,9 +741,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
if training_args.validation_guidance_scale:
|
||||
sampling_param.guidance_scale = float(
|
||||
training_args.validation_guidance_scale)
|
||||
assert self.seed is not None
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
|
||||
@@ -510,8 +510,7 @@ def load_checkpoint(transformer,
|
||||
return 0
|
||||
|
||||
# Extract step number from checkpoint path
|
||||
step = int(
|
||||
os.path.basename(os.path.normpath(checkpoint_path)).split('-')[-1])
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
|
||||
if rank == 0:
|
||||
logger.info("Loading checkpoint from step %s", step)
|
||||
|
||||
@@ -48,26 +48,23 @@ class Worker:
|
||||
# This env var set by Ray causes exceptions with graph building.
|
||||
os.environ.pop("NCCL_ASYNC_ERROR_HANDLING", None)
|
||||
|
||||
# Set environment variables BEFORE calling get_local_torch_device()
|
||||
# so that each worker uses the correct device
|
||||
if self.fastvideo_args.distributed_executor_backend == "mp":
|
||||
os.environ["LOCAL_RANK"] = str(self.local_rank)
|
||||
os.environ["RANK"] = str(self.rank)
|
||||
os.environ["WORLD_SIZE"] = str(self.fastvideo_args.num_gpus)
|
||||
|
||||
# Platform-agnostic device initialization
|
||||
self.device = get_local_torch_device()
|
||||
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
# Set the CUDA device BEFORE any CUDA calls
|
||||
# _check_if_gpu_supports_dtype(self.model_config.dtype)
|
||||
if current_platform.is_cuda_alike():
|
||||
torch.cuda.set_device(self.device)
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0]
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
|
||||
else:
|
||||
# For MPS, we can't get memory info the same way
|
||||
self.init_gpu_memory = 0
|
||||
|
||||
if self.fastvideo_args.distributed_executor_backend == "mp":
|
||||
os.environ["LOCAL_RANK"] = str(self.local_rank)
|
||||
os.environ["RANK"] = str(self.rank)
|
||||
os.environ["WORLD_SIZE"] = str(self.fastvideo_args.num_gpus)
|
||||
|
||||
# Initialize the distributed environment.
|
||||
maybe_init_distributed_environment_and_model_parallel(
|
||||
self.fastvideo_args.tp_size, self.fastvideo_args.sp_size,
|
||||
|
||||
@@ -467,10 +467,6 @@ class WorkerMultiprocProc:
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"logging_info": logging_info
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(
|
||||
method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
else:
|
||||
result = self.worker.execute_method(method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
|
||||
+4
-16
@@ -14,7 +14,6 @@ edit_uri: edit/main/docs/
|
||||
# Configuration
|
||||
theme:
|
||||
name: material
|
||||
favicon: assets/logos/icon_simple.svg
|
||||
palette:
|
||||
- scheme: default
|
||||
toggle:
|
||||
@@ -47,18 +46,11 @@ plugins:
|
||||
hooks:
|
||||
on_pre_build: "docs.generate_examples:on_pre_build_hook"
|
||||
- autorefs
|
||||
# - awesome-nav
|
||||
# - glightbox
|
||||
- git-revision-date-localized:
|
||||
# exclude autogenerated files
|
||||
exclude:
|
||||
- examples/*
|
||||
- api-autonav:
|
||||
modules: ["fastvideo"]
|
||||
modules: ["fastvideo"]
|
||||
api_root_uri: "api"
|
||||
exclude:
|
||||
- "re:fastvideo\\._.*"
|
||||
- "fastvideo.third_party"
|
||||
- "re:fastvideo\\._.*"
|
||||
- mkdocstrings:
|
||||
handlers:
|
||||
python:
|
||||
@@ -83,10 +75,9 @@ plugins:
|
||||
inventories:
|
||||
- https://docs.python.org/3/objects.inv
|
||||
|
||||
|
||||
|
||||
# Markdown extensions
|
||||
markdown_extensions:
|
||||
- admonition
|
||||
- pymdownx.highlight:
|
||||
anchor_linenums: true
|
||||
line_spans: __span
|
||||
@@ -112,10 +103,8 @@ markdown_extensions:
|
||||
- pymdownx.tasklist:
|
||||
custom_checkbox: true
|
||||
- pymdownx.tilde
|
||||
# For in page [TOC] (not sidebar)
|
||||
- toc:
|
||||
permalink: true
|
||||
- mdx_truly_sane_lists
|
||||
|
||||
# Page tree
|
||||
nav:
|
||||
@@ -162,7 +151,6 @@ nav:
|
||||
- Index: contributing/developer_env/index.md
|
||||
- Docker: contributing/developer_env/docker.md
|
||||
- RunPod: contributing/developer_env/runpod.md
|
||||
- Testing: contributing/testing.md
|
||||
- Profiling: contributing/profiling.md
|
||||
- API Reference:
|
||||
- FastVideo: api/fastvideo.md
|
||||
@@ -178,4 +166,4 @@ extra:
|
||||
|
||||
# Custom CSS
|
||||
extra_css:
|
||||
- assets/custom.css
|
||||
- assets/custom.css
|
||||
@@ -7,8 +7,7 @@ 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
|
||||
|
||||
@@ -1,68 +0,0 @@
|
||||
# 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
|
||||
@@ -1,370 +0,0 @@
|
||||
"""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()
|
||||
@@ -1,358 +0,0 @@
|
||||
#!/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()
|
||||
@@ -1,339 +0,0 @@
|
||||
"""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()
|
||||
@@ -1,141 +0,0 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user