"""Run P4D-HSSD navigation episodes with belief or closed-loop Agent control.""" from __future__ import annotations import argparse import json from pathlib import Path import numpy as np from evolvingnav_paper.backend import HabitatInspectionBackend from evolvingnav_paper.agent import Agent, AgentConfig from evolvingnav_paper.calibration import DetectionCalibrator from evolvingnav_paper.controller import LunaToolController from evolvingnav_paper.evaluate import evaluate_search from evolvingnav_paper.memory import VersionedMemory from evolvingnav_paper.perception import GroundedSAMInspector from evolvingnav_paper.policy import load_belief, model_input_batch, pack_public_query, predict_public from evolvingnav_paper.transition import IdentityTransition from evolvingnav_paper.transition_model import NeuralTransition, TransitionHead from evolvingnav_paper.world import HabitatAgentWorld CODE_ROOT = Path(__file__).resolve().parents[1] def rows(path: Path): with path.open(encoding="utf-8") as handle: for line in handle: if line.strip(): yield json.loads(line) def arguments(argv: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--task", choices=("n1", "n2", "n3", "n4"), default="n3") parser.add_argument("--agent", action="store_true", help="Run the event-driven Agent for N1/N2 as well") parser.add_argument("--controller", choices=("utility", "luna"), default="utility") parser.add_argument("--limit", type=int, default=5) parser.add_argument("--world", choices=("routine", "random", "static"), default="routine") parser.add_argument("--dataset", type=Path, required=True) parser.add_argument("--tasks", type=Path, required=True) parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--transition-checkpoint", type=Path) parser.add_argument("--inspection", choices=("semantic-oracle", "grounded-sam"), default="grounded-sam") parser.add_argument("--hssd-root", type=Path, required=True) parser.add_argument("--navmesh-root", type=Path, required=True) parser.add_argument("--grounding-dino-model", default="IDEA-Research/grounding-dino-tiny") parser.add_argument("--sam2-model", default="facebook/sam2.1-hiera-tiny") parser.add_argument("--perception-config", type=Path, default=CODE_ROOT / "configs/perception.yaml") parser.add_argument("--calibration", type=Path) parser.add_argument("--output", type=Path, required=True) args = parser.parse_args(argv) if args.limit < 1: parser.error("--limit must be positive") if args.task == "n4" and args.transition_checkpoint is None: parser.error("N4 requires --transition-checkpoint") return args def main() -> int: args = arguments() if args.output.exists(): raise FileExistsError(f"output already exists: {args.output}") episodes = [] for episode in rows(args.tasks / f"public/episodes_{args.task}.jsonl"): if episode["world_variant"] == args.world: episodes.append(episode) if len(episodes) == args.limit: break if len(episodes) != args.limit: raise ValueError(f"only found {len(episodes)} matching episodes") wanted_queries = {episode["query_id"] for episode in episodes} queries = {row["query_id"]: row for row in rows(args.tasks / "public/query_inputs.jsonl") if row["query_id"] in wanted_queries} active_checkpoint = args.transition_checkpoint if args.task == "n4" else args.checkpoint model, schema = load_belief(active_checkpoint, args.dataset) transition_head = None if args.task == "n4": import torch checkpoint = torch.load(active_checkpoint, map_location="cpu", weights_only=True) transition_head = TransitionHead(checkpoint["model_config"]["hidden_dim"]) transition_head.load_state_dict(checkpoint["transition_head"]) transition_head.eval() with np.load(args.dataset / "records/packed/train.npz", allow_pickle=False) as train: features = {key: train[key] for key in ( "candidate_region_category_id", "candidate_receptacle_category_id", "candidate_center_xyz", "candidate_is_unknown", )} catalog = json.loads((args.tasks / "catalogs/candidate_states_navigation.json").read_text()) public_viewpoints = { int(row["state_id"]): row["navigation_viewpoint"] for row in catalog["states"] if row.get("navigation_eligible") } public_goals = {state: row["position_xyz"] for state, row in public_viewpoints.items()} all_centers = { int(row["state_id"]): row["state_center"] for row in catalog["states"] } state_centers = { int(row["state_id"]): row["state_center"] for row in catalog["states"] if row.get("navigation_eligible") } surface_points = { int(row["state_id"]): [slot["point"] for slot in row.get("sampled_place_points", [])] for row in rows(args.tasks / "catalogs/receptacles.jsonl") } objects = { row["instance_uuid"]: row for row in rows(args.tasks / "catalogs/object_instances.jsonl") } decisions = [] query_batches = [] for episode in episodes: query = queries[episode["query_id"]] packed = pack_public_query(query, schema, features) query_batches.append(model_input_batch(packed, schema)) belief = predict_public(model, schema, packed) candidates = { int(state): belief[int(state)] for state in episode["public_refs"]["candidate_state_ids"] } decisions.append({ "base_episode_id": episode["base_episode_id"], "query_id": episode["query_id"], "task": args.task, "belief": candidates, }) wanted = {episode["base_episode_id"] for episode in episodes} private = { row["base_episode_id"]: row["evaluation_private"] for row in rows(args.tasks / "private/evaluation_gt.jsonl") if row["base_episode_id"] in wanted } scene_ids = {episode["scene_id"] for episode in episodes} if len(scene_ids) != 1: raise ValueError("--tasks must select episodes from one scene") scene_id = next(iter(scene_ids)) navmesh = args.navmesh_root / f"{scene_id}.navmesh" inspector = ( GroundedSAMInspector( args.perception_config, dino_model=args.grounding_dino_model, sam_model=args.sam2_model, ) if args.inspection == "grounded-sam" else None ) backend = HabitatInspectionBackend( args.hssd_root, scene_id, navmesh, public_viewpoints, detector=inspector, ) calibrator = DetectionCalibrator.load(args.calibration) if args.calibration else None controller = LunaToolController() if args.controller == "luna" else None scores = [] try: for episode, decision in zip(episodes, decisions, strict=True): truth = private[episode["base_episode_id"]] backend.prepare( truth, objects[episode["target"]["object_id"]], [state for state in decision["belief"] if state in public_goals], dynamic=args.task == "n4", ) if args.task in {"n3", "n4"} or args.agent: query = queries[episode["query_id"]] target_id = query["input"]["target"]["instance_uuid"] memory = VersionedMemory() for index, event in enumerate(query["input"].get("target_history", [])): if event["event_type"] == "positive_observation": observed_state = int(event["observed_state_id"]) memory.observe( target_id, observed_state, float(event["timestamp_s"]), float(event["detector_confidence"]) * float(event["instance_match_confidence"]), f"{query['query_id']}:history:{index}", np.asarray(all_centers[observed_state]), ) motion_schedule = None if args.task == "n4": motion_schedule = truth.get("target_motion_schedule") if motion_schedule is None: raise ValueError("N4 private evaluation record requires target_motion_schedule") required_motion = { "time_s", "target_position_xyz", "current_state_id", "valid_goal_viewpoints", } if any(required_motion - set(event) for event in motion_schedule): raise ValueError("N4 target_motion_schedule has incomplete motion events") world = HabitatAgentWorld( backend, public_viewpoints, state_centers, episode["agent_start"]["position_xyz"], episode["agent_start"]["rotation_xyzw"], calibrator=calibrator, known_states=set(decision["belief"]), motion_schedule=motion_schedule, surface_points=surface_points, ) unknown_ids = [int(i) for i, flag in enumerate(features["candidate_is_unknown"]) if flag] can_explore = "EXPLORE" in episode["public_refs"].get("action_space", []) transition = (NeuralTransition(model, transition_head, query_batches[len(scores)]) if args.task == "n4" else IdentityTransition()) agent = Agent( decision["belief"], public_goals, world, transition, AgentConfig( max_inspections=1 if args.task == "n1" else int( episode["episode_budget"]["max_candidate_inspections"]), max_path_m=float(episode["episode_budget"]["max_path_length_m"]), chunk_m=2.0, unknown_state=unknown_ids[0] if can_explore and unknown_ids else None, ), sample_count={state: 25 for state in decision["belief"]}, controller=controller, memory=memory, target_id=target_id, time_origin_s=float(query["input"]["query"]["query_time_s"]), ) result = agent.run() final = world.private_inspections[-1] if world.private_inspections else None success = bool( result.found and final is not None and final["distance_to_valid_goal_m"] <= float(episode["success_spec"]["max_geodesic_distance_m"]) and final["visible_fraction"] >= 0.20 ) oracle_m = float(truth["oracle_shortest_path_m"]) score = { "base_episode_id": episode["base_episode_id"], "task": args.task, "success": success, "inspections": len(result.inspections), "inspection_order": result.inspections, "inspection_evidence": world.private_inspections, "actions": result.actions, "path_m": round(result.path_m, 6), "spl": round(float(success) * oracle_m / max(result.path_m, oracle_m, 1e-9), 6), "true_state_id": int(backend.truth["current_state_id"]), "posterior": result.posterior, "termination": result.termination, "evidence_trace": result.evidence_trace, } else: score = evaluate_search( episode, truth, decision["belief"], start=episode["agent_start"]["position_xyz"], goals=public_goals, distance=backend.distance, inspect=backend.inspect, task=args.task, ) scores.append(score) backend.clear() finally: backend.close() if inspector is not None: inspector.close() args.output.mkdir(parents=True) with (args.output / "policy.jsonl").open("w", encoding="utf-8") as handle: for decision in decisions: handle.write(json.dumps(decision, ensure_ascii=False) + "\n") with (args.output / "scores.jsonl").open("w", encoding="utf-8") as handle: for score in scores: handle.write(json.dumps(score, ensure_ascii=False) + "\n") summary = { "task": args.task, "world": args.world, "episodes": len(scores), "successes": sum(row["success"] for row in scores), "sr": sum(row["success"] for row in scores) / len(scores), "spl": sum(row["spl"] for row in scores) / len(scores), "track": f"high_level_{'event_agent' if args.task in {'n3', 'n4'} or args.agent else 'ranked'}_{args.inspection}", "min_visible_fraction": 0.20, "dataset": str(args.tasks), "checkpoint": str(args.checkpoint), } (args.output / "summary.json").write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8") print(json.dumps(summary, ensure_ascii=False, indent=2)) return 0 if __name__ == "__main__": raise SystemExit(main())