Download scripts/benchmark_dememwm_multiview_selection.py from BonanDing/DeMemWM: direct link, hf CLI and curl.
- Browser
- Download file 9.36 kB
-
https://huggingface.co/BonanDing/DeMemWM/resolve/main/scripts/benchmark_dememwm_multiview_selection.py
- Command line
-
hf download hf://BonanDing/DeMemWM/scripts/benchmark_dememwm_multiview_selection.py
-
curl -L -o benchmark_dememwm_multiview_selection.py https://huggingface.co/BonanDing/DeMemWM/resolve/main/scripts/benchmark_dememwm_multiview_selection.py
9.36 kB
| """Benchmark DeMemWM dynamic multiview memory selectors on synthetic poses.""" | |
| from __future__ import annotations | |
| import argparse | |
| import importlib.util | |
| import statistics | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| REPO_ROOT = Path(__file__).resolve().parents[1] | |
| if str(REPO_ROOT) not in sys.path: | |
| sys.path.insert(0, str(REPO_ROOT)) | |
| SELECTORS = ("fov_greedy", "pose_plucker_fps") | |
| def _load_memory_selection_module(): | |
| module_path = REPO_ROOT / "datasets" / "video" / "memory_selection.py" | |
| spec = importlib.util.spec_from_file_location("dememwm_memory_selection", module_path) | |
| if spec is None or spec.loader is None: | |
| raise ImportError(f"could not load memory selection module from {module_path}") | |
| module = importlib.util.module_from_spec(spec) | |
| sys.modules[spec.name] = module | |
| spec.loader.exec_module(module) | |
| return module | |
| memory_selection = _load_memory_selection_module() | |
| def _parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--num-frames", type=int, required=True) | |
| parser.add_argument("--target-start", type=int, required=True) | |
| parser.add_argument("--target-len", type=int, required=True) | |
| parser.add_argument("--num-iters", type=int, required=True) | |
| parser.add_argument("--pose-preselect-topk", type=int, required=True) | |
| parser.add_argument("--candidate-chunk-size", type=int, required=True) | |
| parser.add_argument("--selectors", nargs="+", default=list(SELECTORS)) | |
| parser.add_argument("--write-report", type=Path, default=None) | |
| return parser.parse_args() | |
| def _synthetic_poses(num_frames: int) -> np.ndarray: | |
| frame = np.arange(num_frames, dtype=np.float32) | |
| poses = np.zeros((num_frames, 5), dtype=np.float32) | |
| poses[:, 0] = 0.03 * frame + 24.0 * np.sin(frame * 0.031) | |
| poses[:, 1] = 4.0 * np.cos(frame * 0.019) | |
| poses[:, 2] = 0.015 * frame + 24.0 * np.cos(frame * 0.027) | |
| poses[:, 3] = 18.0 * np.sin(frame * 0.017) | |
| poses[:, 4] = np.remainder(2.7 * frame + 30.0 * np.sin(frame * 0.011) + 180.0, 360.0) - 180.0 | |
| return poses | |
| def _target_positions(target_start: int, target_len: int, num_frames: int) -> np.ndarray: | |
| stop = target_start + target_len | |
| if target_start < 0 or target_len <= 0 or stop > num_frames: | |
| raise ValueError( | |
| f"target window [{target_start}, {stop}) must be non-empty and inside num_frames={num_frames}" | |
| ) | |
| return np.arange(target_start, stop, dtype=np.int64) | |
| def _selection_cfg(selector: str, args: argparse.Namespace) -> dict: | |
| return { | |
| "enabled": True, | |
| "causal": True, | |
| "max_anchor_frames": 0, | |
| "max_dynamic_frames": args.target_len, | |
| "max_revisit_frames": 0, | |
| "pose_similarity_threshold": 0.0, | |
| "training_use_plucker": True, | |
| "training_plucker_weight": 1.0, | |
| "fov_overlap_threshold": 0.6, | |
| "min_total_selected_coverage": 0.1, | |
| "local_context_exclusion_frames": 8, | |
| "plucker_moment_radius": 30.0, | |
| "anchor_diverse_selection": True, | |
| "pose_preselect_topk": args.pose_preselect_topk, | |
| "candidate_chunk_size": args.candidate_chunk_size, | |
| "dynamic": { | |
| "selection_policy": "multiview", | |
| "multiview_selector": selector, | |
| }, | |
| } | |
| def _percentile(values: list[float], fraction: float) -> float: | |
| if not values: | |
| return 0.0 | |
| ordered = sorted(values) | |
| index = min(len(ordered) - 1, int(np.ceil(fraction * len(ordered))) - 1) | |
| return ordered[index] | |
| def _base_candidates(poses: np.ndarray, target_positions: np.ndarray, cfg: dict) -> np.ndarray: | |
| return memory_selection._memory_candidate_frames( | |
| len(poses), | |
| target_positions, | |
| cfg, | |
| "training", | |
| min_candidate_frame=0, | |
| ) | |
| def _fov_candidate_count(poses: np.ndarray, target_positions: np.ndarray, cfg: dict) -> int: | |
| candidates = _base_candidates(poses, target_positions, cfg) | |
| poses_t = torch.as_tensor(poses, dtype=torch.float32) | |
| preselected = memory_selection._pose_preselect(candidates, poses_t, target_positions, cfg) | |
| return int(len(preselected)) | |
| def _pose_plucker_candidate_count(poses: np.ndarray, target_positions: np.ndarray, cfg: dict) -> int: | |
| candidates = _base_candidates(poses, target_positions, cfg) | |
| ranked_ids, _, _ = memory_selection._rank_pose_plucker_candidates(poses, candidates, target_positions, cfg) | |
| topk = memory_selection.cfg_get(cfg, "pose_preselect_topk", 64) | |
| if topk is not None and int(topk) > 0: | |
| return int(min(int(topk), ranked_ids.numel())) | |
| return int(ranked_ids.numel()) | |
| def _run_once(selector: str, poses: np.ndarray, target_positions: np.ndarray, cfg: dict, count: int) -> np.ndarray: | |
| if selector == "fov_greedy": | |
| candidates = _base_candidates(poses, target_positions, cfg) | |
| pool = memory_selection._build_fov_candidate_pool( | |
| poses, | |
| candidates, | |
| target_positions, | |
| cfg, | |
| use_plucker=True, | |
| ) | |
| return memory_selection._select_dynamic_multiview( | |
| poses, | |
| target_positions, | |
| cfg, | |
| count, | |
| split="training", | |
| fov_pool=pool, | |
| ) | |
| if selector == "pose_plucker_fps": | |
| return memory_selection._select_dynamic_multiview( | |
| poses, | |
| target_positions, | |
| cfg, | |
| count, | |
| split="training", | |
| ) | |
| raise ValueError(f"unknown selector {selector!r}") | |
| def _benchmark_selector(selector: str, poses: np.ndarray, target_positions: np.ndarray, args: argparse.Namespace) -> dict: | |
| cfg = _selection_cfg(selector, args) | |
| count = int(args.target_len) | |
| if selector == "fov_greedy": | |
| candidate_count = _fov_candidate_count(poses, target_positions, cfg) | |
| fov_pool_reuse = True | |
| else: | |
| candidate_count = _pose_plucker_candidate_count(poses, target_positions, cfg) | |
| fov_pool_reuse = False | |
| selected = _run_once(selector, poses, target_positions, cfg, count) | |
| timings_ms = [] | |
| for _ in range(args.num_iters): | |
| start = time.perf_counter() | |
| selected = _run_once(selector, poses, target_positions, cfg, count) | |
| timings_ms.append((time.perf_counter() - start) * 1000.0) | |
| return { | |
| "selector": selector, | |
| "mean_ms": statistics.fmean(timings_ms), | |
| "median_ms": statistics.median(timings_ms), | |
| "p90_ms": _percentile(timings_ms, 0.90), | |
| "selected_count": int(len(selected)), | |
| "candidate_count_after_pose_preselection": candidate_count, | |
| "fov_pool_reuse": fov_pool_reuse, | |
| "device": str(torch.device("cpu")), | |
| } | |
| def _format_results(results: list[dict]) -> str: | |
| lines = [ | |
| "selector mean_ms median_ms p90_ms selected_count candidate_count_after_pose_preselection fov_pool_reuse device" | |
| ] | |
| for row in results: | |
| lines.append( | |
| "{selector} {mean_ms:.3f} {median_ms:.3f} {p90_ms:.3f} {selected_count} " | |
| "{candidate_count_after_pose_preselection} {fov_pool_reuse} {device}".format(**row) | |
| ) | |
| return "\n".join(lines) | |
| def _write_report(path: Path, args: argparse.Namespace, results: list[dict]) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| lines = [ | |
| "# DeMemWM Multiview Selection Speed Report", | |
| "", | |
| "This benchmark used deterministic synthetic poses only. It is not a substitute for a real dataset sampling benchmark.", | |
| "", | |
| "```text", | |
| "python " + " ".join(sys.argv), | |
| "```", | |
| "", | |
| "| selector | mean ms | median ms | p90 ms | selected frames | pose-preselected candidates | FOV pool reuse | device |", | |
| "| --- | ---: | ---: | ---: | ---: | ---: | --- | --- |", | |
| ] | |
| for row in results: | |
| lines.append( | |
| "| {selector} | {mean_ms:.3f} | {median_ms:.3f} | {p90_ms:.3f} | {selected_count} | " | |
| "{candidate_count_after_pose_preselection} | {fov_pool_reuse} | {device} |".format(**row) | |
| ) | |
| lines.extend( | |
| [ | |
| "", | |
| f"Synthetic frames: {args.num_frames}", | |
| f"Target window: [{args.target_start}, {args.target_start + args.target_len})", | |
| f"Iterations: {args.num_iters}", | |
| f"pose_preselect_topk: {args.pose_preselect_topk}", | |
| f"candidate_chunk_size: {args.candidate_chunk_size}", | |
| ] | |
| ) | |
| path.write_text("\n".join(lines) + "\n", encoding="utf-8") | |
| def main() -> int: | |
| args = _parse_args() | |
| unknown = [selector for selector in args.selectors if selector not in SELECTORS] | |
| if unknown: | |
| valid = ", ".join(SELECTORS) | |
| print(f"unknown selector(s): {', '.join(unknown)}; valid selectors: {valid}", file=sys.stderr) | |
| return 2 | |
| poses = _synthetic_poses(args.num_frames) | |
| target_positions = _target_positions(args.target_start, args.target_len, args.num_frames) | |
| results = [_benchmark_selector(selector, poses, target_positions, args) for selector in args.selectors] | |
| print(_format_results(results)) | |
| if args.write_report is not None: | |
| _write_report(args.write_report, args, results) | |
| print(f"wrote report: {args.write_report}") | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |