DeMemWM / scripts /benchmark_dememwm_multiview_selection.py
BonanDing's picture
Add DeMemWM multiview selector benchmark
56f5ccc
Raw
History Blame Contribute Delete
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())