# 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()