Skip to content
AI-grafen
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

  1. tokenization — CharTokenizer(text) with encode/decode and vocab_size.
  2. the attention block — CausalSelfAttention(d, heads) returning (output, attention weights) and Block(d, heads).
  3. model + training — TinyGPT(vocab, d, heads, layers, ctx); train(model, data, steps, lr) → loss history.
  4. 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 server
import 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 account

Expected 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).