Self-Forcing-part-2 / scripts /compare_video_metrics.py
Cccccz's picture
Add files using upload-large-folder tool
2847d0b verified
Raw
History Blame Contribute Delete
7.35 kB
#!/usr/bin/env python3
import argparse
import json
import os
import time
from pathlib import Path
import lpips
import numpy as np
import torch
import torch.distributed as dist
from decord import VideoReader, cpu
from skimage.metrics import structural_similarity
def parse_args():
parser = argparse.ArgumentParser(
description="Compare paired videos frame by frame with PSNR, SSIM, and LPIPS."
)
parser.add_argument("--reference_dir", type=Path, required=True)
parser.add_argument("--comparison_dir", type=Path, required=True)
parser.add_argument("--output_json", type=Path, required=True)
parser.add_argument("--batch_size", type=int, default=8)
return parser.parse_args()
def init_distributed():
if "LOCAL_RANK" not in os.environ:
return 0, 0, 1, torch.device("cuda")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
return dist.get_rank(), local_rank, dist.get_world_size(), torch.device(
f"cuda:{local_rank}"
)
def list_video_pairs(reference_dir, comparison_dir):
reference = {path.name: path for path in reference_dir.glob("*.mp4")}
comparison = {path.name: path for path in comparison_dir.glob("*.mp4")}
if reference.keys() != comparison.keys():
missing = sorted(reference.keys() - comparison.keys())
extra = sorted(comparison.keys() - reference.keys())
raise ValueError(f"Video sets do not match: missing={missing}, extra={extra}")
if not reference:
raise ValueError("No paired MP4 videos found")
return [(reference[name], comparison[name]) for name in sorted(reference)]
@torch.inference_mode()
def compare_video(reference_path, comparison_path, lpips_model, device, batch_size):
reference_video = VideoReader(str(reference_path), ctx=cpu(0))
comparison_video = VideoReader(str(comparison_path), ctx=cpu(0))
if len(reference_video) != len(comparison_video):
raise ValueError(
f"Frame-count mismatch for {reference_path.name}: "
f"{len(reference_video)} vs. {len(comparison_video)}"
)
psnr_values = []
ssim_values = []
lpips_values = []
for start in range(0, len(reference_video), batch_size):
end = min(start + batch_size, len(reference_video))
indices = list(range(start, end))
reference_frames = reference_video.get_batch(indices).asnumpy()
comparison_frames = comparison_video.get_batch(indices).asnumpy()
if reference_frames.shape != comparison_frames.shape:
raise ValueError(
f"Frame-shape mismatch for {reference_path.name}: "
f"{reference_frames.shape} vs. {comparison_frames.shape}"
)
reference_tensor = (
torch.from_numpy(np.ascontiguousarray(reference_frames))
.permute(0, 3, 1, 2)
.to(device=device, dtype=torch.float32)
.div_(255.0)
)
comparison_tensor = (
torch.from_numpy(np.ascontiguousarray(comparison_frames))
.permute(0, 3, 1, 2)
.to(device=device, dtype=torch.float32)
.div_(255.0)
)
mse = torch.mean(
(reference_tensor - comparison_tensor) ** 2, dim=(1, 2, 3)
)
batch_psnr = -10.0 * torch.log10(mse)
if not torch.isfinite(batch_psnr).all():
raise ValueError(f"Non-finite PSNR encountered in {reference_path.name}")
psnr_values.extend(batch_psnr.cpu().tolist())
batch_lpips = lpips_model(
reference_tensor.mul(2.0).sub(1.0),
comparison_tensor.mul(2.0).sub(1.0),
)
lpips_values.extend(batch_lpips.flatten().cpu().tolist())
for reference_frame, comparison_frame in zip(
reference_frames, comparison_frames
):
ssim_values.append(
float(
structural_similarity(
reference_frame,
comparison_frame,
channel_axis=2,
data_range=255,
)
)
)
return {
"video": reference_path.name,
"frames": len(reference_video),
"psnr": float(np.mean(psnr_values)),
"ssim": float(np.mean(ssim_values)),
"lpips_alex": float(np.mean(lpips_values)),
}
def main():
args = parse_args()
if args.batch_size <= 0:
raise ValueError("--batch_size must be positive")
rank, local_rank, world_size, device = init_distributed()
started_at = time.perf_counter()
pairs = list_video_pairs(args.reference_dir, args.comparison_dir)
local_pairs = pairs[rank::world_size]
lpips_model = lpips.LPIPS(net="alex").eval().to(device)
local_results = []
for index, (reference_path, comparison_path) in enumerate(local_pairs, start=1):
local_results.append(
compare_video(
reference_path,
comparison_path,
lpips_model,
device,
args.batch_size,
)
)
print(
f"[rank {rank}] {index}/{len(local_pairs)} {reference_path.name}",
flush=True,
)
if dist.is_initialized():
gathered = [None] * world_size if rank == 0 else None
dist.gather_object(local_results, gathered, dst=0)
if rank == 0:
all_results = [
result for rank_results in gathered for result in rank_results
]
else:
all_results = None
else:
all_results = local_results
if rank == 0:
all_results.sort(key=lambda item: item["video"])
frame_count = sum(item["frames"] for item in all_results)
weighted_sums = {
metric: sum(item[metric] * item["frames"] for item in all_results)
for metric in ("psnr", "ssim", "lpips_alex")
}
aggregate = {
metric: weighted_sums[metric] / frame_count for metric in weighted_sums
}
payload = {
"reference_dir": str(args.reference_dir.resolve()),
"comparison_dir": str(args.comparison_dir.resolve()),
"video_count": len(all_results),
"frame_count": frame_count,
"aggregation": "arithmetic mean over paired frames",
"method": {
"psnr": "RGB, data_range=1, per-frame",
"ssim": (
"skimage.metrics.structural_similarity, RGB channel_axis=2, "
"data_range=255, default window"
),
"lpips": "lpips 0.1.4, AlexNet, RGB in [-1, 1], full resolution",
},
"aggregate": aggregate,
"elapsed_seconds": time.perf_counter() - started_at,
"per_video": all_results,
}
args.output_json.parent.mkdir(parents=True, exist_ok=True)
args.output_json.write_text(json.dumps(payload, indent=2) + "\n")
print(json.dumps(payload["aggregate"], indent=2), flush=True)
print(f"Saved metrics to {args.output_json}", flush=True)
if dist.is_initialized():
dist.barrier(device_ids=[local_rank])
dist.destroy_process_group()
if __name__ == "__main__":
main()