File size: 5,473 Bytes
dc9f917 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 | """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()
|