BonanDing's picture
Add isolated Minecraft and RE10K baseline evaluation suite
59630ba verified
Raw History Blame Contribute Delete
7.37 kB
"""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()