Skip to content

Commit 311ceee

Browse files
committed
feat(evaluation): report the reasoning-branch intervention delta as an eval metric
Wire reasoning_intervention_delta (#109) into _evaluate_open_loop so every checkpoint reports eval/reasoning_intervention_delta next to ADE/FDE. Additive and a no-op when there is no reasoning head. Thread a fixed initial_noise and the navigation inputs through both forwards so the delta is measured at the real operating point rather than being noise-dominated. Refs #123, #109. Signed-off-by: GABRIELA CORDOVA <100548769@alumnos.uc3m.es>
1 parent 16d60f0 commit 311ceee

4 files changed

Lines changed: 189 additions & 2 deletions

File tree

‎Model/evaluation/faithfulness.py‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -50,13 +50,20 @@ def reasoning_intervention_delta(
5050
projection: Optional[Any] = None,
5151
geometry_type: Optional[str] = None,
5252
image_transform: Optional[Any] = None,
53+
**forward_kwargs: Any,
5354
) -> dict[str, float]:
5455
"""Mean trajectory L2 between the reasoning-coupled and bypassed runs.
5556
5657
Args:
5758
model: an ``AutoE2E`` built with ``enable_reasoning=True``.
5859
camera_tiles / map_input / visual_history / egomotion_history: one batch.
5960
projection / geometry_type / image_transform: current geometry ABI.
61+
forward_kwargs: extra forward inputs threaded identically into both runs —
62+
navigation (``route_mask`` / ``map_valid`` / ``route_valid``) so the
63+
delta is read at the checkpoint's real operating point, and a fixed
64+
``initial_noise`` so a stochastic planner uses the SAME noise in the
65+
coupled and bypassed runs (otherwise the delta is noise, not the
66+
intervention).
6067
6168
Returns:
6269
``{"trajectory_l2": float}`` — 0.0 while the coupling gate is untrained.
@@ -76,7 +83,7 @@ def reasoning_intervention_delta(
7683
model.eval()
7784
restore_buffer = _snapshot_buffer(model)
7885
fwd = dict(projection=projection, geometry_type=geometry_type,
79-
image_transform=image_transform, mode="infer")
86+
image_transform=image_transform, mode="infer", **forward_kwargs)
8087

8188
try:
8289
with torch.no_grad():
@@ -108,6 +115,7 @@ def horizon_intervention_delta(
108115
projection: Optional[Any] = None,
109116
geometry_type: Optional[str] = None,
110117
image_transform: Optional[Any] = None,
118+
**forward_kwargs: Any,
111119
) -> dict[str, float]:
112120
"""Trajectory delta under a targeted intervention on the horizon tokens.
113121
@@ -155,7 +163,7 @@ def _perturb(pred):
155163
model.eval()
156164
restore_buffer = _snapshot_buffer(model)
157165
fwd = dict(projection=projection, geometry_type=geometry_type,
158-
image_transform=image_transform, mode="infer")
166+
image_transform=image_transform, mode="infer", **forward_kwargs)
159167

