"""Validate complete benchmark coverage and pool saved frame metrics/FID features.""" import argparse import json from pathlib import Path import sys sys.path.insert(0, str(Path(__file__).resolve().parent / "metric_vendor")) import numpy as np import torch import torchmetrics from torchmetrics.image.fid import _compute_fid from protocol import case_key, case_seed, select_cases from score_case import metric_ranges def feature_statistics(case_dirs, name, start, end, device): count, total, product = 0, None, None for directory in case_dirs: feature = torch.from_numpy(np.load(directory / f"{name}_features.npy")[start:end]).to(device, torch.float64) if total is None: total = torch.zeros(feature.shape[1], dtype=torch.float64, device=device) product = torch.zeros((feature.shape[1], feature.shape[1]), dtype=torch.float64, device=device) total += feature.sum(0) product += feature.T @ feature count += len(feature) mean = total / count covariance = (product - total[:, None] @ total[None, :] / count) / (count - 1) return mean, covariance, count def pooled_fid(case_dirs, start, end, device): real_mean, real_cov, count = feature_statistics(case_dirs, "real", start, end, device) fake_mean, fake_cov, _ = feature_statistics(case_dirs, "fake", start, end, device) # This is the same legacy computation called by TorchMetrics0.11.4.compute, # including unbiased covariance and final float32 output conversion. value = _compute_fid(real_mean, real_cov, fake_mean, fake_cov).float() return float(value.cpu()), count def validate_coverage(output, rows, dataset, seed, frames): directories = [Path(output) / "cases" / case_key(row) for row in rows] actual = {path.name for path in (Path(output) / "cases").iterdir() if path.is_dir()} expected = {directory.name for directory in directories} errors, summaries = [], [] if len(expected) != len(rows): errors.append("Duplicate case identities in the selected manifest") if actual != expected: errors.append(f"Case coverage mismatch: missing={sorted(expected - actual)}, extra={sorted(actual - expected)}") if torchmetrics.__version__ != "0.11.4": errors.append(f"Expected torchmetrics0.11.4, found {torchmetrics.__version__}") for row, directory in zip(rows, directories): required = [directory / name for name in ("metrics.json", "frame_metrics.npz", "real_features.npy", "fake_features.npy", "prediction.mp4", "reference.mp4")] missing = [str(path) for path in required if not path.is_file()] if missing: errors.append(f"Incomplete case {directory.name}: {missing}") continue summary = json.loads((directory / "metrics.json").read_text()) summaries.append(summary) wanted = (dict(row), dataset, case_seed(row, seed), frames) found = (summary["case"], summary["dataset"], summary["sample_seed"], summary["frames"]) if found != wanted: errors.append(f"Case identity/dataset/seed/horizon mismatch: {directory.name}") with np.load(directory / "frame_metrics.npz") as metrics: names = {"mse", "psnr", "lpips", "ssim"} if dataset == "re10k" else {"mse", "psnr", "lpips"} if set(metrics.files) != names or any(metrics[name].shape != (frames,) or not np.isfinite(metrics[name]).all() for name in metrics.files): errors.append(f"Invalid per-frame metric population: {directory.name}") for name in ("real", "fake"): feature = np.load(directory / f"{name}_features.npy", mmap_mode="r") if feature.shape != (frames, 2048) or not np.isfinite(feature).all(): errors.append(f"Invalid {name} FID feature population: {directory.name}") contracts = [{key: summary[key] for key in ("metric_resolution", "torchmetrics", "lpips", "psnr_epsilon", "fid", "video")} for summary in summaries] if contracts and any(contract != contracts[0] for contract in contracts): errors.append("Metric contracts differ between cases") if errors: raise ValueError("\n".join(errors)) return directories, contracts[0] def aggregate(output, manifest, dataset, seed=0, scope=None, fid_manifest=None, expected_frames=None, limit=None, device="cuda"): frames = expected_frames or (300 if dataset == "re10k" else 500) rows = select_cases(manifest, scope, limit=limit) directories, contract = validate_coverage(output, rows, dataset, seed, frames) ranges = metric_ranges(dataset, frames) report = {"dataset": dataset, "seed": seed, "manifest": str(manifest), "scope": scope, "cases": len(rows), "generated_frames_per_case": frames, "complete_full_horizon": frames == (300 if dataset == "re10k" else 500), "limited_selection": limit is not None, "metric_contract": contract, "frame_metrics": {}, "fid": {}} all_values = {} for directory in directories: with np.load(directory / "frame_metrics.npz") as values: for name in values.files: all_values.setdefault(name, []).append(values[name]) all_values = {name: np.stack(parts).astype(np.float64) for name, parts in all_values.items()} for name, start, end in ranges: report["frame_metrics"][name] = {metric: float(value[:, start:end].mean()) for metric, value in all_values.items()} populations = {"all_cases": directories} if fid_manifest is not None: subset = select_cases(fid_manifest) requested = {case_key(row) for row in subset} available = {directory.name for directory in directories} if not requested <= available: raise ValueError(f"FID manifest cases absent from evaluated population: {sorted(requested - available)}") populations["official_fid50"] = [Path(output) / "cases" / case_key(row) for row in subset] for population, selected in populations.items(): report["fid"][population] = {} for name, start, end in ranges: value, count = pooled_fid(selected, start, end, torch.device(device)) report["fid"][population][name] = {"fid": value, "frames": count, "cases": len(selected)} Path(output, "aggregate.json").write_text(json.dumps(report, indent=2) + "\n") return report def main(): parser = argparse.ArgumentParser() parser.add_argument("--output", type=Path, required=True) parser.add_argument("--manifest", type=Path, required=True) parser.add_argument("--dataset", choices=["minecraft", "re10k"], required=True) parser.add_argument("--scope") parser.add_argument("--seed", type=int, default=0) parser.add_argument("--fid-manifest", type=Path) parser.add_argument("--expected-frames", type=int) parser.add_argument("--limit", type=int) parser.add_argument("--device", default="cuda") args = parser.parse_args() result = aggregate(args.output, args.manifest, args.dataset, args.seed, args.scope, args.fid_manifest, args.expected_frames, args.limit, args.device) print(json.dumps({"dataset": result["dataset"], "cases": result["cases"], "overall": result["frame_metrics"]["overall"]}), flush=True) if __name__ == "__main__": main()