BonanDing's picture
Add isolated Minecraft and RE10K baseline evaluation suite
59630ba verified
Raw History Blame Contribute Delete
4.72 kB
"""Fixed WorldMem benchmark identities, RGB frames, and model-independent I/O."""
import csv
import json
import re
from pathlib import Path
import av
import numpy as np
import torch
def select_cases(manifest, scope=None, rank=0, world_size=1, limit=None):
with open(manifest, newline="") as stream:
rows = list(csv.DictReader(stream))
if scope is not None and "scope" in rows[0]:
rows = [row for row in rows if row["scope"] == scope]
if limit is not None:
rows = rows[:limit]
if not rows or not 0 <= rank < world_size:
raise ValueError("Empty benchmark selection or invalid worker shard")
return rows[rank::world_size]
def read_frames(path, indices):
indices = torch.as_tensor(indices, dtype=torch.long)
first, last = int(indices.min()), int(indices.max())
frames = []
with av.open(str(path)) as container:
for ordinal, frame in enumerate(container.decode(video=0)):
if ordinal > last:
break
if ordinal >= first:
frames.append(torch.from_numpy(frame.to_ndarray(format="rgb24")))
if len(frames) != last - first + 1:
raise ValueError(f"Missing requested frames in {path}")
video = torch.stack(frames).index_select(0, indices - first)
return video.permute(0, 3, 1, 2).float().div_(255.0).contiguous()
def _scene_file(root, tree, scene_id, suffix):
directory = Path(root) / tree
direct = directory / f"{scene_id}{suffix}"
if direct.is_file():
return direct
matches = list(directory.glob(f"*/{scene_id}{suffix}"))
if len(matches) != 1:
raise ValueError(f"Expected exactly one {tree}/{scene_id}{suffix}")
return matches[0]
def load_re10k_case(row, root):
start = int(row.get("history_start", 0))
indices = torch.cat((torch.arange(200), torch.arange(199, -1, -1))) + start
video_path = _scene_file(root, "test_256", row["scene_id"], ".mp4")
pose_path = _scene_file(root, "test_poses", row["scene_id"], ".pt")
poses = torch.load(pose_path, map_location="cpu", weights_only=True)[indices].float()
return {
"rgb": read_frames(video_path, indices),
"raw_poses": poses,
"cameras": torch.cat((poses[:, :4], poses[:, 6:]), dim=-1),
"history_frames": 100,
"history": 100,
"source_indices": indices,
"video_path": str(video_path),
"row": dict(row),
}
def minecraft_actions(actions):
# Preserve the existing WorldMem 25-channel conversion exactly.
actions = torch.as_tensor(actions)
result = torch.zeros((len(actions), 25), dtype=torch.float32)
result[actions[:, 0] == 1, 11] = 1
result[actions[:, 0] == 2, 12] = 1
result[actions[:, 4] == 11, 16] = -1
result[actions[:, 4] == 13, 16] = 1
result[actions[:, 3] == 11, 15] = -1
result[actions[:, 3] == 13, 15] = 1
result[(actions[:, 5] == 6) | (actions[:, 5] == 1), 24] = 1
result[actions[:, 1] == 1, 13] = 1
result[actions[:, 1] == 2, 14] = 1
result[actions[:, 7] == 1, 2] = 1
return result
def load_minecraft_case(row, root):
indices = torch.arange(int(row["context_start"]), int(row["target_end"]))
path = Path(root) / row["relative_path"]
with np.load(path.with_suffix(".npz")) as annotation:
actions = minecraft_actions(annotation["actions"])[indices]
poses = torch.from_numpy(annotation["poses"].copy()).float()[indices]
return {
"rgb": read_frames(path, indices),
"actions": actions,
"poses": poses,
"history_frames": int(row["context_frames"]),
"history": int(row["context_frames"]),
"source_indices": indices,
"video_path": str(path),
"row": dict(row),
}
def case_key(row):
if "scene_id" in row:
return row["scene_id"]
clean = re.sub(r"[^A-Za-z0-9._-]+", "__", row["clip_id"]).strip("_")
return f"g{int(row['global_test_index']):06d}_{clean}"
def case_seed(row, seed):
if "scene_id" in row:
# The existing dense reverse-loop evaluator uses path ordinal 7.
return int(seed) * 1_000_000 + 1_000_003 * int(row["metadata_index"]) + 7
return int(seed) * 1_000_000 + int(row["global_test_index"])
def save_case(output, row, prediction, metadata):
"""Stage exact float predictions; the scorer removes them after video/metric export."""
directory = Path(output) / "cases" / case_key(row)
directory.mkdir(parents=True, exist_ok=True)
torch.save(prediction.detach().float().cpu().contiguous(), directory / "prediction.pt")
(directory / "case.json").write_text(
json.dumps({"case": dict(row), **metadata}, indent=2, default=str) + "\n"
)
return directory