PidokuInfra

The Transformer

Foundations Intermediate 2h Difficulty 3/5 Topic 09 of 12

Prerequisites 04, 05, 06, 07, 08


1. What is it?#

A transformer is attention and an MLP, stacked with residual connections and normalization, repeated L times.

                    ┌──────────────────────────┐
   tokens → embed → │  × L transformer blocks  │ → norm → lm_head → logits
                    └──────────────────────────┘

   One block:
        x ──┬─► RMSNorm ─► Attention ─┐
            └───────────────(+)◄──────┘
            │
            ├─► RMSNorm ─► FFN ───────┐
            └───────────────(+)◄──────┘

That is the complete architecture of GPT, Llama, Mistral, Qwen, DeepSeek, and essentially every production LLM. The variations between them are small and enumerable.


2. Why does it exist?#

Because it is the first architecture that is simultaneously:

  • Expressive — attention lets any token influence any other.
  • Parallelizable — all positions computed at once during training and prefill.
  • Scalable — performance improves predictably with parameters and data.

Recurrent networks had (1) partially, lacked (2), and scaled poorly. Convolutions had (2) but needed many layers to relate distant tokens. The transformer got all three, and the field converged on it within about three years.


3. Simple analogy#

A committee that meets repeatedly.

Each round has two phases:

  1. Discussion (attention): every member looks at what every other member currently thinks and updates based on who’s relevant to them.
  2. Private reflection (FFN): each member independently processes what they just heard.

After each phase, members add their update to their previous position rather than replacing it (residual connection) and normalize their confidence so no one gets disproportionately loud (normalization).

Repeat 32-80 times. The final positions are the answer.

The “residual stream” framing is more than an analogy: it’s how interpretability researchers actually think about transformers. Each layer reads from the stream, computes something, and adds it back.


4. Tiny example#

A complete, runnable transformer block in ~40 lines:

Go
// transformer.go — a complete decoder-only transformer forward pass, standard library only.
package main

import (
	"fmt"
	"math"
	"math/rand"
)

// Mat is a weight matrix of shape (Out, In), stored row by row.
type Mat struct {
	Out, In int
	W       []float32
}

// Apply computes y = W·x for one vector.
func (m Mat) Apply(x []float32) []float32 {
	y := make([]float32, m.Out)
	for i := range y {
		row := m.W[i*m.In : (i+1)*m.In]
		var s float32
		for j, w := range row {
			s += w * x[j]
		}
		y[i] = s
	}
	return y
}

type Layer struct {
	H, HKV, D                       int // query heads, key/value heads, head dimension
	AttnNorm, FFNNorm               []float32
	Wq, Wk, Wv, Wo, Wgate, Wup, Wdn Mat
}

type Model struct {
	Embed     [][]float32 // (V, d)
	Layers    []Layer
	FinalNorm []float32
	LMHead    Mat // (V, d)
}

func rmsnorm(x, weight []float32) []float32 {
	var ss float32
	for _, v := range x {
		ss += v * v
	}
	inv := float32(1 / math.Sqrt(float64(ss)/float64(len(x))+1e-6))
	y := make([]float32, len(x))
	for i, v := range x {
		y[i] = v * inv * weight[i]
	}
	return y
}

// rope rotates each pair of dimensions of one head's vector by an angle that depends on position.
func rope(x []float32, pos int) {
	D := len(x)
	for i := 0; i < D; i += 2 {
		angle := float64(pos) / math.Pow(10000, float64(i)/float64(D))
		sin, cos := math.Sincos(angle)
		a, b := x[i], x[i+1]
		x[i] = a*float32(cos) - b*float32(sin)
		x[i+1] = a*float32(sin) + b*float32(cos)
	}
}

func silu(x float32) float32 { return x / (1 + float32(math.Exp(float64(-x)))) }

