Spaces:
Running on Zero
Running on Zero
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()
|