Leplanner / code /scripts /probe_attention.py
nottygian's picture
Push scripts
dc9f917 verified
Raw
History Blame Contribute Delete
8.11 kB
"""E4 -- attention readout probe (exp7 explainability battery, sec 5).
Temporarily swaps each transformer Block's bound ``forward`` for a version
that also asks its ``MultiheadAttention`` for weights (``need_weights=True``,
which the default ``need_weights=False`` forward never computes), for the
duration of one probe pass, then restores the originals -- so the
controller's real forward path (and every existing checkpoint/test) is
untouched.
Reports mass from plan-slot queries onto four groups -- self, other plan
slots, context frames, goal token -- per layer and per refinement. Works
unchanged on the split architecture (two named nets, ``consequence_net``
then ``refine_net`` per refine() call) and the fused one (one ``net`` per
call), so the "one coherent map vs. smeared across two networks" comparison
in sec 5/E4 is literally the same script pointed at two checkpoints.
"""
import argparse
import json
import sys
import types
from pathlib import Path
import numpy as np
import torch
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 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=128)
p.add_argument('--horizon', type=int, default=5)
p.add_argument('--tag', default='main')
p.add_argument('--out', default='data/runs/probe_attention')
return p.parse_args()
def _capturing_forward(cap_list):
"""A ``Block.forward`` that also records self-attention weights."""
def forward(self, x):
h = self.norm1(x)
attn_out, w = self.attn(
h, h, h, need_weights=True, average_attn_weights=True
)
cap_list.append(w.detach()) # (B, L, L), already averaged over heads
x = x + attn_out
return x + self.mlp(self.norm2(x))
return forward
class AttentionCapture:
"""Patches every ``Block`` under ``root`` to log attention weights in
call order, for the lifetime of the ``with`` block. Blocks that never
run during the wrapped call (e.g. ``consequence_net`` during
``initial_plan``) simply never append -- no separate bookkeeping needed
to know which net actually fired.
"""
def __init__(self, root):
self.blocks = [m for m in root.modules() if type(m).__name__ == 'Block']
self.log = []
self._originals = []
def __enter__(self):
for block in self.blocks:
self._originals.append((block, block.forward))
block.forward = types.MethodType(_capturing_forward(self.log), block)
return self
def __exit__(self, *exc):
for block, orig in self._originals:
block.forward = orig
return False
def mass_by_group(weights, horizon, num_context):
"""``weights``: (B, L, L) self-attention, L = horizon + num_context + 1,
sequence laid out as [plan slots, context frames, goal] (matching
``condition()`` / ``_run_refine`` / ``_run_fused``). Returns group mass
averaged over batch and over the ``horizon`` plan-slot query rows.
"""
H = horizon
rows = weights[:, :H] # (B, H, L): only plan-slot queries
diag = rows[:, torch.arange(H), torch.arange(H)].mean()
plan_total = rows[:, :, :H].sum(-1).mean()
ctx_mass = rows[:, :, H:H + num_context].sum(-1).mean()
goal_mass = rows[:, :, H + num_context].mean()
return {
'self': float(diag),
'other_plan': float(plan_total - diag),
'context': float(ctx_mass),
'goal': float(goal_mass),
}
def label_layers(log, fused, depth):
"""Positionally names each captured layer by which net produced it.
Call order is deterministic from ``controller.py``: ``initial_plan``
only runs the "refine" net; ``refine()`` runs consequence-then-refine
(split) or the single fused net (fused) -- so position alone identifies
the source net, no encoder-identity bookkeeping required.
"""
if fused:
return [('net', i) for i in range(len(log))]
if len(log) == depth: # this was an initial_plan() call
return [('refine_net', i) for i in range(len(log))]
return (
[('consequence_net', i) for i in range(depth)]
+ [('refine_net', i) for i in range(depth)]
)
@torch.no_grad()
def run_probe(controller, model, ctx, past, goal, depth):
"""Replays ``IterativeController.forward``'s loop, capturing attention
inside every net call at every refinement. Returns a list (per
iteration k = 0..K) of ``{(net_name, layer): group_mass_dict}``.
"""
H, N = controller.horizon, controller.num_context
cond = controller.condition(ctx, goal)
per_iter = []
with AttentionCapture(controller) as cap:
tokens = controller.initial_plan(cond)
layers = label_layers(cap.log, controller.fused, depth)
per_iter.append({
layer: mass_by_group(w, H, N) for layer, w in zip(layers, cap.log)
})
for k in range(controller.refinements):
actions = controller.to_actions(tokens)
pred = rollout_plan(model, ctx, past, actions)
with AttentionCapture(controller) as cap:
delta = controller.refine(tokens, cond, pred, goal)
layers = label_layers(cap.log, controller.fused, depth)
per_iter.append({
layer: mass_by_group(w, H, N) for layer, w in zip(layers, cap.log)
})
idx = min(k, controller.step_logit.numel() - 1)
tokens = tokens + torch.sigmoid(controller.step_logit[idx]) * delta
return per_iter
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()
depth = ckpt['args']['depth']
print(f'controller step {ckpt["step"]}, K={controller.refinements}, '
f'fused={controller.fused}, depth={depth}')
_, 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')
}
per_iter = run_probe(
controller, model, batch['context'], batch['past_actions'], batch['goal'],
depth,
)
print('\n=== E4: attention mass from plan-slot queries ===')
print(f'{"k":>3}{"net":>16}{"layer":>6}{"self":>8}{"plan":>8}'
f'{"ctx":>8}{"goal":>8}')
for k, layers in enumerate(per_iter):
for (net, layer), g in layers.items():
print(f'{k:>3}{net:>16}{layer:>6}{g["self"]:>8.3f}'
f'{g["other_plan"]:>8.3f}{g["context"]:>8.3f}{g["goal"]:>8.3f}')
report = {
'tag': args.tag,
'checkpoint': args.controller,
'step': int(ckpt['step']),
'fused': bool(controller.fused),
'per_iteration': [
{f'{net}.{layer}': g for (net, layer), g in layers.items()}
for layers in per_iter
],
}
out = Path(args.out)
out.mkdir(parents=True, exist_ok=True)
with (out / 'probe_attention.jsonl').open('a') as f:
f.write(json.dumps(report) + '\n')
print(f'\nwrote {out / "probe_attention.jsonl"}')
if __name__ == '__main__':
main()