Transformer v kódu

Celý blok v padesáti řádcích, včetně rozměrů tenzorů.

Co se naučíš: Přečteš si celý transformer v padesáti řádcích včetně rozměrů všech tenzorů.

12 min čtení + cvičeníNavazuje na:🏗️ Anatomie modelu

Nejrychlejší způsob, jak si architekturu ověřit, je přečíst ji jako kód. Celý model se vejde na jednu obrazovku, vážně. Tohle je zjednodušený, ale úplný transformerový blok v PyTorchu: nic tu není vynechané, jen jsou vypuštěné optimalizace. U každého řádku je poznámka o rozměrech.


✂ Referenční kód v PyTorchi, ne spustitelný skript. Slouží ke čtení a k porovnání s tím, co jsi počítal ručně. Spustitelnou verzi bez frameworku najdeš v postav si vlastní model.


Co budeme stavět

tokeny → embedding → N × blok → norm → logits
                      └ RMSNorm → attention (GQA + RoPE) → +residual
                        RMSNorm → SwiGLU FFN            → +residual

Notace rozměrů: B dávka, T délka sekvence, C = d_model.


Normalizace

class RMSNorm(nn.Module):
    def __init__(self, dim, eps=1e-6):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(dim))   # naučené škálování
        self.eps = eps

    def forward(self, x):                    # x: [B, T, C]
        rms = x.pow(2).mean(-1, keepdim=True).sqrt()   # [B, T, 1]
        return self.weight * (x / (rms + self.eps))    # [B, T, C]

Žádné odečítání průměru, žádný bias, přesně jak popisuje normalizace a aktivace. Osm řádků.


Attention (s GQA a kauzální maskou)

class Attention(nn.Module):
    def __init__(self, dim, n_heads, n_kv_heads):
        super().__init__()
        self.n_heads, self.n_kv_heads = n_heads, n_kv_heads
        self.head_dim = dim // n_heads
        self.rep = n_heads // n_kv_heads          # kolikrát se K/V zopakuje

        self.wq = nn.Linear(dim, n_heads * self.head_dim, bias=False)
        self.wk = nn.Linear(dim, n_kv_heads * self.head_dim, bias=False)   # menší!
        self.wv = nn.Linear(dim, n_kv_heads * self.head_dim, bias=False)   # menší!
        self.wo = nn.Linear(dim, dim, bias=False)

    def forward(self, x, freqs, cache=None):     # x: [B, T, C]
        B, T, C = x.shape

        q = self.wq(x).view(B, T, self.n_heads,    self.head_dim)  # [B,T,32,128]
        k = self.wk(x).view(B, T, self.n_kv_heads, self.head_dim)  # [B,T, 8,128]
        v = self.wv(x).view(B, T, self.n_kv_heads, self.head_dim)

        q, k = apply_rope(q, freqs), apply_rope(k, freqs)   # pozice = rotace

        if cache is not None:                       # KV cache při generování
            k, v = cache.append(k, v)               # k,v: [B, T_celkem, 8, 128]

        k = k.repeat_interleave(self.rep, dim=2)    # 8 → 32 hlav (GQA)
        v = v.repeat_interleave(self.rep, dim=2)

        q, k, v = (t.transpose(1, 2) for t in (q, k, v))   # [B, hlavy, T, hd]

        scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim)  # [B,32,T,T]
        scores = scores.masked_fill(causal_mask(T), float('-inf'))   # do budoucna ne
        weights = scores.softmax(dim=-1)                             # [B,32,T,T]

        out = weights @ v                          # [B, 32, T, 128]
        out = out.transpose(1, 2).reshape(B, T, C) # spoj hlavy → [B, T, C]
        return self.wo(out)                        # výstupní projekce

Všimni si tří věcí:

  1. wk a wv jsou menší než wq, to je celé GQA. Klíče a hodnoty se pak jen zopakují, aby jich bylo tolik jako dotazů.
  2. repeat_interleave nekopíruje do cache, do cache jde jen osm hlav, kopie vzniká až při výpočtu. Proto je KV cache čtvrtinová.
  3. masked_fill před softmaxem, přesně jak jsme počítali v attention ručně.

Feed-forward se SwiGLU

class FFN(nn.Module):
    def __init__(self, dim, hidden):
        super().__init__()
        self.w1 = nn.Linear(dim, hidden, bias=False)   # brána
        self.w2 = nn.Linear(dim, hidden, bias=False)   # obsah
        self.w3 = nn.Linear(hidden, dim, bias=False)   # zpět dolů

    def forward(self, x):                       # [B, T, C]
        return self.w3(F.silu(self.w1(x)) * self.w2(x))

Tři matice místo dvou, proto se hidden volí kolem ⅔ × 4 × dim, aby počet parametrů seděl. Násobení * je ta „brána“: model se učí, kolik z čeho propustit.


Blok a celý model

class Block(nn.Module):
    def __init__(self, dim, n_heads, n_kv_heads, hidden):
        super().__init__()
        self.attn_norm, self.attn = RMSNorm(dim), Attention(dim, n_heads, n_kv_heads)
        self.ffn_norm,  self.ffn  = RMSNorm(dim), FFN(dim, hidden)

    def forward(self, x, freqs, cache=None):
        x = x + self.attn(self.attn_norm(x), freqs, cache)   # pre-norm + residual
        x = x + self.ffn(self.ffn_norm(x))                   # pre-norm + residual
        return x


