BonanDing's picture
Add isolated Minecraft and RE10K baseline evaluation suite
59630ba verified
Raw History Blame Contribute Delete
2.38 kB
"""Check every requested case and fixed Minecraft reference cache before full evaluation."""
import argparse
import json
from pathlib import Path
import sys
import numpy as np
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "shared"))
from protocol import _scene_file, select_cases
parser = argparse.ArgumentParser()
parser.add_argument("--minecraft-root", type=Path, required=True)
parser.add_argument("--re10k-root", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
manifests = Path(__file__).resolve().parent / "manifests"
errors = []
mc = select_cases(manifests / "minecraft_test_300.csv")
for row in mc:
relative = Path(row["relative_path"])
video = args.minecraft_root / relative
feature = args.minecraft_root / "vae_features" / relative.parent / f"{relative.stem}_vae_feature.npy"
required = [video, video.with_suffix(".npz"), feature]
missing = [str(path) for path in required if not path.is_file()]
if missing:
errors.extend(missing)
continue
with np.load(video.with_suffix(".npz")) as data:
if min(len(data["actions"]), len(data["poses"])) < int(row["target_end"]):
errors.append(f"Short annotation: {video}")
z = np.load(feature, mmap_mode="r")
if z.dtype != np.float16 or z.shape[1:] != (16, 18, 32) or len(z) < int(row["target_end"]):
errors.append(f"Incompatible fixed reference cache: {feature}")
re = select_cases(manifests / "re10k_real_trajectory_h100_f399_scopes_v1.csv", "final100")
for row in re:
try:
_scene_file(args.re10k_root, "test_256", row["scene_id"], ".mp4")
path = _scene_file(args.re10k_root, "test_poses", row["scene_id"], ".pt")
pose = torch.load(path, map_location="cpu", weights_only=True)
if pose.ndim != 2 or pose.shape[1] != 18 or len(pose) < int(row["history_start"]) + 200 or not torch.isfinite(pose).all():
errors.append(f"Invalid camera annotation: {path}")
except ValueError as error:
errors.append(str(error))
result = {"minecraft_cases": len(mc), "re10k_cases": len(re), "errors": errors}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(result, indent=2) + "\n")
print(json.dumps(result, indent=2))
sys.exit(1 if errors or len(mc) != 300 or len(re) != 100 else 0)