EUniversityProject· about 180 min· server sandbox
Project 3: Dissect a language model
Train a small character-level GPT on a short text, then open it up: tokenization, embeddings, attention maps as images, activations layer by layer.
Theory
A decoder-only transformer: token embedding + position embedding → N blocks (causal self-attention + MLP with residuals and layer norm) → logits. Training = next-character prediction. After training every part can be inspected.
Sub-tasks
- tokenization —
CharTokenizer(text)withencode/decodeandvocab_size. - the attention block —
CausalSelfAttention(d, heads)returning (output, attention weights) andBlock(d, heads). - model + training —
TinyGPT(vocab, d, heads, layers, ctx);train(model, data, steps, lr)→ loss history. - dissect —
attention_map(model, tokens)→ matrix for the first layer/head;save_attention_png(mat, tokenizer, tokens)writes out/attention.png.
Passes when: loss <= 2.2
The starter code
runs in an isolated sandbox on the serverimport torch
import torch.nn as nn
import torch.nn.functional as F
class CharTokenizer:
def __init__(self, text):
# TODO: sorterat teckenförråd; stoi/itos
...
@property
def vocab_size(self):
...
def encode(self, s):
...
def decode(self, ids):
...
class CausalSelfAttention(nn.Module):
def __init__(self, d, heads):
super().__init__()
assert d % heads == 0
self.heads, self.d = heads, d
self.qkv = nn.Linear(d, 3 * d)
self.proj = nn.Linear(d, d)
def forward(self, x):
"""x: (B, T, d). Returnerar (out, weights) där weights: (B, heads, T, T)."""
# TODO: qkv → split → (B, heads, T, hd); scores/sqrt(hd); kausal mask (-inf); softmax; @ v; merge; proj
...
class Block(nn.Module):
def __init__(self, d, heads):
super().__init__()
self.ln1, self.ln2 = nn.LayerNorm(d), nn.LayerNorm(d)
self.attn = CausalSelfAttention(d, heads)
self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
def forward(self, x):
# TODO: pre-norm residualblock; returnera (x, weights)
...
class TinyGPT(nn.Module):
def __init__(self, vocab, d=64, heads=4, layers=2, ctx=64):
super().__init__()
self.ctx = ctx
self.tok = nn.Embedding(vocab, d)
self.pos = nn.Embedding(ctx, d)
self.blocks = nn.ModuleList([Block(d, heads) for _ in range(layers)])
self.ln = nn.LayerNorm(d)
self.head = nn.Linear(d, vocab)
def forward(self, idx):
"""idx: (B, T). Returnerar (logits (B, T, V), lista med attention-vikter per lager)."""
# TODO
...
@torch.no_grad()
def generate(self, idx, n):
for _ in range(n):
logits, _ = self(idx[:, -self.ctx:])
nxt = torch.multinomial(F.softmax(logits[:, -1], -1), 1)
idx = torch.cat([idx, nxt], 1)
return idx
def batches(data, ctx, batch, gen):
ix = torch.randint(0, len(data) - ctx - 1, (batch,), generator=gen)
x = torch.stack([data[i:i + ctx] for i in ix]); y = torch.stack([data[i + 1:i + ctx + 1] for i in ix])
return x, y
def train(model, data, steps=600, lr=3e-3, batch=32, seed=0):
gen = torch.Generator().manual_seed(seed)
opt = torch.optim.AdamW(model.parameters(), lr=lr)
hist = []
# TODO: loop: x, y = batches(...); logits, _ = model(x); loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1)); steg; hist
...
return hist
def attention_map(model, tokens):
# TODO: kör modellen på tokens (1, T); returnera weights[0][0, 0] (lager 0, huvud 0) som (T, T)-tensor
...
def save_attention_png(mat, tokenizer, tokens, path="out/attention.png"):
import os
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
os.makedirs("out", exist_ok=True)
chars = [tokenizer.decode([int(t)]) for t in tokens]
plt.figure(figsize=(6, 5)); plt.imshow(mat.detach().numpy(), cmap="viridis")
plt.xticks(range(len(chars)), chars); plt.yticks(range(len(chars)), chars); plt.colorbar(); plt.tight_layout(); plt.savefig(path); plt.close()
return path
You write the code; tests you cannot see decide whether it holds up. Create a free account to run the lab.
Try the diagnosticCreate a free accountExpected results
Loss from ≈ 3.5 to ≤ 2.2 in 600 steps (CPU, ~1 min); attention.png shows clear diagonal/local structure; generated text has word-like sequences.
Common mistakes
- Causal mask missing → the model cheats and the loss becomes implausibly low.
- Position embedding forgotten → the model cannot tell order apart.
- Loss on the wrong shape: logits should be (B·T, V).