File size: 16,237 Bytes
23a59ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# interactive_four_corr.py
"""Interactive demo: four hallucination signals vs. true one-step error.

Extends the three-predictor correlation UI (`interactive_uncertainty_corr.py`)
with the inverse-dynamics rollout-divergence signal, and reports both the
CUMULATIVE correlation (all steps since reset) and an INSTANTANEOUS one
(sliding window over the last W steps) for each signal.

Signals, all measured at the same step against the same live env transition:

  u_r        motion-normalized tokenizer round-trip residual   [future-free]
  u_f        denoising-trajectory instability                  [future-free]
  u_s        motion-normalized inter-seed variance             [future-free]
  WAV_H      IDM-inferred-action rollout distance              [uses the real frames]

Targets (one per row, see "Horizons" below):

  err_h1 / err_h2 = mean RMS(z_sim - z_real) over an open-loop rollout of that
                    many steps under the GROUND-TRUTH actions.  UNNORMALIZED.

The per-step `true_error = RMS(z_pred - z_env)` is still streamed and shown in
the readout line, but it is no longer what the panels correlate against.

The targets are deliberately NOT motion-normalized. u_r and u_s are themselves
divided by a motion term; correlating them against a motion-normalized target
shares that 1/motion factor and manufactures correlation. Measured offline over
120 trajectories, switching the target from normalized to unnormalized moved
u_r from 0.715 to 0.394 and u_s from 0.818 to 0.490, while u_f *rose* from
0.591 to 0.705 -- it reversed the ranking. This page therefore scores against
the raw error, and additionally streams the normalized variant so the artifact
is visible live rather than hidden.

WAV is not a peer of the u_* signals and is drawn apart in the UI: it
reads the real next frame (it infers the action from it), so it is an offline
audit measurement, not a runtime predictor. Its offline correlation with the
raw one-step error was 0.890 -- against a ceiling of 1.0, since using the true
action instead of the inferred one reproduces the target exactly.

Horizons
--------
Both correlation rows score open-loop rollouts, at two horizons: `--wav_horizon`
(main row, default 4) and `--wav_horizon_long` (second row, default 8). Because
the live history is teacher-forced, a multi-step rollout can only be scored once
the real frames it should be compared against have arrived, so each is restarted
that many steps back and evaluated with a matching lag.

Each row is correlated against the error at ITS OWN horizon -- `err_h1` for the
main row, `err_h2` for the second -- both being the ground-truth-action rollout
distance over the same steps, unnormalized. Pairing a 4-step WAV against a
1-step error would be a mismatch; that is why the row target moves with the row.

Offline over 120 trajectories, discrimination is flat in horizon (a random-action
rollout is 1.52x the ground-truth one at H=1 and still 1.77x at H=16) while the
u_* signals degrade with it (u_f 0.705 -> 0.524, u_s 0.490 -> 0.409 against the
raw error). Comparing the two rows is where that shows up live.

Run from ``src/``:
  ./run_interactive_four_corr.sh combined
or directly:
  python interactive_four_corr.py --tokenizer_ckpt ... --dynamics_ckpt ... --idm_ckpt ...
"""
from __future__ import annotations

import argparse
import math
from typing import Any, Dict, Optional, Tuple

import torch
from aiohttp import web

from model import unpack_spatial_to_bottleneck
from idm_hallucination import IDMRolloutDivergence, load_idm_tokenizer_head
from interactive_uncertainty import (
    SessionState,
    _as_2d_packed,
    build_action_from_keys,
    classify_uncertainty,
    decode_single_packed_frame,
    frame_to_jpeg_bytes,
    frame_to_uint8_hwc,
    reward_from_reward_head_output,
    sample_one_timestep_packed,
)
from interactive_uncertainty_corr import CorrInteractiveServer, build_parser as _corr_parser


def _rms(a: torch.Tensor, b: torch.Tensor) -> float:
    return float((a.float() - b.float()).pow(2).mean().sqrt().item())


