354 lines
13 KiB
Python
354 lines
13 KiB
Python
import hashlib
|
|
import itertools
|
|
import json
|
|
import threading
|
|
from collections import OrderedDict
|
|
|
|
import torch
|
|
from jlens.lens import JacobianLens
|
|
|
|
import config
|
|
|
|
GEN_STORE_MAX = 4
|
|
|
|
MASKS_DIR = config.DATA_DIR / "masks"
|
|
|
|
# Last range of layers captured per lens: {lens key: [layers]}.
|
|
# Avoids re-entering the range on every reload (user request).
|
|
LENS_PREFS_PATH = config.DATA_DIR / "lens_prefs.json"
|
|
|
|
|
|
def _load_lens_prefs():
|
|
try:
|
|
return json.loads(LENS_PREFS_PATH.read_text(encoding="utf-8"))
|
|
except Exception:
|
|
return {}
|
|
|
|
|
|
def _save_lens_pref(key, layers):
|
|
prefs = _load_lens_prefs()
|
|
prefs[key] = [int(l) for l in layers]
|
|
LENS_PREFS_PATH.parent.mkdir(parents=True, exist_ok=True)
|
|
LENS_PREFS_PATH.write_text(json.dumps(prefs, indent=1), encoding="utf-8")
|
|
|
|
|
|
def _lens_pref_key(source):
|
|
if source.get("path"):
|
|
return f"path:{source['path']}"
|
|
return f"hub:{source['repo_id']}:{source['filename']}@{source.get('revision') or 'main'}"
|
|
|
|
|
|
class ActivationCatcher:
|
|
def __init__(self, layers, indices):
|
|
self.acts = {}
|
|
self._handles = [
|
|
layers[i].register_forward_hook(self._make(i)) for i in indices
|
|
]
|
|
|
|
def _make(self, index):
|
|
def hook(module, inputs, output):
|
|
tensor = output[0] if isinstance(output, tuple) else output
|
|
self.acts[index] = tensor.detach()
|
|
|
|
return hook
|
|
|
|
def close(self):
|
|
for handle in self._handles:
|
|
handle.remove()
|
|
self._handles = []
|
|
|
|
|
|
def _vocab_fingerprint(tokenizer):
|
|
payload = json.dumps(sorted(tokenizer.get_vocab().items()), ensure_ascii=False)
|
|
return hashlib.sha1(payload.encode()).hexdigest()[:16]
|
|
|
|
|
|
def _wordlike(raw):
|
|
s = raw.strip()
|
|
if len(s) < 1 or "<|" in s or (s.startswith("<") and s.endswith(">")):
|
|
return False
|
|
if s.isascii():
|
|
return (
|
|
raw.startswith(" ")
|
|
and len(s) > 2
|
|
and s[0].isalpha()
|
|
and all(c.isalpha() or c in "'-" for c in s)
|
|
)
|
|
return all(ch.isalnum() for ch in s)
|
|
|
|
|
|
def display_token_mask(tokenizer, vocab_size):
|
|
MASKS_DIR.mkdir(parents=True, exist_ok=True)
|
|
path = MASKS_DIR / f"{_vocab_fingerprint(tokenizer)}_{vocab_size}.pt"
|
|
if path.exists():
|
|
return torch.load(path, weights_only=True)
|
|
mask = torch.zeros(vocab_size, dtype=torch.bool)
|
|
n_decodable = min(vocab_size, len(tokenizer))
|
|
decoded = tokenizer.batch_decode(
|
|
[[tid] for tid in range(n_decodable)], clean_up_tokenization_spaces=False
|
|
)
|
|
for tid, raw in enumerate(decoded):
|
|
mask[tid] = _wordlike(raw)
|
|
torch.save(mask, path)
|
|
return mask
|
|
|
|
|
|
class LensManager:
|
|
def __init__(self):
|
|
self._lock = threading.Lock()
|
|
self.lens = None
|
|
self.meta = None
|
|
self.layers = []
|
|
self.k = 8
|
|
self.mask = None
|
|
self._J = None
|
|
self._tok_strs = {}
|
|
self.gen_store = OrderedDict()
|
|
self._gen_counter = itertools.count(1)
|
|
self._pref_key = None
|
|
|
|
def load(self, model_manager, *, repo_id=None, filename="lens.pt", revision=None,
|
|
path=None, layers=None, k=8):
|
|
with self._lock:
|
|
if model_manager.hf_model is None:
|
|
raise ValueError("load a model first")
|
|
if path:
|
|
lens = JacobianLens.from_pretrained(path)
|
|
source = {"path": path, "repo_id": None, "filename": None, "revision": None}
|
|
else:
|
|
lens = JacobianLens.from_pretrained(
|
|
repo_id, filename=filename, revision=revision
|
|
)
|
|
source = {"path": None, "repo_id": repo_id, "filename": filename, "revision": revision}
|
|
|
|
model_meta = model_manager.meta
|
|
if lens.d_model != model_meta["d_model"]:
|
|
raise ValueError(
|
|
f"lens d_model ({lens.d_model}) != model ({model_meta['d_model']})"
|
|
)
|
|
n_layers = model_meta["n_layers"]
|
|
fitted = lens.source_layers
|
|
if fitted[-1] >= n_layers:
|
|
raise ValueError(
|
|
f"the lens covers layer {fitted[-1]}, outside a model with {n_layers} layers"
|
|
)
|
|
pref_key = _lens_pref_key(source)
|
|
if layers:
|
|
tapped = sorted(set(layers) & set(fitted))
|
|
if not tapped:
|
|
raise ValueError(
|
|
f"no requested layer is fitted (fitted: {fitted[0]}..{fitted[-1]})"
|
|
)
|
|
else:
|
|
# last range used for THIS lens, otherwise all the fitted layers
|
|
# (= the selection made when the fit was created; max range for a
|
|
# downloaded lens)
|
|
saved = _load_lens_prefs().get(pref_key)
|
|
tapped = (sorted(set(saved) & set(fitted)) if saved else None) or list(fitted)
|
|
_save_lens_pref(pref_key, tapped)
|
|
self._pref_key = pref_key
|
|
|
|
device = model_manager.jl.input_device
|
|
stacked = torch.stack([lens.jacobians[l].float() for l in tapped]).to(device)
|
|
tokenizer = model_manager.tokenizer
|
|
vocab_size = model_manager.hf_model.get_output_embeddings().weight.shape[0]
|
|
mask = display_token_mask(tokenizer, vocab_size).to(device)
|
|
|
|
warnings = []
|
|
if model_meta.get("quant"):
|
|
warnings.append(
|
|
f"model loaded in {model_meta['quant']}: the lens was probably "
|
|
"fitted on the unquantized weights, the readouts may drift"
|
|
)
|
|
if model_meta["model_id"].startswith("local/"):
|
|
warnings.append(
|
|
"local model: cannot verify that the lens matches these exact weights"
|
|
)
|
|
|
|
self.lens = lens
|
|
self.layers = tapped
|
|
self.k = int(k)
|
|
self.mask = mask
|
|
self._J = stacked
|
|
self._tok_strs = {}
|
|
self.meta = {
|
|
**source,
|
|
"model_id": model_meta["model_id"],
|
|
"model_revision": model_meta.get("revision"),
|
|
"d_model": lens.d_model,
|
|
"n_prompts": lens.n_prompts,
|
|
"fitted_layers": [int(fitted[0]), int(fitted[-1])],
|
|
"fitted_layers_all": [int(l) for l in fitted],
|
|
"tapped_layers": [int(l) for l in tapped],
|
|
"k": self.k,
|
|
"warnings": warnings,
|
|
}
|
|
return self.meta
|
|
|
|
def set_layers(self, model_manager, layers, k=None):
|
|
with self._lock:
|
|
if self.lens is None:
|
|
raise ValueError("no lens loaded")
|
|
fitted = self.lens.source_layers
|
|
tapped = sorted(set(layers) & set(fitted))
|
|
if not tapped:
|
|
raise ValueError(
|
|
f"no requested layer is fitted (fitted: {fitted[0]}..{fitted[-1]})"
|
|
)
|
|
device = model_manager.jl.input_device
|
|
self.layers = tapped
|
|
self._J = torch.stack(
|
|
[self.lens.jacobians[l].float() for l in tapped]
|
|
).to(device)
|
|
if k:
|
|
self.k = int(k)
|
|
self.meta = dict(self.meta, tapped_layers=[int(l) for l in tapped], k=self.k)
|
|
if getattr(self, "_pref_key", None):
|
|
_save_lens_pref(self._pref_key, tapped)
|
|
return self.meta
|
|
|
|
def unload(self):
|
|
with self._lock:
|
|
self.lens = None
|
|
self.meta = None
|
|
self.layers = []
|
|
self.mask = None
|
|
self._J = None
|
|
self._tok_strs = {}
|
|
self.gen_store.clear()
|
|
torch.cuda.empty_cache()
|
|
return {"unloaded": True}
|
|
|
|
def start_gen(self):
|
|
gen_id = next(self._gen_counter)
|
|
self.gen_store[gen_id] = {
|
|
"layers": list(self.layers),
|
|
"residuals": {l: [] for l in self.layers},
|
|
"positions": [],
|
|
"token_ids": [],
|
|
"phases": [],
|
|
}
|
|
while len(self.gen_store) > GEN_STORE_MAX:
|
|
self.gen_store.popitem(last=False)
|
|
return gen_id
|
|
|
|
@torch.no_grad()
|
|
def pin_ranks(self, gen_id, token_ids, jl, chunk=32):
|
|
store = self.gen_store.get(gen_id)
|
|
if store is None:
|
|
raise ValueError("unknown generation (residual store expired)")
|
|
layers = store["layers"]
|
|
device = self._J.device
|
|
tids = torch.tensor(token_ids, dtype=torch.long, device=device)
|
|
pins = {
|
|
int(t): {"ranks": [], "p": []} for t in token_ids
|
|
}
|
|
for layer in layers:
|
|
residuals = torch.cat(store["residuals"][layer]).to(device).float()
|
|
J = self.lens.jacobians[layer].float().to(device)
|
|
layer_ranks = {int(t): [] for t in token_ids}
|
|
layer_p = {int(t): [] for t in token_ids}
|
|
for start in range(0, residuals.shape[0], chunk):
|
|
h = residuals[start : start + chunk]
|
|
logits = jl.unembed(h @ J.T).float()
|
|
probs = torch.softmax(logits, -1)
|
|
sel = logits[:, tids]
|
|
rank = (logits.unsqueeze(-1) > sel.unsqueeze(1)).sum(1)
|
|
p_sel = probs[:, tids]
|
|
rank_l, p_l = rank.tolist(), p_sel.tolist()
|
|
for ti, t in enumerate(token_ids):
|
|
layer_ranks[int(t)].extend(row[ti] for row in rank_l)
|
|
layer_p[int(t)].extend(round(row[ti], 6) for row in p_l)
|
|
for t in token_ids:
|
|
pins[int(t)]["ranks"].append(layer_ranks[int(t)])
|
|
pins[int(t)]["p"].append(layer_p[int(t)])
|
|
return {
|
|
"gen_id": gen_id,
|
|
"layers": [int(l) for l in layers],
|
|
"positions": store["positions"],
|
|
"phases": store["phases"],
|
|
"tokens": self._strs(jl.tokenizer, store["token_ids"]),
|
|
"pins": pins,
|
|
}
|
|
|
|
def _strs(self, tokenizer, ids):
|
|
out = []
|
|
for tid in ids:
|
|
s = self._tok_strs.get(tid)
|
|
if s is None:
|
|
s = tokenizer.decode([tid], clean_up_tokenization_spaces=False)
|
|
self._tok_strs[tid] = s
|
|
out.append(s)
|
|
return out
|
|
|
|
@torch.no_grad()
|
|
def compute_frames(self, acts, positions, phase, jl, token_ids, gen_id=None,
|
|
abs_positions=None, chunk=None):
|
|
tokenizer = jl.tokenizer
|
|
if chunk is None:
|
|
chunk = max(1, 96 // max(1, len(self.layers)))
|
|
if abs_positions is None:
|
|
abs_positions = positions
|
|
frames = [
|
|
{
|
|
"type": "frame",
|
|
"phase": phase,
|
|
"pos": int(pos),
|
|
"token_id": int(tid),
|
|
"tok": self._strs(tokenizer, [tid])[0],
|
|
"gen": gen_id,
|
|
"layers": {},
|
|
}
|
|
for pos, tid in zip(abs_positions, token_ids)
|
|
]
|
|
store = self.gen_store.get(gen_id) if gen_id is not None else None
|
|
if store is not None:
|
|
store["positions"].extend(int(p) for p in abs_positions)
|
|
store["token_ids"].extend(int(t) for t in token_ids)
|
|
store["phases"].extend(phase for _ in abs_positions)
|
|
device = self._J.device
|
|
for start in range(0, len(positions), chunk):
|
|
batch_positions = positions[start : start + chunk]
|
|
gathered = []
|
|
for layer in self.layers:
|
|
full = acts[layer][0]
|
|
gathered.append(full[list(batch_positions)].float().to(device))
|
|
h = torch.stack(gathered)
|
|
if store is not None:
|
|
for li, layer in enumerate(self.layers):
|
|
store["residuals"][layer].append(h[li].half().cpu())
|
|
# L2 norm of the residual per layer/position ("Activations" view)
|
|
h_norms = h.norm(dim=-1).tolist()
|
|
transported = torch.einsum("lij,lpj->lpi", self._J, h)
|
|
logits = jl.unembed(transported).float()
|
|
lse = logits.logsumexp(-1, keepdim=True)
|
|
raw_v, raw_ids = logits.topk(self.k)
|
|
raw_p = (raw_v - lse).exp()
|
|
m_v, m_ids = logits.masked_fill(~self.mask, float("-inf")).topk(self.k)
|
|
m_p = (m_v - lse).exp()
|
|
sel = logits.gather(-1, m_ids)
|
|
# rank of each top-k token in the full distribution. We loop over k
|
|
# rather than materializing a boolean [L, P, k, V] (≈760 MB at k=32 /
|
|
# 32 layers → OOM): each iteration only touches [L, P, V].
|
|
m_rank = torch.empty_like(m_ids)
|
|
for ki in range(m_ids.shape[-1]):
|
|
m_rank[..., ki] = (logits > sel[..., ki : ki + 1]).sum(-1)
|
|
del logits
|
|
raw_ids_l, raw_p_l = raw_ids.tolist(), raw_p.tolist()
|
|
m_ids_l, m_p_l, m_rank_l = m_ids.tolist(), m_p.tolist(), m_rank.tolist()
|
|
for li, layer in enumerate(self.layers):
|
|
for pi in range(len(batch_positions)):
|
|
ids = raw_ids_l[li][pi]
|
|
mids = m_ids_l[li][pi]
|
|
frames[start + pi]["layers"][str(layer)] = {
|
|
"ids": ids,
|
|
"p": [round(v, 5) for v in raw_p_l[li][pi]],
|
|
"strs": self._strs(tokenizer, ids),
|
|
"m_ids": mids,
|
|
"m_p": [round(v, 5) for v in m_p_l[li][pi]],
|
|
"m_rank": m_rank_l[li][pi],
|
|
"m_strs": self._strs(tokenizer, mids),
|
|
"h_norm": round(h_norms[li][pi], 2),
|
|
}
|
|
return frames
|