Repository navigation
Expand file tree
/
Copy pathtest_m7_controller.py
More file actions
304 lines (258 loc) · 13.3 KB
/
Copy pathtest_m7_controller.py
File metadata and controls
304 lines (258 loc) · 13.3 KB
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
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
"""M7 demo: the seven milestone scenarios, all driven through the controller.
Every query goes through ``AgentController.run`` -- no tool is ever called
directly -- because the controller and its trace are what the problem statement
evaluates. Queries 3 and 6 print their full JSON trace.
Usage (Colab, with the RS adapter):
python scripts/test_m7_controller.py \
--adapter /content/drive/MyDrive/SatQueryAI/adapters/.../checkpoint-1590 \
--patches /content/data/processed/ben_patches \
--cdvqa /content/data/raw/cdvqa
Without --adapter it runs on a stub backend, which still exercises routing,
validation, chaining and the trace contract (but not real model answers).
"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
import numpy as np
from PIL import Image
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from agent.planner import AgentController, TaskType # noqa: E402
SEP = "=" * 78
class StubBackend:
"""Deterministic stand-in so the pipeline is testable without weights."""
model_id = "stub/no-model"
adapter_path = None
rs_adapted = False
def generate(self, images, prompt, params=None):
if not isinstance(images, (list, tuple)):
images = [images] if images is not None else []
p = prompt.lower()
if "classify this remote-sensing question" in p:
text = "vqa"
elif "earlier" in p and "later" in p:
text = "The built-up area increased in the north-east near the river."
elif "optical" in p and "sar" in p:
text = ("Conclusion: built-up and water both present. "
"Optical evidence: bright roofs and a dark surface. "
"SAR evidence: strong returns over buildings, specular water.")
elif "bounding box" in p or "box" in p:
text = "The region is at [0.32, 0.41, 0.63, 0.77]."
else:
text = "Agricultural land with a river and scattered built-up areas."
return {"text": text, "model_id": self.model_id, "rs_adapted": False,
"duration_ms": 1.0, "backend": "stub", "warnings": []}
def patch_embeddings(self, image):
arr = np.asarray(image.convert("L").resize((16, 16), Image.BILINEAR),
dtype=np.float32)
return np.stack([arr, 255.0 - arr, np.full_like(arr, 7.0)], -1), (16, 16)
def build_inputs(work: Path, patches: Path | None, cdvqa: Path | None,
demo_assets: str | None = None) -> dict:
"""Prefer real data when present; fall back to georeferenced synthetics."""
work.mkdir(parents=True, exist_ok=True)
inputs: dict[str, object] = {"source": "synthetic",
"_demo_assets": demo_assets or ""}
# Real optical + SAR from the M6 BigEarthNet patches, if available.
if patches and (patches / "selection.json").exists():
try:
from agent.ben_io import (iter_patches, load_patch_stack,
rgb_from_stack, sar_preview_from_stack)
rec = next(iter(iter_patches(patches, split="test")))
stack = load_patch_stack(rec["s2_dir"], rec["s1_dir"])
o = work / "real_s2_rgb.png"
s = work / "real_s1_vv.png"
Image.fromarray(rgb_from_stack(stack)).save(o)
Image.fromarray(sar_preview_from_stack(stack)).save(s)
inputs.update(optical_real=o, sar_real=s, band_stack=stack,
source="bigearthnet")
except Exception as exc: # noqa: BLE001
print(f"[m7] could not load BigEarthNet patch ({exc}); using synthetics")
# Real bi-temporal pair from CDVQA, if available.
if cdvqa and cdvqa.exists():
try:
from scripts.test_m5_cdvqa import pick_samples
sample = pick_samples(cdvqa, 1)[0]
a = work / "real_before.png"
b = work / "real_after.png"
sample["img0"].save(a)
sample["img1"].save(b)
inputs.update(before_real=a, after_real=b)
except Exception as exc: # noqa: BLE001
print(f"[m7] could not load CDVQA pair ({exc}); using synthetics")
# Real Sentinel-2/Sentinel-1 assets, when they have been fetched. These are
# what the demo and the screenshots use; the synthetic fixtures below stay
# for the offline/deterministic path.
demo = Path(inputs.pop("_demo_assets", "") or "demo_assets")
manifest = demo / "manifest.json"
if manifest.exists():
m = json.loads(manifest.read_text(encoding="utf-8"))
a = m["assets"]
def _f(key):
f = (a.get(key) or {}).get("file")
return demo / f if f else None
optical, early, late = _f("optical_single"), _f("bitemporal_early"), _f("bitemporal_late")
sar, other = _f("sar"), _f("crs_mismatch_pair")
if all(p is not None and p.exists() for p in (optical, early, late, other)):
inputs.update(optical=optical, before=early, after=late,
wrong_crs=other, source="real_sentinel")
if sar is not None and sar.exists():
inputs["sar"] = sar
else:
# No real SAR: fall back to the synthetic one and SAY SO, rather
# than quietly mixing a fake modality into a "real" run.
from tests.fixtures.synthetic import build_all as _ba
inputs["sar"] = _ba(work)["sar_dual"]
inputs["source"] = "real_sentinel (SAR=SYNTHETIC FALLBACK)"
return inputs
# Georeferenced synthetics, from the one fixture generator the tests and
# the UI screenshots also use. They are physically ordered scenes -- fields,
# a river, a settlement -- not noise, so an answer about them is meaningful
# and the CRS-mismatch case is available regardless of what real data is on
# disk.
from tests.fixtures.synthetic import build_all
fixtures = build_all(work)
inputs["optical"] = fixtures["scene_rgb"]
inputs["before"] = fixtures["pair_before"]
inputs["after"] = fixtures["pair_after"]
inputs["sar"] = fixtures["sar_dual"]
inputs["wrong_crs"] = fixtures["pair_other_crs"]
return inputs
def show(trace, label: str, full_json: bool = False) -> None:
print(f"\n{SEP}\n{label}\n{SEP}")
print(f"query : {trace.query}")
print(f"classified_task : {trace.classified_task.value} "
f"(confidence {trace.task_confidence}, {trace.classification_method})")
v = trace.input_validation
if v:
print(f"validation : ok={v.ok} images={v.n_images} "
f"modalities={v.modalities}"
+ (f" ERROR={v.error_code}" if v.error_code else ""))
print(f"plan : {[s.tool_name for s in trace.plan.steps]}"
f"{' (CHAIN)' if trace.plan.is_chain else ''}")
for s in trace.steps:
print(f" step {s.order}: {s.tool_name} [{s.status.value}] "
f"{s.duration_ms}ms conf={s.confidence} adapter={s.adapter_id}")
if s.output_summary:
print(f" -> {s.output_summary[:160]}")
if s.error:
print(f" !! {s.error.get('code')}: {s.error.get('message')}")
print(f"final_answer : {trace.final_answer[:300]}")
print(f"evidence : {trace.visual_evidence_paths}")
print(f"overall_conf : {trace.overall_confidence} "
f"[{trace.confidence_breakdown.get('weakest_component', '?')} is weakest]")
print(f"status : {trace.status}")
if full_json:
print(f"\n--- FULL JSON TRACE ---\n{trace.to_json()}")
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None,
help="local adapter directory (implies --backend local)")
ap.add_argument("--backend", choices=("stub", "local", "remote"), default=None,
help="stub (default), local weights, or the Colab server "
"behind SATQUERY_REMOTE_URL")
ap.add_argument("--patches", default=None)
ap.add_argument("--cdvqa", default=None)
ap.add_argument("--fusion-head", default="models/fusion/fusion_head.pt")
ap.add_argument("--demo-assets", default="demo_assets",
help="real Sentinel assets from "
"scripts/fetch_real_demo_assets.py; falls back "
"to synthetic fixtures when absent")
ap.add_argument("--work", default="outputs/m7")
ap.add_argument("--out", default="reports/m7")
args = ap.parse_args()
work, out_dir = Path(args.work), Path(args.out)
inputs = build_inputs(work, Path(args.patches) if args.patches else None,
Path(args.cdvqa) if args.cdvqa else None,
demo_assets=args.demo_assets)
print(f"[m7] input source: {inputs['source']}")
kind = args.backend or ("local" if args.adapter else "stub")
adapter_for_trace = args.adapter
if kind == "local":
from models.serving.local import LocalBackend
backend = LocalBackend(adapter_path=args.adapter)
backend._get_backend()
print(f"[m7] LocalBackend ready (rs_adapted={backend.rs_adapted})")
elif kind == "remote":
from models.serving.remote import RemoteBackend
backend = RemoteBackend()
health = backend.health()
if not health.get("reachable"):
print(f"[m7] remote backend unreachable: {health.get('error')}")
return 4
print(f"[m7] RemoteBackend ready: {backend.model_id} "
f"rs_adapted={backend.rs_adapted} "
f"adapter={backend.adapter_path} "
f"caps={health.get('remote', {}).get('capabilities')}")
adapter_for_trace = adapter_for_trace or backend.adapter_path
else:
backend = StubBackend()
print("[m7] no backend selected: using StubBackend (routing/validation/"
"trace still exercised; answers are not real)")
head_path = Path(args.fusion_head)
ctrl = AgentController(
backend=backend, report_dir=out_dir, adapter_path=adapter_for_trace,
fusion_head_path=head_path if head_path.exists() else None)
optical = inputs.get("optical_real") or inputs["optical"]
before = inputs.get("before_real") or inputs["before"]
after = inputs.get("after_real") or inputs["after"]
sar = inputs.get("sar_real") or inputs["sar"]
opt_for_sar = inputs.get("optical_real") or inputs["optical"]
traces: list[tuple[str, object]] = []
t1 = ctrl.run("Describe the land-cover and major objects visible in this image.",
[optical])
show(t1, "QUERY 1 - captioning [1 optical]")
traces.append(("q1_captioning", t1))
t2 = ctrl.run("Highlight the water body referred to in the query.", [optical])
show(t2, "QUERY 2 - grounding [1 optical]")
traces.append(("q2_grounding", t2))
t3 = ctrl.run("What changed between these two dates, and where did the change "
"occur?", [before, after])
show(t3, "QUERY 3 - change + WHERE [bi-temporal] -> 2-TOOL CHAIN",
full_json=True)
traces.append(("q3_change_chain", t3))
t4 = ctrl.run("Use the optical and SAR images together to identify built-up "
"and water-covered regions.", [opt_for_sar, sar])
show(t4, "QUERY 4 - cross-modal [optical + SAR]")
traces.append(("q4_cross_modal", t4))
t5 = ctrl.run("Has the built-up area increased, decreased, or remained "
"unchanged?", [before, after])
show(t5, "QUERY 5 - direction question [bi-temporal]")
traces.append(("q5_direction", t5))
t6 = ctrl.run("What changed between these two dates, and where did the change "
"occur?", [optical])
show(t6, "QUERY 6 - FAILURE: change query with ONE image", full_json=True)
traces.append(("q6_too_few_images", t6))
t7 = ctrl.run("What changed between these two dates?",
[inputs["before"], inputs["wrong_crs"]])
show(t7, "QUERY 7 - FAILURE: mismatched CRS -> NotCoRegistered")
traces.append(("q7_crs_mismatch", t7))
out_dir.mkdir(parents=True, exist_ok=True)
for name, tr in traces:
(out_dir / f"trace_{name}.json").write_text(tr.to_json())
print(f"\n{SEP}\nSUMMARY\n{SEP}")
ok = True
expectations = {
"q1_captioning": (TaskType.CAPTIONING, None),
"q2_grounding": (TaskType.GROUNDING, None),
"q3_change_chain": (TaskType.CHANGE_ANALYSIS, None),
"q4_cross_modal": (TaskType.CROSS_MODAL, None),
"q5_direction": (TaskType.CHANGE_ANALYSIS, None),
"q6_too_few_images": (None, "too_few_images"),
"q7_crs_mismatch": (None, "not_co_registered"),
}
for name, tr in traces:
want_task, want_err = expectations[name]
got_err = tr.input_validation.error_code if tr.input_validation else None
task_ok = want_task is None or tr.classified_task is want_task
err_ok = want_err is None or got_err == want_err
ok &= task_ok and err_ok
print(f" {'PASS' if task_ok and err_ok else 'FAIL'} {name:20s} "
f"task={tr.classified_task.value:16s} status={tr.status:8s} "
f"conf={tr.overall_confidence:<5} err={got_err}")
print(f"\nchain executed on q3: {t3.plan.is_chain} "
f"({len(t3.steps)} steps: {[s.tool_name for s in t3.steps]})")
print(f"traces written to {out_dir}")
return 0 if ok else 5
if __name__ == "__main__":
raise SystemExit(main())