Leplanner / code /scripts /probe_alignment.py
nottygian's picture
Push scripts
dc9f917 verified
Raw
History Blame Contribute Delete
5.47 kB
"""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()