Download shared/aggregate.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 7.37 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/shared/aggregate.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/shared/aggregate.py
-
curl -L -o aggregate.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/shared/aggregate.py
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() | |