160168
original_forward = head.forward
161169

Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
"""Demonstration test for the faithfulness `forward_kwargs` extension.
2+
3+
Proves the two gaps found when trying to wire `reasoning_intervention_delta` into
4+
the route-conditioned, flow-matching eval loop, and that threading `forward_kwargs`
5+
(a fixed `initial_noise` + the navigation inputs) fixes both:
6+
7+
1. Without a fixed `initial_noise`, a stochastic planner draws different noise in
8+
the coupled and bypassed runs, so the "intervention delta" is dominated by
9+
noise rather than by the reasoning intervention (which here is 1e-3).
10+
2. Passing the navigation inputs (`route_mask`) must not raise — the pre-extension
11+
signature had no `**kwargs`, so the eval loop could not thread them at all.
12+
13+
Self-contained: a minimal stochastic stub, no fixture / GPU / network.
14+
"""
15+
16+
from __future__ import annotations
17+
18+
import pytest
19+
import torch
20+
from torch import nn
21+
22+
from evaluation.faithfulness import reasoning_intervention_delta
23+
24+
_T, _D = 8, 2
25+
_BUMP = 1e-3 # the reasoning intervention's effect on the trajectory
26+
27+
28+
class _Reactive(nn.Module):
29+
def __init__(self) -> None:
30+
super().__init__()
31+
self.ReasoningHead = nn.Identity() # present => reasoning coupled
32+
33+
34+
class _StochasticStub(nn.Module):
35+
"""trajectory = reasoning_bump + route_term + noise.
36+
37+
- reasoning_bump: `_BUMP` while `Reactive_E2E.ReasoningHead` is not None (the
38+
intervention sets it to None to bypass);
39+
- noise: uses `initial_noise` if given (fixed across runs), else fresh randn;
40+
- route_term: shifts the operating point when `route_mask` is provided.
41+
"""
42+
43+
def __init__(self) -> None:
44+
super().__init__()
45+
self.Reactive_E2E = _Reactive()
46+
47+
def forward(self, camera, map_input, vis_hist, ego, *,
48+
projection=None, geometry_type=None, image_transform=None,
49+
route_mask=None, initial_noise=None, mode="infer", **_):
50+
b = camera.shape[0]
51+
reasoning_on = self.Reactive_E2E.ReasoningHead is not None
52+
bump = _BUMP if reasoning_on else 0.0
53+
route_term = 0.0 if route_mask is None else 0.5
54+
noise = initial_noise if initial_noise is not None else torch.randn(b, _T, _D)
55+
return torch.zeros(b, _T, _D) + bump + route_term + noise
56+
57+
58+
def _inputs(b: int = 2):
59+
return (torch.randn(b, 7, 3, 256, 256), torch.randn(b, 3, 256, 256),
60+
torch.randn(b, 896), torch.randn(b, 256))
61+
62+
63+
def test_gap_naive_call_is_noise_dominated():
64+
torch.manual_seed(0)
65+
out = reasoning_intervention_delta(_StochasticStub(), *_inputs())
66+
# noise ~ N(0,1) in each run -> delta on the order of 1, not the 1e-3 signal.
67+
assert out["trajectory_l2"] > 0.1
68+
69+
70+
def test_fix_fixed_noise_recovers_the_intervention():
71+
torch.manual_seed(0)
72+
b = 2
73+
fixed = torch.zeros(b, _T, _D) # same noise threaded into both runs
74+
out = reasoning_intervention_delta(_StochasticStub(), *_inputs(b),
75+
initial_noise=fixed)
76+
# noise cancels; only the reasoning bump survives: ||[1e-3, 1e-3]|| = 1e-3*sqrt(2)
77+
assert out["trajectory_l2"] == pytest.approx(_BUMP * 2 ** 0.5, abs=2e-4)
78+
79+
80+
def test_fix_navigation_inputs_are_threaded():
81+
b = 2
82+
fixed = torch.zeros(b, _T, _D)
83+
route = torch.ones(b, 10) # the pre-extension signature would TypeError on this
84+
out = reasoning_intervention_delta(_StochasticStub(), *_inputs(b),
85+
route_mask=route, initial_noise=fixed)
86+
# route is identical in both runs -> it cancels, leaving the intervention only.
87+
assert out["trajectory_l2"] == pytest.approx(_BUMP * 2 ** 0.5, abs=2e-4)

‎Model/tests/test_workflow_training_lifecycle.py‎

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1688,3 +1688,64 @@ def test_resume_load_keeps_rng_tensors_on_cpu():
16881688
keywords = {item.arg: item.value for item in resume_load.keywords}
16891689
assert ast.literal_eval(keywords["map_location"]) == "cpu"
16901690
assert ast.literal_eval(keywords["weights_only"]) is False
1691+
1692+
1693+
class _ReasoningMetricModel(_MetricModel):
1694+
"""`_MetricModel` with a bypassable reasoning head. The trajectory picks up a
1695+
fixed bump while ``Reactive_E2E.ReasoningHead`` is present, so the intervention
1696+
(which sets the head to None) moves the trajectory by exactly that bump — the
1697+
quantity the gate metric must report."""
1698+
1699+
_BUMP = 1e-3
1700+
1701+
def __init__(self):
1702+
super().__init__()
1703+
self.Reactive_E2E = SimpleNamespace(ReasoningHead=object())
1704+
1705+
def __call__(self, visual, *args, **kwargs):
1706+
out = super().__call__(visual, *args, **kwargs)
1707+
if self.Reactive_E2E.ReasoningHead is not None:
1708+
out = out + self._BUMP
1709+
return out
1710+
1711+
1712+
def test_intervention_delta_reported_when_opted_in():
1713+
model = _ReasoningMetricModel()
1714+
loader = [
1715+
(_validation_batch(["sample-b", "sample-a"]), None, "pseudo")
1716+
]
1717+
1718+
metrics = workflows._evaluate_open_loop(
1719+
model, loader, torch.device("cpu"), report_intervention=True
1720+
)
1721+
1722+
assert "reasoning_intervention_delta" in metrics
1723+
# a fixed initial_noise is threaded into both runs, so the delta reflects only
1724+
# the reasoning bump (1e-3 across 128 signals): ||1e-3||_2 = 1e-3 * sqrt(128).
1725+
assert metrics["reasoning_intervention_delta"] == pytest.approx(
1726+
_ReasoningMetricModel._BUMP * 128 ** 0.5, abs=1e-4
1727+
)
1728+
assert model.training is True # eval mode restored after the intervention
1729+
1730+
1731+
def test_no_intervention_delta_without_reasoning_head():
1732+
model = _MetricModel() # no Reactive_E2E.ReasoningHead
1733+
loader = [(_validation_batch(["sample-a"]), None, "pseudo")]
1734+
1735+
metrics = workflows._evaluate_open_loop(
1736+
model, loader, torch.device("cpu"), report_intervention=True
1737+
)
1738+
1739+
assert "reasoning_intervention_delta" not in metrics
1740+
1741+
1742+
def test_no_intervention_delta_by_default():
1743+
model = _ReasoningMetricModel()
1744+
loader = [(_validation_batch(["sample-a"]), None, "pseudo")]
1745+
1746+
# report_intervention defaults to False -> no cost, no metric (training path).
1747+
metrics = workflows._evaluate_open_loop(
1748+
model, loader, torch.device("cpu")
1749+
)
1750+
1751+
assert "reasoning_intervention_delta" not in metrics

