EvolvingNav / code /scripts /collect_calibration.py
fengnian1678's picture
Mirror ZJU4EmbodiedAI/EvolvingNav at dfa6872
ad91e86 verified
Raw History Blame Contribute Delete
5.87 kB
#!/usr/bin/env python3
"""Collect held-out RGB-D detector observations for logistic recall calibration."""
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.coverage import (
camera_forward, candidate_surface_samples, depth_quality, visible_sample_ids,
)
from evolvingnav_paper.perception import GroundedSAMInspector
from evolvingnav_paper.run import rows
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--dataset", type=Path, required=True)
parser.add_argument("--tasks", type=Path, required=True)
parser.add_argument("--hssd-root", type=Path, required=True)
parser.add_argument("--navmesh-root", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--limit", type=int, default=16)
parser.add_argument("--dino-model", default="IDEA-Research/grounding-dino-tiny")
parser.add_argument("--sam-model", default="facebook/sam2.1-hiera-tiny")
parser.add_argument("--perception-config", type=Path,
default=Path(__file__).resolve().parents[1] / "configs/perception.yaml")
args = parser.parse_args()
if args.output.exists():
raise FileExistsError(args.output)
catalog = json.loads((args.tasks / "catalogs/candidate_states_navigation.json").read_text())
viewpoints = {int(row["state_id"]): row["navigation_viewpoint"]
for row in catalog["states"] if row.get("navigation_eligible")}
centers = {int(row["state_id"]): np.asarray(row["state_center"], dtype=float)
for row in catalog["states"]}
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")}
selected = []
for record in rows(args.dataset / "records/val_queries.jsonl"):
state = int(record["supervision"]["current_state_id"])
if state in viewpoints and record["input"]["target"]["instance_uuid"] in objects:
selected.append(record)
if len(selected) >= args.limit:
break
if len(selected) < args.limit:
raise ValueError(f"only {len(selected)} validation records have a public viewpoint")
inspector = GroundedSAMInspector(
args.perception_config, dino_model=args.dino_model, sam_model=args.sam_model
)
backend = HabitatInspectionBackend(
args.hssd_root, selected[0]["scene_id"],
args.navmesh_root / f"{selected[0]['scene_id']}.navmesh",
viewpoints, detector=inspector,
)
observations = []
try:
for record in selected:
target = record["input"]["target"]
state = int(record["supervision"]["current_state_id"])
center = centers[state]
nearest = sorted(viewpoints, key=lambda candidate: float(np.linalg.norm(
np.asarray(viewpoints[candidate]["position_xyz"]) - center
)))
selected_views = list(dict.fromkeys(
[state, *nearest[:3], nearest[len(nearest) // 2], nearest[-1]]
))
truth = {
"target_position_xyz": record["supervision"]["current_position"],
"current_state_id": state,
"valid_goal_viewpoints": [viewpoints[state]],
}
backend.prepare(truth, objects[target["instance_uuid"]], selected_views)
samples = candidate_surface_samples(center, place_points=surface_points.get(state))
for view_state in selected_views:
viewpoint = viewpoints[view_state]
backend.inspect(view_state, viewpoint["position_xyz"])
depth = backend.last_observation["depth"]
covered = visible_sample_ids(
samples, viewpoint["position_xyz"], viewpoint["rotation_xyzw"],
depth, 79.0,
)
semantic = backend.last_semantic == int(objects[target["instance_uuid"]]["semantic_instance_id"])
target_pixels = int(np.count_nonzero(semantic))
overlaps = [int(np.count_nonzero(detection.mask & semantic))
for detection in backend.last_detections]
position = np.asarray(viewpoint["position_xyz"], dtype=float)
direction = center - position
distance = float(np.linalg.norm(direction))
observations.append({
"coverage": len(covered) / len(samples),
"range_m": distance,
"angle_cos": max(0.0, float(np.dot(
direction / max(distance, 1e-9),
camera_forward(viewpoint["rotation_xyzw"])))),
"projected_pixels": len(covered),
"depth_quality": depth_quality(depth),
"category_recall": 0.8,
"detected": bool(target_pixels >= 20 and max(overlaps, default=0) / target_pixels >= 0.1),
})
backend.clear()
finally:
backend.close()
inspector.close()
args.output.parent.mkdir(parents=True, exist_ok=True)
with args.output.open("w", encoding="utf-8") as handle:
for row in observations:
handle.write(json.dumps(row) + "\n")
print(json.dumps({
"observations": len(observations),
"detected": sum(row["detected"] for row in observations),
"output": str(args.output),
}, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())