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