‎Platform/pipelines/workflows.py‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -870,6 +870,7 @@ def _evaluate_open_loop(
870870
route_swap_counterfactual: bool = False,
871871
include_navigation_records: bool = False,
872872
include_rollout_selector_records: bool = False,
873+
report_intervention: bool = False,
873874
) -> dict:
874875
"""Evaluate one fixed loader and return finite ADE/FDE plus its UID digest."""
875876
import hashlib
@@ -902,6 +903,11 @@ def _evaluate_open_loop(
902903
route_swap_records: list[dict] = []
903904
rollout_selector_records: list[dict] = []
904905
route_cache: dict[str, dict] = {}
906+
_reactive = getattr(model, "Reactive_E2E", None)
907+
report_delta = report_intervention and (
908+
getattr(_reactive, "ReasoningHead", None) is not None
909+
)
910+
intervention_deltas: list[float] = []
905911
model.eval()
906912
try:
907913
with torch.no_grad():
@@ -1160,6 +1166,22 @@ def _evaluate_open_loop(
11601166
horizon_fde[label].append(
11611167
float(horizon_errors[-1])
11621168
)
1169+
if report_delta and len(intervention_deltas) < 50:
1170+
from evaluation.faithfulness import (
1171+
reasoning_intervention_delta,
1172+
)
1173+
intervention_deltas.append(
1174+
reasoning_intervention_delta(
1175+
model, visual, map_context, vis_hist,
1176+
ego_hist, projection=projection,
1177+
geometry_type=geometry_type,
1178+
route_mask=route_mask, map_valid=map_valid,
1179+
route_valid=route_valid,
1180+
history_frames=history_frames,
1181+
future_frames=future_frames,
1182+
initial_noise=initial_noise,
1183+
)["trajectory_l2"]
1184+
)
11631185
if navigation_geometry is not None:
11641186
from evaluation.navigation_metrics import (
11651187
ROUTE_QUALITY_FIELDS,
@@ -1313,6 +1335,10 @@ def _evaluate_open_loop(
13131335
for label in horizon_steps
13141336
},
13151337
}
1338+
if intervention_deltas:
1339+
result["reasoning_intervention_delta"] = float(
1340+
np.mean(intervention_deltas)
1341+
)
13161342
if navigation_geometry is not None:
13171343
from evaluation.navigation_metrics import (
13181344
summarize_navigation_metrics,
@@ -6455,6 +6481,7 @@ def _run_evaluation(
64556481
device,
64566482
training_policy=training_policy,
64576483
navigation_geometry=navigation_geometry,
6484+
report_intervention=True,
64586485
route_swap_counterfactual=(navigation_geometry is not None),
64596486
include_navigation_records=(
64606487
navigation_records_output is not None
@@ -6755,6 +6782,10 @@ def _run_evaluation(
67556782
for key, value in navigation_metrics.items()
67566783
if value is not None
67576784
})
6785+
if evaluation.get("reasoning_intervention_delta") is not None:
6786+
logged_metrics["eval/reasoning_intervention_delta"] = float(
6787+
evaluation["reasoning_intervention_delta"]
6788+
)
67586789
mlflow.log_metrics(logged_metrics)
67596790

67606791
# Artifacts

0 commit comments

Comments
 (0)