World-Action-Verifier / src /idm_hallucination.py
ffeng1017's picture
WAV interactive demo (ZeroGPU)
23a59ea verified
Raw
History Blame Contribute Delete
15.7 kB
# idm_hallucination.py
"""
Hallucination detection by inverse-dynamics rollout divergence.
A standalone predictor. It deliberately does not extend the u_r / u_f / u_s
scorers in `uncertainty.py`: those score a *single predicted latent* per
candidate action and are label-free at the level of one step, whereas this
method needs a whole observed trajectory to compare against. Their
`(predictions_KN, z_prev_K)` scorer signature cannot express that, so this
lives on its own rather than being bent to fit it.
Method
------
1. Infer the action behind every observed transition with a trained inverse
dynamics model (`idm_tokenizer.py`):
a_hat_t = IDM(z_t, z_{t+1})
This needs no action labels -- only the video -- which is what makes the
method applicable to arbitrary unlabeled trajectories.
2. Re-simulate the trajectory with the world model, open-loop, conditioned on
a_hat. Each predicted latent is fed back in, so errors accumulate the way
they do in a real rollout.
3. Compare the simulated latents against the real ones, step by step:
dist_t = RMS(z_sim_t - z_real_t)
dist_norm_t = dist_t / RMS(z_real_t - z_real_{t-1})
score = mean_t dist_norm_t
A world model that has genuinely learned this region reproduces the real
trajectory under the actions that demonstrably generated it. A hallucinating
one invents its own continuation and drifts away.
Keep the horizon short
----------------------
Measured on walker-run with the released `base` checkpoint, the divergence of
a ground-truth-action rollout versus a random-action rollout is 1.55x at step
2 and decays to ~1.2x by step 6: an open-loop rollout reaches the distance
between two *unrelated* real latents (0.587 there) within about four steps,
after which every action sequence looks equally wrong and further steps only
add noise. Aggregating over a 13-step rollout collapsed that separation to
1.07x. Hence `horizon=4` by default.
Validated over 16 rollouts across 4 tasks (walker-run, cheetah-run,
mw-assembly, pygame-cowboy) with `horizon=3`:
ground-truth actions 0.578 (1.00x)
IDM-inferred actions 0.678 (1.17x)
shuffled GT actions 1.073 (1.86x)
random actions 1.148 (1.99x)
with ground-truth below random in 16/16 rollouts and IDM-inferred below random
in 16/16. The inferred actions land far closer to ground truth than to wrong
actions, which is the point: the label-free signal retains most of what the
labeled one gives you.
Usage
-----
from idm_hallucination import IDMRolloutDivergence
det = IDMRolloutDivergence.from_checkpoints(
tokenizer_ckpt="./checkpoints/base/tokenizer.pt",
dynamics_ckpt="./checkpoints/base/dynamics.pt",
idm_ckpt="./logs/idm_tokenizer_ckpts/base_run3/step_040000.pt",
device=torch.device("cuda"),
)
z = det.encode(obs_u8) # (1, T+1, 3, H, W) uint8 -> packed latents
out = det(z, act_mask=mask, lang_emb=lang)
print(out["score"], out["dist"])
Or from the command line, scoring trajectories of one task:
python idm_hallucination.py --task walker-run --idm_ckpt <path> --n_traj 4
"""
from typing import Any, Dict, Optional
import torch
from model import (
Dynamics, pack_bottleneck_to_spatial, temporal_patchify,
unpack_spatial_to_bottleneck,
)
from train_dynamics import sample_one_timestep_packed
def load_idm_tokenizer_head(ckpt_path: str, *, device, n_latents: int, d_bottleneck: int):
"""Load a trained tokenizer-arm inverse dynamics head (`idm_tokenizer.py`).
The architecture is recovered from the checkpoint's own tensor shapes, so
any `--d_hidden` loads without being told which one was used.
Only the tokenizer arm applies here. It maps a pair of latents
(z_t, z_{t+1}) to an action, which is exactly what this method has at each
step. The backbone arm (`idm_backbone.py`) reads contextual features out of
a Dynamics forward over an observed window, which is a different input
contract.
"""
from idm_tokenizer import InverseDynamicsHeadTokenizer
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
sd = ck["model"] if isinstance(ck, dict) and "model" in ck else ck
d_hidden, d_in3 = sd["in_proj.weight"].shape
d_in = d_in3 // 3
expected = n_latents * d_bottleneck
if d_in != expected:
raise ValueError(
f"IDM checkpoint expects flattened latents of {d_in}, but the tokenizer "
f"gives n_latents*d_bottleneck = {expected}. Checkpoint/tokenizer mismatch."
)
head = InverseDynamicsHeadTokenizer(
n_latents=n_latents, d_bottleneck=d_bottleneck, d_hidden=int(d_hidden),
)
head.load_state_dict(sd)
head.to(device).eval()
for p in head.parameters():
p.requires_grad_(False)
meta = {"step": (ck.get("step") if isinstance(ck, dict) else None),
"d_hidden": int(d_hidden), "d_in": int(d_in)}
return head, meta
class IDMRolloutDivergence:
"""Hallucination score from inverse-dynamics rollout divergence.
Holds the frozen world model + inverse dynamics model and the rollout
configuration; call the instance on a trajectory's packed latents.
"""
def __init__(
self,
dyn: Dynamics,
idm,
*,
k_max: int,
sched: Dict[str, Any],
packing_factor: int,
encoder=None,
patch: Optional[int] = None,
n_latents: Optional[int] = None,
n_context: int = 8,
horizon: Optional[int] = 4,
max_ctx: int = 8,
tau_ctx: float = 0.0,
motion_eps: float = 1e-3,
use_kv_cache: bool = False,
):
self.dyn = dyn
self.idm = idm
self.k_max = int(k_max)
self.sched = sched
self.packing_factor = int(packing_factor)
self.encoder = encoder
self.patch = patch
self.n_latents = n_latents
self.n_context = int(n_context)
self.horizon = horizon
self.max_ctx = int(max_ctx)
self.tau_ctx = float(tau_ctx)
self.motion_eps = float(motion_eps)
# The rollout is the dominant cost in a live demo: H steps x 2 action
# sources, each re-attending over the whole context. Caching the
# context K,V is an exact optimisation -- same numbers, less compute.
self.use_kv_cache = bool(use_kv_cache)
@classmethod
def from_checkpoints(
cls,
*,
tokenizer_ckpt: str,
dynamics_ckpt: str,
idm_ckpt: str,
device,
packing_factor: int = 2,
schedule: str = "shortcut",
eval_d: float = 0.25,
**kwargs,
) -> "IDMRolloutDivergence":
"""Build everything from checkpoint paths (tokenizer, dynamics, IDM)."""
from train_dynamics import load_frozen_tokenizer_from_pt_ckpt, make_tau_schedule
from interactive import load_dynamics_from_ckpt
enc, _dec, tok_args = load_frozen_tokenizer_from_pt_ckpt(tokenizer_ckpt, device=device)
patch = int(tok_args.get("patch", 14))
n_latents = int(tok_args.get("n_latents", 64))
d_bottleneck = int(tok_args.get("d_bottleneck", 64))
dyn, _rew, _pol, dyn_meta = load_dynamics_from_ckpt(
dynamics_ckpt, device=device, d_bottleneck=d_bottleneck,
n_latents=n_latents, packing_factor=packing_factor,
)
idm, _meta = load_idm_tokenizer_head(
idm_ckpt, device=device, n_latents=n_latents, d_bottleneck=d_bottleneck,
)
sched = make_tau_schedule(
k_max=dyn_meta["k_max"], schedule=schedule,
d=(eval_d if schedule == "shortcut" else None),
)
return cls(
dyn, idm, k_max=dyn_meta["k_max"], sched=sched,
packing_factor=packing_factor, encoder=enc, patch=patch,
n_latents=n_latents, **kwargs,
)
@torch.no_grad()
def encode(self, obs_u8: torch.Tensor) -> torch.Tensor:
"""(B, T+1, 3, H, W) uint8 -> packed latents (B, T+1, Sz, Dz)."""
if self.encoder is None or self.patch is None or self.n_latents is None:
raise RuntimeError("encode() needs an encoder; build via from_checkpoints().")
z = self.encoder(temporal_patchify(obs_u8.float() / 255.0, self.patch))[0]
return pack_bottleneck_to_spatial(
z, n_spatial=self.n_latents // self.packing_factor, k=self.packing_factor,
)
@torch.no_grad()
def infer_actions(self, z_real_packed: torch.Tensor) -> torch.Tensor:
"""Packed latents (1, T+1, Sz, Dz) -> inferred actions (1, T, A).
`a_hat[:, t]` is the action for the transition z_t -> z_{t+1}.
"""
dtype = next(self.idm.parameters()).dtype
zu = unpack_spatial_to_bottleneck(z_real_packed, k=self.packing_factor)
return self.idm(zu[:, :-1].to(dtype), zu[:, 1:].to(dtype)).float()
@torch.no_grad()
def __call__(
self,
z_real_packed: torch.Tensor, # (1, T+1, Sz, Dz)
*,
act_mask: Optional[torch.Tensor] = None,
lang_emb: Optional[torch.Tensor] = None,
actions_override: Optional[torch.Tensor] = None, # (1, T, A), bypass the IDM
n_context: Optional[int] = None,
horizon: Optional[int] = -1, # -1 = use the instance default
) -> Dict[str, torch.Tensor]:
"""Score one trajectory.
`actions_override` runs the same rollout under known or deliberately
perturbed actions instead of inferred ones -- the control used to
validate the signal.
Returns:
"dist": (H,) RMS(z_sim_t - z_real_t) per simulated step
"dist_norm": (H,) same, divided by the real per-step motion
"motion": (H,) RMS(z_real_t - z_real_{t-1})
"actions_hat": (T, A) the inferred (or overridden) actions
"score": scalar mean of dist_norm -- the hallucination score
"""
assert z_real_packed.dim() == 4 and z_real_packed.shape[0] == 1, \
"z_real_packed must be (1, T+1, Sz, Dz)"
device = z_real_packed.device
dtype = next(self.dyn.parameters()).dtype
Tp1 = z_real_packed.shape[1]
T = Tp1 - 1
assert T >= 1, "need at least one transition"
n_ctx = self.n_context if n_context is None else int(n_context)
n_ctx = max(1, min(n_ctx, T))
h = self.horizon if horizon == -1 else horizon
# 1. Infer the action sequence from the real transitions.
if actions_override is not None:
a_hat = actions_override.to(device).float()
else:
a_hat = self.infer_actions(z_real_packed)
A = a_hat.shape[-1]
# led-to convention: actions_led_to[i] is the action that led to frame i,
# so a_hat_t (for z_t -> z_{t+1}) sits at index t+1 and index 0 is zero.
actions_led_to = torch.zeros((1, Tp1, A), device=device, dtype=torch.float32)
actions_led_to[:, 1:] = a_hat
z_real = z_real_packed.float()
history = [z_real[:, i].to(dtype) for i in range(n_ctx)]
last = Tp1 if h is None else min(Tp1, n_ctx + int(h))
dists, norms, motions = [], [], []
for m in range(n_ctx, last):
z_win = history[-self.max_ctx:] if (self.max_ctx > 0 and len(history) > self.max_ctx) else history
ctx_len = len(z_win)
past = torch.stack(z_win, dim=1) # (1, ctx_len, Sz, Dz)
acts = torch.cat(
[actions_led_to[:, m - ctx_len:m], actions_led_to[:, m:m + 1]], dim=1,
).to(dtype) # (1, ctx_len+1, A)
z_next = sample_one_timestep_packed(
self.dyn,
past_packed=past,
k_max=self.k_max, sched=self.sched,
actions=acts, act_mask=act_mask, tau_ctx=self.tau_ctx,
lang_emb=lang_emb, use_kv_cache=self.use_kv_cache,
) # (1, Sz, Dz)
z_sim = z_next.float()[0]
d = (z_sim - z_real[0, m]).pow(2).mean().sqrt()
mo = (z_real[0, m] - z_real[0, m - 1]).pow(2).mean().sqrt()
dists.append(d)
motions.append(mo)
norms.append(d / mo.clamp(min=self.motion_eps))
history.append(z_next.to(dtype)[0].unsqueeze(0)) # open-loop feedback
dist_norm = torch.stack(norms)
return {
"dist": torch.stack(dists),
"dist_norm": dist_norm,
"motion": torch.stack(motions),
"actions_hat": a_hat[0],
"score": dist_norm.mean(),
}
if __name__ == "__main__":
import argparse
from wm_dataset import WMDataset, collate_batch
p = argparse.ArgumentParser(description=__doc__.split("\n")[1])
p.add_argument("--task", type=str, default="walker-run")
p.add_argument("--tokenizer_ckpt", type=str, default="./checkpoints/base/tokenizer.pt")
p.add_argument("--dynamics_ckpt", type=str, default="./checkpoints/base/dynamics.pt")
p.add_argument("--idm_ckpt", type=str,
default="./logs/idm_tokenizer_ckpts/base_run3/step_040000.pt")
p.add_argument("--data_dir", type=str, default="/data5/expert")
p.add_argument("--frame_dir", type=str, default="/data5/expert-shards")
p.add_argument("--tasks_json", type=str, default="../tasks.json")
p.add_argument("--seq_len", type=int, default=16)
p.add_argument("--n_traj", type=int, default=4)
p.add_argument("--n_context", type=int, default=8)
p.add_argument("--horizon", type=int, default=3)
p.add_argument("--packing_factor", type=int, default=2)
p.add_argument("--controls", action="store_true",
help="also score ground-truth / shuffled / random actions as controls")
args = p.parse_args()
dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
det = IDMRolloutDivergence.from_checkpoints(
tokenizer_ckpt=args.tokenizer_ckpt, dynamics_ckpt=args.dynamics_ckpt,
idm_ckpt=args.idm_ckpt, device=dev, packing_factor=args.packing_factor,
n_context=args.n_context, horizon=args.horizon,
)
lang_dim = det.dyn.lang_dim
ds = WMDataset(
data_dir=[args.data_dir], frames_dir=[args.frame_dir], seq_len=args.seq_len,
img_size=224, action_dim=16, lang_dim=lang_dim, tasks_json=args.tasks_json,
tasks=[args.task], verbose=False, cache_mb=512, ddp_partition=False,
iid_sampling=False, strict_tasks=False,
)
hdr = f"{'traj':>4s} {'score':>9s}"
if args.controls:
hdr += f" {'GT':>9s} {'shuffled':>9s} {'random':>9s}"
print(f"task={args.task} n_context={args.n_context} horizon={args.horizon}")
print(hdr)
print("-" * len(hdr))
for j in range(args.n_traj):
b = collate_batch([ds[(len(ds) // (args.n_traj + 1)) * (j + 1)]])
obs, am = b["obs"].to(dev), b["act_mask"].to(dev)
a_gt = (b["act"].to(dev).clamp(-1, 1)) * am
lang = b["lang_emb"].to(dev) if lang_dim > 0 else None
z = det.encode(obs)
kw = dict(act_mask=am[0, 0], lang_emb=lang)
row = f"{j:>4d} {float(det(z, **kw)['score']):>9.4f}"
if args.controls:
torch.manual_seed(j)
perm = torch.randperm(a_gt.shape[1], device=dev)
for a in (a_gt, a_gt[:, perm], (torch.rand_like(a_gt) * 2 - 1) * am):
row += f" {float(det(z, actions_override=a, **kw)['score']):>9.4f}"
print(row)