File size: 7,352 Bytes
2847d0b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 | #!/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()
|