lap-steering-demo / lap_compute.py
Jayadev Billa
Initial commit: LAP entity redirect demo
13b3e15
Raw
History Blame Contribute Delete
3.4 kB
"""
A_lin: the Linear Accessibility Profile metric from the LAP paper.
A_lin(l) = fraction of prompts where argmax(W_U · LayerNorm(h_l)) == target token.
This is the logit lens applied to intermediate hidden states — no training required.
It measures how "output-aligned" the concept is at each layer.
"""
import numpy as np
import torch
from typing import Dict, List, Tuple
from steering import get_final_norm, get_lm_head, get_layers, resolve_target_token
def compute_alin(
model,
tokenizer,
prompts: List[str],
target_token: str,
device: str = "cpu",
) -> List[float]:
"""
Compute logit-lens accuracy at every layer for `prompts` with `target_token`.
For each prompt and each layer l:
- Take hidden state h_l at the last token position
- Apply final LayerNorm + LM head (the model's own unembedding projection)
- Check if argmax == target_token_id
A_lin(l) = fraction of prompts that pass the check.
This matches the paper's Eq. (1) exactly. No training required.
"""
target_id = resolve_target_token(tokenizer, target_token)
print(f" Target '{target_token}' → token id {target_id} "
f"('{tokenizer.decode([target_id])}')")
final_norm = get_final_norm(model)
lm_head = get_lm_head(model)
layers = get_layers(model)
n_layers = len(layers)
hits = [0] * n_layers
model.eval()
with torch.no_grad():
for prompt in prompts:
buf: Dict[int, torch.Tensor] = {}
def make_hook(idx):
def hook(module, inp, out):
h = out[0] if isinstance(out, tuple) else out
buf[idx] = h[0, -1, :].detach()
return hook
handles = [layers[i].register_forward_hook(make_hook(i)) for i in range(n_layers)]
try:
inputs = tokenizer(
prompt, return_tensors="pt", truncation=True, max_length=128
).to(device)
model(**inputs)
finally:
for h in handles:
h.remove()
for i in range(n_layers):
h = buf[i].to(device)
h_normed = final_norm(h.unsqueeze(0)).squeeze(0)
logits = lm_head(h_normed)
if logits.argmax().item() == target_id:
hits[i] += 1
return [h / len(prompts) for h in hits]
# ---------------------------------------------------------------------------
# Layer selection
# ---------------------------------------------------------------------------
def select_optimal_layer(alin: List[float]) -> int:
"""
Layer with highest A_lin, excluding the final layer.
The last layer sits immediately before the LM head; A_lin there is
trivially high (it's what the model actually predicts) but steering
at that point disrupts generation rather than redirecting it.
"""
n = len(alin)
candidates = list(range(n - 1)) # exclude last layer
return candidates[int(np.argmax([alin[i] for i in candidates]))]
def select_comparison_layer(n_layers: int) -> int:
"""
Return the middle layer — the standard practitioner heuristic.
The paper explicitly contrasts LAP-recommended vs middle-layer steering,
showing that the middle layer fails for entity redirect (A_lin = 0 there).
"""
return n_layers // 2