World-Action-Verifier / src /interactive_uncertainty_corr.py
ffeng1017's picture
WAV interactive demo (ZeroGPU)
23a59ea verified
Raw
History Blame Contribute Delete
12 kB
# interactive_uncertainty_corr.py
"""Interactive world-model UI that also steps the live env to expose true
one-step prediction error, for correlating u_r / u_f / u_s with ground truth.
Separate from interactive_uncertainty.py (original open-loop vis unchanged).
True error matches collect_data.py:
true_error = RMS(z_pred - encode(env_next))
and the latent history is teacher-forced with the env encoding so each step
is a clean one-step prediction (not open-loop accumulation).
Run from ``src/``:
python interactive_uncertainty_corr.py --tokenizer_ckpt ... --dynamics_ckpt ...
"""
from __future__ import annotations
import argparse
import math
from typing import Any, Dict, Optional, Tuple
import numpy as np
import torch
from aiohttp import web
from interactive_uncertainty import (
InteractiveServer,
SessionState,
_as_2d_packed,
build_action_from_keys,
build_parser as _base_parser,
classify_uncertainty,
decode_single_packed_frame,
env_obs_to_frame_chw01,
frame_to_jpeg_bytes,
frame_to_uint8_hwc,
reward_from_reward_head_output,
sample_one_timestep_packed,
)
def rms_latent_error(z_a: torch.Tensor, z_b: torch.Tensor) -> float:
return float((z_a.float() - z_b.float()).pow(2).mean().sqrt().item())
class CorrInteractiveServer(InteractiveServer):
"""Same interactive server, but each stepped frame also env.steps and
reports true_error + raw u_* for a correlation histogram frontend."""
def _step_env_and_encode(self, st: SessionState, a: torch.Tensor) -> Tuple[torch.Tensor, bool]:
"""Apply action to the live env; return (z_gt packed, done)."""
env = self._get_or_make_env(st.task)
act_dim = max(0, int(st.act_dim))
if act_dim <= 0:
# No controllable dims — still need a zero step for some envs.
a_np = np.zeros(env.action_space.shape, dtype=np.float32)
else:
real_dim = int(env.action_space.shape[0])
a_np = a.detach().float().cpu().numpy()[:real_dim].astype(np.float32)
if a_np.shape[0] < real_dim:
pad = np.zeros(real_dim, dtype=np.float32)
pad[: a_np.shape[0]] = a_np
a_np = pad
obs, _reward, terminated, truncated, _info = env.step(a_np)
done = bool(terminated or truncated)
frame = env_obs_to_frame_chw01(obs, H=self.H, W=self.W).to(self.device)
z_gt = self._encode_frame_to_packed(frame)
return _as_2d_packed(z_gt), done
def _render_step_sync(self, st: SessionState) -> Tuple[Optional[bytes], Dict[str, Any]]:
if st.reset_requested:
self._reset_session(st)
# New episode: clear the cumulative error accumulator.
st.cum_true_error = 0.0
st.n_true_error = 0
# Ensure corr-only fields exist (SessionState is a plain dataclass).
if not hasattr(st, "last_true_error"):
st.last_true_error = float("nan")
if not hasattr(st, "env_done"):
st.env_done = False
if not hasattr(st, "cum_true_error"):
st.cum_true_error = 0.0
if not hasattr(st, "n_true_error"):
st.n_true_error = 0
a_raw = build_action_from_keys(
st.keys_down, act_dim=st.act_dim, A=16
).to(self.device)
a_raw = (a_raw.clamp(-1, 1) * st.act_mask_1d).to(torch.float32)
beta = float(st.action_beta)
if beta > 0.0:
beta = min(max(beta, 0.0), 0.999)
st.a_smooth = (beta * st.a_smooth + (1.0 - beta) * a_raw).to(torch.float32)
a = st.a_smooth
else:
a = a_raw
frame_cur: Optional[torch.Tensor] = None
stepped: bool = False
env_done = False
if not st.paused and st.act_dim >= 0:
stepped = True
st.a_hist.append(a)
past, actions_local, actmask_local = self._build_local_window(st)
need_h = self.rew_head is not None
N = self.n_samples_u
z_prev_1 = st.z_hist[-1]
past_N = past.expand(N, -1, -1, -1).contiguous()
actions_N = actions_local.expand(N, -1, -1).contiguous()
actmask_N = actmask_local.expand(N, -1, -1).contiguous()
z_prev_N = (
z_prev_1.unsqueeze(0).expand(N, -1, -1).contiguous()
if self.tau_init > 0.0 else None
)
lang_N = None if st.lang_emb is None else st.lang_emb.expand(N, -1).contiguous()
result = sample_one_timestep_packed(
self.dyn,
past_packed=past_N,
k_max=self.k_max,
sched=self.sched,
actions=actions_N,
act_mask=actmask_N,
use_amp=self.use_amp,
return_h=need_h,
tau_ctx=self.tau_ctx,
lang_emb=lang_N,
z_prev=z_prev_N,
tau_init=self.tau_init,
use_kv_cache=self.use_kv_cache,
)
if need_h:
z_next_N, h_N, instability = result
else:
z_next_N, instability = result
st.last_u_f = float(instability)
z_mean_N = z_next_N.float().mean(dim=0)
u_s_raw = z_next_N.float().var(dim=0).mean().clamp(min=0).sqrt().item()
motion = (z_mean_N - z_prev_1.float()).pow(2).mean().sqrt().item()
st.last_u_s = u_s_raw / max(motion, 1e-3)
z_next = z_next_N[0]
h = h_N[0:1] if need_h else None
# Live env one-step target (same action), then teacher-force history.
z_gt, env_done = self._step_env_and_encode(st, a)
st.last_true_error = rms_latent_error(z_next, z_gt)
st.env_done = env_done
# Cumulative (running sum) of the per-step error since last reset.
st.cum_true_error = float(st.cum_true_error) + st.last_true_error
st.n_true_error = int(st.n_true_error) + 1
# History uses GT latent so the next predictor call is one-step.
st.z_hist.append(_as_2d_packed(z_gt.detach()))
st.step += 1
cap = int(st.ctx_window) + 1
if len(st.z_hist) > cap:
st.z_hist = st.z_hist[-cap:]
st.a_hist = st.a_hist[-cap:]
if self.args.uncertainty_overlay and (st.step % max(1, int(self.args.u_every)) == 0):
# Round-trip on the *predicted* latent (what u_r scores), not GT.
frame_cur = decode_single_packed_frame(
self.decoder,
z_packed=_as_2d_packed(z_next.detach()),
H=self.H, W=self.W, C=self.C, patch=self.patch,
packing_factor=self.args.packing_factor,
d_bottleneck=self.d_bottleneck,
)
z_recon = self._encode_frame_to_packed(frame_cur)
diff = z_next.to(torch.float32) - z_recon
st.last_u_r = float(diff.pow(2).mean().sqrt().item())
if not st.calib_done:
st.calib_f_samples.append(st.last_u_f)
st.calib_r_samples.append(st.last_u_r)
st.calib_s_samples.append(st.last_u_s)
if len(st.calib_f_samples) >= int(self.args.calibration_steps):
st.calib_done = True
if need_h:
logits_btlk, centers = self.rew_head(h[:, -1:])
st.last_reward_pred = reward_from_reward_head_output(logits_btlk[0, 0], centers)
st.cum_reward += st.last_reward_pred
# Display the WM prediction being scored (not the GT frame).
if frame_cur is None:
frame_cur = decode_single_packed_frame(
self.decoder,
z_packed=_as_2d_packed(z_next.detach()),
H=self.H, W=self.W, C=self.C, patch=self.patch,
packing_factor=self.args.packing_factor,
d_bottleneck=self.d_bottleneck,
)
if env_done:
# Episode ended in the real env — reseed on the next tick.
st.reset_requested = True
frame_id = st.step
need_encode = (st.cached_jpeg is None) or (st.cached_frame_id != frame_id)
jpeg: Optional[bytes] = None
if need_encode:
if frame_cur is None:
frame_cur = decode_single_packed_frame(
self.decoder,
z_packed=st.z_hist[-1],
H=self.H, W=self.W, C=self.C, patch=self.patch,
packing_factor=self.args.packing_factor,
d_bottleneck=self.d_bottleneck,
)
st.cached_jpeg = frame_to_jpeg_bytes(frame_cur, quality=int(self.args.jpeg_quality))
st.cached_frame_id = frame_id
jpeg = st.cached_jpeg
if self.args.record and stepped:
st.recorded_frames.append(frame_to_uint8_hwc(frame_cur))
else:
jpeg = None
u_r_state = u_f_state = u_s_state = "off"
u_suffix = ""
if self.args.uncertainty_overlay:
if not st.calib_done:
u_r_state = u_f_state = u_s_state = "calibrating"
u_suffix = f" [cal {len(st.calib_f_samples)}/{int(self.args.calibration_steps)}]"
else:
u_r_state = classify_uncertainty(st.last_u_r, st.calib_r_samples)
u_f_state = classify_uncertainty(st.last_u_f, st.calib_f_samples)
u_s_state = classify_uncertainty(st.last_u_s, st.calib_s_samples)
te = float(getattr(st, "last_true_error", float("nan")))
te_str = "nan" if not math.isfinite(te) else f"{te:.3f}"
cum = float(getattr(st, "cum_true_error", 0.0))
n_te = int(getattr(st, "n_true_error", 0))
mean_te = (cum / n_te) if n_te else None
status = {
"type": "status",
"task": st.task,
"paused": bool(st.paused),
"act_dim": int(st.act_dim),
"step": int(st.step),
"u_r": float(st.last_u_r),
"u_f": float(st.last_u_f),
"u_s": float(st.last_u_s),
"u_r_state": u_r_state,
"u_f_state": u_f_state,
"u_s_state": u_s_state,
"true_error": te if math.isfinite(te) else None,
"cum_error": cum,
"mean_error": mean_te,
"text": (
f"step={st.step} | "
f"err={te_str} | "
f"cum={cum:.2f} | "
f"u_r={st.last_u_r:.3f} u_f={st.last_u_f:.3f} u_s={st.last_u_s:.3f}{u_suffix} | "
f"r={st.last_reward_pred:+.2f} R={st.cum_reward:+.2f}"
),
}
return jpeg, status
def build_parser() -> argparse.ArgumentParser:
p = _base_parser()
# Retarget help defaults for this entrypoint.
for action in p._actions:
if action.dest == "html":
action.default = "interactive_uncertainty_corr.html"
if action.dest == "port":
action.default = 7862
if action.dest == "uncertainty_overlay":
# Always on for this tool; keep flag for CLI compatibility.
pass
return p
def main():
args = build_parser().parse_args()
# Correlation UI always wants the three predictors.
args.uncertainty_overlay = True
server = CorrInteractiveServer(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] corr UI on http://{args.host}:{args.port} "
f"(task={args.task}; true_error = RMS(z_pred - z_env))"
)
web.run_app(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()