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í:
wkawvjsou menší nežwq, to je celé GQA. Klíče a hodnoty se pak jen zopakují, aby jich bylo tolik jako dotazů.repeat_interleavenekopíruje do cache, do cache jde jen osm hlav, kopie vzniká až při výpočtu. Proto je KV cache čtvrtinová.masked_fillpř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í
- 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. - Co se stane, když z
Block.forwardodstraníšx +(tedy residual)? Proč by se takový model nedal natrénovat do hloubky? - V
generatese 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
- 8,03 miliardy. Na vrstvu: attention
2 × 4096² + 2 × 4096 × 1024 = 41 943 040(Q a O jsou plné, K a V jen8 × 128 = 1024široké kvůli GQA), SwiGLU FFN3 × 4096 × 14336 = 176 160 768, normalizace2 × 4096. Vrstva tedy218 112 000, krát 32 vrstev je6,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 %. - 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řesf, 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. - 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/wvarepeat_interleaveaž 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í.