class FourCorrServer(CorrInteractiveServer):
    """Three runtime predictors + the IDM rollout-divergence measurement."""

    def __init__(self, args):
        super().__init__(args)
        self.idm, idm_meta = load_idm_tokenizer_head(
            args.idm_ckpt, device=self.device,
            n_latents=self.n_latents, d_bottleneck=self.d_bottleneck,
        )
        print(f"[four-corr] IDM: {args.idm_ckpt} (step={idm_meta['step']}, "
              f"d_hidden={idm_meta['d_hidden']})", flush=True)

        # Two horizons, one rollout. The first h1 steps of an h2-step open-loop
        # rollout ARE the h1-step rollout, so both rows come from a single pair
        # of rollouts via prefix means -- the same trick the offline horizon
        # sweep used. Reuses the offline detector verbatim, so the live numbers
        # are the same quantities that were measured offline.
        self.wav_h1 = int(args.wav_horizon)
        self.wav_h2 = int(args.wav_horizon_long)
        self.wav_ctx = int(args.wav_context)
        self.wav_every = max(1, int(args.wav_every))
        if self.wav_h2 < self.wav_h1:
            raise ValueError("--wav_horizon_long must be >= --wav_horizon")
        self.long_wav = IDMRolloutDivergence(
            self.dyn, self.idm,
            k_max=self.k_max, sched=self.sched,
            packing_factor=int(args.packing_factor),
            n_context=self.wav_ctx, horizon=self.wav_h2,
            max_ctx=self.wav_ctx, tau_ctx=self.tau_ctx,
            use_kv_cache=self.use_kv_cache,
        )
        print(f"[four-corr] WAV horizons: H={self.wav_h1} (main row), "
              f"H={self.wav_h2} (second row), context={self.wav_ctx}, "
              f"every {self.wav_every} step(s)", flush=True)
        # The session keeps only ctx_window+1 latents; without enough of them the
        # lagged rollout can never be scored and both rows stay empty.
        need = self.wav_ctx + self.wav_h2
        if int(args.ctx_window) + 1 < need:
            print(f"[four-corr] WARNING: --ctx_window {args.ctx_window} keeps only "
                  f"{int(args.ctx_window) + 1} latents but the rollout needs {need} "
                  f"(--wav_context {self.wav_ctx} + --wav_horizon_long {self.wav_h2}). "
                  f"The correlation rows will stay empty; raise --ctx_window to >= {need - 1}.")

    @torch.no_grad()
    def _long_horizon(self, st: SessionState):
        """Lagged H-step open-loop rollout over the most recent real frames.

        The live history is teacher-forced, so a multi-step rollout can only be
        scored after the real frames it should be compared against have
        arrived. We therefore look back: take the last (context + H) real
        latents, restart an open-loop rollout at the frame H steps ago, and
        score it against what actually happened. The value is exact but lags
        the display by H steps.

        One rollout is run to the longer horizon; the shorter one is its prefix.
        Returns (wav_h1, err_h1, wav_h2, err_h2), all UNNORMALIZED mean
        distances over the simulated steps -- wav_* under IDM-inferred actions,
        err_* under the ground-truth actions. Their ratio is the live analogue
        of the offline IDM/GT column.
        """
        need = self.wav_ctx + self.wav_h2
        if len(st.z_hist) < need or len(st.a_hist) < need:
            return None
        frames = st.z_hist[-need:]
        acts = st.a_hist[-need:]                      # acts[i] produced frames[i]

        z_seq = torch.stack(frames, dim=0).unsqueeze(0)             # (1, need, Sz, Dz)
        # Transition frames[k] -> frames[k+1] was produced by acts[k+1].
        a_gt = torch.stack(acts[1:], dim=0).unsqueeze(0)            # (1, need-1, 16)

        # Paired seeding: the sampler starts from torch.randn, so without this the
        # two rollouts differ by sampler noise as well as by action source and the
        # wav_H/err_H ratio becomes meaningless per sample.
        kw = dict(act_mask=st.act_mask_1d, lang_emb=st.lang_emb)
        seed = int(st.step)
        torch.manual_seed(seed)
        d_idm = self.long_wav(z_seq, **kw)["dist"]                        # inferred actions
        torch.manual_seed(seed)
        d_gt = self.long_wav(z_seq, actions_override=a_gt, **kw)["dist"]  # true actions
        # Return the PER-STEP distances, not a pair of means. Any horizon H is the
        # mean of the first H entries, so the client can choose H itself and
        # recompute instantly over samples it already holds -- no server round
        # trip, and changing H never discards accumulated statistics.
        return [float(v) for v in d_idm], [float(v) for v in d_gt]

    # ---- IDM helpers -----------------------------------------------------

    @torch.no_grad()
    def _infer_action(self, z_prev: torch.Tensor, z_next: torch.Tensor) -> torch.Tensor:
        """Packed (Sz,Dz) x2 -> inferred action (16,) in [-1,1]."""
        dtype = next(self.idm.parameters()).dtype
        k = int(self.args.packing_factor)
        zp = unpack_spatial_to_bottleneck(z_prev.unsqueeze(0).unsqueeze(0), k=k)
        zn = unpack_spatial_to_bottleneck(z_next.unsqueeze(0).unsqueeze(0), k=k)
        return self.idm(zp.to(dtype), zn.to(dtype))[0, 0].float()

    def _window_with_action(self, st: SessionState, a_override: torch.Tensor):
        """Rebuild the pre-step context window but with `a_override` as the
        action that produces the new frame. Mirrors `_build_local_window`,
        which cannot be reused here because history has already advanced."""
        g = len(st.z_hist) - 1                      # index the new frame occupied
        s = max(0, g - int(st.ctx_window))
        past_list = st.z_hist[s:g]
        if not past_list:
            return None, None
        past = torch.stack(past_list, dim=0).unsqueeze(0)          # (1,t,Sz,Dz)
        t = past.shape[1]
        actions = torch.zeros((1, t + 1, 16), device=self.device, dtype=torch.float32)
        actions[0, 0:t] = torch.stack(st.a_hist[s: s + t], dim=0)
        actions[0, t] = a_override
        return past, actions

    # ---- per-step --------------------------------------------------------

    def _render_step_sync(self, st: SessionState) -> Tuple[Optional[bytes], Dict[str, Any]]:
        prev_step = int(getattr(st, "step", 0))
        # z_prev must be captured before the parent teacher-forces the history.
        z_prev = st.z_hist[-1] if st.z_hist else None

        jpeg, status = super()._render_step_sync(st)

        stepped = int(status.get("step", prev_step)) != prev_step
        if not hasattr(st, "last_wav"):
            st.last_wav = float("nan")
            st.last_u_r_norm = float("nan")
            st.last_u_s_raw = float("nan")
            st.last_h = None                    # (wav_steps, err_steps) per-step distances
            st.wav_h_fresh = False

        if stepped and z_prev is not None and len(st.z_hist) >= 2:
            z_gt = st.z_hist[-1]                       # teacher-forced real latent
            motion_real = _rms(z_gt, z_prev)

            # WAV: infer the action that explains the REAL transition,
            # re-simulate that one step, and measure the distance to reality.
            a_hat = self._infer_action(z_prev, z_gt)
            past, actions = self._window_with_action(st, a_hat)
            if past is not None:
                actmask = st.act_mask_1d.view(1, 1, -1).expand(1, actions.shape[1], -1)
                res = sample_one_timestep_packed(
                    self.dyn,
                    past_packed=past,
                    k_max=self.k_max,
                    sched=self.sched,
                    actions=actions,
                    act_mask=actmask,
                    use_amp=self.use_amp,
                    tau_ctx=self.tau_ctx,
                    lang_emb=st.lang_emb,
                    use_kv_cache=self.use_kv_cache,
                )
                z_sim = res[0] if isinstance(res, tuple) else res
                st.last_wav = _rms(_as_2d_packed(z_sim), z_gt)

            # The parent stores u_r raw; expose the motion-normalized variant too
            # so the normalization artifact is visible side by side.
            st.last_u_r_norm = float(st.last_u_r) / max(motion_real, 1e-3)

            # Long-horizon pair, throttled: each computation is 2 rollouts of H
            # sampler calls, which would otherwise dominate the frame budget.
            st.wav_h_fresh = False
            if int(status.get("step", 0)) % self.wav_every == 0:
                pair = self._long_horizon(st)
                if pair is not None:
                    st.last_h = pair
                    st.wav_h_fresh = True

        status["wav"] = (None if not math.isfinite(float(st.last_wav))
                         else float(st.last_wav))
        status["u_r_norm"] = (None if not math.isfinite(float(st.last_u_r_norm))
                              else float(st.last_u_r_norm))
        # Only emit on steps where the rollout was recomputed, so the client does
        # not fold the same stale sample into its correlation many times over.
        fresh = bool(getattr(st, "wav_h_fresh", False))
        steps = st.last_h if fresh else None
        status["wav_steps"] = steps[0] if steps else None   # per-step, IDM-inferred actions
        status["err_steps"] = steps[1] if steps else None   # per-step, ground-truth actions
        status["wav_max_horizon"] = self.wav_h2             # length of those arrays
        # Kept for the older four_corr page, which reads fixed-horizon scalars.
        for name, arr, h in (("wav_h1", 0, self.wav_h1), ("err_h1", 1, self.wav_h1),
                             ("wav_h2", 0, self.wav_h2), ("err_h2", 1, self.wav_h2)):
            status[name] = (sum(steps[arr][:h]) / h) if steps else None
        status["wav_horizon1"] = self.wav_h1
        status["wav_horizon2"] = self.wav_h2

        fmt = lambda v: "nan" if v is None or not math.isfinite(v) else f"{v:.3f}"
        status["text"] = status.get("text", "") + (
            f" | wav={fmt(float(st.last_wav))}"
            f" wav{self.wav_h1}={fmt(status['wav_h1'])}"
            f" wav{self.wav_h2}={fmt(status['wav_h2'])}"
        )
        return jpeg, status


