File size: 5,869 Bytes
ad91e86
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
#!/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())