ProCreations/llmviz / code /capture.py
ProCreations's picture
download
raw
6.64 kB
"""Run MiniCPM5-2B for real and record EVERYTHING the visualisation needs.
Records, for every token position (prompt + generated), at every layer:
embedding output, RMSNorm outputs, q/k/v projections, attention weights
(all 16 heads over the whole context), attention output, MLP intermediate
(silu(gate)*up), MLP output, residual stream after the layer, final norm,
and the full next-token probability distribution for every generated token.
"""
import json, os, sys, time
import numpy as np
import torch
MODEL = os.environ.get("MODEL", "openbmb/MiniCPM5-2B")
OUT = os.environ.get("OUT", "acts.npz")
N_NEW = int(os.environ.get("N_NEW", "96"))
SEED = int(os.environ.get("SEED", "0"))
PROMPT = os.environ.get(
"PROMPT",
"In a few sentences, describe what happens inside a language model as it generates each word.",
)
from transformers import AutoModelForCausalLM, AutoTokenizer
from transformers.cache_utils import DynamicCache
if torch.cuda.is_available():
dev, dtype = "cuda", torch.bfloat16
elif torch.backends.mps.is_available():
dev, dtype = "mps", torch.bfloat16
else:
dev, dtype = "cpu", torch.float32
print("device", dev, dtype, flush=True)
tok = AutoTokenizer.from_pretrained(MODEL)
model = AutoModelForCausalLM.from_pretrained(MODEL, dtype=dtype, attn_implementation="eager")
model.to(dev).eval()
cfg = model.config
L, H, KVH, D, I, HD, V = (cfg.num_hidden_layers, cfg.num_attention_heads, cfg.num_key_value_heads,
cfg.hidden_size, cfg.intermediate_size, cfg.head_dim, cfg.vocab_size)
print(f"layers={L} heads={H} kv={KVH} hidden={D} inter={I} head_dim={HD} vocab={V}", flush=True)
messages = [{"role": "user", "content": PROMPT}]
try:
ids = tok.apply_chat_template(messages, add_generation_prompt=True, enable_thinking=False,
return_tensors="pt", return_dict=False)
except TypeError:
ids = tok.apply_chat_template(messages, add_generation_prompt=True, return_tensors="pt")
if isinstance(ids, dict) or hasattr(ids, "input_ids"):
ids = ids["input_ids"]
ids = ids.to(dev)
P = ids.shape[1]
T = P + N_NEW
print("prompt tokens", P, "total", T, flush=True)
# ---- storage (fp16) ------------------------------------------------------
st = {
"emb": np.zeros((T, D), np.float16),
"ln1": np.zeros((L, T, D), np.float16),
"q": np.zeros((L, T, H * HD), np.float16),
"k": np.zeros((L, T, KVH * HD), np.float16),
"v": np.zeros((L, T, KVH * HD), np.float16),
"attn": np.zeros((L, H, T, T), np.float16),
"attn_out": np.zeros((L, T, D), np.float16),
"ln2": np.zeros((L, T, D), np.float16),
"inter": np.zeros((L, T, I), np.float16),
"mlp_out": np.zeros((L, T, D), np.float16),
"resid": np.zeros((L, T, D), np.float16),
"final_norm": np.zeros((T, D), np.float16),
"probs": np.zeros((N_NEW + 1, V), np.float16), # distribution that produced token P+i
}
pos = {"start": 0, "end": P} # rows being written this forward
def put(key, layer, val):
a, b = pos["start"], pos["end"]
v = val.detach().float().cpu().numpy().astype(np.float16)
if layer is None:
st[key][a:b] = v
else:
st[key][layer, a:b] = v
hooks = []
m = model.model
hooks.append(m.embed_tokens.register_forward_hook(lambda mod, i, o: put("emb", None, o[0])))
hooks.append(m.norm.register_forward_hook(lambda mod, i, o: put("final_norm", None, o[0])))
for li, layer in enumerate(m.layers):
def mk(key, l, which="out"):
def h(mod, inp, out):
if which == "out":
t = out[0] if isinstance(out, tuple) else out
put(key, l, t[0])
else:
put(key, l, inp[0][0])
return h
hooks.append(layer.input_layernorm.register_forward_hook(mk("ln1", li)))
hooks.append(layer.self_attn.q_proj.register_forward_hook(mk("q", li)))
hooks.append(layer.self_attn.k_proj.register_forward_hook(mk("k", li)))
hooks.append(layer.self_attn.v_proj.register_forward_hook(mk("v", li)))
hooks.append(layer.self_attn.o_proj.register_forward_hook(mk("attn_out", li)))
hooks.append(layer.post_attention_layernorm.register_forward_hook(mk("ln2", li)))
hooks.append(layer.mlp.down_proj.register_forward_hook(mk("inter", li, "in")))
hooks.append(layer.mlp.down_proj.register_forward_hook(mk("mlp_out", li)))
hooks.append(layer.register_forward_hook(mk("resid", li)))
def attn_hook(mod, inp, out, l=li):
w = out[1]
if w is None:
raise RuntimeError("attention weights not returned; need eager attention")
a, b = pos["start"], pos["end"]
st["attn"][l, :, a:b, :b] = w[0].detach().float().cpu().numpy().astype(np.float16)
hooks.append(layer.self_attn.register_forward_hook(attn_hook))
gen = torch.Generator(device="cpu").manual_seed(SEED)
temperature, top_p = 1.0, 0.95
def sample(logits):
logits = logits.float().cpu() / temperature
probs = torch.softmax(logits, -1)
sp, si = torch.sort(probs, descending=True)
cum = torch.cumsum(sp, 0)
keep = cum - sp < top_p
sp = sp * keep
sp = sp / sp.sum()
j = torch.multinomial(sp, 1, generator=gen).item()
return si[j].item(), probs
tokens = ids[0].tolist()
gen_ids = []
t0 = time.time()
with torch.no_grad():
cache = DynamicCache()
out = model(input_ids=ids, past_key_values=cache, use_cache=True)
for i in range(N_NEW):
nxt, probs = sample(out.logits[0, -1])
st["probs"][i] = probs.numpy().astype(np.float16)
gen_ids.append(nxt)
tokens.append(nxt)
pos["start"], pos["end"] = P + i, P + i + 1
out = model(input_ids=torch.tensor([[nxt]], device=dev), past_key_values=cache, use_cache=True)
if (i + 1) % 16 == 0:
print(f" gen {i+1}/{N_NEW} {time.time()-t0:.1f}s", flush=True)
nxt, probs = sample(out.logits[0, -1])
st["probs"][N_NEW] = probs.numpy().astype(np.float16)
text = tok.decode(gen_ids, skip_special_tokens=False)
print("GENERATED:", repr(text), flush=True)
meta = {
"model": MODEL, "prompt": PROMPT, "prompt_len": P, "n_new": N_NEW, "seed": SEED,
"temperature": temperature, "top_p": top_p,
"layers": L, "heads": H, "kv_heads": KVH, "hidden": D, "inter": I, "head_dim": HD, "vocab": V,
"tokens": tokens, "token_strs": [tok.decode([t]) for t in tokens], "generated_text": text,
"device": dev, "seconds": time.time() - t0,
}
np.savez(OUT, tokens=np.array(tokens, np.int32), meta=json.dumps(meta), **st)
with open(os.path.splitext(OUT)[0] + ".json", "w") as f:
json.dump(meta, f, indent=1)
print("saved", OUT, f"{os.path.getsize(OUT)/1e6:.0f} MB", flush=True)

Xet Storage Details

Size:
6.64 kB
·
Xet hash:
e08222a56434f2d493cae40ddf2288655a894b6d14f3ee99e21a855155ea2c58

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.