// block runs one transformer layer over a sequence x of shape (S, d).
func block(x [][]float32, p *Layer) [][]float32 {
	S, d := len(x), len(x[0])
	h, hkv, D := p.H, p.HKV, p.D

	// ---- attention ----
	q, k, v := make([][]float32, S), make([][]float32, S), make([][]float32, S)
	for s := range x {
		xn := rmsnorm(x[s], p.AttnNorm)
		q[s], k[s], v[s] = p.Wq.Apply(xn), p.Wk.Apply(xn), p.Wv.Apply(xn) // (h·D), (hkv·D), (hkv·D)
		for i := 0; i < h; i++ {
			rope(q[s][i*D:(i+1)*D], s)
		}
		for i := 0; i < hkv; i++ {
			rope(k[s][i*D:(i+1)*D], s)
		}
	}
	scale := float32(1 / math.Sqrt(float64(D)))
	for s := 0; s < S; s++ {
		out := make([]float32, d)
		for head := 0; head < h; head++ {
			kv := head / (h / hkv) // GQA: several query heads share one key/value head
			qs := q[s][head*D : (head+1)*D]

			// scores against every position t <= s (the causal mask: the future is never read)
			att := make([]float32, s+1)
			maxScore := float32(math.Inf(-1))
			for t := 0; t <= s; t++ {
				kt := k[t][kv*D : (kv+1)*D]
				var dot float32
				for i := range qs {
					dot += qs[i] * kt[i]
				}
				att[t] = dot * scale
				maxScore = max(maxScore, att[t])
			}
			var sum float32
			for t := range att { // softmax, stable form
				att[t] = float32(math.Exp(float64(att[t] - maxScore)))
				sum += att[t]
			}
			for t := range att { // weighted sum of the values
				vt, w := v[t][kv*D:(kv+1)*D], att[t]/sum
				for i := range vt {
					out[head*D+i] += w * vt[i]
				}
			}
		}
		o := p.Wo.Apply(out)
		y := make([]float32, d)
		for i := range y {
			y[i] = x[s][i] + o[i] // residual
		}
		x[s] = y
	}

	// ---- feed-forward ----
	for s := range x {
		xn := rmsnorm(x[s], p.FFNNorm)
		gate, up := p.Wgate.Apply(xn), p.Wup.Apply(xn)
		for i := range gate {
			gate[i] = silu(gate[i]) * up[i]
		}
		dn := p.Wdn.Apply(gate)
		for i := range dn {
			x[s][i] += dn[i] // residual
		}
	}
	return x
}

// forward returns the logits for every position: shape (S, V).
func forward(ids []int, m *Model) [][]float32 {
	x := make([][]float32, len(ids))
	for s, id := range ids {
		x[s] = append([]float32{}, m.Embed[id]...) // (S, d)
	}
	for i := range m.Layers {
		x = block(x, &m.Layers[i])
	}
	logits := make([][]float32, len(ids))
	for s := range x {
		logits[s] = m.LMHead.Apply(rmsnorm(x[s], m.FinalNorm)) // (V,)
	}
	return logits
}

// ---- random weights so the program runs on its own ----

func randMat(out, in int) Mat {
	m := Mat{out, in, make([]float32, out*in)}
	for i := range m.W {
		m.W[i] = float32(rand.NormFloat64()) * 0.05
	}
	return m
}

func ones(n int) []float32 {
	v := make([]float32, n)
	for i := range v {
		v[i] = 1
	}
	return v
}

func main() {
	const V, d, h, hkv, D, ff, L = 100, 64, 8, 2, 8, 128, 2
	m := &Model{FinalNorm: ones(d), LMHead: randMat(V, d)}
	for i := 0; i < V; i++ {
		m.Embed = append(m.Embed, randMat(1, d).W)
	}
	for i := 0; i < L; i++ {
		m.Layers = append(m.Layers, Layer{h, hkv, D, ones(d), ones(d),
			randMat(h*D, d), randMat(hkv*D, d), randMat(hkv*D, d), randMat(d, h*D),
			randMat(ff, d), randMat(ff, d), randMat(d, ff)})
	}
	logits := forward([]int{5, 17, 42, 8}, m)
	fmt.Println("logits shape:", len(logits), "x", len(logits[0])) // 4 x 100
}

