"""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