Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9d95b1c10 | ||
|
|
8b55e9706c |
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -1,103 +0,0 @@
|
||||
# FVD (Fréchet Video Distance) Benchmark
|
||||
|
||||
Evaluate generated video quality using FVD with the I3D feature extractor.
|
||||
|
||||
## Quick Start
|
||||
|
||||
**Run the benchmark:**
|
||||
|
||||
```bash
|
||||
bash benchmarks/scripts/run.sh
|
||||
```
|
||||
|
||||
That's it! The script auto-installs dependencies and runs the benchmark.
|
||||
|
||||
**To customize:** Edit `benchmarks/fvd/run_fvd.py` to change:
|
||||
- Video paths (`real_dir`, `gen_dir`)
|
||||
- Number of videos, frames, sampling strategy
|
||||
- Device, batch size, caching, etc.
|
||||
|
||||
## Advanced Usage (CLI)
|
||||
|
||||
For more control without editing Python files, use the CLI.
|
||||
|
||||
**First-time setup** (one-time per pod/environment):
|
||||
|
||||
```bash
|
||||
bash benchmarks/scripts/setup_fvd.sh
|
||||
```
|
||||
|
||||
Then run any configuration you want:
|
||||
|
||||
```bash
|
||||
# Custom configuration
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--num-videos 1024 \
|
||||
--num-frames 32 \
|
||||
--clip-strategy random \
|
||||
--batch-size 32 \
|
||||
--seed 42
|
||||
```
|
||||
|
||||
**Standard protocols:**
|
||||
|
||||
```bash
|
||||
# Use predefined protocols
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
|
||||
```
|
||||
|
||||
**Feature caching** (speed up repeated evaluations):
|
||||
|
||||
```bash
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--protocol fvd2048_16f \
|
||||
--cache-real-features cache/real # Directory path (will save/load cache/real/real_features.pkl)
|
||||
```
|
||||
|
||||
Run `python -m benchmarks.fvd.cli --help` for all options.
|
||||
|
||||
## Available Protocols
|
||||
|
||||
- `fvd2048_16f` - Standard (2048 videos, 16 frames)
|
||||
- `fvd2048_128f` - Long videos (128 frames)
|
||||
- `fvd2048_128f_subsample8` - Subsampled long videos
|
||||
- `quick_test` - Fast testing (10 videos)
|
||||
|
||||
## Configuration Options
|
||||
|
||||
Key options in `FVDConfig`:
|
||||
|
||||
```python
|
||||
num_videos=2048, # Videos to evaluate
|
||||
num_frames_per_clip=16, # Frames per clip
|
||||
clip_strategy='beginning', # beginning|random|uniform|middle|sliding
|
||||
frame_stride=1, # Frame subsampling
|
||||
batch_size=32, # GPU batch size
|
||||
device='cuda', # cuda|cpu
|
||||
cache_real_features=None, # Cache path for speed
|
||||
seed=42, # Reproducibility
|
||||
```
|
||||
|
||||
## Programmatic Usage
|
||||
|
||||
```python
|
||||
from benchmarks.fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
config = FVDConfig.fvd2048_16f() # or custom config
|
||||
results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
print(f"FVD: {results['fvd']:.2f}")
|
||||
```
|
||||
|
||||
## Notes
|
||||
|
||||
- I3D model auto-downloads from Hugging Face on first run
|
||||
- Requires minimum 10 frames per clip
|
||||
- Supports both video files (.mp4, .avi, etc.) and frame directories
|
||||
- `--cache-real-features` expects a **directory path** (e.g., `cache/real`), it will automatically create/load `real_features.pkl` inside that directory
|
||||
@@ -1,35 +0,0 @@
|
||||
"""
|
||||
FastVideo Frechet Video Distance (FVD) Benchmark Module.
|
||||
>>> from fastvideo.benchmarks.fvd import compute_fvd_with_config, FVDConfig
|
||||
>>> config = FVDConfig.fvd2048_16f() # Standard protocol
|
||||
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
>>> print(f"FVD: {results['fvd']:.2f}")
|
||||
"""
|
||||
|
||||
from .fvd import (
|
||||
compute_fvd,
|
||||
compute_fvd_with_config,
|
||||
compute_frechet_distance,
|
||||
compute_statistics,
|
||||
FVDConfig,
|
||||
)
|
||||
from .i3d_model import I3DFeatureExtractor
|
||||
from .video_utils import (
|
||||
load_video_auto,
|
||||
sample_clips_from_video,
|
||||
load_video_clips_streaming,
|
||||
ClipSamplingStrategy,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
'compute_fvd',
|
||||
'compute_fvd_with_config',
|
||||
'compute_frechet_distance',
|
||||
'compute_statistics',
|
||||
'FVDConfig',
|
||||
'I3DFeatureExtractor',
|
||||
'load_video_auto',
|
||||
'sample_clips_from_video',
|
||||
'load_video_clips_streaming',
|
||||
'ClipSamplingStrategy',
|
||||
]
|
||||
@@ -1,185 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from .fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Compute Fréchet Video Distance (FVD)',
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
epilog="""
|
||||
Examples:
|
||||
# Standard FVD2048_16f protocol
|
||||
python -m fastvideo.benchmarks.fvd.cli \\
|
||||
--real-path data/real/ \\
|
||||
--gen-path outputs/gen/ \\
|
||||
--protocol fvd2048_16f
|
||||
|
||||
# Custom configuration
|
||||
python -m fastvideo.benchmarks.fvd.cli \\
|
||||
--real-path data/real/ \\
|
||||
--gen-path outputs/gen/ \\
|
||||
--num-videos 1024 \\
|
||||
--num-frames 32 \\
|
||||
--clip-strategy random \\
|
||||
--frame-stride 2
|
||||
""")
|
||||
|
||||
# Required arguments
|
||||
parser.add_argument('--real-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to real videos directory')
|
||||
parser.add_argument('--gen-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to generated videos directory')
|
||||
|
||||
# Reproducibility
|
||||
parser.add_argument(
|
||||
'--seed',
|
||||
type=int,
|
||||
default=None,
|
||||
help='Random seed for reproducibility (np.random, random, torch)')
|
||||
|
||||
# Protocol presets
|
||||
parser.add_argument('--protocol',
|
||||
type=str,
|
||||
default=None,
|
||||
choices=[
|
||||
'fvd2048_16f', 'fvd2048_128f',
|
||||
'fvd2048_128f_subsample8', 'quick_test'
|
||||
],
|
||||
help='Use standard protocol (overrides other settings)')
|
||||
|
||||
# Video selection
|
||||
parser.add_argument('--num-videos',
|
||||
type=int,
|
||||
default=2048,
|
||||
help='Number of videos to use (default: 2048)')
|
||||
|
||||
# Clip sampling
|
||||
parser.add_argument('--num-frames',
|
||||
type=int,
|
||||
default=16,
|
||||
help='Number of frames per clip (default: 16)')
|
||||
parser.add_argument('--num-clips',
|
||||
type=int,
|
||||
default=1,
|
||||
help='Number of clips per video (default: 1)')
|
||||
parser.add_argument(
|
||||
'--clip-strategy',
|
||||
type=str,
|
||||
default='beginning',
|
||||
choices=['beginning', 'random', 'uniform', 'middle', 'sliding', 'all'],
|
||||
help='Clip sampling strategy (default: beginning)')
|
||||
parser.add_argument(
|
||||
'--frame-stride',
|
||||
type=int,
|
||||
default=1,
|
||||
help='Frame stride for FPS subsampling (default: 1, no subsampling)')
|
||||
parser.add_argument('--temporal-stride',
|
||||
type=int,
|
||||
default=1,
|
||||
help='Temporal stride for sliding window (default: 1)')
|
||||
|
||||
# Data processing
|
||||
parser.add_argument('--no-frame-dirs',
|
||||
action='store_true',
|
||||
help='Disable frame directory support')
|
||||
|
||||
# Computation
|
||||
parser.add_argument('--batch-size',
|
||||
type=int,
|
||||
default=32,
|
||||
help='Batch size for feature extraction (default: 32)')
|
||||
parser.add_argument('--device',
|
||||
type=str,
|
||||
default='cuda',
|
||||
choices=['cuda', 'cpu'],
|
||||
help='Device to use (default: cuda)')
|
||||
|
||||
# Caching
|
||||
parser.add_argument('--cache-real-features',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Path to cache real video features')
|
||||
parser.add_argument('--i3d-model-path',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Custom cache path for I3D model')
|
||||
|
||||
# Output
|
||||
parser.add_argument('--output',
|
||||
type=str,
|
||||
default='fvd_results.json',
|
||||
help='Output JSON file (default: fvd_results.json)')
|
||||
parser.add_argument('--quiet',
|
||||
action='store_true',
|
||||
help='Suppress progress output')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create config
|
||||
if args.protocol:
|
||||
protocol_map = {
|
||||
'fvd2048_16f': FVDConfig.fvd2048_16f,
|
||||
'fvd2048_128f': FVDConfig.fvd2048_128f,
|
||||
'fvd2048_128f_subsample8': FVDConfig.fvd2048_128f_subsample8,
|
||||
'quick_test': FVDConfig.quick_test,
|
||||
}
|
||||
config = protocol_map[args.protocol]()
|
||||
|
||||
# Override device and caching from args
|
||||
config.device = args.device
|
||||
config.cache_real_features = args.cache_real_features
|
||||
config.i3d_model_path = args.i3d_model_path
|
||||
config.batch_size = args.batch_size
|
||||
config.seed = args.seed
|
||||
else:
|
||||
# Custom config from args
|
||||
config = FVDConfig(num_videos=args.num_videos,
|
||||
num_frames_per_clip=args.num_frames,
|
||||
num_clips_per_video=args.num_clips,
|
||||
clip_strategy=args.clip_strategy,
|
||||
frame_stride=args.frame_stride,
|
||||
temporal_stride=args.temporal_stride,
|
||||
support_frame_dirs=not args.no_frame_dirs,
|
||||
batch_size=args.batch_size,
|
||||
device=args.device,
|
||||
cache_real_features=args.cache_real_features,
|
||||
i3d_model_path=args.i3d_model_path,
|
||||
seed=args.seed)
|
||||
|
||||
# Compute FVD
|
||||
try:
|
||||
results = compute_fvd_with_config(real_videos=args.real_path,
|
||||
gen_videos=args.gen_path,
|
||||
config=config,
|
||||
verbose=not args.quiet)
|
||||
|
||||
# Save results
|
||||
output_path = Path(args.output)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
with open(output_path, 'w') as f:
|
||||
json.dump(results, f, indent=2)
|
||||
|
||||
print(f"\nResults saved to {output_path}")
|
||||
print(f"FVD: {results['fvd']:.2f}")
|
||||
print(f"Protocol: {results['protocol']}")
|
||||
|
||||
return 0
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error: {e}", file=sys.stderr)
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
sys.exit(main())
|
||||
@@ -1,447 +0,0 @@
|
||||
import numpy as np
|
||||
import scipy.linalg
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from collections.abc import Iterator
|
||||
import pickle
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from .i3d_model import I3DFeatureExtractor
|
||||
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
|
||||
|
||||
|
||||
def compute_statistics(features: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Compute mean and covariance."""
|
||||
mu = np.mean(features, axis=0)
|
||||
sigma = np.cov(features, rowvar=False)
|
||||
return mu, sigma
|
||||
|
||||
|
||||
def compute_frechet_distance(mu1: np.ndarray,
|
||||
sigma1: np.ndarray,
|
||||
mu2: np.ndarray,
|
||||
sigma2: np.ndarray,
|
||||
eps: float = 1e-6) -> float:
|
||||
"""
|
||||
Compute Fréchet distance between two Gaussians.
|
||||
"""
|
||||
sigma1 = sigma1 + eps * np.eye(sigma1.shape[0])
|
||||
sigma2 = sigma2 + eps * np.eye(sigma2.shape[0])
|
||||
|
||||
diff = mu1 - mu2
|
||||
mean_distance = np.sum(diff**2)
|
||||
|
||||
trace_sum = np.trace(sigma1 + sigma2)
|
||||
|
||||
covmean = scipy.linalg.sqrtm(sigma1 @ sigma2)
|
||||
|
||||
if np.iscomplexobj(covmean):
|
||||
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
|
||||
print(
|
||||
f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
|
||||
covmean = covmean.real
|
||||
|
||||
trace_product = np.trace(covmean)
|
||||
|
||||
fvd = mean_distance + trace_sum - 2 * trace_product
|
||||
|
||||
return float(fvd)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FVDConfig:
|
||||
# default configuration for FVD computation:
|
||||
|
||||
# Video selection
|
||||
num_videos: int = 2048
|
||||
|
||||
# Clip sampling
|
||||
num_frames_per_clip: int = 16
|
||||
num_clips_per_video: int = 1
|
||||
clip_strategy: str | ClipSamplingStrategy = 'beginning'
|
||||
|
||||
# Temporal subsampling
|
||||
frame_stride: int = 1 # 1=no subsampling, 2=every 2nd, 8=every 8th
|
||||
temporal_stride: int = 1 # For sliding window clips
|
||||
|
||||
# Data processing
|
||||
video_extensions: list[str] = field(
|
||||
default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
|
||||
support_frame_dirs: bool = True
|
||||
|
||||
# Computation
|
||||
batch_size: int = 32
|
||||
device: str = 'cuda'
|
||||
|
||||
use_streaming: bool = True
|
||||
resize_before_extraction: bool = True
|
||||
|
||||
# Caching
|
||||
cache_real_features: str | None = None
|
||||
i3d_model_path: str | None = None
|
||||
|
||||
# Reproducibility
|
||||
seed: int | None = None
|
||||
|
||||
@classmethod
|
||||
def fvd2048_16f(cls) -> 'FVDConfig':
|
||||
"""
|
||||
Standard FVD protocol: 2048 videos, 16 frames, beginning clip.
|
||||
|
||||
most common FVD configuration used in papers
|
||||
"""
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def fvd2048_128f(cls) -> 'FVDConfig':
|
||||
"""Long video protocol: 2048 videos, 128 frames."""
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=128,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def fvd2048_128f_subsample8(cls) -> 'FVDConfig':
|
||||
"""
|
||||
Long video with FPS subsampling: 2048 videos, 128 frames (every 8th).
|
||||
Used for very long videos - samples every 8th frame
|
||||
"""
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=16,
|
||||
frame_stride=8,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def quick_test(cls) -> 'FVDConfig':
|
||||
"""Quick test config: 100 videos, 16 frames."""
|
||||
return cls(num_videos=100,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning')
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Export config to dict for logging"""
|
||||
return {
|
||||
'num_videos': self.num_videos,
|
||||
'num_frames_per_clip': self.num_frames_per_clip,
|
||||
'num_clips_per_video': self.num_clips_per_video,
|
||||
'clip_strategy': str(self.clip_strategy),
|
||||
'frame_stride': self.frame_stride,
|
||||
'temporal_stride': self.temporal_stride,
|
||||
'batch_size': self.batch_size,
|
||||
'device': self.device,
|
||||
'seed': self.seed,
|
||||
'use_streaming': self.use_streaming,
|
||||
}
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""Human-readable protocol name"""
|
||||
desc = f"FVD{self.num_videos}_{self.num_frames_per_clip}f"
|
||||
if self.frame_stride > 1:
|
||||
desc += f"_subsample{self.frame_stride}"
|
||||
if self.num_clips_per_video > 1:
|
||||
desc += f"_{self.num_clips_per_video}clips"
|
||||
if self.clip_strategy != 'beginning':
|
||||
desc += f"_{self.clip_strategy}"
|
||||
return desc
|
||||
|
||||
|
||||
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
|
||||
extractor: I3DFeatureExtractor,
|
||||
batch_size: int = 32,
|
||||
max_clips: int | None = None,
|
||||
verbose: bool = True) -> np.ndarray:
|
||||
"""
|
||||
Extract features from a video clip generator using streaming.
|
||||
|
||||
Args:
|
||||
video_generator: Iterator yielding clips [T, C, H, W]
|
||||
extractor: I3D feature extractor
|
||||
batch_size: Batch size for processing
|
||||
max_clips: Maximum clips to process (for validation)
|
||||
verbose: Show progress
|
||||
|
||||
Returns:
|
||||
features: [N, 400] numpy array
|
||||
"""
|
||||
all_features = []
|
||||
batch = []
|
||||
clip_count = 0
|
||||
|
||||
if verbose:
|
||||
print(f"Extracting features with batch_size={batch_size}...")
|
||||
|
||||
for clip_count, clip in enumerate(video_generator):
|
||||
batch.append(clip)
|
||||
|
||||
# Process batch when full
|
||||
if len(batch) == batch_size:
|
||||
batch_tensor = torch.stack(batch).to(extractor.device)
|
||||
features = extractor.extract_features(batch_tensor,
|
||||
batch_size=batch_size,
|
||||
verbose=False)
|
||||
all_features.append(features.cpu().numpy())
|
||||
|
||||
batch = [] # Clear batch
|
||||
|
||||
if verbose and clip_count % (batch_size * 10) == 0:
|
||||
print(f"Processed {clip_count} clips...")
|
||||
|
||||
# Stop if we've reached max_clips
|
||||
if max_clips is not None and clip_count >= max_clips:
|
||||
break
|
||||
|
||||
# Process remaining clips
|
||||
if len(batch) > 0:
|
||||
batch_tensor = torch.stack(batch).to(extractor.device)
|
||||
features = extractor.extract_features(batch_tensor,
|
||||
batch_size=len(batch),
|
||||
verbose=False)
|
||||
all_features.append(features.cpu().numpy())
|
||||
|
||||
if len(all_features) == 0:
|
||||
raise RuntimeError("No features extracted - check video loading")
|
||||
|
||||
features = np.concatenate(all_features, axis=0)
|
||||
|
||||
if verbose:
|
||||
print(f"Extracted {len(features)} feature vectors")
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def load_or_compute_features(videos: str | Path | torch.Tensor,
|
||||
extractor: I3DFeatureExtractor,
|
||||
config: FVDConfig,
|
||||
cache_path: str | None = None,
|
||||
cache_name: str = "real_features") -> np.ndarray:
|
||||
"""Load features from cache or compute (with streaming support)"""
|
||||
|
||||
if cache_path is not None:
|
||||
cache_file = Path(cache_path) / f"{cache_name}.pkl"
|
||||
if cache_file.exists():
|
||||
print(f"Loading cached features from {cache_file}")
|
||||
with open(cache_file, 'rb') as f:
|
||||
features = pickle.load(f)
|
||||
|
||||
# Validate and limit based on config
|
||||
max_features = config.num_videos * config.num_clips_per_video
|
||||
if len(features) < max_features:
|
||||
print(
|
||||
f"WARNING: Cache has {len(features)} features but need {max_features}"
|
||||
)
|
||||
print("Recomputing features...")
|
||||
elif len(features) > max_features:
|
||||
features = features[:max_features]
|
||||
return features
|
||||
else:
|
||||
return features
|
||||
|
||||
# Compute features
|
||||
if isinstance(videos, str | Path):
|
||||
target_size = (224, 224) if config.resize_before_extraction else None
|
||||
|
||||
video_generator = load_video_clips_streaming(
|
||||
videos,
|
||||
num_frames=config.num_frames_per_clip,
|
||||
max_videos=config.num_videos,
|
||||
clip_strategy=config.clip_strategy,
|
||||
frame_stride=config.frame_stride,
|
||||
num_clips_per_video=config.num_clips_per_video,
|
||||
video_extensions=config.video_extensions,
|
||||
support_frame_dirs=config.support_frame_dirs,
|
||||
target_size=target_size,
|
||||
verbose=True)
|
||||
|
||||
max_clips = config.num_videos * config.num_clips_per_video
|
||||
features = extract_features_streaming(video_generator,
|
||||
extractor,
|
||||
batch_size=config.batch_size,
|
||||
max_clips=max_clips,
|
||||
verbose=True)
|
||||
|
||||
else:
|
||||
# Already a tensor
|
||||
print(f"Extracting features from {len(videos)} video tensors...")
|
||||
features = extractor.extract_features(videos,
|
||||
batch_size=config.batch_size,
|
||||
verbose=True)
|
||||
features = features.numpy()
|
||||
|
||||
# Validate feature count
|
||||
expected_count = config.num_videos * config.num_clips_per_video
|
||||
if len(features) < expected_count:
|
||||
raise ValueError(
|
||||
f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
|
||||
f"Found fewer videos than expected. Check your video directory.")
|
||||
elif len(features) > expected_count:
|
||||
print(f"Truncating {len(features)} features to {expected_count}")
|
||||
features = features[:expected_count]
|
||||
|
||||
# Cache features if requested
|
||||
if cache_path is not None:
|
||||
cache_dir = Path(cache_path)
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
cache_file = cache_dir / f"{cache_name}.pkl"
|
||||
print(f"Caching features to {cache_file}")
|
||||
with open(cache_file, 'wb') as f:
|
||||
pickle.dump(features, f)
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def compute_fvd(real_videos: str | Path | torch.Tensor,
|
||||
gen_videos: str | Path | torch.Tensor,
|
||||
num_frames: int = 16,
|
||||
batch_size: int = 32,
|
||||
device: str = 'cuda',
|
||||
num_videos: int | None = 2048,
|
||||
cache_real_features: str | None = None,
|
||||
i3d_model_path: str | None = None,
|
||||
seed: int | None = None,
|
||||
verbose: bool = True) -> float:
|
||||
"""
|
||||
Compute Fréchet Video Distance (FVD)
|
||||
|
||||
For advanced control, use compute_fvd_with_config() instead.
|
||||
|
||||
Args:
|
||||
real_videos: Path to real videos or tensor [N, T, C, H, W]
|
||||
gen_videos: Path to generated videos or tensor [N, T, C, H, W]
|
||||
num_frames: Frames per video (default: 16)
|
||||
batch_size: Batch size (default: 32)
|
||||
device: 'cuda' or 'cpu' (default: 'cuda')
|
||||
num_videos: Max videos (default: 2048)
|
||||
cache_real_features: Cache path for real features
|
||||
i3d_model_path: Custom I3D model cache path
|
||||
seed: Random seed for reproducibility
|
||||
verbose: Print progress
|
||||
|
||||
Returns:
|
||||
FVD score (float). Lower is better.
|
||||
"""
|
||||
num_videos = num_videos if num_videos is not None else 2048
|
||||
|
||||
config = FVDConfig(
|
||||
num_videos=num_videos,
|
||||
num_frames_per_clip=num_frames,
|
||||
batch_size=batch_size,
|
||||
device=device,
|
||||
cache_real_features=cache_real_features,
|
||||
i3d_model_path=i3d_model_path,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
|
||||
return result['fvd']
|
||||
|
||||
|
||||
def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
|
||||
gen_videos: str | Path | torch.Tensor,
|
||||
config: FVDConfig,
|
||||
verbose: bool = True) -> dict:
|
||||
"""
|
||||
Compute FVD using a standardized configuration.
|
||||
|
||||
This is the recommended way to compute FVD for reproducibility.
|
||||
|
||||
Args:
|
||||
real_videos: Path or tensors
|
||||
gen_videos: Path or tensors
|
||||
config: FVDConfig specifying protocol
|
||||
verbose: Print progress
|
||||
|
||||
Returns:
|
||||
results: Dictionary with:
|
||||
- 'fvd': FVD score (float)
|
||||
- 'protocol': Protocol name (str)
|
||||
- 'config': Configuration dict
|
||||
|
||||
Example:
|
||||
>>> config = FVDConfig.fvd2048_16f()
|
||||
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
>>> print(f"FVD: {results['fvd']:.2f}")
|
||||
>>> print(f"Protocol: {results['protocol']}") # "FVD2048_16f"
|
||||
"""
|
||||
# Seed for reproducibility
|
||||
if config.seed is not None:
|
||||
import random as _rnd
|
||||
_rnd.seed(config.seed)
|
||||
np.random.seed(config.seed)
|
||||
torch.manual_seed(config.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(config.seed)
|
||||
|
||||
if verbose:
|
||||
print("=" * 70)
|
||||
print(f"Computing FVD with protocol: {config}")
|
||||
print("=" * 70)
|
||||
print("\nConfiguration:")
|
||||
for key, value in config.to_dict().items():
|
||||
print(f" {key}: {value}")
|
||||
print()
|
||||
|
||||
# Initialize I3D
|
||||
if verbose:
|
||||
print(f"\nInitializing I3D model on {config.device}...")
|
||||
|
||||
extractor = I3DFeatureExtractor(device=config.device,
|
||||
cache_dir=config.i3d_model_path)
|
||||
|
||||
# Extract features
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Extracting REAL video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
real_features = load_or_compute_features(
|
||||
videos=real_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=config.cache_real_features,
|
||||
cache_name="real_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Extracting GENERATED video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
gen_features = load_or_compute_features(videos=gen_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=None,
|
||||
cache_name="gen_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\nReal videos/clips: {len(real_features)}")
|
||||
print(f"Generated videos/clips: {len(gen_features)}")
|
||||
print(f"\n{'='*70}")
|
||||
print("Computing statistics...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
mu_real, sigma_real = compute_statistics(real_features)
|
||||
mu_gen, sigma_gen = compute_statistics(gen_features)
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Computing Fréchet distance...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
fvd = compute_frechet_distance(mu_real, sigma_real, mu_gen, sigma_gen)
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print(f"FVD Score: {fvd:.4f}")
|
||||
print(f"Protocol: {config}")
|
||||
print(f"{'='*70}\n")
|
||||
|
||||
results = {
|
||||
'fvd': fvd,
|
||||
'protocol': str(config),
|
||||
'config': config.to_dict(),
|
||||
}
|
||||
|
||||
return results
|
||||
@@ -1,142 +0,0 @@
|
||||
"""I3D Feature Extractor for FVD Computation"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pathlib import Path
|
||||
from huggingface_hub import hf_hub_download
|
||||
from tqdm import tqdm
|
||||
from contextlib import suppress
|
||||
|
||||
|
||||
class I3DFeatureExtractor(nn.Module):
|
||||
"""
|
||||
I3D feature extractor for FVD computation.
|
||||
Extracts 400-dimensional features from videos using I3D model
|
||||
trained on Kinetics-400.
|
||||
"""
|
||||
|
||||
REPO_ID = 'flateon/FVD-I3D-torchscript'
|
||||
MODEL_FILENAME = 'i3d_torchscript.pt'
|
||||
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
cache_dir: str | Path | None = None):
|
||||
super().__init__()
|
||||
|
||||
self.device_str = device
|
||||
if device == 'cuda' and not torch.cuda.is_available():
|
||||
print(
|
||||
"Warning: CUDA requested but not available – falling back to CPU"
|
||||
)
|
||||
self.device = torch.device('cpu')
|
||||
else:
|
||||
self.device = torch.device(device)
|
||||
|
||||
self.cache_dir: str | None
|
||||
if cache_dir is not None:
|
||||
self.cache_dir = str(Path(cache_dir).resolve())
|
||||
else:
|
||||
self.cache_dir = None # Use HF default cache
|
||||
|
||||
self.model = self._load_model()
|
||||
self.model.eval()
|
||||
|
||||
with suppress(Exception):
|
||||
self.model.to(self.device)
|
||||
|
||||
def _load_model(self) -> torch.nn.Module:
|
||||
"""Download and load I3D TorchScript model from Hugging Face Hub."""
|
||||
print(f"Loading I3D model from Hugging Face Hub ({self.REPO_ID})...")
|
||||
|
||||
try:
|
||||
# Download model from Hugging Face Hub
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID,
|
||||
filename=self.MODEL_FILENAME,
|
||||
cache_dir=self.cache_dir)
|
||||
|
||||
# Load directly to chosen device
|
||||
model = torch.jit.load(model_path, map_location=self.device)
|
||||
print("I3D model loaded successfully")
|
||||
return model
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Preprocess videos for I3D.
|
||||
|
||||
Args:
|
||||
videos: [B, T, C, H, W], values in [0, 255]
|
||||
|
||||
Returns:
|
||||
Preprocessed videos [B, C, T, 224, 224] (normalized and resized)
|
||||
"""
|
||||
B, T, C, H, W = videos.shape
|
||||
|
||||
if T < 10:
|
||||
raise ValueError(f"I3D requires at least 10 frames, got {T}")
|
||||
|
||||
# Normalize to [0, 1] if needed
|
||||
if videos.max() > 1.0:
|
||||
videos = videos / 255.0
|
||||
|
||||
# Resize to 224x224 if needed
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.reshape(B * T, C, H, W)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.reshape(B, T, C, 224, 224)
|
||||
|
||||
# Convert to [B, C, T, H, W] format
|
||||
videos = videos.permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
return videos
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_features(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32,
|
||||
verbose: bool = True) -> torch.Tensor:
|
||||
"""
|
||||
Extract I3D features
|
||||
|
||||
Args:
|
||||
videos: [N, T, C, H, W], values in [0, 255]
|
||||
batch_size: Batch size for processing
|
||||
verbose: Show progress bar
|
||||
|
||||
Returns:
|
||||
Features [N, 400]
|
||||
"""
|
||||
N = len(videos)
|
||||
all_features = []
|
||||
|
||||
iterator = range(0, N, batch_size)
|
||||
if verbose:
|
||||
iterator = tqdm(iterator, desc="Extracting I3D features")
|
||||
|
||||
for i in iterator:
|
||||
batch = videos[i:i + batch_size].to(self.device)
|
||||
batch = self.preprocess(batch) # Now returns [B, C, T, H, W]
|
||||
|
||||
# Use the HF model without rescale/resize (we handle it in preprocess)
|
||||
features = self.model(batch,
|
||||
rescale=False,
|
||||
resize=False,
|
||||
return_features=True)
|
||||
|
||||
all_features.append(features.cpu())
|
||||
|
||||
return torch.cat(all_features, dim=0)
|
||||
|
||||
def __call__(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32) -> torch.Tensor:
|
||||
return self.extract_features(videos, batch_size=batch_size)
|
||||
@@ -1,34 +0,0 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config
|
||||
|
||||
root_dir = Path(__file__).parent.parent.parent
|
||||
sys.path.insert(0, str(root_dir))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# Get script directory
|
||||
script_dir = Path(__file__).parent.resolve()
|
||||
|
||||
clip_strategy = 'beginning' # Options: 'uniform', 'random', 'beginning', 'end', 'all'
|
||||
cfg = FVDConfig(
|
||||
num_videos=650,
|
||||
num_frames_per_clip=16,
|
||||
num_clips_per_video=1,
|
||||
clip_strategy=clip_strategy,
|
||||
frame_stride=1,
|
||||
batch_size=32,
|
||||
device='cuda',
|
||||
seed=42,
|
||||
cache_real_features=str(script_dir / f'fvd-cache/{clip_strategy}'),
|
||||
)
|
||||
|
||||
real_dir = "benchmarks/data/real_videos"
|
||||
gen_dir = "benchmarks/data/generated_videos"
|
||||
|
||||
results = compute_fvd_with_config(real_dir, gen_dir, cfg, verbose=True)
|
||||
print(f"FVD = {results['fvd']:.2f}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,97 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
import sys
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import random
|
||||
from fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
script_path = Path(__file__).resolve()
|
||||
fastvideo_root = script_path.parent.parent.parent
|
||||
sys.path.insert(0, str(fastvideo_root))
|
||||
|
||||
|
||||
def split_videos(video_dir: Path, n_per_subset: int = 128, seed: int = 42):
|
||||
subset_a = video_dir.parent / 'bair_full_subset_A'
|
||||
subset_b = video_dir.parent / 'bair_full_subset_B'
|
||||
|
||||
if subset_a.exists():
|
||||
shutil.rmtree(subset_a)
|
||||
if subset_b.exists():
|
||||
shutil.rmtree(subset_b)
|
||||
|
||||
subset_a.mkdir(parents=True)
|
||||
subset_b.mkdir(parents=True)
|
||||
|
||||
videos = sorted(video_dir.glob('*.mp4'))
|
||||
|
||||
random.seed(seed)
|
||||
shuffled = list(videos)
|
||||
random.shuffle(shuffled)
|
||||
|
||||
needed = n_per_subset * 2
|
||||
if len(shuffled) > needed:
|
||||
shuffled = shuffled[:needed]
|
||||
|
||||
mid = len(shuffled) // 2
|
||||
|
||||
print(f"\nSplitting {len(shuffled)} BAIR FULL videos:")
|
||||
print(f" Subset A: {mid} videos")
|
||||
print(f" Subset B: {len(shuffled) - mid} videos")
|
||||
|
||||
for v in shuffled[:mid]:
|
||||
shutil.copy2(v, subset_a / v.name)
|
||||
|
||||
for v in shuffled[mid:]:
|
||||
shutil.copy2(v, subset_b / v.name)
|
||||
|
||||
return subset_a, subset_b, mid
|
||||
|
||||
|
||||
def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
|
||||
config = FVDConfig(num_videos=num_videos,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning',
|
||||
batch_size=8,
|
||||
device='cuda',
|
||||
seed=42)
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST 1: Identity Test")
|
||||
print("=" * 70)
|
||||
|
||||
result1 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_a),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_identity = result1['fvd']
|
||||
print(f"\nIdentity FVD: {fvd_identity:.2f}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST 2: Real vs Real")
|
||||
print("=" * 70)
|
||||
|
||||
result2 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_b),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_real = result2['fvd']
|
||||
print(f"\nReal vs Real FVD: {fvd_real:.2f}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("RESULTS")
|
||||
print("=" * 70)
|
||||
print(f"Identity: {fvd_identity:.2f}")
|
||||
print(f"Real vs Real: {fvd_real:.2f}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
bair_dir = Path('benchmarks/data/bair_full_videos')
|
||||
|
||||
subset_a, subset_b, count = split_videos(bair_dir,
|
||||
n_per_subset=128,
|
||||
seed=42)
|
||||
validate_fvd(subset_a, subset_b, count)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,490 +0,0 @@
|
||||
import torch
|
||||
import cv2
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from collections.abc import Iterator
|
||||
from tqdm import tqdm
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ClipSamplingStrategy(Enum):
|
||||
"""Clip sampling strategies for FVD evaluation."""
|
||||
BEGINNING = 'beginning' # Take first N frames (most common)
|
||||
RANDOM = 'random' # Random N consecutive frames
|
||||
UNIFORM = 'uniform' # Uniformly spaced frames across video
|
||||
MIDDLE = 'middle' # Middle N frames
|
||||
SLIDING = 'sliding' # Multiple sliding windows
|
||||
ALL = 'all' # All possible clips
|
||||
|
||||
|
||||
def _load_video_cv2(video_path: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform') -> torch.Tensor:
|
||||
"""
|
||||
Load video from video file using OpenCV.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file (MP4, AVI, MOV, MKV)
|
||||
num_frames: Number of frames to extract
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
video_path = str(video_path)
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
|
||||
if not cap.isOpened():
|
||||
raise RuntimeError(f"Cannot open video: {video_path}")
|
||||
|
||||
frames = []
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
|
||||
if num_frames is None:
|
||||
# Read all available frames
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
cap.release()
|
||||
if len(frames) == 0:
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
return frames
|
||||
|
||||
if total_frames == 0:
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
# Determine frame indices for sampling
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(
|
||||
total_frames)) + [total_frames - 1] * (num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0, total_frames - 1, num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
# Extract frames
|
||||
for idx in frame_indices:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
||||
ret, frame = cap.read()
|
||||
|
||||
if not ret:
|
||||
if len(frames) > 0:
|
||||
frames.append(frames[-1].copy())
|
||||
else:
|
||||
h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
frames.append(np.zeros((h, w, 3), dtype=np.uint8))
|
||||
continue
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
cap.release()
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _load_video_from_frames(
|
||||
frame_dir: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform',
|
||||
frame_extensions: list[str] | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Load video from directory of frame images.
|
||||
|
||||
Args:
|
||||
frame_dir: Directory containing frames
|
||||
num_frames: Number of frames to sample
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
frame_extensions: Image file extensions to look for
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
if frame_extensions is None:
|
||||
frame_extensions = ['.jpg', '.png', '.jpeg', '.bmp']
|
||||
|
||||
frame_dir = Path(frame_dir)
|
||||
|
||||
if not frame_dir.exists():
|
||||
raise FileNotFoundError(f"Frame directory not found: {frame_dir}")
|
||||
|
||||
# Find all frames
|
||||
frame_files: list[Path] = []
|
||||
for ext in frame_extensions:
|
||||
frame_files.extend(frame_dir.glob(f"*{ext}"))
|
||||
|
||||
if len(frame_files) == 0:
|
||||
raise ValueError(
|
||||
f"No frames found in {frame_dir} with extensions {frame_extensions}"
|
||||
)
|
||||
|
||||
frame_files = sorted(frame_files, key=lambda x: x.name)
|
||||
total_frames = len(frame_files)
|
||||
|
||||
# Determine frame indices
|
||||
if num_frames is None:
|
||||
frame_indices = list(range(total_frames))
|
||||
else:
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(total_frames)) + [total_frames - 1] * (
|
||||
num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0,
|
||||
total_frames - 1,
|
||||
num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
# Load frames
|
||||
frames = []
|
||||
for idx in frame_indices:
|
||||
frame_path = frame_files[idx]
|
||||
frame = cv2.imread(str(frame_path))
|
||||
|
||||
if frame is None:
|
||||
raise RuntimeError(f"Failed to load frame: {frame_path}")
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
# Stack and convert to tensor
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _detect_video_format(path: str | Path) -> str:
|
||||
"""
|
||||
Detect if path is a video file or frame directory.
|
||||
|
||||
Returns:
|
||||
'video_file', 'frame_directory', or 'unknown'
|
||||
"""
|
||||
path = Path(path)
|
||||
|
||||
if path.is_file():
|
||||
return 'video_file'
|
||||
elif path.is_dir():
|
||||
# Check if contains image files
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
|
||||
for ext in image_extensions:
|
||||
if list(path.glob(f"*{ext}")):
|
||||
return 'frame_directory'
|
||||
return 'unknown'
|
||||
else:
|
||||
raise ValueError(f"Path does not exist: {path}")
|
||||
|
||||
|
||||
def load_video_auto(video_path: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform') -> torch.Tensor:
|
||||
"""
|
||||
Automatically detect format and load video.
|
||||
|
||||
Supports:
|
||||
- Video files (MP4, AVI, MOV, MKV)
|
||||
- Frame directories (JPG, PNG)
|
||||
|
||||
Args:
|
||||
video_path: Path to video file or frame directory
|
||||
num_frames: Number of frames to extract
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
format_type = _detect_video_format(video_path)
|
||||
|
||||
if format_type == 'video_file':
|
||||
return _load_video_cv2(video_path, num_frames, sample_strategy)
|
||||
elif format_type == 'frame_directory':
|
||||
return _load_video_from_frames(video_path, num_frames, sample_strategy)
|
||||
else:
|
||||
raise ValueError(f"Unknown video format at {video_path}")
|
||||
|
||||
|
||||
def sample_clips_from_video(
|
||||
video: torch.Tensor,
|
||||
num_frames_per_clip: int = 16,
|
||||
num_clips: int = 1,
|
||||
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
|
||||
frame_stride: int = 1,
|
||||
temporal_stride: int = 1) -> list[torch.Tensor]:
|
||||
"""
|
||||
Sample clips from a video with various strategies.
|
||||
|
||||
Args:
|
||||
video: [T, C, H, W] full video
|
||||
num_frames_per_clip: Frames per clip
|
||||
num_clips: Number of clips to extract
|
||||
strategy: ClipSamplingStrategy or string ('beginning', 'random', etc.)
|
||||
frame_stride: Skip frames (FPS control: 1=all, 2=every 2nd, 8=every 8th)
|
||||
temporal_stride: Stride between clips for sliding window
|
||||
|
||||
Returns:
|
||||
List of clips, each [num_frames_per_clip, C, H, W]
|
||||
|
||||
Examples:
|
||||
>>> # Beginning clip (most common for FVD)
|
||||
>>> clips = sample_clips_from_video(video, 16, strategy='beginning')
|
||||
|
||||
>>> # Multiple random clips
|
||||
>>> clips = sample_clips_from_video(video, 16, num_clips=4, strategy='random')
|
||||
|
||||
>>> # Subsample FPS by 2x (every 2nd frame)
|
||||
>>> clips = sample_clips_from_video(video, 16, frame_stride=2)
|
||||
|
||||
>>> # Sliding window with overlap
|
||||
>>> clips = sample_clips_from_video(video, 16, strategy='sliding', temporal_stride=8)
|
||||
"""
|
||||
# Convert string to enum if needed
|
||||
if isinstance(strategy, str):
|
||||
strategy = ClipSamplingStrategy(strategy)
|
||||
|
||||
T, C, H, W = video.shape
|
||||
|
||||
# Apply frame stride (FPS subsampling)
|
||||
if frame_stride > 1:
|
||||
video = video[::frame_stride]
|
||||
T = len(video)
|
||||
|
||||
effective_clip_length = num_frames_per_clip
|
||||
|
||||
# Handle videos shorter than clip length
|
||||
if effective_clip_length > T:
|
||||
pad_length = effective_clip_length - T
|
||||
last_frame = video[-1:].repeat(pad_length, 1, 1, 1)
|
||||
video = torch.cat([video, last_frame], dim=0)
|
||||
T = len(video)
|
||||
|
||||
clips = []
|
||||
|
||||
if strategy == ClipSamplingStrategy.BEGINNING:
|
||||
# Take first clip (most common for FVD evaluation)
|
||||
clip = video[:effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.MIDDLE:
|
||||
# Take middle clip
|
||||
start = (T - effective_clip_length) // 2
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.RANDOM:
|
||||
# Sample N random clips
|
||||
for _ in range(num_clips):
|
||||
if effective_clip_length == T:
|
||||
start = 0
|
||||
else:
|
||||
start = np.random.randint(0, T - effective_clip_length + 1)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.UNIFORM:
|
||||
# Uniformly spaced clips
|
||||
if num_clips == 1:
|
||||
# Single clip from middle
|
||||
start = (T - effective_clip_length) // 2
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
else:
|
||||
# Multiple uniformly spaced clips
|
||||
step = (T - effective_clip_length) / (num_clips -
|
||||
1) if num_clips > 1 else 0
|
||||
for i in range(num_clips):
|
||||
start = int(i * step)
|
||||
start = min(start, T - effective_clip_length)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.SLIDING:
|
||||
# Sliding window with stride
|
||||
for start in range(0, T - effective_clip_length + 1, temporal_stride):
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
if len(clips) >= num_clips:
|
||||
break
|
||||
|
||||
elif strategy == ClipSamplingStrategy.ALL:
|
||||
# All possible clips (overlapping)
|
||||
for start in range(T - effective_clip_length + 1):
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {strategy}")
|
||||
|
||||
return clips
|
||||
|
||||
|
||||
def load_video_clips_streaming(directory: str | Path,
|
||||
num_frames: int = 16,
|
||||
max_videos: int | None = None,
|
||||
clip_strategy: str
|
||||
| ClipSamplingStrategy = 'beginning',
|
||||
frame_stride: int = 1,
|
||||
num_clips_per_video: int = 1,
|
||||
video_extensions: list[str] | None = None,
|
||||
support_frame_dirs: bool = True,
|
||||
target_size: tuple[int, int] | None = (224, 224),
|
||||
verbose: bool = True) -> Iterator[torch.Tensor]:
|
||||
"""
|
||||
This generator yields clips one-by-one instead of loading all videos into RAM.
|
||||
Perfect for large datasets where memory is limited.
|
||||
|
||||
Args:
|
||||
directory: Path to directory with videos
|
||||
num_frames: Frames per clip
|
||||
max_videos: Max videos to load
|
||||
clip_strategy: 'beginning', 'random', 'uniform', etc.
|
||||
frame_stride: Frame skip (1=all, 2=every 2nd, 8=every 8th)
|
||||
num_clips_per_video: Number of clips per video
|
||||
video_extensions: Video file extensions
|
||||
support_frame_dirs: Also load frame directories
|
||||
target_size: Resize clips to (H, W). If None, keep original size.
|
||||
verbose: Show progress
|
||||
|
||||
Yields:
|
||||
clip: [T, C, H, W] individual clips
|
||||
|
||||
Example:
|
||||
>>> for clip in load_video_clips_streaming('data/videos/', num_frames=16):
|
||||
>>> features = model.extract_features(clip.unsqueeze(0))
|
||||
>>> # Process one clip at a time - low memory usage!
|
||||
"""
|
||||
if video_extensions is None:
|
||||
video_extensions = ['.mp4', '.avi', '.mov', '.mkv']
|
||||
|
||||
directory = Path(directory)
|
||||
|
||||
if not directory.exists():
|
||||
raise FileNotFoundError(f"Directory not found: {directory}")
|
||||
|
||||
# Find video paths
|
||||
video_paths: list[Path] = []
|
||||
|
||||
# Find video files
|
||||
for ext in video_extensions:
|
||||
video_paths.extend(directory.glob(f"**/*{ext}"))
|
||||
|
||||
# Find frame directories if enabled
|
||||
if support_frame_dirs:
|
||||
for subdir in directory.iterdir():
|
||||
if subdir.is_dir():
|
||||
# Check if it contains frames
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
|
||||
for ext in image_extensions:
|
||||
if list(subdir.glob(f"*{ext}")):
|
||||
video_paths.append(subdir)
|
||||
break
|
||||
|
||||
if len(video_paths) == 0:
|
||||
raise ValueError(f"No videos found in {directory}")
|
||||
|
||||
video_paths = sorted(video_paths)
|
||||
|
||||
if max_videos is not None:
|
||||
video_paths = video_paths[:max_videos]
|
||||
|
||||
if verbose:
|
||||
print(f"Found {len(video_paths)} videos in {directory}")
|
||||
if num_clips_per_video > 1:
|
||||
print(f"Extracting {num_clips_per_video} clips per video...")
|
||||
if frame_stride > 1:
|
||||
print(f"Subsampling frames with stride {frame_stride}...")
|
||||
if target_size:
|
||||
print(f"Resizing clips to {target_size}...")
|
||||
|
||||
# Track statistics
|
||||
failed_count = 0
|
||||
total_clips = 0
|
||||
|
||||
iterator = tqdm(video_paths,
|
||||
desc="Loading videos") if verbose else video_paths
|
||||
|
||||
for video_path in iterator:
|
||||
try:
|
||||
# Load full video
|
||||
video = load_video_auto(video_path,
|
||||
num_frames=None,
|
||||
sample_strategy='uniform')
|
||||
|
||||
# Sample clips from video
|
||||
clips = sample_clips_from_video(video,
|
||||
num_frames_per_clip=num_frames,
|
||||
num_clips=num_clips_per_video,
|
||||
strategy=clip_strategy,
|
||||
frame_stride=frame_stride)
|
||||
|
||||
if target_size is not None:
|
||||
resized_clips = []
|
||||
for clip in clips:
|
||||
T, C, H, W = clip.shape
|
||||
if target_size != (H, W):
|
||||
# Resize to target size
|
||||
clip = clip.contiguous(
|
||||
) # Fix non-contiguous tensors first
|
||||
clip_flat = clip.view(T * C, H,
|
||||
W).unsqueeze(0) # [1, T*C, H, W]
|
||||
clip_resized = torch.nn.functional.interpolate(
|
||||
clip_flat,
|
||||
size=target_size,
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
clip = clip_resized.squeeze(0).view(
|
||||
T, C, target_size[0],
|
||||
target_size[1]) # Back to [T, C, H, W]
|
||||
resized_clips.append(clip)
|
||||
clips = resized_clips
|
||||
|
||||
# Yield clips one by one
|
||||
for clip in clips:
|
||||
yield clip
|
||||
total_clips += 1
|
||||
|
||||
# Free memory
|
||||
del video, clips
|
||||
|
||||
except Exception as e:
|
||||
failed_count += 1
|
||||
if verbose:
|
||||
print(f"\nWarning: Failed to load {video_path}: {e}")
|
||||
continue
|
||||
|
||||
# Validate
|
||||
if total_clips == 0:
|
||||
raise RuntimeError(f"Failed to load any videos from {directory}")
|
||||
|
||||
failure_rate = failed_count / len(video_paths)
|
||||
if failure_rate > 0.1: # More than 10% failed
|
||||
print(
|
||||
f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})"
|
||||
)
|
||||
|
||||
if verbose:
|
||||
print(
|
||||
f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos"
|
||||
)
|
||||
@@ -1,7 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
@@ -1,4 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
@@ -247,12 +247,11 @@ def _attn_bwd_dq(dq, q, K, V, #
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + q_blk)
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
kv_idx = tl.load(kv_ptr + blk_idx//2).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx) - (blk_idx % 2) * step_n
|
||||
block_sparse_offset = (kv_idx*2 + blk_idx%2) * step_n * stride_tok
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -1,68 +0,0 @@
|
||||
# LoRA Extraction and Merging
|
||||
|
||||
Tools for extracting and merging LoRA adapters for FastVideo models.
|
||||
|
||||
## Extract LoRA Adapter
|
||||
|
||||
```bash
|
||||
python extract_lora.py \
|
||||
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--out adapter_r32.safetensors \
|
||||
--rank 32
|
||||
```
|
||||
|
||||
**Options:**
|
||||
- `--base`: Base model (HuggingFace ID or local path)
|
||||
- `--ft`: Fine-tuned model (HuggingFace ID or local path)
|
||||
- `--out`: Output adapter file
|
||||
- `--rank`: LoRA rank (16, 32, 64, 128)
|
||||
- `--full-rank`: Extract full-rank adapter (optional)
|
||||
|
||||
## Merge Adapter
|
||||
|
||||
```bash
|
||||
python merge_lora.py \
|
||||
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
|
||||
--adapter adapter_r32.safetensors \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--output merged_model
|
||||
```
|
||||
|
||||
**Options:**
|
||||
- `--base`: Base model (HuggingFace ID or local path)
|
||||
- `--adapter`: LoRA adapter file (.safetensors)
|
||||
- `--ft`: Fine-tuned model (for configuration)
|
||||
- `--output`: Output directory
|
||||
|
||||
## Validate Quality (Optional)
|
||||
|
||||
```bash
|
||||
python lora_inference_comparison.py \
|
||||
--base merged_model \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--adapter NONE \
|
||||
--output-dir results \
|
||||
--prompt "A cat sitting on a windowsill" \
|
||||
--seed 42 \
|
||||
--height 480 \
|
||||
--width 480 \
|
||||
--num-frames 49 \
|
||||
--num-inference-steps 32 \
|
||||
--compute-ssim \
|
||||
--compute-lpips
|
||||
```
|
||||
|
||||
**Options:**
|
||||
- `--base`: Merged model or base model path
|
||||
- `--ft`: Fine-tuned model (reference)
|
||||
- `--adapter`: Path to adapter or NONE
|
||||
- `--output-dir`: Output directory
|
||||
- `--prompt`: Text prompt (default: "A cat sitting on a windowsill")
|
||||
- `--seed`: Random seed (default: 42)
|
||||
- `--height`: Video height (default: 480)
|
||||
- `--width`: Video width (default: 832)
|
||||
- `--num-frames`: Number of frames (default: 49)
|
||||
- `--num-inference-steps`: Inference steps (default: 32)
|
||||
- `--compute-ssim`: Compute SSIM metric
|
||||
- `--compute-lpips`: Compute LPIPS metric
|
||||
@@ -1,370 +0,0 @@
|
||||
"""Extract FastVideo-style LoRA adapters from a fine-tuned model by SVDing (FT - base).
|
||||
|
||||
Usage:
|
||||
python scripts/lora_extraction/extract_lora.py \
|
||||
--base <base_model> --ft <fine_tuned_model> --out adapter.safetensors --rank 16
|
||||
|
||||
example: python extract_lora.py \
|
||||
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--out fastvideo_adapter.safetensors \
|
||||
--rank 16
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
# Set distributed env BEFORE any fastvideo imports
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29500")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
# Optional safetensors support
|
||||
_HAVE_SAFETENSORS = True
|
||||
try:
|
||||
from safetensors.torch import save_file as safetensors_save # type: ignore
|
||||
except Exception:
|
||||
_HAVE_SAFETENSORS = False
|
||||
|
||||
# Configure minimal logging
|
||||
LOG = logging.getLogger("extract_lora")
|
||||
|
||||
|
||||
def configure_logging(level: str = "INFO") -> None:
|
||||
handler = logging.StreamHandler()
|
||||
fmt = "%(asctime)s %(levelname)s %(message)s"
|
||||
handler.setFormatter(logging.Formatter(fmt, datefmt="%Y-%m-%d %H:%M:%S"))
|
||||
LOG.addHandler(handler)
|
||||
LOG.setLevel(level)
|
||||
|
||||
|
||||
def get_pipeline_class_for_model(model_path: str):
|
||||
"""Return appropriate FastVideo Pipeline class for the model."""
|
||||
from fastvideo.utils import maybe_download_model_index # local import
|
||||
from fastvideo.pipelines.pipeline_registry import get_pipeline_registry, PipelineType
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
|
||||
config = maybe_download_model_index(model_path)
|
||||
pipeline_name = config.get("_class_name")
|
||||
if pipeline_name is None:
|
||||
raise ValueError(f"Model config for {model_path} missing _class_name (diffusers format expected).")
|
||||
|
||||
pipeline_registry = get_pipeline_registry(PipelineType.BASIC)
|
||||
pipeline_cls = pipeline_registry.resolve_pipeline_cls(pipeline_name, PipelineType.BASIC, WorkloadType.T2V)
|
||||
return pipeline_cls
|
||||
|
||||
|
||||
def load_transformer_state_dict_from_model(
|
||||
model_path: str,
|
||||
num_gpus: int = 1,
|
||||
dit_cpu_offload: bool = True,
|
||||
vae_cpu_offload: bool = True,
|
||||
text_encoder_cpu_offload: bool = True,
|
||||
pin_cpu_memory: bool = True,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Load pipeline and extract transformer.state_dict as CPU tensors."""
|
||||
pipeline_cls = get_pipeline_class_for_model(model_path)
|
||||
pipeline = pipeline_cls.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=num_gpus,
|
||||
inference_mode=True,
|
||||
dit_cpu_offload=dit_cpu_offload,
|
||||
vae_cpu_offload=vae_cpu_offload,
|
||||
text_encoder_cpu_offload=text_encoder_cpu_offload,
|
||||
pin_cpu_memory=pin_cpu_memory,
|
||||
)
|
||||
|
||||
# Try to locate transformer in several typical attributes
|
||||
transformer = getattr(pipeline, "transformer", None)
|
||||
if transformer is None:
|
||||
modules = getattr(pipeline, "modules", None)
|
||||
if isinstance(modules, dict):
|
||||
transformer = modules.get("transformer")
|
||||
if transformer is None:
|
||||
pipeline_attr = getattr(pipeline, "pipeline", None)
|
||||
transformer = getattr(pipeline_attr, "transformer", None) if pipeline_attr else None
|
||||
if transformer is None:
|
||||
raise RuntimeError("Transformer not found in pipeline. Expected pipeline.transformer or pipeline.modules['transformer'].")
|
||||
|
||||
state_dict = transformer.state_dict()
|
||||
|
||||
# DTensor safe handling
|
||||
try:
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
_HAS_DTENSOR = True
|
||||
except Exception:
|
||||
DTensor = None # type: ignore
|
||||
_HAS_DTENSOR = False
|
||||
|
||||
state_dict_cpu: Dict[str, torch.Tensor] = {}
|
||||
for k, v in state_dict.items():
|
||||
if _HAS_DTENSOR and isinstance(v, DTensor): # type: ignore
|
||||
state_dict_cpu[k] = v.to_local().detach().cpu().contiguous()
|
||||
else:
|
||||
state_dict_cpu[k] = v.detach().cpu().contiguous()
|
||||
|
||||
# cleanup
|
||||
try:
|
||||
del pipeline, transformer
|
||||
except Exception:
|
||||
pass
|
||||
torch.cuda.empty_cache()
|
||||
return state_dict_cpu
|
||||
|
||||
|
||||
def is_extractable_weight(key: str) -> bool:
|
||||
"""Return True if key represents a weight suitable for LoRA extraction."""
|
||||
if not key.endswith("weight"):
|
||||
return False
|
||||
low = key.lower()
|
||||
for skip in ("norm", "bias", "embedding"):
|
||||
if skip in low:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def save_adapter_state(adapter_state: Dict[str, torch.Tensor], out_path: Path) -> None:
|
||||
"""Save adapter state dict to safetensors (if available) or torch.save."""
|
||||
cleaned = {k: v.detach().cpu().contiguous() for k, v in adapter_state.items()}
|
||||
out_str = str(out_path)
|
||||
if out_path.suffix == ".safetensors" and _HAVE_SAFETENSORS:
|
||||
safetensors_save(cleaned, out_str)
|
||||
else:
|
||||
torch.save(cleaned, out_str)
|
||||
|
||||
|
||||
def build_adapter_from_states(
|
||||
base_sd: Dict[str, torch.Tensor],
|
||||
ft_sd: Dict[str, torch.Tensor],
|
||||
rank: int,
|
||||
full_rank: bool,
|
||||
min_delta: float,
|
||||
checkpoint_interval: int,
|
||||
checkpoint_path: Optional[Path],
|
||||
resume_from: int = 0,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Compute low-rank LoRA adapters by SVD on (ft - base) for extractable weights."""
|
||||
# DTensor detection
|
||||
try:
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
_HAS_DTENSOR = True
|
||||
except Exception:
|
||||
DTensor = None # type: ignore
|
||||
_HAS_DTENSOR = False
|
||||
|
||||
keys = sorted(ft_sd.keys())
|
||||
adapter_state: Dict[str, torch.Tensor] = {}
|
||||
processed = 0
|
||||
mean_deltas = []
|
||||
|
||||
for idx, key in enumerate(tqdm(keys, desc="scanning keys", unit="keys")):
|
||||
if idx < resume_from:
|
||||
continue
|
||||
if not is_extractable_weight(key):
|
||||
continue
|
||||
if key not in base_sd:
|
||||
continue
|
||||
|
||||
Wb_raw = base_sd[key]
|
||||
Wf_raw = ft_sd[key]
|
||||
|
||||
# Convert DTensor if present
|
||||
if _HAS_DTENSOR and isinstance(Wb_raw, DTensor): # type: ignore
|
||||
Wb = Wb_raw.to_local().detach().cpu().to(torch.float32).contiguous()
|
||||
else:
|
||||
Wb = Wb_raw.detach().cpu().to(torch.float32).contiguous()
|
||||
|
||||
if _HAS_DTENSOR and isinstance(Wf_raw, DTensor): # type: ignore
|
||||
Wf = Wf_raw.to_local().detach().cpu().to(torch.float32).contiguous()
|
||||
else:
|
||||
Wf = Wf_raw.detach().cpu().to(torch.float32).contiguous()
|
||||
|
||||
if Wb.shape != Wf.shape:
|
||||
continue
|
||||
|
||||
delta = (Wf - Wb).contiguous()
|
||||
mean_abs = float(delta.abs().mean().item())
|
||||
mean_deltas.append(mean_abs)
|
||||
if mean_abs < min_delta:
|
||||
continue
|
||||
|
||||
# SVD (CPU)
|
||||
try:
|
||||
U, S, Vh = torch.linalg.svd(delta, full_matrices=False)
|
||||
except RuntimeError:
|
||||
# skip layers that fail SVD
|
||||
continue
|
||||
|
||||
max_rank = S.numel()
|
||||
chosen_rank = max_rank if full_rank or rank <= 0 else min(rank, max_rank)
|
||||
if chosen_rank == 0:
|
||||
continue
|
||||
|
||||
S_sqrt = torch.sqrt(S[:chosen_rank].to(torch.float32))
|
||||
U_r = U[:, :chosen_rank].to(torch.float32) # (out, r)
|
||||
Vh_r = Vh[:chosen_rank, :].to(torch.float32) # (r, in)
|
||||
|
||||
lora_B = (U_r * S_sqrt.unsqueeze(0)).contiguous() # (out, r)
|
||||
tmp = (Vh_r.T * S_sqrt.unsqueeze(0)).contiguous() # (in, r)
|
||||
lora_A = tmp.T.contiguous() # (r, in)
|
||||
|
||||
base_name = key[:-len(".weight")]
|
||||
a_key = f"{base_name}.lora_A.weight"
|
||||
b_key = f"{base_name}.lora_B.weight"
|
||||
rank_key = f"{base_name}.lora_rank"
|
||||
alpha_key = f"{base_name}.lora_alpha"
|
||||
|
||||
adapter_state[a_key] = lora_A.cpu()
|
||||
adapter_state[b_key] = lora_B.cpu()
|
||||
adapter_state[rank_key] = torch.tensor([chosen_rank], dtype=torch.int32)
|
||||
adapter_state[alpha_key] = torch.tensor([float(chosen_rank)], dtype=torch.float32)
|
||||
|
||||
processed += 1
|
||||
|
||||
# checkpoint periodically
|
||||
if checkpoint_path and checkpoint_interval > 0 and (idx + 1) % checkpoint_interval == 0:
|
||||
try:
|
||||
torch.save({"index": idx + 1, "adapter": adapter_state}, str(checkpoint_path))
|
||||
except Exception:
|
||||
# non-fatal; continue
|
||||
pass
|
||||
|
||||
# free local large tensors
|
||||
del delta, U, S, Vh, U_r, Vh_r, tmp, lora_A, lora_B
|
||||
|
||||
# final checkpoint
|
||||
if checkpoint_path:
|
||||
try:
|
||||
torch.save({"index": len(keys), "adapter": adapter_state}, str(checkpoint_path))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
avg_delta = (sum(mean_deltas) / len(mean_deltas)) if mean_deltas else 0.0
|
||||
LOG.info("Extraction complete: processed_keys=%d, extracted_layers=%d, avg_abs_delta=%.6e",
|
||||
len(keys), processed, avg_delta)
|
||||
return adapter_state
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description="Extract FastVideo-style LoRA adapter (CPU SVD).")
|
||||
p.add_argument("--base", required=True, help="Base model id or local path")
|
||||
p.add_argument("--ft", required=True, help="Fine-tuned model id or local path")
|
||||
p.add_argument("--out", default="fastvideo_adapter.safetensors", help="Output adapter file (.safetensors or .pt)")
|
||||
p.add_argument("--rank", type=int, default=16, help="Truncated SVD rank; <=0 for full rank")
|
||||
p.add_argument("--full-rank", action="store_true", help="Use full SVD rank for every layer")
|
||||
p.add_argument("--min-delta", type=float, default=1e-8, help="Minimum mean abs delta to consider a layer changed")
|
||||
p.add_argument("--checkpoint", default="extract_lora_checkpoint.pt", help="Checkpoint path to resume/save progress")
|
||||
p.add_argument("--resume", action="store_true", help="Resume from checkpoint if available")
|
||||
p.add_argument("--log-level", default="INFO", choices=["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"])
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def extract_lora_adapter(
|
||||
base: str,
|
||||
ft: str,
|
||||
out: str,
|
||||
rank: int = 32,
|
||||
full_rank: bool = False,
|
||||
min_delta: float = 1e-6,
|
||||
checkpoint: Optional[str] = None,
|
||||
resume: bool = False,
|
||||
log_level: str = "INFO",
|
||||
) -> None:
|
||||
"""Extract LoRA adapter from fine-tuned model.
|
||||
|
||||
Args:
|
||||
base: Base model path or HuggingFace ID
|
||||
ft: Fine-tuned model path or HuggingFace ID
|
||||
out: Output adapter file path
|
||||
rank: LoRA rank (default: 32)
|
||||
full_rank: Extract full-rank adapter
|
||||
min_delta: Minimum delta for extraction
|
||||
checkpoint: Checkpoint file path
|
||||
resume: Resume from checkpoint
|
||||
log_level: Logging level
|
||||
"""
|
||||
configure_logging(log_level)
|
||||
|
||||
# ensure fastvideo import
|
||||
try:
|
||||
import fastvideo # noqa: F401
|
||||
except Exception as exc:
|
||||
LOG.error("Failed to import fastvideo: %s", exc)
|
||||
sys.exit(2)
|
||||
|
||||
out_path = Path(out)
|
||||
checkpoint_path = Path(checkpoint) if checkpoint else None
|
||||
|
||||
# load state_dicts
|
||||
LOG.info("Loading base model: %s", base)
|
||||
base_sd = load_transformer_state_dict_from_model(base)
|
||||
|
||||
LOG.info("Loading fine-tuned model: %s", ft)
|
||||
ft_sd = load_transformer_state_dict_from_model(ft)
|
||||
|
||||
resume_idx = 0
|
||||
adapter_existing: Dict[str, torch.Tensor] = {}
|
||||
if resume and checkpoint_path and checkpoint_path.exists():
|
||||
try:
|
||||
ck = torch.load(str(checkpoint_path), map_location="cpu")
|
||||
adapter_existing = ck.get("adapter", {}) or {}
|
||||
resume_idx = int(ck.get("index", 0) or 0)
|
||||
LOG.info("Resuming from checkpoint index=%d with %d existing entries", resume_idx, len(adapter_existing))
|
||||
except Exception:
|
||||
adapter_existing = {}
|
||||
|
||||
adapter_state = dict(adapter_existing) if adapter_existing else {}
|
||||
new_adapter = build_adapter_from_states(
|
||||
base_sd=base_sd,
|
||||
ft_sd=ft_sd,
|
||||
rank=rank,
|
||||
full_rank=full_rank,
|
||||
min_delta=min_delta,
|
||||
checkpoint_interval=50,
|
||||
checkpoint_path=checkpoint_path,
|
||||
resume_from=resume_idx,
|
||||
)
|
||||
adapter_state.update(new_adapter)
|
||||
|
||||
# final save
|
||||
save_adapter_state(adapter_state, out_path)
|
||||
|
||||
# cleanup checkpoint if present
|
||||
if checkpoint_path and checkpoint_path.exists():
|
||||
try:
|
||||
checkpoint_path.unlink()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
LOG.info("Saved adapter to %s (entries=%d)", str(out_path), len(adapter_state) // 4)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""CLI wrapper for extract_lora_adapter."""
|
||||
args = parse_args()
|
||||
extract_lora_adapter(
|
||||
base=args.base,
|
||||
ft=args.ft,
|
||||
out=args.out,
|
||||
rank=args.rank,
|
||||
full_rank=args.full_rank,
|
||||
min_delta=args.min_delta,
|
||||
checkpoint=args.checkpoint,
|
||||
resume=args.resume,
|
||||
log_level=args.log_level,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,358 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Compare Fine-Tuned vs Merged/Base+LoRA inference outputs.
|
||||
|
||||
Generates two videos with the same seed and computes SSIM.
|
||||
|
||||
Usage examples:
|
||||
python lora_inference_comparison.py \
|
||||
--base ./merged_model \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--adapter NONE \
|
||||
--output-dir ./inference_comparison \
|
||||
--compute-ssim \
|
||||
--seed 41
|
||||
|
||||
or
|
||||
|
||||
python lora_inference_comparison.py \
|
||||
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--adapter adapter.safetensors \
|
||||
--output-dir ./inference_comparison \
|
||||
--compute-ssim \
|
||||
--seed 41
|
||||
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional, Dict, Any
|
||||
|
||||
import logging
|
||||
|
||||
# minimal distributed env defaults (kept for compatibility)
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29500")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
|
||||
# allow running from repo root where fastvideo is located
|
||||
_FASTVIDEO_PATH = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "fastvideo_pr", "FastVideo"))
|
||||
if _FASTVIDEO_PATH not in sys.path:
|
||||
sys.path.insert(0, _FASTVIDEO_PATH)
|
||||
|
||||
logger = logging.getLogger("inference_comparison")
|
||||
|
||||
|
||||
def configure_logging(level: str = "INFO") -> None:
|
||||
handler = logging.StreamHandler()
|
||||
fmt = "%(asctime)s %(levelname)s %(message)s"
|
||||
handler.setFormatter(logging.Formatter(fmt, datefmt="%Y-%m-%d %H:%M:%S"))
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(level)
|
||||
|
||||
|
||||
def _validate_adapter(adapter: Optional[str]) -> Optional[str]:
|
||||
if not adapter:
|
||||
return None
|
||||
if adapter.upper() == "NONE":
|
||||
return None
|
||||
p = Path(adapter).expanduser()
|
||||
if not p.exists():
|
||||
raise FileNotFoundError(f"Adapter not found: {p}")
|
||||
|
||||
# Accept both files and directories (FastVideo expects directories for HF-style adapters)
|
||||
if p.is_file():
|
||||
if p.suffix != ".safetensors":
|
||||
raise ValueError(f"Adapter file must be .safetensors, got: {p.suffix}")
|
||||
if p.stat().st_size == 0:
|
||||
raise ValueError(f"Adapter file is empty: {p}")
|
||||
elif p.is_dir():
|
||||
# Check if directory contains at least one .safetensors file
|
||||
safetensors_files = list(p.glob("*.safetensors"))
|
||||
if not safetensors_files:
|
||||
raise ValueError(f"Adapter directory contains no .safetensors files: {p}")
|
||||
else:
|
||||
raise ValueError(f"Adapter must be a file or directory: {p}")
|
||||
|
||||
return str(p.resolve())
|
||||
|
||||
|
||||
def generate_with_model(
|
||||
model_path: str,
|
||||
output_dir: str,
|
||||
output_name: str,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
lora_path: Optional[str],
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
num_inference_steps: int,
|
||||
guidance_scale: float,
|
||||
flow_shift: Optional[float] = None,
|
||||
embedded_guidance_scale: Optional[float] = None,
|
||||
) -> str:
|
||||
"""Produce a video with VideoGenerator.from_pretrained; returns video path."""
|
||||
try:
|
||||
from fastvideo import VideoGenerator # lazy import
|
||||
except Exception as exc:
|
||||
raise RuntimeError(f"Failed to import fastvideo.VideoGenerator: {exc}") from exc
|
||||
|
||||
init_kwargs: Dict[str, Any] = {
|
||||
"num_gpus": 1,
|
||||
"dit_cpu_offload": True,
|
||||
"vae_cpu_offload": True,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"pin_cpu_memory": True,
|
||||
}
|
||||
if lora_path:
|
||||
init_kwargs["lora_path"] = lora_path
|
||||
init_kwargs["lora_nickname"] = "extracted"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_path, **init_kwargs)
|
||||
|
||||
gen_kwargs = {
|
||||
"height": height,
|
||||
"width": width,
|
||||
"num_frames": num_frames,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"guidance_scale": guidance_scale,
|
||||
"seed": seed,
|
||||
"output_path": output_dir,
|
||||
"output_video_name": output_name,
|
||||
"save_video": True,
|
||||
}
|
||||
if flow_shift is not None:
|
||||
gen_kwargs["flow_shift"] = flow_shift
|
||||
if embedded_guidance_scale is not None:
|
||||
gen_kwargs["embedded_guidance_scale"] = embedded_guidance_scale
|
||||
|
||||
result = generator.generate_video(prompt, **gen_kwargs)
|
||||
|
||||
# best-effort cleanup of internal executors
|
||||
try:
|
||||
if hasattr(generator, "executor") and hasattr(generator.executor, "shutdown"):
|
||||
generator.executor.shutdown()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# determine saved video path
|
||||
expected = Path(output_dir) / f"{output_name}.mp4"
|
||||
if expected.exists():
|
||||
return str(expected)
|
||||
# fallback: check result dict
|
||||
if isinstance(result, dict) and "video_path" in result:
|
||||
return str(result["video_path"])
|
||||
raise FileNotFoundError(f"Video not found at expected path: {expected}")
|
||||
|
||||
|
||||
def compute_metrics(output_dir: str, ft_video: str, other_video: str, num_inference_steps: int, prompt: str, compute_ssim: bool, compute_lpips: bool) -> dict:
|
||||
results = {}
|
||||
|
||||
if compute_ssim:
|
||||
try:
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results # type: ignore
|
||||
ssim_values = compute_video_ssim_torchvision(ft_video, other_video, use_ms_ssim=True)
|
||||
results["mean_ssim"] = float(ssim_values[0])
|
||||
write_ssim_results(output_dir, ssim_values, ft_video, other_video, num_inference_steps, prompt)
|
||||
except Exception as e:
|
||||
logger.warning(f"SSIM computation failed: {e}")
|
||||
|
||||
if compute_lpips:
|
||||
try:
|
||||
import torch
|
||||
import lpips
|
||||
import torchvision.io as tv_io
|
||||
|
||||
loss_fn = lpips.LPIPS(net='alex')
|
||||
|
||||
# Load videos
|
||||
vid1, _, _ = tv_io.read_video(ft_video, pts_unit='sec')
|
||||
vid2, _, _ = tv_io.read_video(other_video, pts_unit='sec')
|
||||
|
||||
# Normalize to [-1, 1]
|
||||
vid1 = (vid1.float() / 127.5 - 1.0).permute(0, 3, 1, 2) # (T, C, H, W)
|
||||
vid2 = (vid2.float() / 127.5 - 1.0).permute(0, 3, 1, 2)
|
||||
|
||||
lpips_scores = []
|
||||
with torch.no_grad():
|
||||
for frame1, frame2 in zip(vid1, vid2):
|
||||
score = loss_fn(frame1.unsqueeze(0), frame2.unsqueeze(0))
|
||||
lpips_scores.append(float(score.item()))
|
||||
|
||||
results["mean_lpips"] = sum(lpips_scores) / len(lpips_scores)
|
||||
|
||||
# Write LPIPS results
|
||||
import json
|
||||
lpips_file = Path(output_dir) / f"steps{num_inference_steps}_{prompt.replace(' ', '_')[:30]}_lpips.json"
|
||||
with open(lpips_file, 'w') as f:
|
||||
json.dump({"mean_lpips": results["mean_lpips"], "lpips_per_frame": lpips_scores}, f, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"LPIPS computation failed: {e}")
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description="Compare Fine-Tuned vs Merged/Base+LoRA inference outputs")
|
||||
p.add_argument("--base", required=True, help="Base model ID or merged model path")
|
||||
p.add_argument("--ft", required=True, help="Fine-tuned model ID or path (reference)")
|
||||
p.add_argument("--adapter", default="NONE", help="Path to .safetensors adapter, or NONE to use merged model")
|
||||
p.add_argument("--output-dir", default="./inference_comparison")
|
||||
p.add_argument("--prompt", default="A cat sitting on a windowsill")
|
||||
p.add_argument("--seed", type=int, default=42)
|
||||
p.add_argument("--height", type=int, default=480)
|
||||
p.add_argument("--width", type=int, default=832)
|
||||
p.add_argument("--num-frames", type=int, default=49)
|
||||
p.add_argument("--num-inference-steps", type=int, default=32)
|
||||
p.add_argument("--guidance-scale", type=float, default=6.0)
|
||||
p.add_argument("--compute-ssim", action="store_true")
|
||||
p.add_argument("--compute-lpips", action="store_true")
|
||||
p.add_argument("--flow-shift", type=float, default=None)
|
||||
p.add_argument("--embedded-guidance-scale", type=float, default=None)
|
||||
p.add_argument("--log-level", default="INFO")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def compare_inference(
|
||||
base: str,
|
||||
ft: str,
|
||||
adapter: Optional[str],
|
||||
output_dir: str,
|
||||
prompt: str = "A cat sitting on a windowsill",
|
||||
seed: int = 42,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 49,
|
||||
num_inference_steps: int = 32,
|
||||
guidance_scale: float = 5.0,
|
||||
flow_shift: Optional[float] = None,
|
||||
embedded_guidance_scale: Optional[float] = None,
|
||||
compute_ssim: bool = False,
|
||||
compute_lpips: bool = False,
|
||||
log_level: str = "INFO",
|
||||
) -> dict:
|
||||
"""Compare inference between fine-tuned model and merged/base+adapter model.
|
||||
|
||||
Args:
|
||||
base: Base or merged model ID/path
|
||||
ft: Fine-tuned model ID/path
|
||||
adapter: LoRA adapter path (or NONE for merged model)
|
||||
output_dir: Output directory for videos
|
||||
prompt: Generation prompt
|
||||
seed: Random seed
|
||||
height: Video height
|
||||
width: Video width
|
||||
num_frames: Number of frames
|
||||
num_inference_steps: Inference steps
|
||||
guidance_scale: CFG scale
|
||||
flow_shift: Flow shift
|
||||
embedded_guidance_scale: Embedded guidance scale
|
||||
compute_ssim: Compute SSIM metric
|
||||
compute_lpips: Compute LPIPS metric
|
||||
log_level: Logging level
|
||||
|
||||
Returns:
|
||||
Dictionary with metric results
|
||||
"""
|
||||
configure_logging(log_level)
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
try:
|
||||
adapter_path = _validate_adapter(adapter)
|
||||
except Exception as exc:
|
||||
logger.error("Adapter validation failed: %s", exc)
|
||||
sys.exit(2)
|
||||
|
||||
# 1) generate with fine-tuned model (reference)
|
||||
logger.info("Generating reference (fine-tuned): %s", ft)
|
||||
try:
|
||||
ft_video = generate_with_model(
|
||||
model_path=ft,
|
||||
output_dir=output_dir,
|
||||
output_name="fine_tuned",
|
||||
prompt=prompt,
|
||||
seed=seed,
|
||||
lora_path=None,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
flow_shift=flow_shift,
|
||||
embedded_guidance_scale=embedded_guidance_scale,
|
||||
)
|
||||
logger.info("Reference video saved: %s", ft_video)
|
||||
except Exception as exc:
|
||||
logger.error("Reference generation failed: %s", exc)
|
||||
sys.exit(3)
|
||||
|
||||
# 2) generate with merged model OR base + adapter
|
||||
use_merged = adapter_path is None
|
||||
mode = "merged model" if use_merged else "base+adapter"
|
||||
logger.info("Generating target (%s): %s", mode, base)
|
||||
try:
|
||||
target_video = generate_with_model(
|
||||
model_path=base,
|
||||
output_dir=output_dir,
|
||||
output_name="merged_model" if use_merged else "base_plus_lora",
|
||||
prompt=prompt,
|
||||
seed=seed,
|
||||
lora_path=None if use_merged else adapter_path,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
flow_shift=flow_shift,
|
||||
embedded_guidance_scale=embedded_guidance_scale,
|
||||
)
|
||||
logger.info("Target video saved: %s", target_video)
|
||||
except Exception as exc:
|
||||
logger.error("Target generation failed: %s", exc)
|
||||
sys.exit(4)
|
||||
|
||||
# 3) compute metrics
|
||||
results = {}
|
||||
if compute_ssim or compute_lpips:
|
||||
results = compute_metrics(output_dir, ft_video, target_video, num_inference_steps, prompt, compute_ssim, compute_lpips)
|
||||
if results.get("mean_ssim") is not None:
|
||||
logger.info("Mean SSIM: %.4f", results["mean_ssim"])
|
||||
if results.get("mean_lpips") is not None:
|
||||
logger.info("Mean LPIPS: %.4f", results["mean_lpips"])
|
||||
|
||||
logger.info("Comparison complete. Videos in: %s", output_dir)
|
||||
return results
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""CLI wrapper for compare_inference."""
|
||||
args = parse_args()
|
||||
compare_inference(
|
||||
base=args.base,
|
||||
ft=args.ft,
|
||||
adapter=args.adapter,
|
||||
output_dir=args.output_dir,
|
||||
prompt=args.prompt,
|
||||
seed=args.seed,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
flow_shift=args.flow_shift,
|
||||
embedded_guidance_scale=args.embedded_guidance_scale,
|
||||
compute_ssim=args.compute_ssim,
|
||||
compute_lpips=args.compute_lpips,
|
||||
log_level=args.log_level,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,339 +0,0 @@
|
||||
"""Merge LoRA adapter into base model weights.
|
||||
|
||||
Usage:
|
||||
python merge_lora_updated.py \
|
||||
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
|
||||
--adapter adapter.safetensors \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
|
||||
--output ./merged_model
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import shutil
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29500")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
|
||||
_FASTVIDEO_PATH = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "fastvideo_pr", "FastVideo"))
|
||||
if _FASTVIDEO_PATH not in sys.path:
|
||||
sys.path.insert(0, _FASTVIDEO_PATH)
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
from extract_lora import load_transformer_state_dict_from_model, get_pipeline_class_for_model
|
||||
from fastvideo.training.training_utils import custom_to_hf_state_dict
|
||||
from fastvideo.models.loader.utils import get_param_names_mapping, hf_to_custom_state_dict
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def configure_logging(level: str = "INFO"):
|
||||
handler = logging.StreamHandler()
|
||||
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
|
||||
LOG.addHandler(handler)
|
||||
LOG.setLevel(level)
|
||||
|
||||
|
||||
def fix_adapter_naming(adapter: dict) -> dict:
|
||||
fixed = {}
|
||||
renamed = 0
|
||||
|
||||
for key, tensor in adapter.items():
|
||||
new_key = key
|
||||
|
||||
if ".lora_A.weight" in key:
|
||||
base = key.replace(".lora_A.weight", "")
|
||||
if not base.endswith(".weight"):
|
||||
new_key = base + ".weight.lora_A.weight"
|
||||
renamed += 1
|
||||
elif ".lora_B.weight" in key:
|
||||
base = key.replace(".lora_B.weight", "")
|
||||
if not base.endswith(".weight"):
|
||||
new_key = base + ".weight.lora_B.weight"
|
||||
renamed += 1
|
||||
elif ".lora_rank" in key:
|
||||
base = key.replace(".lora_rank", "")
|
||||
if not base.endswith(".weight"):
|
||||
new_key = base + ".weight.lora_rank"
|
||||
renamed += 1
|
||||
elif ".lora_alpha" in key:
|
||||
base = key.replace(".lora_alpha", "")
|
||||
if not base.endswith(".weight"):
|
||||
new_key = base + ".weight.lora_alpha"
|
||||
renamed += 1
|
||||
|
||||
fixed[new_key] = tensor
|
||||
|
||||
if renamed > 0:
|
||||
LOG.info(f"Fixed {renamed} adapter key names")
|
||||
|
||||
return fixed
|
||||
|
||||
|
||||
def load_adapter(adapter_path: str) -> dict:
|
||||
abs_path = os.path.abspath(adapter_path)
|
||||
if not os.path.exists(abs_path):
|
||||
raise FileNotFoundError(f"Adapter file not found: {abs_path}")
|
||||
if not abs_path.endswith('.safetensors'):
|
||||
raise ValueError(f"Adapter must be .safetensors: {abs_path}")
|
||||
|
||||
LOG.info(f"Loading adapter: {abs_path}")
|
||||
adapter = load_file(abs_path)
|
||||
file_size_mb = os.path.getsize(abs_path) / (1024 * 1024)
|
||||
LOG.info(f"Loaded {len(adapter)} tensors ({file_size_mb:.1f} MB)")
|
||||
|
||||
return fix_adapter_naming(adapter)
|
||||
|
||||
|
||||
def group_adapter_keys(adapter: dict) -> dict:
|
||||
grouped = defaultdict(dict)
|
||||
|
||||
for key, tensor in adapter.items():
|
||||
if key.endswith(".lora_A.weight"):
|
||||
grouped[key.replace(".lora_A.weight", "")]["A"] = tensor
|
||||
elif key.endswith(".lora_B.weight"):
|
||||
grouped[key.replace(".lora_B.weight", "")]["B"] = tensor
|
||||
elif key.endswith(".lora_rank"):
|
||||
grouped[key.replace(".lora_rank", "")]["rank"] = tensor
|
||||
elif key.endswith(".lora_alpha"):
|
||||
grouped[key.replace(".lora_alpha", "")]["alpha"] = tensor
|
||||
|
||||
LOG.info(f"Grouped {len(grouped)} LoRA layers")
|
||||
return grouped
|
||||
|
||||
|
||||
def get_reverse_param_mapping(base_model_path: str):
|
||||
LOG.info("Loading base model for parameter mapping")
|
||||
|
||||
pipeline_cls = get_pipeline_class_for_model(base_model_path)
|
||||
pipeline = pipeline_cls.from_pretrained(
|
||||
base_model_path,
|
||||
num_gpus=1,
|
||||
inference_mode=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
transformer = None
|
||||
if hasattr(pipeline, "transformer"):
|
||||
transformer = pipeline.transformer
|
||||
elif hasattr(pipeline, "modules") and isinstance(pipeline.modules, dict):
|
||||
if "transformer" in pipeline.modules:
|
||||
transformer = pipeline.modules["transformer"]
|
||||
|
||||
if transformer is None:
|
||||
raise RuntimeError("Could not find transformer in pipeline")
|
||||
|
||||
if hasattr(transformer, "reverse_param_names_mapping"):
|
||||
reverse_mapping = transformer.reverse_param_names_mapping
|
||||
elif hasattr(transformer, "config") and hasattr(transformer.config, "arch_config"):
|
||||
arch_config = transformer.config.arch_config
|
||||
if hasattr(arch_config, "reverse_param_names_mapping"):
|
||||
reverse_mapping = arch_config.reverse_param_names_mapping
|
||||
else:
|
||||
param_mapping = arch_config.param_names_mapping
|
||||
param_names_mapping_fn = get_param_names_mapping(param_mapping)
|
||||
|
||||
from diffusers import DiffusionPipeline
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
if not os.path.exists(base_model_path) or not os.path.isdir(base_model_path):
|
||||
model_path = snapshot_download(
|
||||
repo_id=base_model_path,
|
||||
ignore_patterns=["*.onnx", "*.msgpack"]
|
||||
)
|
||||
else:
|
||||
model_path = base_model_path
|
||||
|
||||
hf_pipeline = DiffusionPipeline.from_pretrained(model_path, torch_dtype=torch.float32)
|
||||
hf_transformer = hf_pipeline.transformer
|
||||
hf_sd = hf_transformer.state_dict()
|
||||
|
||||
_, reverse_mapping = hf_to_custom_state_dict(hf_sd, param_names_mapping_fn)
|
||||
|
||||
del hf_pipeline
|
||||
del hf_transformer
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
raise RuntimeError("Could not find reverse_param_names_mapping in transformer or config")
|
||||
|
||||
del pipeline
|
||||
del transformer
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return reverse_mapping
|
||||
|
||||
|
||||
def merge_lora_into_base(base_sd: dict, adapter: dict) -> dict:
|
||||
LOG.info("Merging LoRA into base weights")
|
||||
|
||||
adapter_layers = group_adapter_keys(adapter)
|
||||
merged_sd = dict(base_sd)
|
||||
|
||||
merged_count = 0
|
||||
skipped_count = 0
|
||||
|
||||
for base_name, parts in adapter_layers.items():
|
||||
weight_key = base_name if base_name.endswith(".weight") else base_name + ".weight"
|
||||
|
||||
if weight_key not in base_sd:
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
if "A" not in parts or "B" not in parts:
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
lora_A = parts["A"].to(torch.float32)
|
||||
lora_B = parts["B"].to(torch.float32)
|
||||
base_weight = base_sd[weight_key].to(torch.float32)
|
||||
|
||||
out_dim, in_dim = base_weight.shape
|
||||
if lora_B.shape[0] != out_dim or lora_A.shape[1] != in_dim or lora_B.shape[1] != lora_A.shape[0]:
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
delta = lora_B @ lora_A
|
||||
|
||||
rank = int(parts.get("rank", torch.tensor([lora_A.shape[0]])).item()) if "rank" in parts else lora_A.shape[0]
|
||||
alpha = float(parts.get("alpha", torch.tensor([rank])).item()) if "alpha" in parts else float(rank)
|
||||
|
||||
if rank != 0 and alpha != rank:
|
||||
delta = delta * (alpha / float(rank))
|
||||
|
||||
merged_weight = base_weight + delta
|
||||
merged_sd[weight_key] = merged_weight.to(base_sd[weight_key].dtype)
|
||||
merged_count += 1
|
||||
|
||||
LOG.info(f"Merged {merged_count} layers, skipped {skipped_count}")
|
||||
return merged_sd
|
||||
|
||||
|
||||
def save_merged_model(merged_sd: dict, base_model_path: str, ft_model_path: str, output_dir: str, reverse_mapping: dict):
|
||||
output_path = Path(output_dir)
|
||||
output_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
transformer_dir = output_path / "transformer"
|
||||
transformer_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
LOG.info("Converting to HuggingFace format")
|
||||
hf_merged_sd = custom_to_hf_state_dict(merged_sd, reverse_mapping)
|
||||
LOG.info(f"Converted {len(hf_merged_sd)} parameters")
|
||||
|
||||
base_path = Path(base_model_path)
|
||||
if not base_path.exists() or not base_path.is_dir():
|
||||
from huggingface_hub import snapshot_download
|
||||
base_path = Path(snapshot_download(
|
||||
repo_id=base_model_path,
|
||||
ignore_patterns=["*.onnx", "*.msgpack"]
|
||||
))
|
||||
|
||||
LOG.info("Copying model components")
|
||||
for component in ["scheduler", "text_encoder", "tokenizer", "vae"]:
|
||||
src = base_path / component
|
||||
if src.exists():
|
||||
dst = output_path / component
|
||||
if dst.exists():
|
||||
shutil.rmtree(dst)
|
||||
shutil.copytree(src, dst)
|
||||
|
||||
LOG.info("Copying finetuned model config")
|
||||
ft_path = Path(ft_model_path)
|
||||
if not ft_path.exists() or not ft_path.is_dir():
|
||||
from huggingface_hub import snapshot_download
|
||||
ft_path = Path(snapshot_download(
|
||||
repo_id=ft_model_path,
|
||||
allow_patterns=["model_index.json"],
|
||||
ignore_patterns=["*.onnx", "*.msgpack"]
|
||||
))
|
||||
|
||||
ft_index = ft_path / "model_index.json"
|
||||
if ft_index.exists():
|
||||
shutil.copy2(ft_index, output_path / "model_index.json")
|
||||
else:
|
||||
LOG.warning("Finetuned model_index.json not found, using base")
|
||||
src_index = base_path / "model_index.json"
|
||||
if src_index.exists():
|
||||
shutil.copy2(src_index, output_path / "model_index.json")
|
||||
|
||||
weight_path = transformer_dir / "diffusion_pytorch_model.safetensors"
|
||||
LOG.info(f"Saving merged weights to {weight_path}")
|
||||
to_save_hf = {k: v.detach().cpu() for k, v in hf_merged_sd.items()}
|
||||
save_file(to_save_hf, str(weight_path))
|
||||
|
||||
config_src = base_path / "transformer" / "config.json"
|
||||
if config_src.exists():
|
||||
shutil.copy2(config_src, transformer_dir / "config.json")
|
||||
|
||||
file_size_mb = weight_path.stat().st_size / (1024 * 1024)
|
||||
LOG.info(f"Saved to {output_dir} ({file_size_mb:.0f} MB, {len(hf_merged_sd)} params)")
|
||||
|
||||
|
||||
def merge_lora(
|
||||
base: str,
|
||||
adapter: str,
|
||||
ft: str,
|
||||
output: str,
|
||||
log_level: str = "INFO",
|
||||
) -> None:
|
||||
"""Merge LoRA adapter into base model.
|
||||
|
||||
Args:
|
||||
base: Base model ID or path
|
||||
adapter: LoRA adapter .safetensors file
|
||||
ft: Finetuned model ID (for config)
|
||||
output: Output directory
|
||||
log_level: Logging level
|
||||
"""
|
||||
configure_logging(log_level)
|
||||
|
||||
LOG.info(f"Base: {base}")
|
||||
LOG.info(f"Adapter: {adapter}")
|
||||
LOG.info(f"Output: {output}")
|
||||
|
||||
reverse_mapping = get_reverse_param_mapping(base)
|
||||
|
||||
LOG.info(f"Loading base model: {base}")
|
||||
base_sd = load_transformer_state_dict_from_model(base)
|
||||
LOG.info(f"Loaded { len(base_sd)} parameters")
|
||||
|
||||
adapter_sd = load_adapter(adapter)
|
||||
merged_sd = merge_lora_into_base(base_sd, adapter_sd)
|
||||
|
||||
save_merged_model(merged_sd, base, ft, output, reverse_mapping)
|
||||
LOG.info("Merge complete")
|
||||
|
||||
|
||||
def main():
|
||||
"""CLI wrapper for merge_lora."""
|
||||
parser = argparse.ArgumentParser(description="Merge LoRA adapter into base model")
|
||||
parser.add_argument("--base", required=True, help="Base model ID or path")
|
||||
parser.add_argument("--adapter", required=True, help="LoRA adapter .safetensors file")
|
||||
parser.add_argument("--ft", required=True, help="Finetuned model ID (for config)")
|
||||
parser.add_argument("--output", required=True, help="Output directory")
|
||||
parser.add_argument("--log-level", default="INFO", help="Logging level")
|
||||
args = parser.parse_args()
|
||||
|
||||
merge_lora(
|
||||
base=args.base,
|
||||
adapter=args.adapter,
|
||||
ft=args.ft,
|
||||
output=args.output,
|
||||
log_level=args.log_level,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,141 +0,0 @@
|
||||
"""Verify merged LoRA model matches finetuned model numerically.
|
||||
|
||||
Usage:
|
||||
python verify_lora.py \
|
||||
--merged merged_model \
|
||||
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def configure_logging(level: str = "INFO"):
|
||||
handler = logging.StreamHandler()
|
||||
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
|
||||
LOG.addHandler(handler)
|
||||
LOG.setLevel(level)
|
||||
|
||||
|
||||
def load_transformer_weights(model_path: str | Path) -> dict:
|
||||
"""Load transformer weights from model directory."""
|
||||
model_path = Path(model_path)
|
||||
|
||||
if not model_path.exists() or not model_path.is_dir():
|
||||
from huggingface_hub import snapshot_download
|
||||
LOG.info(f"Downloading {model_path} from HuggingFace Hub...")
|
||||
model_path = Path(snapshot_download(
|
||||
repo_id=str(model_path),
|
||||
ignore_patterns=["*.onnx", "*.msgpack"]
|
||||
))
|
||||
|
||||
transformer_dir = model_path / "transformer"
|
||||
|
||||
if not transformer_dir.exists():
|
||||
raise FileNotFoundError(f"Transformer directory not found: {transformer_dir}")
|
||||
|
||||
weight_files = sorted(transformer_dir.glob("*.safetensors"))
|
||||
if not weight_files:
|
||||
raise FileNotFoundError(f"No safetensors files in {transformer_dir}")
|
||||
|
||||
LOG.info(f"Loading {len(weight_files)} file(s) from {transformer_dir}")
|
||||
|
||||
state_dict = {}
|
||||
for f in weight_files:
|
||||
if "custom" in f.name:
|
||||
continue
|
||||
state_dict.update(load_file(str(f)))
|
||||
|
||||
return state_dict
|
||||
|
||||
|
||||
def compare_models(merged_sd: dict, ft_sd: dict) -> dict:
|
||||
"""Compare merged and finetuned model weights."""
|
||||
|
||||
common_keys = set(merged_sd.keys()) & set(ft_sd.keys())
|
||||
merged_only = set(merged_sd.keys()) - set(ft_sd.keys())
|
||||
ft_only = set(ft_sd.keys()) - set(merged_sd.keys())
|
||||
|
||||
LOG.info(f"Common keys: {len(common_keys)}")
|
||||
if merged_only:
|
||||
LOG.warning(f"Keys only in merged: {len(merged_only)}")
|
||||
if ft_only:
|
||||
LOG.warning(f"Keys only in finetuned: {len(ft_only)}")
|
||||
|
||||
results = []
|
||||
for key in sorted(common_keys):
|
||||
merged_param = merged_sd[key]
|
||||
ft_param = ft_sd[key]
|
||||
|
||||
if merged_param.shape != ft_param.shape:
|
||||
LOG.error(f"{key}: shape mismatch {merged_param.shape} vs {ft_param.shape}")
|
||||
continue
|
||||
|
||||
diff = (merged_param.float() - ft_param.float()).abs()
|
||||
max_abs = diff.max().item()
|
||||
mean_abs = diff.mean().item()
|
||||
|
||||
merged_norm = merged_param.float().norm().item()
|
||||
rel_mean = (mean_abs / merged_norm * 100) if merged_norm > 0 else 0
|
||||
|
||||
results.append({
|
||||
"key": key,
|
||||
"shape": tuple(merged_param.shape),
|
||||
"max_abs": max_abs,
|
||||
"mean_abs": mean_abs,
|
||||
"rel_mean": rel_mean
|
||||
})
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Verify merged LoRA model matches finetuned model")
|
||||
parser.add_argument("--merged", required=True, help="Merged model directory")
|
||||
parser.add_argument("--ft", required=True, help="Finetuned model ID or path")
|
||||
parser.add_argument("--log-level", default="INFO", help="Logging level")
|
||||
args = parser.parse_args()
|
||||
|
||||
configure_logging(args.log_level)
|
||||
|
||||
LOG.info(f"Loading merged model: {args.merged}")
|
||||
merged_sd = load_transformer_weights(args.merged)
|
||||
LOG.info(f"Loaded {len(merged_sd)} parameters")
|
||||
|
||||
LOG.info(f"Loading finetuned model: {args.ft}")
|
||||
ft_sd = load_transformer_weights(args.ft)
|
||||
LOG.info(f"Loaded {len(ft_sd)} parameters")
|
||||
|
||||
LOG.info("Comparing models...")
|
||||
results = compare_models(merged_sd, ft_sd)
|
||||
|
||||
results.sort(key=lambda x: x["max_abs"], reverse=True)
|
||||
|
||||
LOG.info(f"\nTop 10 mismatches by max_abs_error:")
|
||||
for i, r in enumerate(results[:10], 1):
|
||||
LOG.info(f"{i:2d}. {r['key']}")
|
||||
LOG.info(f" shape={r['shape']}, max_abs={r['max_abs']:.3e}, mean_abs={r['mean_abs']:.3e}, rel_mean={r['rel_mean']:.4f}%")
|
||||
|
||||
overall_mean = sum(r["mean_abs"] for r in results) / len(results)
|
||||
overall_max = max(r["max_abs"] for r in results)
|
||||
|
||||
LOG.info(f"\nOverall metrics:")
|
||||
LOG.info(f" Layers compared: {len(results)}")
|
||||
LOG.info(f" Mean(mean_abs): {overall_mean:.3e}")
|
||||
LOG.info(f" Max(max_abs): {overall_max:.3e}")
|
||||
|
||||
if overall_mean < 1e-4:
|
||||
LOG.info("\nVerification PASSED: Merge is numerically accurate")
|
||||
else:
|
||||
LOG.warning(f"\nVerification WARNING: Mean error {overall_mean:.3e} > 1e-4")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user