class Transformer(nn.Module):
    def __init__(self, vocab, dim, n_layers, n_heads, n_kv_heads, hidden):
        super().__init__()
        self.embed  = nn.Embedding(vocab, dim)
        self.blocks = nn.ModuleList([Block(dim, n_heads, n_kv_heads, hidden)
                                     for _ in range(n_layers)])
        self.norm   = RMSNorm(dim)
        self.head   = nn.Linear(dim, vocab, bias=False)   # unembedding

    def forward(self, ids, cache=None):          # ids: [B, T] celá čísla
        x = self.embed(ids)                      # [B, T, C]
        freqs = rope_freqs(x.shape[1])
        for i, block in enumerate(self.blocks):
            x = block(x, freqs, cache[i] if cache else None)
        x = self.norm(x)
        return self.head(x)                      # [B, T, vocab] = logits

To je celý model. Padesát řádků. Rozdíl mezi tímhle a modelem, který běží v produkci, je v optimalizacích (FlashAttention, kvantizace, paged cache), ne v myšlence.


Generační smyčka

@torch.no_grad()
def generate(model, ids, max_new=100, temperature=0.7, top_p=0.9):
    cache = KVCache(model.n_layers)

    logits = model(ids, cache)[:, -1]        # PREFILL: celý prompt naráz
    for _ in range(max_new):
        probs = softmax(logits / temperature)
        probs = top_p_filter(probs, top_p)   # viz kapitola o vzorkování
        next_id = torch.multinomial(probs, 1)

        if next_id == EOS: break
        ids = torch.cat([ids, next_id], dim=1)

        logits = model(next_id, cache)[:, -1] # DECODE: jen JEDEN token
    return ids

Tady je vidět rozdíl, o kterém mluví inference stack: prefill zpracuje celý prompt jedním voláním, decode pak posílá do modelu vždycky jen jeden nový token, zbytek je v cache. Proto je první token pomalý a další rychlé.


Co se do padesáti řádků nevešlo

Poctivě: tenhle kód je funkčně úplný, ale produkční implementace navíc řeší

  • FlashAttention místo naivního q @ k.T (jinak ti dojde paměť u dlouhých sekvencí),
  • paged KV cache místo prostého připojování,
  • kvantizaci vah a cache,
  • paralelismus přes víc GPU,
  • RoPE scaling pro dlouhý kontext,
  • u MoE navíc router místo jednoho FFN.

Žádná z těch věcí ale nemění to, co jsi právě četl. Mění jen, jak rychle to poběží.


Cvičení

  1. Spočítej z kódu počet parametrů pro dim=4096, n_layers=32, n_heads=32, n_kv_heads=8, hidden=14336, vocab=128256. Porovnej s anatomií modelu.
  2. Co se stane, když z Block.forward odstraníš x + (tedy residual)? Proč by se takový model nedal natrénovat do hloubky?
  3. V generate se při decode posílá do modelu jen jeden token. Proč to funguje, kde jsou předchozí?
Náčrt řešení: rozbal, až si cvičení zkusíš sám
  1. 8,03 miliardy. Na vrstvu: attention 2 × 4096² + 2 × 4096 × 1024 = 41 943 040 (Q a O jsou plné, K a V jen 8 × 128 = 1024 široké kvůli GQA), SwiGLU FFN 3 × 4096 × 14336 = 176 160 768, normalizace 2 × 4096. Vrstva tedy 218 112 000, krát 32 vrstev je 6,98 mld. K tomu embedding a unembedding, každý 128256 × 4096 = 525 336 576. Všimni si poměru: FFN je 81 % parametrů vrstvy, attention jen 19 %.
  2. Zmizí identická cesta od vstupu k výstupu. S residualem je výstup bloku x + f(x), takže gradient teče zpět jednak přes f, jednak přímo. Bez něj musí projít úplně každou maticí a každou nelinearitou, takže se násobí desítkami Jakobiánů za sebou. Součin mnoha čísel menších než jedna jde k nule, součin čísel větších než jedna k nekonečnu. Do hloubky pár vrstev to ještě jde, do osmdesáti ne. Viz normalizace a aktivace.
  3. Jsou v KV cache. Při decode se nový token musí porovnat s klíči všech předchozích, ale ty se nepočítají znovu: jejich K a V leží uložené z minulých kroků. Do modelu proto stačí poslat jediný token, spočítat jeho Q, K a V, K a V přilepit do cache a pozornost počítat proti celé cache. Odtud plyne, že decode je omezený propustností paměti, ne výpočtem.

Shrnutí

  • Celý transformer se vejde na obrazovku: RMSNorm, attention, SwiGLU FFN, residual, opakovat.
  • GQA je vidět v kódu jako menší matice wk/wv a repeat_interleave až při výpočtu.
  • Prefill zpracuje celý prompt naráz, decode posílá jeden token a zbytek bere z KV cache.
  • Produkční implementace přidávají optimalizace, ne novou myšlenku.
V kódu jsou matice wk a wv menší než wq. Co to znamená a co tím model získá?

Je to GQA: klíčů a hodnot se počítá méně hlav než dotazů a před výpočtem se zopakují. Do KV cache se ukládá jen ten menší počet, takže cache je několikanásobně menší při prakticky nezměněné kvalitě.

Proč se při generování posílá do modelu jen jeden token, když má odpověď navazovat na celý prompt?

Protože klíče a hodnoty všech předchozích tokenů jsou uložené v KV cache. Nový token potřebuje spočítat jen svoje Q, K a V a porovnat se s cache, přepočítávat celý kontext by bylo zbytečné.

Co by se stalo, kdybys z bloku odstranil residual spojení (řádek x = x + …)?

Gradient by musel při zpětném průchodu procházet všemi transformacemi a u hluboké sítě by cestou zanikl. Model by se prakticky nenatrénoval. Residual navíc drží průběžnou reprezentaci tokenu, do které jednotlivé vrstvy jen přidávají.