That is a Llama block. Everything else is stacking, embedding, and the final head. Project 04 turns this into a working inference engine.


5. Technical explanation#

The full model#

Go
// The whole model is this loop (from the program above):
func forward(ids []int, m *Model) [][]float32 {
	x := embed(ids, m)                  // (S, d)   one table lookup per token
	for i := range m.Layers {
		x = block(x, &m.Layers[i])      // (S, d)   same shape in, same shape out
	}
	return lmHead(finalNorm(x, m), m)   // (S, V)   one score per vocabulary entry
}

Eight lines. Every LLM you have used is this, with a specific set of weights and a few architectural choices.

Pre-norm vs post-norm#

Post-norm (original 2017):  x = Norm(x + Sublayer(x))
Pre-norm (everything since): x = x + Sublayer(Norm(x))

Pre-norm won because it trains stably at depth — the residual path is a clean identity, so signal flows through 80 layers without normalization interfering.

Inference consequence: with pre-norm the residual stream is never normalized, so its magnitude grows through the network. Activations in late layers can be large. This matters for FP16 (range limits) and for activation quantization (Section VII.04).

The architecture variation table#

Everything that differs between modern LLMs:

ChoiceOptionsInference impact
NormalizationLayerNorm / RMSNormminor: kernel cost
Norm positionpre / postactivation magnitudes
Position encodinglearned / sinusoidal / RoPE / ALiBiRoPE constrains prefix caching
AttentionMHA / GQA / MQA / MLAmajor: KV cache size
FFN activationGELU / SiLU / gatedminor
FFN typedense / MoEmajor: active vs total params
Weight tyingyes / nomemory
Biasespresent / absentminor
Attention windowfull / sliding / hybridmajor: KV cache growth
QK normalizationyes / nonumerical stability

The two rows in bold determine your serving economics. When someone says “we’re switching to model X,” those are the first two questions.

Sliding-window attention#

Mistral popularized limiting attention to the last W tokens (e.g. 4096):

Full attention:      token 10000 attends to tokens 0..10000     KV cache grows forever
Sliding window (W):  token 10000 attends to tokens 5905..10000  KV cache CAPPED at W

Enormous serving benefit: KV cache memory is bounded regardless of context length. Cost: the model can only access distant information through the “chain” of intermediate layers (information propagates W tokens per layer, so L layers reach L×W tokens indirectly).

Modern models often interleave: some layers full attention, some sliding window (Gemma 2, Llama 4-era designs). This bounds most of the KV cache while retaining some global attention.

Counting everything#

Parameters ≈ L · (4d² · (adjusted for GQA) + 3·d·d_ff) + 2·V·d
FLOPs/token ≈ 2·P + 4·L·S·d
KV bytes/token ≈ 2 · L · h_kv · d_head · bytes

Apply these three formulas to any model config and you know its serving profile before you download it. Practice until it’s automatic.


6. Under the hood#

Trace a single token through Llama-3-8B during decode with 2,000 tokens of context:

1.  input_ids = [15496]                                   1 token
2.  embed lookup → x (1, 1, 4096)                         8 KB read from a 1 GB table
3.  For each of 32 layers:
      a. RMSNorm            read/write 8 KB
      b. q_proj  (1,4096)@(4096,4096)   33.6 MFLOP, 33.6 MB weight read
      c. k_proj  (1,4096)@(4096,1024)    8.4 MFLOP,  8.4 MB
      d. v_proj                          8.4 MFLOP,  8.4 MB
      e. RoPE on q,k        negligible
      f. append k,v to cache            0.25 MB write
      g. attention vs 2001 cached keys  8.2 MFLOP, 16.4 MB KV read
      h. o_proj                        33.6 MFLOP, 33.6 MB
      i. residual add
      j. RMSNorm
      k. gate,up,down                  352 MFLOP, 352 MB weight read
      l. residual add
