The idea in one minute#
A large language model is a function from a sequence of token IDs to a score for every possible next token. Inside, each token is a vector that passes through a stack of identical layers. Every layer does two things: attention, where the token gathers information from all the tokens before it, and an MLP, where it processes what it gathered. The final vector is compared with every vocabulary entry to produce logits; pick a token from them, append it, and repeat.
Generating token by token would mean redoing the whole sequence for each new token. The KV cache avoids that: each token’s attention keys and values are computed once and stored, so a new token only computes its own. That single idea turns quadratic work into linear, and is the reason LLM serving is a memory-management problem.
An analogy#
A meeting where each person speaks once, in order. Before speaking, each person reviews what everyone earlier said (attention), then thinks (MLP), then adds their conclusion to the shared notes. Without minutes, every new speaker would make all the previous speakers repeat themselves from the start. The KV cache is the minutes: what each earlier speaker contributed is written down once and simply read by everyone after.
A picture#
flowchart TB
TOK["token id at position t"] --> EMB["x = token embedding + position embedding"]
EMB --> N1["normalize"]
N1 --> QKV["q = Wq h k = Wk h v = Wv h"]
QKV -->|"store k, v at position t"| CACHE[("KV cache<br/>keys and values of positions 0..t<br/>per layer")]
CACHE --> ATT["scores = q . k_i / sqrt(d) for i = 0..t<br/>weights = softmax(scores)<br/>out = sum of weights x v_i"]
QKV --> ATT
ATT --> WO["x = x + Wo out"]
WO --> N2["normalize"]
N2 --> MLP["x = x + W2 gelu(W1 h)"]
MLP -->|"repeat for every layer"| N1
MLP --> FIN["final normalize"]
FIN --> LOG["logits = embedding table x h<br/>one score per vocabulary entry"]
LOG --> SAMP["sample: greedy, temperature, top-k"]
SAMP -->|"next token"| TOK
class TOK,SAMP neutral
class EMB,N1,N2,FIN,QKV,WO,MLP,LOG compute
class CACHE memory
class ATT queueHow it really works#
The residual stream#
Each position carries one vector x of size d_model through the whole network. Layers do not
replace it; they add to it (x = x + f(x)). That running sum is the residual stream, and
it is why information — and gradients, during training — flows through dozens of layers.
Embedding#
x = TokEmb[token] + PosEmb[position]. The first term says what the token is, the second
where. (Current models encode position differently — rotary embeddings applied inside
attention — but the role is the same.)
Attention#
For the token at position t, in each layer:
h = normalize(x)
q = Wq·h "what am I looking for?"
k = Wk·h "what do I contain?" ← stored in the cache
v = Wv·h "what do I pass on?" ← stored in the cache
score_i = (q · k_i) / √d_head for every position i ≤ t
weight_i = softmax(score)_i
out = Σ weight_i · v_i
x = x + Wo·out- Causal: a token attends only to itself and earlier positions. This is what makes left-to-right generation valid, and what makes the cache possible — earlier tokens’ keys and values never depend on later tokens, so they never change.
- Heads: the vector is split into several independent slices (“heads”), each with its own scores, so a token can attend to different things for different purposes.
- Softmax turns scores into weights that are positive and sum to one. Subtracting the maximum score first prevents overflow.
The MLP#
x = x + W2·gelu(W1·normalize(x)), with a hidden width of 4·d_model. Two-thirds of a
transformer’s weights live here. It acts on each position independently.
Logits and sampling#
After the last layer, logits = TokEmb · normalize(x): a dot product with every vocabulary
row. Then choose:
| Strategy | How | Effect |
|---|---|---|
| Greedy | Take the largest | Deterministic; can loop |
Temperature T | Divide logits by T before softmax | T < 1 sharpens, T > 1 flattens |
| Top-k | Keep only the k largest | Removes the long tail of bad options |
| Top-p (nucleus) | Keep the smallest set whose probability sums to p | Adapts to how confident the model is |
Why the KV cache matters#
Without a cache, producing token n means running positions 0..n−1 through every layer
again: the total work to generate n tokens grows as n². With the cache, each new token runs
once and reads the stored keys and values: total work grows as n.
The price is memory, and it is exact:
KV bytes per token = 2 × layers × d_model × bytes per valueFor an 8B-class model at 16-bit precision that is about half a megabyte per token, per sequence. A thousand concurrent conversations of 4,000 tokens is two terabytes — more than the weights by two orders of magnitude. Deciding which sequences get cache space, sharing common prefixes, and evicting are what an inference engine’s scheduler spends its effort on (Inference Engineering V.05, V.10).
Prefill and decode#
The same function is used in two regimes:
- Prefill: the prompt’s tokens are all known, so they can be processed together as one matrix operation. Compute-bound; determines time to first token.
- Decode: one new token per step. Every weight is read to produce a single vector: memory-bound; determines tokens per second.
The program below processes the prompt one token at a time for simplicity. A real engine batches the prefill across positions, and batches decode across sequences.
The Go that matters here#
- Weights are flat
[]float32, read-only after loading: shared by every goroutine with no lock (IV.05) and free for the garbage collector (III.04). - All per-sequence state — the KV cache and scratch vectors — lives in a
Stateallocated once.Stepallocates nothing (III.05). - The inner kernel is the unrolled dot product from V.02.
That is the layout of a real engine in miniature: immutable shared weights, per-request state, an allocation-free step.
Code#
Weights are random — training a language model is out of scope — so the output tokens are meaningless. Everything else is real: the arithmetic, the cache, the sampling.
// transformer.go — a decoder-only transformer forward pass with a KV cache and sampling.
package main
import (
"fmt"
"math"
"math/rand"
"sort"
"testing"
"time"
)
type Config struct{ Vocab, Ctx, D, Heads, Layers int }
type Layer struct {
G1, G2 []float32 // RMSNorm gains
Wq, Wk, Wv, Wo []float32 // D x D each
W1 []float32 // 4D x D
W2 []float32 // D x 4D
}
type Model struct {
Config
TokEmb, PosEmb, GF []float32
L []Layer
}
// State is everything one sequence needs: its KV cache and scratch buffers. Nothing in a
// decode step allocates; a server keeps one State per active request.
type State struct {
K, V [][]float32 // per layer: Ctx x D, filled up to Len
Len int
x, h, q, att, mlp, scores []float32
Logits []float32
}
func randn(rng *rand.Rand, n int, scale float64) []float32 {
w := make([]float32, n)
for i := range w {
w[i] = float32(rng.NormFloat64() * scale)
}
return w
}
func ones(n int) []float32 {
g := make([]float32, n)
for i := range g {
g[i] = 1
}
return g
}
func NewModel(c Config, seed int64) *Model {
rng := rand.New(rand.NewSource(seed))
s := 1 / math.Sqrt(float64(c.D))
m := &Model{Config: c, TokEmb: randn(rng, c.Vocab*c.D, s), PosEmb: randn(rng, c.Ctx*c.D, s), GF: ones(c.D)}
for l := 0; l < c.Layers; l++ {
m.L = append(m.L, Layer{
G1: ones(c.D), G2: ones(c.D),
Wq: randn(rng, c.D*c.D, s), Wk: randn(rng, c.D*c.D, s), Wv: randn(rng, c.D*c.D, s), Wo: randn(rng, c.D*c.D, s),
W1: randn(rng, 4*c.D*c.D, s), W2: randn(rng, 4*c.D*c.D, s/2),
})
}
return m
}
func (m *Model) NewState() *State {
s := &State{
x: make([]float32, m.D), h: make([]float32, m.D), q: make([]float32, m.D),
att: make([]float32, m.D), mlp: make([]float32, 4*m.D), scores: make([]float32, m.Ctx),
Logits: make([]float32, m.Vocab),
}
for l := 0; l < m.Layers; l++ {
s.K = append(s.K, make([]float32, m.Ctx*m.D))
s.V = append(s.V, make([]float32, m.Ctx*m.D))
}
return s
}
func dot(a, b []float32) float32 {
b = b[:len(a)]
var s0, s1, s2, s3 float32
i := 0
for ; i+4 <= len(a); i += 4 {
s0 += a[i] * b[i]
s1 += a[i+1] * b[i+1]
s2 += a[i+2] * b[i+2]
s3 += a[i+3] * b[i+3]
}
for ; i < len(a); i++ {
s0 += a[i] * b[i]
}
return s0 + s1 + s2 + s3
}
// matVec: dst = W x, with W stored row-major as len(dst) rows of len(x).
func matVec(dst, w, x []float32) {
n := len(x)
for r := range dst {
dst[r] = dot(x, w[r*n:(r+1)*n])
}
}
func rmsNorm(dst, x, gain []float32) {
var ss float32
for _, v := range x {
ss += v * v
}
inv := float32(1 / math.Sqrt(float64(ss)/float64(len(x))+1e-5))
for i, v := range x {
dst[i] = v * inv * gain[i]
}
}
// Step feeds one token at position s.Len, updates the KV cache, and fills s.Logits.
func (m *Model) Step(s *State, token int) {
D, pos := m.D, s.Len
dh := D / m.Heads
for i := 0; i < D; i++ {
s.x[i] = m.TokEmb[token*D+i] + m.PosEmb[pos*D+i]
}
for l := range m.L {
ly := &m.L[l]
// --- attention: this token looks at every earlier token, and itself
rmsNorm(s.h, s.x, ly.G1)
matVec(s.q, ly.Wq, s.h)
matVec(s.K[l][pos*D:(pos+1)*D], ly.Wk, s.h) // cache this position's key...
matVec(s.V[l][pos*D:(pos+1)*D], ly.Wv, s.h) // ...and value: computed once, reused forever
for i := range s.att {
s.att[i] = 0
}
scale := float32(1 / math.Sqrt(float64(dh)))
for h := 0; h < m.Heads; h++ {
qh := s.q[h*dh : (h+1)*dh]
maxScore := float32(math.Inf(-1))
for t := 0; t <= pos; t++ {
sc := dot(qh, s.K[l][t*D+h*dh:t*D+(h+1)*dh]) * scale
s.scores[t] = sc
if sc > maxScore {
maxScore = sc
}
}
var sum float32
for t := 0; t <= pos; t++ { // softmax, shifted by the max for numerical stability
s.scores[t] = float32(math.Exp(float64(s.scores[t] - maxScore)))
sum += s.scores[t]
}
out := s.att[h*dh : (h+1)*dh]
for t := 0; t <= pos; t++ {
w := s.scores[t] / sum
vh := s.V[l][t*D+h*dh : t*D+(h+1)*dh]
for i := range out {
out[i] += w * vh[i]
}
}
}
matVec(s.h, ly.Wo, s.att)
for i := range s.x {
s.x[i] += s.h[i] // residual connection
}
// --- MLP
rmsNorm(s.h, s.x, ly.G2)
matVec(s.mlp, ly.W1, s.h)
for i, v := range s.mlp { // GELU (tanh approximation)
x := float64(v)
s.mlp[i] = float32(0.5 * x * (1 + math.Tanh(0.7978845608*(x+0.044715*x*x*x))))
}
matVec(s.h, ly.W2, s.mlp)
for i := range s.x {
s.x[i] += s.h[i]
}
}
rmsNorm(s.h, s.x, m.GF)
matVec(s.Logits, m.TokEmb, s.h) // output projection shares the embedding table
s.Len++
}
func argmax(x []float32) int {
best := 0
for i, v := range x {
if v > x[best] {
best = i
}
}
return best
}
// sample draws from the top-k tokens after dividing logits by the temperature.
func sample(rng *rand.Rand, logits []float32, temperature float64, k int) int {
idx := make([]int, len(logits))
for i := range idx {
idx[i] = i
}
sort.Slice(idx, func(a, b int) bool { return logits[idx[a]] > logits[idx[b]] })
idx = idx[:k]
p := make([]float64, k)
sum := 0.0
for i, id := range idx {
p[i] = math.Exp(float64(logits[id]-logits[idx[0]]) / temperature)
sum += p[i]
}
r := rng.Float64() * sum
for i, id := range idx {
if r -= p[i]; r <= 0 {
return id
}
}
return idx[k-1]
}
func main() {
m := NewModel(Config{Vocab: 512, Ctx: 96, D: 128, Heads: 4, Layers: 4}, 7)
prompt := []int{11, 42, 7, 300}
const newTokens = 64
// 1. With the KV cache: each new token costs one Step.
s := m.NewState()
start := time.Now()
for _, t := range prompt {
m.Step(s, t)
}
cached := make([]int, 0, newTokens)
for i := 0; i < newTokens; i++ {
next := argmax(s.Logits)
cached = append(cached, next)
m.Step(s, next)
}
tCached := time.Since(start)
// 2. Without it: to produce each token, re-run the whole sequence so far from scratch.
start = time.Now()
seq := append([]int{}, prompt...)
var plain []int
steps := 0
for i := 0; i < newTokens; i++ {
fresh := m.NewState()
for _, t := range seq {
m.Step(fresh, t)
steps++
}
next := argmax(fresh.Logits)
plain = append(plain, next)
seq = append(seq, next)
}
tPlain := time.Since(start)
same := true
for i := range cached {
same = same && cached[i] == plain[i]
}
fmt.Println("first tokens generated:", cached[:12], "(random weights: meaningless, but deterministic)")
fmt.Println("cached and uncached outputs identical:", same)
fmt.Printf("with KV cache: %4d steps %7.1f ms %6.0f tokens/s\n", len(prompt)+newTokens,
float64(tCached.Microseconds())/1000, newTokens/tCached.Seconds())
fmt.Printf("without KV cache: %4d steps %7.1f ms %6.0f tokens/s (%.0fx slower)\n", steps,
float64(tPlain.Microseconds())/1000, newTokens/tPlain.Seconds(), float64(tPlain)/float64(tCached))
kvBytes := 2 * m.Layers * m.D * 4
fmt.Printf("KV cache: %d bytes per token, %d kB for this sequence\n", kvBytes, kvBytes*s.Len/1024)
// 3. The decode step does not allocate.
s2 := m.NewState()
m.Step(s2, 1)
allocs := testing.AllocsPerRun(20, func() {
if s2.Len < m.Ctx {
m.Step(s2, 5)
}
})
fmt.Printf("allocations per decode step: %.0f\n", allocs)
// 4. Sampling: temperature and top-k change which token is picked from the same logits.
rng := rand.New(rand.NewSource(1))
fmt.Print("greedy: ", argmax(s.Logits), " temperature 1.0, top-k 40:")
for i := 0; i < 8; i++ {
fmt.Print(" ", sample(rng, s.Logits, 1.0, 40))
}
fmt.Print("\n temperature 0.1, top-k 40:")
for i := 0; i < 8; i++ {
fmt.Print(" ", sample(rng, s.Logits, 0.1, 40))
}
fmt.Println()
}Remember this#
- Per layer: attention (gather from earlier tokens) then an MLP (process), each added to a running vector.
- Causal attention means earlier tokens’ keys and values never change — so cache them.
- The KV cache turns
n²work inton, and costs2 × layers × d_model × bytesper token per sequence. - Prefill is compute-bound; decode is memory-bound.
- Shared read-only weights, per-request state, no allocation per step.
Try it#
- Run
transformer.go. Double the number of generated tokens (raiseCtxto fit). How does the gap between cached and uncached change, and why? - Add top-p sampling.
- Run several sequences concurrently, one goroutine and one
Stateeach, sharing the model. Confirm with-racethat there is no data race, and measure total tokens per second against the number of goroutines.
Check yourself#
- What do the query, key and value of a token represent?
- Why does causal attention make the KV cache valid?
- How many bytes of KV cache does one token cost, and why does that dominate serving?