-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcontroller.py
More file actions
459 lines (414 loc) · 21 KB
/
Copy pathcontroller.py
File metadata and controls
459 lines (414 loc) · 21 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
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
"""The agentic controller: classify -> validate -> plan -> execute -> combine.
One entry point, :meth:`AgentController.run`, and one output, an
:class:`~agent.planner.trace.ExecutionTrace`. Every stage writes into the trace
as it goes, so a run that is rejected at validation still produces a complete,
inspectable record instead of an exception.
Parameter binding is deliberately restrictive: a caller may pass anything, but
only names in the selected tool's ``ToolSpec.permitted_params`` are forwarded.
Everything else is dropped and recorded in ``rejected_params``, because the
problem statement scores whether the system respected permitted parameters.
The one real chain: a change query that also asks *where* runs
``ChangeAnalysis`` and then feeds the change-map evidence PNG into
``Grounding``, so the second tool localises the change the first one found.
That is a genuine dependency -- step 2 consumes step 1's output, not the
original input.
"""
from __future__ import annotations
import hashlib
import os
import time
from pathlib import Path
from typing import Any, Callable
from ..errors import SatQueryError
from ..tools import (
ChangeAnalysisTool,
OpticalSARFusionTool,
SceneCaptioningTool,
SingleImageVQATool,
TextGuidedGroundingTool,
)
from ..tools.base import ToolResult
from ..validators import validate
from .classifier import Classification, classify
from .confidence import RULE_TEXT, aggregate
from .trace import (
ExecutionStep,
ExecutionTrace,
Plan,
PlanStep,
StepStatus,
TaskType,
TraceBuilder,
)
#: Task -> tool class. One place, so routing and execution cannot disagree.
TASK_TOOLS: dict[TaskType, type] = {
TaskType.VQA: SingleImageVQATool,
TaskType.CAPTIONING: SceneCaptioningTool,
TaskType.GROUNDING: TextGuidedGroundingTool,
TaskType.CHANGE_ANALYSIS: ChangeAnalysisTool,
TaskType.CROSS_MODAL: OpticalSARFusionTool,
}
#: A change query that also asks "where" earns the two-tool chain.
WHERE_MARKERS = ("where", "locate", "highlight", "which part", "which area",
"what part", "pinpoint", "show me")
MAX_SUMMARY_CHARS = 400
def _summarise(result: ToolResult) -> str:
"""A short factual summary of a tool's output.
Never chain-of-thought: this is the answer or a description of the
artefact, not the model's reasoning about how it got there.
"""
out = result.output or {}
for key in ("answer", "caption", "text"):
if out.get(key):
return str(out[key])[:MAX_SUMMARY_CHARS]
if out.get("boxes") is not None:
n = len(out["boxes"])
return f"{n} region(s) localised" + (
f"; best box {out['boxes'][0]}" if n else "")
if out.get("quantitative"):
q = out["quantitative"]
classes = q.get("classes", {})
return ("fusion(" + q.get("method", "?") + "): " + ", ".join(
f"{c}={'present' if v.get('present') else 'absent'}"
f"({v.get('coverage')})" for c, v in classes.items()))[:MAX_SUMMARY_CHARS]
return str(out)[:MAX_SUMMARY_CHARS] if out else ""
def _adapter_fingerprint(adapter_path: str | None) -> tuple[str | None, str | None]:
"""(adapter_id, sha256-of-adapter-weights). Evidence of RS adaptation.
The id is recorded on every query even when hashing is unavailable, since
it is what shows the answer came from the adapted model.
"""
if not adapter_path:
return None, None
p = Path(adapter_path)
adapter_id = p.name or str(p)
weights = p / "adapter_model.safetensors"
if not weights.exists():
return adapter_id, None
try:
h = hashlib.sha256()
with open(weights, "rb") as fh:
for chunk in iter(lambda: fh.read(1 << 20), b""):
h.update(chunk)
return adapter_id, h.hexdigest()
except OSError:
return adapter_id, None
class AgentController:
"""Routes one query to the right tool(s) and records the whole run."""
def __init__(self, backend=None, report_dir: str | Path | None = None,
fusion_head=None, fusion_head_path: str | Path | None = None,
adapter_path: str | None = None) -> None:
self._backend = backend
self._report_dir = Path(report_dir) if report_dir else None
self._fusion_head = fusion_head
self._fusion_head_path = fusion_head_path
self._adapter_path = adapter_path or getattr(backend, "adapter_path", None) \
or os.environ.get("ADAPTER_PATH")
self._tools: dict[TaskType, Any] = {}
# -- tool construction -------------------------------------------------
def _tool_for(self, task: TaskType):
"""Construct a tool, passing only the kwargs its __init__ accepts.
Tools differ in what they take (only the ones that write evidence
accept ``report_dir``), so bind by introspection rather than keeping a
second list of which tool wants what -- that list would drift.
"""
if task not in self._tools:
import inspect
cls = TASK_TOOLS[task]
available = {
"backend": self._backend,
"report_dir": self._report_dir,
"fusion_head": self._fusion_head,
"fusion_head_path": self._fusion_head_path,
}
accepted = set(inspect.signature(cls.__init__).parameters)
kwargs = {k: v for k, v in available.items()
if k in accepted and v is not None}
kwargs.setdefault("backend", self._backend)
self._tools[task] = cls(**kwargs)
return self._tools[task]
# -- planning ----------------------------------------------------------
def _build_plan(self, task: TaskType, query: str) -> Plan:
tool_cls = TASK_TOOLS[task]
spec = tool_cls.spec
steps = [PlanStep(order=0, tool_name=spec.name, task=spec.task,
reason=f"query classified as {task.value}")]
is_chain = False
rationale = f"single-tool plan for {task.value}"
if task is TaskType.CHANGE_ANALYSIS and self._asks_where(query):
g = TextGuidedGroundingTool.spec
steps.append(PlanStep(
order=1, tool_name=g.name, task=g.task,
reason="query asks WHERE the change occurred; localise it on the "
"change map produced by step 0",
consumes="step_0.change_map_evidence"))
is_chain = True
rationale = ("two-tool chain: change analysis produces a spatial "
"change map, then grounding localises the change on it")
return Plan(steps=steps, is_chain=is_chain, rationale=rationale)
@staticmethod
def _asks_where(query: str) -> bool:
q = (query or "").lower()
return any(m in q for m in WHERE_MARKERS)
# -- parameter binding -------------------------------------------------
@staticmethod
def _bind_params(tool, params: dict[str, Any]
) -> tuple[dict[str, Any], list[str]]:
"""Keep only ToolSpec-permitted names; report what was dropped."""
permitted = set(tool.spec.permitted_params)
bound = {k: v for k, v in (params or {}).items() if k in permitted}
rejected = sorted(set(params or {}) - permitted)
return bound, rejected
def _record_step(self, order: int, tool, result: ToolResult,
bound: dict[str, Any], rejected: list[str],
duration_ms: float) -> ExecutionStep:
model = result.model or {}
warnings = list(result.warnings or [])
if rejected:
warnings.append(
f"parameters not permitted for {tool.spec.name} and dropped: "
f"{rejected}")
step = ExecutionStep(
order=order,
tool_name=tool.spec.name,
task=tool.spec.task,
model_id=model.get("model_id"),
adapter_id=self._adapter_id,
rs_adapted=model.get("rs_adapted"),
params=bound,
duration_ms=round(duration_ms, 1),
status=StepStatus.SUCCESS if result.ok else StepStatus.FAILED,
confidence=result.confidence,
output_summary=_summarise(result) if result.ok else "",
evidence_paths=list((result.output or {}).get("evidence_paths") or []),
error=result.error,
warnings=warnings,
)
if rejected:
step.params = {**bound, "_rejected_params": rejected}
return step
# -- execution ---------------------------------------------------------
def run(self, query: str, image_paths: list[str | Path],
params: dict[str, Any] | None = None,
progress: Callable[[str], None] | None = None) -> ExecutionTrace:
"""Answer one query and return the full execution trace.
``progress`` is an optional callback invoked with a short stage name
at each controller stage (classify, validate, plan, execute:<tool>,
combine). It is purely observability for callers who show live
progress; passing None (the default) changes nothing about how the
run executes or what the trace records.
"""
paths = [Path(p) for p in (image_paths or [])]
_progress = progress or (lambda _stage: None)
self._adapter_id, adapter_sha = _adapter_fingerprint(self._adapter_path)
builder = TraceBuilder(query, [str(p) for p in paths],
manifest={"tool_registry": sorted(
t.spec.name for t in TASK_TOOLS.values())})
builder.set_provenance(
model_id=getattr(self._backend, "model_id", None),
adapter_id=self._adapter_id, adapter_sha256=adapter_sha)
# -- 1. classify ---------------------------------------------------
# Peek at modalities first: an optical+SAR pair is evidence for routing.
peek = validate(paths, TaskType.UNSUPPORTED) if paths else None
modalities = peek.modalities if peek else []
cls: Classification = classify(query, n_images=len(paths),
modalities=modalities,
backend=self._backend)
builder.set_classification(cls.task, cls.confidence, cls.method)
_progress("classify")
builder.trace.confidence_breakdown = {"classification": cls.to_dict()}
if cls.task is TaskType.UNSUPPORTED or cls.task not in TASK_TOOLS:
# The query asked for a real task the inputs cannot support:
# validate against the *intended* task so the trace carries the
# specific typed error (too_few_images, modality_mismatch) rather
# than a vague "unsupported".
intended = cls.intended_task
v = validate(paths, intended or TaskType.UNSUPPORTED)
builder.set_validation(v)
if intended is not None and not v.ok:
return builder.fail(
v.error_code or "validation_failed",
v.error_message or "input validation failed",
final_answer=(
f"This query asks for {intended.value}, which cannot run "
f"on the supplied input: {v.error_message}"))
return builder.fail(
"unsupported_query",
f"No supported task matches this query with {len(paths)} image(s). "
f"Routing notes: {'; '.join(cls.notes) or 'none'}",
final_answer="This query is not supported by the available tools.")
# -- 2. validate ---------------------------------------------------
_progress("validate")
validation = validate(paths, cls.task)
builder.set_validation(validation)
if not validation.ok:
return builder.fail(validation.error_code or "validation_failed",
validation.error_message or "input validation failed")
# -- 3. plan -------------------------------------------------------
_progress("plan")
plan = self._build_plan(cls.task, query)
builder.set_plan(plan)
# -- 4. execute ----------------------------------------------------
step_confidences: list[float] = []
failed = 0
primary: ToolResult | None = None
rs_adapted_seen = True
tool = self._tool_for(cls.task)
_progress(f"execute:{tool.spec.name}")
bound, rejected = self._bind_params(tool, params or {})
t0 = time.perf_counter()
try:
primary = self._invoke(tool, cls.task, paths, query, bound)
except SatQueryError as exc:
primary = ToolResult(tool=tool.spec.name, task=tool.spec.task, ok=False,
error=exc.to_dict())
except Exception as exc: # noqa: BLE001 - controller never raises
primary = ToolResult(tool=tool.spec.name, task=tool.spec.task, ok=False,
error={"code": "tool_crashed",
"message": str(exc),
"context": {"type": type(exc).__name__}})
dt = (time.perf_counter() - t0) * 1000
step = builder.add_step(
self._record_step(0, tool, primary, bound, rejected, dt))
if primary.ok:
step_confidences.append(primary.confidence or 0.0)
rs_adapted_seen = bool((primary.model or {}).get("rs_adapted", True))
else:
failed += 1
# -- 5. chained step, if planned and step 0 produced a map ---------
if plan.is_chain and primary.ok:
chain_step = self._run_chain_step(builder, primary, query, params or {})
if chain_step is not None:
if chain_step.status is StepStatus.SUCCESS:
step_confidences.append(chain_step.confidence or 0.0)
else:
# DEGRADED keeps the failed-step penalty: the chain did not
# deliver what it promised, and the substituted answer is
# weaker than the one the step was planned to produce.
failed += 1
# -- 6. combine ----------------------------------------------------
_progress("combine")
final_answer = self._compose_answer(builder.trace, primary)
overall, breakdown = aggregate(
cls.confidence, step_confidences,
failed_steps=failed, benchmark_mode=validation.benchmark_mode,
rs_adapted=rs_adapted_seen)
breakdown["classification"] = cls.to_dict()
return builder.finish(final_answer, overall, RULE_TEXT, breakdown)
def _invoke(self, tool, task: TaskType, paths: list[Path], query: str,
bound: dict[str, Any]) -> ToolResult:
"""Call one tool with the argument shape its task requires."""
if task in (TaskType.VQA, TaskType.GROUNDING):
return tool.run(paths[0], query, **bound)
if task is TaskType.CAPTIONING:
return tool.run(paths[0], query, **bound)
if task is TaskType.CHANGE_ANALYSIS:
return tool.run(paths[0], paths[1], query, **bound)
if task is TaskType.CROSS_MODAL:
return tool.run(paths[0], paths[1], query, **bound)
raise SatQueryError(f"No invocation defined for task {task.value}")
def _run_chain_step(self, builder: TraceBuilder, primary: ToolResult,
query: str, params: dict[str, Any]) -> ExecutionStep | None:
"""Ground the change on the change map produced by step 0."""
change_map = (primary.output or {}).get("change_map") or {}
# Ground on the single-panel overlay, not the before|after|overlay
# triptych: the triptych is a figure for a human reader, and a model
# trained on satellite scenes cannot localise anything inside a
# three-panel contact sheet. Fall back to the triptych only if the
# single-panel render is missing.
evidence = change_map.get("overlay_path") or change_map.get("evidence_path")
if not evidence or not Path(evidence).exists():
skipped = ExecutionStep(
order=1, tool_name=TextGuidedGroundingTool.spec.name,
task=TextGuidedGroundingTool.spec.task,
adapter_id=self._adapter_id, status=StepStatus.SKIPPED,
output_summary="",
warnings=["step 0 produced no change-map evidence to ground on"])
return builder.add_step(skipped)
tool = self._tool_for(TaskType.GROUNDING)
bound, rejected = self._bind_params(tool, params)
t0 = time.perf_counter()
try:
result = tool.run(evidence, self._chain_query(query), **bound)
except Exception as exc: # noqa: BLE001
result = ToolResult(tool=tool.spec.name, task=tool.spec.task, ok=False,
error={"code": "tool_crashed", "message": str(exc),
"context": {"type": type(exc).__name__}})
dt = (time.perf_counter() - t0) * 1000
step = self._record_step(1, tool, result, bound, rejected, dt)
step.warnings.append(
"grounded on the step-0 change map (chained input), not the raw image")
# Graceful degradation. On a diffuse change map -- which is what real
# bi-temporal imagery produces once seasonal and illumination
# differences enter the pixel channel -- there is no tight blob for the
# grounding model to find, and it correctly returns "not found". That
# is honest, but it throws away the localisation the mask ALREADY
# contains. So instead of leaving a failed step, fall back to the mask
# statistics computed in step 0 and mark the step DEGRADED.
#
# Deliberately not a fabricated box: no `boxes` are emitted and the
# tool name stays what actually ran, so nothing here can be mistaken
# for a grounding result. The confidence penalty is unchanged.
if not result.ok:
fallback = self._mask_fallback(primary)
if fallback is not None:
step.status = StepStatus.DEGRADED
step.output_summary = fallback["summary"]
step.error = None
step.warnings.append(
"grounding returned no bounding box on the change map; "
"location below comes from step-0 mask statistics "
"(largest connected component), not from the grounding model")
step.params = {**step.params, "_fallback": "mask_statistics"}
return builder.add_step(step)
@staticmethod
def _mask_fallback(primary: ToolResult) -> dict[str, Any] | None:
"""Describe the change from step-0 mask statistics, in plain words."""
loc = (primary.output or {}).get("change_location") or {}
comp = loc.get("largest_component")
if not comp:
return None
pct = comp["area_fraction_of_scene"] * 100.0
share = comp["share_of_changed_pixels"] * 100.0
box = comp["bbox_fraction"]
return {
"summary": (
f"Largest change region concentrated in the {comp['quadrant']} "
f"of the scene, covering {pct:.1f}% of the image "
f"({share:.0f}% of all changed pixels); bounding box {box} as "
f"fractions of width/height. Derived from the change mask, not "
f"from the grounding model."),
"component": comp,
}
@staticmethod
def _chain_query(query: str) -> str:
"""Ask grounding for the changed region rather than re-asking the
original bi-temporal question, which grounding cannot answer.
Kept to a short noun phrase: the grounding tool prompts the model with
this as an object name, and a full sentence asks it to find something
literally called "the region that changed between the two
acquisitions".
"""
return "changed area"
@staticmethod
def _compose_answer(trace: ExecutionTrace, primary: ToolResult) -> str:
"""Combine step outputs into the user-facing answer."""
if not primary.ok:
err = primary.error or {}
return f"Could not answer: {err.get('message', 'tool failed')}"
parts = [_summarise(primary)]
for step in trace.steps[1:]:
if not step.output_summary:
continue
if step.status is StepStatus.SUCCESS:
parts.append(f"Localisation: {step.output_summary}")
elif step.status is StepStatus.DEGRADED:
parts.append(step.output_summary)
return " ".join(p for p in parts if p).strip()
def answer(query: str, image_paths: list[str | Path], **kwargs: Any
) -> ExecutionTrace:
"""Convenience wrapper: one query in, one trace out."""
ctrl_kwargs = {k: kwargs.pop(k) for k in
("backend", "report_dir", "fusion_head", "fusion_head_path",
"adapter_path") if k in kwargs}
return AgentController(**ctrl_kwargs).run(query, image_paths,
kwargs.get("params"))