| """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
|
| from lejepa_control.rollout import goal_distance, rollout_plan
|
| from lejepa_control.solver import load_controller
|
| from lejepa_control.world_model import load_lewm
|
|
|
|
|
| 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
|
|
|
| cos = F.cosine_similarity(emitted, true_dir, dim=-1)
|
| cosines.append(cos)
|
|
|
| tokens = tokens + emitted
|
|
|
| return torch.stack(cosines)
|
|
|
|
|
| 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']
|
| )
|
|
|
| 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()
|
|
|