"""E2 -- gradient alignment probe (exp7 explainability battery, sec 5). Per refinement k, cosine between the controller's emitted update ``alpha_k * Delta Y_j`` and the true descent direction ``-grad_{Y_j} sum_i d_i``, obtained by autograd through the *frozen rollout only* (never through the controller's own network). This directly tests the "amortized descent" story -- does the controller behave like an unrolled gradient step? -- and works unchanged on both the split and fused architectures, since it only calls the controller's public forward pieces (``condition`` / ``initial_plan`` / ``refine`` / ``to_actions``). """ import argparse import json import sys from pathlib import Path import numpy as np import torch import torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from lejepa_control.data import LatentGoalDataset, split_episodes # noqa: E402 from lejepa_control.rollout import goal_distance, rollout_plan # noqa: E402 from lejepa_control.solver import load_controller # noqa: E402 from lejepa_control.world_model import load_lewm # noqa: E402 def parse_args(): p = argparse.ArgumentParser() p.add_argument('--controller', default='data/runs/controller/controller.pt') p.add_argument('--latents', default='data/latents') p.add_argument('--samples', type=int, default=256) p.add_argument('--horizon', type=int, default=5) p.add_argument('--tag', default='main') p.add_argument('--out', default='data/runs/probe_alignment') return p.parse_args() def descent_direction(controller, model, tokens, ctx, past, goal): """``-grad_Y sum_i d_i``: true descent, autograd through the frozen rollout only -- ``tokens`` is detached first, so no gradient reaches the controller's own network here, only the path token -> action -> rollout. """ y = tokens.detach().clone().requires_grad_(True) pred = rollout_plan(model, ctx, past, controller.to_actions(y)) loss = goal_distance(pred, goal).sum() (grad,) = torch.autograd.grad(loss, y) return -grad def alignment_trace(controller, model, ctx, past, goal): """Per-refinement cosine between the emitted update and true descent. Mirrors ``IterativeController.forward``'s loop exactly (same ``cond``, same ``tokens`` recursion), but at each step also scores the emitted update against the gradient-through-rollout direction at that same ``Y_k``. Returns a ``(K, B, H)`` cosine tensor. """ with torch.no_grad(): cond = controller.condition(ctx, goal) tokens = controller.initial_plan(cond) cosines = [] for k in range(controller.refinements): with torch.no_grad(): actions = controller.to_actions(tokens) pred = rollout_plan(model, ctx, past, actions) true_dir = descent_direction(controller, model, tokens, ctx, past, goal) with torch.no_grad(): delta = controller.refine(tokens, cond, pred, goal) idx = min(k, controller.step_logit.numel() - 1) step = torch.sigmoid(controller.step_logit[idx]) emitted = step * delta # (B, H, W): the actual applied update cos = F.cosine_similarity(emitted, true_dir, dim=-1) # (B, H) cosines.append(cos) tokens = tokens + emitted return torch.stack(cosines) # (K, B, H) def main(): args = parse_args() device = 'cuda' if torch.cuda.is_available() else 'cpu' torch.manual_seed(0) stats = json.loads((Path(args.latents) / 'stats.json').read_text()) model = load_lewm(device=device) controller, ckpt = load_controller( args.controller, latent_dim=stats['latent_dim'], device=device ) controller.eval() print(f'controller step {ckpt["step"]}, K={controller.refinements}, ' f'fused={controller.fused}, no_latent_proj={controller.no_latent_proj}') _, val_eps = split_episodes(stats['n_episodes']) val = LatentGoalDataset( args.latents, max_offset=5, episodes=val_eps, horizon=args.horizon ) idx = np.random.default_rng(0).choice(len(val), args.samples, replace=False) batch = { k: torch.stack([val[int(i)][k] for i in idx]).to(device) for k in ('context', 'past_actions', 'goal') } cos = alignment_trace( controller, model, batch['context'], batch['past_actions'], batch['goal'] ) # (K, B, H) per_k = cos.mean(dim=(1, 2)) per_k_std = cos.std(dim=(1, 2)) print('\n=== E2: gradient alignment (emitted update vs true descent) ===') print(f'{"k":>3}{"mean cos":>11}{"std":>9}') for k in range(cos.size(0)): print(f'{k:>3}{per_k[k].item():>11.4f}{per_k_std[k].item():>9.4f}') report = { 'tag': args.tag, 'checkpoint': args.controller, 'step': int(ckpt['step']), 'fused': bool(controller.fused), 'no_latent_proj': bool(controller.no_latent_proj), 'mean_cosine_per_k': [round(float(v), 5) for v in per_k], 'std_cosine_per_k': [round(float(v), 5) for v in per_k_std], } out = Path(args.out) out.mkdir(parents=True, exist_ok=True) with (out / 'probe_alignment.jsonl').open('a') as f: f.write(json.dumps(report) + '\n') print(f'\nwrote {out / "probe_alignment.jsonl"}') if __name__ == '__main__': main()