def build_parser() -> argparse.ArgumentParser:
    p = _corr_parser()
    for action in p._actions:
        if action.dest == "html":
            action.default = "interactive_four_corr.html"
        if action.dest == "port":
            action.default = 7863
    p.add_argument("--idm_ckpt", type=str,
                   default="./logs/idm_tokenizer_ckpts/base_run3/step_040000.pt",
                   help="trained tokenizer-arm inverse dynamics head (idm_tokenizer.py)")
    p.add_argument("--wav_horizon", type=int, default=4,
                   help="open-loop rollout length scored in the MAIN row")
    p.add_argument("--wav_horizon_long", type=int, default=8,
                   help="open-loop rollout length scored in the SECOND row. Must be >= "
                        "--wav_horizon; both come from one rollout via prefix means, so "
                        "the extra horizon costs only the additional steps.")
    p.add_argument("--wav_context", type=int, default=8,
                   help="real frames used as context for the rollout")
    p.add_argument("--wav_every", type=int, default=4,
                   help="recompute every N steps. Each computation costs 2 rollouts of "
                        "--wav_horizon_long sampler calls, so 1 would dominate the frame "
                        "budget.")
    return p


def main():
    args = build_parser().parse_args()
    args.uncertainty_overlay = True

    server = FourCorrServer(args)
    app = web.Application()
    app.router.add_get("/", server.index)
    app.router.add_get("/ws", server.ws_handler)
    app.router.add_get("/status", server.status)
    app.router.add_get("/healthz", server.healthz)

    print(f"[web] four-signal corr UI on http://{args.host}:{args.port}  "
          f"(task={args.task}; target = RMS(z_pred - z_env), unnormalized)")
    web.run_app(app, host=args.host, port=args.port)


if __name__ == "__main__":
    main()