4.  Final RMSNorm
5.  lm_head (1,4096)@(4096,128256)     1.05 GFLOP, 1.05 GB weight read
6.  softmax over 128256 logits, sample
──────────────────────────────────────────────────────────────
Total: ~15.6 GFLOPs,  ~16.5 GB read
On an H100: compute 0.016 ms, memory 4.9 ms  →  MEMORY BOUND by 300x

That trace is the whole of Section I made concrete. Note where the bytes are: 16 GB of weights, 0.5 GB of KV. At batch 1 and 2k context, weights dominate utterly. Recompute this trace at batch 64 and 32k context and the balance flips.


7. Performance implications#

  • FFN dominates FLOPs and weight bytes (~65-70%). Quantize it first; MoE it if you’re designing the model.
  • Attention dominates at long context, through KV reads.
  • The LM head is ~7% of decode cost and produces a large tensor.
  • Norms and residuals are ~1% of FLOPs and ~10% of time unfused.
  • Layer count sets kernel-launch count: 80 layers × ~10 kernels = 800 launches per token. At 5 µs each, 4 ms of pure CPU overhead. CUDA graphs eliminate it.

8. Production implications#

  • Read the config before deploying. config.json tells you num_hidden_layers, hidden_size, num_attention_heads, num_key_value_heads, intermediate_size, vocab_size, max_position_embeddings, rope_theta, sliding_window. From those you can compute everything in section 5.
  • Architecture changes break kernels. A new attention variant may not have a fused kernel in your engine yet, silently falling back to a slow path.
  • Sliding window changes your capacity model completely — KV is bounded, so concurrency is much higher.
  • Interleaved attention layers complicate the KV manager: different layers have different cache sizes.

9. Common mistakes#

Assuming all transformers are the same. GQA vs MHA is an 8x difference in your primary capacity constraint.

Ignoring rope_theta. Models fine-tuned for long context change it; using the wrong value produces degraded output at long positions with no error.

Forgetting sliding_window when computing KV memory — you’ll overprovision by 10x (or underprovision if you assume it’s there and it isn’t).

Trusting parameter count as a proxy for serving cost. An MoE with 400B total / 40B active serves like a 40B model for compute and a 400B model for memory.

Not checking whether your engine has an optimized kernel for the architecture.


10. Hands-on exercise#

A. Complete the model. Extend the block in section 4 into a full model: embedding, L layers, final norm, lm_head. Load real Llama weights (a small one) and verify your implementation produces the same logits as Hugging Face’s, to within 1e-3. This is Project 04 and it is the single most valuable exercise in the curriculum.

B. Config archaeology. Download config.json for five different open models. For each, compute: parameters, FLOPs/token, KV bytes/token, and memory for weights+KV at 8k context, batch 32. Build a comparison table. Which is cheapest to serve? Is it the smallest?

C. Trace the bytes. Reproduce the trace in section 6 for a model and batch size of your choice. Verify the total against a measured decode step time.

D. Sliding window. Compute KV memory for a 32k-context request with and without a 4096 sliding window. What is the concurrency difference?

E. Ablate. Take your implementation from A and: remove the causal mask, remove RoPE, remove the residual connections, one at a time. Generate text with each. Describe what breaks and why.


11. Interview questions#

  1. Draw a transformer block and label every operation.
  2. What is pre-norm vs post-norm and why did pre-norm win?
  3. List five architectural choices that differ between modern LLMs, ranked by inference impact.
  4. What is sliding-window attention and what does it do to your capacity model?
  5. Given a config.json, compute parameters, FLOPs/token, and KV bytes/token.
  6. Why is the FFN the first target for quantization?
  7. How many kernel launches per token for an 80-layer model, and why does that matter?

12. Further reading#

  • [FUNDAMENTAL] Vaswani et al., “Attention Is All You Need” (2017)
  • [FUNDAMENTAL] Karpathy, nanoGPT — read model.py end to end
  • [ESTABLISHED] Llama, Mistral, Qwen, DeepSeek technical reports — read the architecture sections
  • [ESTABLISHED] Elhage et al., “A Mathematical Framework for Transformer Circuits” — the residual stream view
  • Next: 10 — Training vs inference: the math

↑↓ navigate↵ openesc close