Compare commits

..
Author SHA1 Message Date
SolitaryThinker e9d95b1c10 wip 2025-12-12 06:30:08 +00:00
SolitaryThinker 8b55e9706c wip 2025-12-10 01:59:07 +00:00
29 changed files with 278 additions and 2939 deletions
-11
View File
@@ -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"
-4
View File
@@ -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
+1 -1
View File
@@ -6,7 +6,7 @@
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/TM8JyJCd" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
-103
View File
@@ -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
-35
View File
@@ -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',
]
-185
View File
@@ -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())
-447
View File
@@ -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
-142
View File
@@ -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)
-34
View File
@@ -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()
-97
View File
@@ -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()
-490
View File
@@ -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"
)
-7
View File
@@ -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
-4
View File
@@ -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)
+20
View File
@@ -83,6 +83,9 @@ class PreprocessConfig:
speed_factor: float = 1.0
drop_short_ratio: float = 1.0
do_temporal_sample: bool = False
enable_smart_resize: bool = False
smart_resize_max_area: int | None = None
hw_aspect_threshold: float = 1.5
# Model configuration
training_cfg_rate: float = 0.0
@@ -184,6 +187,23 @@ class PreprocessConfig:
action=StoreBoolean,
default=PreprocessConfig.do_temporal_sample,
help="Whether to do temporal sampling")
preprocess_args.add_argument(
f"--{prefix_with_dot}enable-smart-resize",
action=StoreBoolean,
default=PreprocessConfig.enable_smart_resize,
help="Whether to enable smart resizing")
preprocess_args.add_argument(
f"--{prefix_with_dot}smart-resize-max-area",
type=int,
default=PreprocessConfig.smart_resize_max_area,
help="Maximum area for smart resizing")
preprocess_args.add_argument(
f"--{prefix_with_dot}hw-aspect-threshold",
type=float,
default=PreprocessConfig.hw_aspect_threshold,
help=
"Height/Width aspect ratio threshold. Allowed range is [1/threshold * target_aspect, threshold * target_aspect]."
)
# Model Training configuration
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
+41
View File
@@ -152,3 +152,44 @@ class TemporalRandomCrop:
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
def best_output_size(
width: int,
height: int,
width_stride: int,
height_stride: int,
max_area: int,
) -> tuple[int, int]:
"""
Calculate the best output size (width, height) given the original dimensions, strides and max area.
The aspect ratio is preserved as much as possible.
Args:
width (int): Original width
height (int): Original height
width_stride (int): Width stride requirement
height_stride (int): Height stride requirement
max_area (int): Maximum allowed area (width * height)
Returns:
tuple[int, int]: (new_width, new_height)
"""
aspect_ratio = width / height
# Scale dimensions if they exceed max_area
current_area = width * height
if current_area > max_area:
scale = (max_area / current_area)**0.5
width = int(width * scale)
height = int(height * scale)
# Round to the nearest multiple of stride
width = round(width / width_stride) * width_stride
height = round(height / height_stride) * height_stride
# Ensure dimensions are at least one stride
width = max(width, width_stride)
height = max(height, height_stride)
return width, height
+2 -2
View File
@@ -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
@@ -5,13 +5,19 @@ from typing import cast
import numpy as np
import torch
import torchvision
import torchvision.transforms.functional as TF
from PIL import Image
from einops import rearrange
from torchvision import transforms
from fastvideo.configs.configs import VideoLoaderType
from fastvideo.dataset.transform import (CenterCropResizeVideo,
TemporalRandomCrop)
TemporalRandomCrop, best_output_size)
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.logger import init_logger
logger = init_logger(__name__)
from fastvideo.pipelines.pipeline_batch_info import (ForwardBatch,
PreprocessBatch)
from fastvideo.pipelines.stages.base import PipelineStage
@@ -41,6 +47,7 @@ class VideoTransformStage(PipelineStage):
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.fps, list)
assert isinstance(batch.num_frames, list)
assert fastvideo_args.preprocess_config is not None
if batch.data_type != "video":
return batch
@@ -49,8 +56,17 @@ class VideoTransformStage(PipelineStage):
raise ValueError("Video loader is not set")
video_pixel_batch = []
pil_image_batch = []
enable_smart_resize = fastvideo_args.preprocess_config.enable_smart_resize
smart_resize_max_area = fastvideo_args.preprocess_config.smart_resize_max_area
if smart_resize_max_area is None:
smart_resize_max_area = 480 * 832
calculated_size = None
for i in range(len(batch.video_loader)):
# logger.info(f"Processing video {i+1}/{len(batch.video_loader)}")
frame_interval = batch.fps[i] / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch.num_frames[i],
@@ -63,8 +79,28 @@ class VideoTransformStage(PipelineStage):
else:
frame_indices = frame_indices[:self.num_frames]
logger.info(
f"Frame indices selected (count={len(frame_indices)}): [{frame_indices[0]}, ..., {frame_indices[-1]}]"
)
if fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
video = batch.video_loader[i].get_frames_at(frame_indices).data
try:
video = batch.video_loader[i].get_frames_at(
frame_indices).data
except Exception as e:
# Try to get filename if available in PreprocessBatch
video_path = "unknown"
print(f"batch: {batch}")
if isinstance(batch, PreprocessBatch) and hasattr(
batch, 'video_file_name') and i < len(
batch.video_file_name):
video_path = batch.video_file_name[i]
logger.error(
f"Failed to load frames for video {video_path}: {e}")
logger.error(
f"Attempting to load frame indices: {frame_indices}")
raise e
elif fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHVISION:
video, _, _ = torchvision.io.read_video(batch.video_loader[i],
output_format="TCHW")
@@ -73,16 +109,75 @@ class VideoTransformStage(PipelineStage):
raise ValueError(
f"Invalid video loader type: {fastvideo_args.preprocess_config.video_loader_type}"
)
video = self.video_transform(video)
video_pixel_batch.append(video)
logger.info(f"Video tensor shape after loading: {video.shape}")
if enable_smart_resize:
if calculated_size is None:
_, _, h_in, w_in = video.shape
# Get config values
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[
2] * vae_stride
ow, oh = best_output_size(w_in, h_in, dw, dh,
smart_resize_max_area)
calculated_size = (oh, ow)
logger.info(
f"Smart resize: input=({h_in}, {w_in}), output=({oh}, {ow})"
)
# Resize video frames using CenterCropResizeVideo (efficient)
processed_video = CenterCropResizeVideo(calculated_size)(video)
logger.info(
f"Processed video shape after resize: {processed_video.shape}"
)
video_pixel_batch.append(processed_video)
# Process pil_image (condition) with high quality Lanczos if I2V
if fastvideo_args.workload_type == WorkloadType.I2V:
# Extract first frame
img_tensor = video[0] # C, H, W
img = TF.to_pil_image(img_tensor)
iw, ih = img.width, img.height
ow, oh = calculated_size[1], calculated_size[0]
# Smart Resize logic for PIL image
scale = max(ow / iw, oh / ih)
resampling = Image.Resampling.LANCZOS if hasattr(
Image, 'Resampling') else Image.LANCZOS
img = img.resize((round(iw * scale), round(ih * scale)),
resampling)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
# to tensor [0, 255] uint8
img_t = torch.from_numpy(np.array(img)).permute(
2, 0, 1).unsqueeze(0)
pil_image_batch.append(img_t)
else:
video = self.video_transform(video)
video_pixel_batch.append(video)
video_pixel_values = torch.stack(video_pixel_batch)
logger.info(
f"Final stacked video batch shape: {video_pixel_values.shape}")
video_pixel_values = rearrange(video_pixel_values,
"b t c h w -> b c t h w")
video_pixel_values = video_pixel_values.to(torch.uint8)
if fastvideo_args.workload_type == WorkloadType.I2V:
batch.pil_image = video_pixel_values[:, :, 0, :, :]
if enable_smart_resize and len(pil_image_batch) > 0:
batch.pil_image = torch.cat(
pil_image_batch, dim=0).to(self.device if hasattr(
self, 'device') else video_pixel_values.device)
else:
batch.pil_image = video_pixel_values[:, :, 0, :, :]
video_pixel_values = video_pixel_values.float() / 255.0
batch.latents = video_pixel_values
@@ -82,7 +82,6 @@ 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,
@@ -93,7 +92,6 @@ class LatentPreparationStage(PipelineStage):
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,
@@ -104,7 +102,6 @@ class LatentPreparationStage(PipelineStage):
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
bcthw_shape = shape
# Validate generator if it's a list
if isinstance(generator, list) and len(generator) != batch_size:
@@ -126,7 +123,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
@@ -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")
-4
View File
@@ -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")
+66 -8
View File
@@ -71,6 +71,7 @@ class PreprocessingDataValidator:
for name, validator in self.validators.items():
if not validator(batch):
logger.info(f"Failed validation for {name}")
self.filter_counts[name] += 1
return False
@@ -87,6 +88,8 @@ class PreprocessingDataValidator:
"""Validate resolution constraints"""
aspect = self.max_height / self.max_width
height = None
width = None
if batch["resolution"] is not None:
height = batch["resolution"].get("height", None)
width = batch["resolution"].get("width", None)
@@ -94,12 +97,15 @@ class PreprocessingDataValidator:
if height is None or width is None:
return False
return self._filter_resolution(
ret = self._filter_resolution(
height,
width,
max_h_div_w_ratio=self.hw_aspect_threshold * aspect,
min_h_div_w_ratio=1 / self.hw_aspect_threshold * aspect,
)
if not ret:
logger.info(f"failed in resolution: {batch['caption']}")
return ret
def _filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
min_h_div_w_ratio: float) -> bool:
@@ -113,14 +119,19 @@ class PreprocessingDataValidator:
if (batch["num_frames"] / batch["fps"]
> self.video_length_tolerance_range *
(self.num_frames / self.train_fps * self.speed_factor)):
logger.info("Failed in 1")
return False
frame_interval = batch["fps"] / self.train_fps
frame_interval = (batch["fps"] / self.train_fps) * self.speed_factor
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch["num_frames"],
frame_interval).astype(int)
return not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
# logger.info("Failed in 2")
result = not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
if not result:
logger.info(f"failed in frame_sampling: {batch['caption']}")
return result
def log_validation_stats(self):
info = ""
@@ -280,9 +291,46 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
dataset = dataset.shard(num_shards=get_world_size(),
index=get_world_rank())
elif preprocess_config.dataset_type == DatasetType.MERGED:
metadata_json_path = os.path.join(preprocess_config.dataset_path,
"videos2caption.json")
video_folder = os.path.join(preprocess_config.dataset_path, "videos")
merge_txt_path = os.path.join(preprocess_config.dataset_path,
"merge.txt")
if os.path.exists(merge_txt_path):
logger.info(f"Found merge.txt at {merge_txt_path}")
with open(merge_txt_path) as f:
line = f.read().strip()
if "," not in line:
raise ValueError(
f"Invalid format in {merge_txt_path}: expected 'video_folder,metadata_json_path'"
)
video_folder, metadata_json_path = line.split(",", 1)
video_folder = video_folder.strip()
metadata_json_path = metadata_json_path.strip()
if not os.path.isabs(video_folder):
video_folder = os.path.join(preprocess_config.dataset_path,
video_folder)
if not os.path.isabs(metadata_json_path):
metadata_json_path = os.path.join(
preprocess_config.dataset_path, metadata_json_path)
else:
logger.info(
f"merge.txt not found at {merge_txt_path}, using default paths")
metadata_json_path = os.path.join(preprocess_config.dataset_path,
"videos2caption.json")
video_folder = os.path.join(preprocess_config.dataset_path,
"videos")
if not os.path.exists(metadata_json_path):
logger.error(f"Metadata file not found: {metadata_json_path}")
raise FileNotFoundError(
f"Metadata file not found: {metadata_json_path}")
if not os.path.exists(video_folder):
logger.error(f"Video folder not found: {video_folder}")
raise FileNotFoundError(f"Video folder not found: {video_folder}")
logger.info(f"Using metadata file: {metadata_json_path}")
logger.info(f"Using video folder: {video_folder}")
dataset = load_dataset("json",
data_files=metadata_json_path,
split=split)
@@ -293,9 +341,18 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
if "path" in column_names:
dataset = dataset.rename_column("path", "name")
dataset = dataset.filter(validator)
print(f"Length of dataset before filtering: {len(dataset)}")
if len(dataset) > 0:
print(f"DEBUG: First item in dataset: {dataset[0]}")
# Disable caching to ensure our print statements run
dataset = dataset.filter(validator, load_from_cache_file=False)
validator.log_validation_stats()
print(f"Length of dataset after filtering: {len(dataset)}")
dataset = dataset.shard(num_shards=get_world_size(),
index=get_world_rank())
print(f"Length of dataset after sharding: {len(dataset)}")
# add video column
def add_video_column(item: dict[str, Any]) -> dict[str, Any]:
@@ -303,6 +360,7 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
return item
dataset = dataset.map(add_video_column)
print(f"Length of dataset after mapping: {len(dataset)}")
if preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
dataset = dataset.cast_column("video", Video())
else:
@@ -40,6 +40,7 @@ class PreprocessWorkflow(WorkflowBase):
video_length_tolerance_range=preprocess_config.
video_length_tolerance_range,
drop_short_ratio=preprocess_config.drop_short_ratio,
hw_aspect_threshold=preprocess_config.hw_aspect_threshold,
)
self.add_component("raw_data_validator", raw_data_validator)
@@ -4,17 +4,37 @@ import os
import random
def generate_merged_validation_json(input_dir, output_file):
# read in video2caption.json
with open(os.path.join(input_dir, "video2caption_replace.json"), "r") as f:
def generate_merged_validation_json(args):
input_file = args.input_file
output_validation_file = args.output_validation_file
if args.output_train_file:
output_train_file = args.output_train_file
else:
base, ext = os.path.splitext(input_file)
output_train_file = f"{base}_train{ext}"
# read in input json
print(f"Reading from {input_file}")
with open(input_file, "r") as f:
video2caption = json.load(f)
# count how many elements are in the list
num_elements = len(video2caption)
print(f"Number of elements in video2caption.json: {num_elements}")
print(f"Number of elements in input file: {num_elements}")
# randomly sample 64 elements from the list
sampled_elements = random.sample(video2caption, 64)
# randomly sample elements from the list
num_sample = min(args.num_elements, num_elements)
indices = set(random.sample(range(num_elements), num_sample))
sampled_elements = []
remaining_elements = []
for i in range(num_elements):
if i in indices:
sampled_elements.append(video2caption[i])
else:
remaining_elements.append(video2caption[i])
# Transform sampled elements into validation.json format
validation_data = []
@@ -23,10 +43,10 @@ def generate_merged_validation_json(input_dir, output_file):
validation_entry = {
"caption": element["cap"],
"video_path": element.get("path", ""),
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
"num_inference_steps": args.num_inference_steps,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames
}
validation_data.append(validation_entry)
@@ -36,18 +56,25 @@ def generate_merged_validation_json(input_dir, output_file):
}
# Write the validation JSON to the output file
with open(output_file, "w") as f:
with open(output_validation_file, "w") as f:
json.dump(validation_json, f, indent=2)
print(f"Generated validation JSON with {len(validation_data)} entries and saved to {output_file}")
print(f"Generated validation JSON with {len(validation_data)} entries and saved to {output_validation_file}")
# Write the remaining JSON to the output train file
with open(output_train_file, "w") as f:
json.dump(remaining_elements, f, indent=2)
print(f"Saved remaining {len(remaining_elements)} entries to {output_train_file}")
def main():
parser = argparse.ArgumentParser()
# dataset_type: "mixkit"
# dataset_type: "merged"
parser.add_argument("--dataset_type", choices=["merged"], required=True)
parser.add_argument("--input_dir", type=str, required=True)
parser.add_argument("--output_file", type=str, required=True)
parser.add_argument("--input_file", type=str, required=True, help="Path to input json file")
parser.add_argument("--output_validation_file", type=str, required=True, help="Path to output validation json file")
parser.add_argument("--output_train_file", type=str, help="Path to output train json file (remaining data). Defaults to {input_filename}_train.json")
parser.add_argument("--num_elements", type=int, default=64)
parser.add_argument("--num_frames", type=int, default=77)
parser.add_argument("--height", type=int, default=480)
@@ -56,8 +83,8 @@ def main():
args = parser.parse_args()
if args.dataset_type == "merged":
generate_merged_validation_json(args.input_dir, args.output_file)
generate_merged_validation_json(args)
if __name__ == "__main__":
main()
main()
-68
View File
@@ -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
-370
View File
@@ -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()
-339
View File
@@ -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()
-141
View File
@@ -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()