Dataflow in Generative LLMs
This is a rough sketch of the inference loop in an LLM, as it predicts one token at a time. I find new research hard to follow unless it’s clear where the change plugs in. Code helps, even pseudo-code that doesn’t run, since it makes the primitives and the control-flow explicit. It’s also been helpful to ground explanations from AI assistants, in the style I prefer.
The snippets below only convey the shape, and aren’t real implementations. For more background on why tranformers are the way they are, there are a lot of good sources. My favorite piece of exposition is on Cosma Shalizi’s blog, which you should go read!
The outer loop
Text in, Text out
A language model is a function from “text so far” to “more text”. You get a sequence of tokens, which the model processes and produces more tokens in the response.
def complete(context: str, n: int = 20) -> str:
...
Break Text into Tokens
To make processing tractable, text is broken into tokens (similar to how humans use words). Split it into pieces, and look each piece up in a fixed vocabulary.
# vocabulary size
VOCAB_DIM: int = 50
def tokenize(text: str) -> list[int]: ...
def detokenize(ids: list[int]) -> str: ...
def complete(context: str, n: int = 20) -> str:
...
ids = tokenize(context)
ids = generate(ids, n)
return detokenize(ids)
So, the model only ever sees sequences of integers.
Generate tokens one at a time
The model doesn’t produce 20 tokens in one shot.
- At each step, it selects a token from the vocabulary, and appends it to the input sequence.
- Then the inference pass is run again, with the appended sequence as the input.
def generate(ids: list[int], n: int) -> list[int]:
for _ in range(n):
next_id = model(ids)
ids.append(next_id)
return ids
The model takes the whole sequence and returns one integer. Rest of the steps unpack its internals.
The residual vector
Token generation is broken into two steps.
- First it reads the context and builds up some state.
- Then that state is used to choose a token from the vocabulary.
MODEL_DIM = 8
def model(ids: list[int]) -> int:
# context → one vector of MODEL_DIM floats
h = process_tokens(ids)
# vector → one token id
return generate_token(h)
h is a fixed-length list of floats, called the residual vector. It is the only thing that crosses from the reading half to the choosing half. generate_token projects it onto the vocabulary and picks a token.
Layers and context
Every token gets a vector
process_tokens starts with a lookup table, one embedding vector per vocabulary entry. Each input token is mapped to its own vector h.
embed_table = [
[random.gauss(0, 0.5) for _ in range(MODEL_DIM)] for _ in range(VOCAB_DIM)
]
def process_tokens(ids: list[int]) -> list[float]:
hs = [embed_table[token] for token in ids]
return ...
Embedding vectors above are context-independent; token T always maps to the same vector, whichever sentence it is put in. The transformer layers following it add context by mixing per-token vectors.
Push tokens through Layers
LLM computation is structured by pushing the whole list of vectors through a stack of “layers”. Each layer takes a list of vectors hs in, and outputs hs out with the same shape.
Each layer refines every token’s h, so more tokens means more rounds of mixing.
NUM_LAYERS = 40
def process_tokens(ids: list[int]) -> list[float]:
hs = [embed_table[token] for token in ids]
for layer_idx in range(NUM_LAYERS):
hs = layer(layer_idx, hs)
return hs[-1]
Only the residual vector corresponding to the last token, in the last layer, is used for generating the next token.
Inside a layer: each token looks back
A layer updates each position using its own h vector and the hs before it. Not after it: position j only sees hs[:j+1]. This slice operation makes the LLM causal.
def mix(layer_idx: int, past: list[list[float]], h: list[float]) -> list[float]:
return ...
def layer(layer_idx: int, hs: list[list[float]]) -> list[list[float]]:
out = []
for j in range(len(hs)):
out.append(mix(layer_idx, hs[:j + 1], hs[j]))
return out
The mix method is where each token interacts with tokens that came before it.
Attention
Mixing as regression
Each token combines with the tokens before it, in a regression-like setup.
- Score how relevant vectors at positions ≤N are to the vector at position N.
- Normalize the scores into weights that sum to 1.
- Add the weighted sum to the vector at position N.
def mix(layer_idx: int, past: list[list[float]], h: list[float]) -> list[float]:
scores = [dot(h, other) for other in past]
weights = softmax(scores)
for wt, other in zip(weights, past):
h = add(h, [wt * x for x in other])
return h
Note that h = add(h, ...) adds the contribution from previous vectors instead of substituting it. So there’s an untransformed path for original vectors to pass through each layer, called residual connection.
Here is what the computational primitives look like:
def add(a: list[float], b: list[float]) -> list[float]:
return [x + y for x, y in zip(a, b)]
def dot(a: list[float], b: list[float]) -> float:
return sum(x * y for x, y in zip(a, b))
def softmax(scores: list[float]) -> list[float]:
exps = [math.exp(s - max(scores)) for s in scores]
return [e / sum(exps) for e in exps]
Softmax subtracts the largest score before computing the exponents to avoid overflow.
One vector, three jobs
Note that there are no trainable parameters in our LLM yet, besides the embedding table. For the model to be useful, it needs the “ability” to learn patterns from training data, so it can generate text when used.
In mixing as regression the same vector h is used in multiple ways.
- Query: Ask: what am I looking for?
- Key: Get matched against: how well do I answer someone else’s question?
- Value: Get retrieved: what do I hand over once matched?
Overloading a single vector to serve all these needs is limiting. For ex, a verb in a sentence might need to find its subject and not just other verbs.
For each of these “functions”, every layer introduces a learnable linear map.
MODEL_TO_QUERY = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_KEY = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_VALUE = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
def mix(layer_idx, past, h):
scores = [dot(h, other) for other in past]
query = MODEL_TO_QUERY[layer_idx](h)
keys = [MODEL_TO_KEY[layer_idx](other) for other in past]
values = [MODEL_TO_VALUE[layer_idx](other) for other in past]
scores = [dot(query, key) for key in keys]
weights = softmax(scores)
for wt, other in zip(weights, past):
h = add(h, [wt * x for x in other])
for wt, value in zip(weights, values):
h = add(h, [wt * v for v in value])
return h
Every MODEL_* variable above is a linear map, that takes a vector and outputs another.
LearnedFn = Callable[[list[float]], list[float]]
def make_learned_fn(in_dim: int, out_dim: int) -> LearnedFn:
# random numbers for illustration. In an actual LLM, these get updated during training
weights = [[random.gauss(0, 0.5) for _ in range(in_dim)] for _ in range(out_dim)]
return lambda vec: [dot(w, vec) for w in weights]
This is called the attention mechanism - weights from query times key, which are then used to scale the values.
(Took me forever to really intuit exactly what attention is meant for. Cosma Shalizi’s blog probably does the best job at this exposition - how attention is basically Kernel regression).
Transform each token separately
Attention moves information between tokens, but it is a weighted average at heart. Transformer layers do some per token computation following it too (empirically, that’s where most learned knowledge resides).
We add another learned map, a squash operation to bring it between -1 and 1. It follows the same residual pattern, add the output, instead of replacing it.
MODEL_MLP = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
def mlp(layer_idx: int, h: list[float]) -> list[float]:
out = MODEL_MLP[layer_idx](h)
return add(h, [math.tanh(x) for x in out])
def layer(layer_idx: int, hs: list[list[float]]) -> list[list[float]]:
out = []
for j in range(len(hs)):
out.append(mix(layer_idx, hs[:j + 1], hs[j]))
return out
return [mlp(layer_idx, vec) for vec in out]
So each layer does two operations per token
- Attention: Tokens talk to each other
- Multi layer perceptron: Transform each token separately
The KV cache
Notice the waste
Looking back at the generate loop (with model inlined) in one token at a time - every iteration calls process_tokens(ids) on the whole sequence, which recomputes every h at every layer, including the keys and values for tokens we saw ten steps ago.
def generate(ids: list[int], n: int) -> list[int]:
for _ in range(n):
# recomputes hs for ALL of ids, every layer
h = process_tokens(ids)
ids.append(generate_token(h))
return ids
However, an earlier token (at position j < sequence length N) contributes only its key and its value to the mix operation.
- Both key and value vectors at
jcan be computed from the inputhvector to the layer atj. - Since position
jnever sees anything after it, thathdoesn’t change when the sequence grows.
The key and value for an old token are final the first time they are computed, so we can persist and reuse them!
The Key-Value cache
For each layer, we keep a list of (key, value) rows, one for every token seen yet. Now process_token can take as input, a single token and the cache entries for tokens before it.
# (key-vector, value-vector)
CacheRow = tuple[list[float], list[float]]
# one list per layer
Cache = list[list[CacheRow]]
def empty_cache() -> Cache:
return [[] for _ in range(NUM_LAYERS)]
def process_tokens(ids: list[int]) -> list[float]:
hs = [embed_table[token] for token in ids]
for layer_idx in range(NUM_LAYERS):
hs = layer(layer_idx, hs)
return hs[-1]
def process_token(cache: Cache, token: int) -> tuple[Cache, list[float]]:
h = embed_table[token]
for layer_idx in range(NUM_LAYERS):
h = layer(layer_idx, cache[layer_idx], h)
return cache, h
The KV cache per layer grows with the sequence, one entry per token. We rewrite layer to take the cache below.
Prefill, then decode
The outer generation loop can also be split into two separate phases
- Prefill: read the prompt into the cache, and only keep the residual vector
hfrom the last token in the context. - Decode: do one read and one prediction, at each step
def generate(ids: list[int], n: int) -> list[int]:
cache = empty_cache()
# prefill
for token in ids:
cache, h = process_token(cache, token)
# decode
for _ in range(n):
# recomputes hs for ALL of ids, every layer
h = process_tokens(ids)
ids.append(generate_token(h))
next_id = generate_token(h)
cache, h = process_token(cache, next_id)
ids.append(next_id)
return ids
Rewriting a Layer: write to cache
Each layer now gets one h vector (from the previous layer) and its cache.
First, we add a write operation: compute the key and values vectors for the new token, and add them to the KV cache. This row is written before any reads, so the token can see itself.
def make_cache_row(layer_idx: int, h: list[float]) -> CacheRow:
return MODEL_TO_KEY[layer_idx](h), MODEL_TO_VALUE[layer_idx](h)
def layer(layer_idx: int, hs: list[list[float]]) -> list[list[float]]:
out = []
for j in range(len(hs)):
out.append(mix(layer_idx, hs[:j + 1], hs[j]))
return [mlp(layer_idx, vec) for vec in out]
def layer(layer_idx: int, layer_cache: list[CacheRow], h: list[float]) -> list[float]:
# write
layer_cache.append(make_cache_row(layer_idx, h))
...
return h
Read from cache, and compute attention
Then, we add the regression from one vector, three jobs, except the keys and values are already there.
- Compute the
queryusing the currenthvector - Score against every cached key vector, and compute weights
- Take weighted sum of the cached value vectors
def mix(layer_idx, past, h):
query = MODEL_TO_QUERY[layer_idx](h)
keys = [MODEL_TO_KEY[layer_idx](other) for other in past]
values = [MODEL_TO_VALUE[layer_idx](other) for other in past]
scores = [dot(query, key) for key in keys]
weights = softmax(scores)
for wt, value in zip(weights, values):
h = add(h, [wt * v for v in value])
return h
def attend(layer_idx: int, layer_cache: list[CacheRow], h: list[float]) -> list[float]:
query = MODEL_TO_QUERY[layer_idx](h)
weights = softmax([dot(query, key) for key, _ in layer_cache])
for wt, (_, value) in zip(weights, layer_cache):
h = add(h, [wt * v for v in value])
return h
def layer(layer_idx, layer_cache, h):
# write
layer_cache.append(make_cache_row(layer_idx, h))
# read
h = attend(layer_idx, layer_cache, h)
...
return h
Note the attention computation above is causal by construction, since the KV cache only stores vectors corresponding to past and current entries.
Add back the MLP op
Add back the MLP operation from the then think alone section. It doesn’t interact with the cache and only depends on the attention output.
def layer(layer_idx, layer_cache, h):
# write
layer_cache.append(make_cache_row(layer_idx, h))
# read
h = attend(layer_idx, layer_cache, h)
...
# think
h = mlp(layer_idx, h)
return h
Putting it together
Implement token generation
The last transformer layer outputs a residual vector, from which we want to generate a token.
- We add another learnable map to resize it from
MODEL_DIMto theVOCAB_DIM, which is the size of the vocabulary. - A softmax function then converts it into a list of probabilities, one for each token in the vocabulary.
- The token with the highest assigned probability is selected.
MODEL_TO_VOCAB = make_learned_fn(MODEL_DIM, VOCAB_DIM)
def generate_token(h: list[float]) -> int:
# 8 floats → 50 floats
scores = MODEL_TO_VOCAB(h)
probs = softmax(scores)
return max(range(VOCAB_DIM), key=lambda i: probs[i])
There are also alternate non-greedy decoding techniques like beam search, which try to maximize probabilities of subsequences over multiple steps. In which case, we keep around K of the highest scoring tokens, and not just the top one.
Everything all at once
import math, random
from typing import Callable
random.seed(42)
NUM_LAYERS, MODEL_DIM, VOCAB_DIM = 40, 8, 50
CacheRow = tuple[list[float], list[float]]
Cache = list[list[CacheRow]]
LearnedFn = Callable[[list[float]], list[float]]
def add(a, b): return [x + y for x, y in zip(a, b)]
def dot(a, b): return sum(x * y for x, y in zip(a, b))
def make_learned_fn(in_dim, out_dim) -> LearnedFn:
weights = [[random.gauss(0, 0.5) for _ in range(in_dim)] for _ in range(out_dim)]
return lambda vec: [dot(w, vec) for w in weights]
embed_table = [[random.gauss(0, 0.5) for _ in range(MODEL_DIM)] for _ in range(VOCAB_DIM)]
MODEL_TO_QUERY = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_KEY = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_VALUE = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_MLP = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_VOCAB = make_learned_fn(MODEL_DIM, VOCAB_DIM)
def softmax(scores):
exps = [math.exp(s - max(scores)) for s in scores]
return [e / sum(exps) for e in exps]
def empty_cache() -> Cache:
return [[] for _ in range(NUM_LAYERS)]
def make_cache_row(layer_idx, h) -> CacheRow:
return MODEL_TO_KEY[layer_idx](h), MODEL_TO_VALUE[layer_idx](h)
def attend(layer_idx, layer_cache, h):
query = MODEL_TO_QUERY[layer_idx](h)
weights = softmax([dot(query, key) for key, _ in layer_cache])
for wt, (_, value) in zip(weights, layer_cache):
h = add(h, [wt * v for v in value])
return h
def mlp(layer_idx, h):
return add(h, [math.tanh(x) for x in MODEL_MLP[layer_idx](h)])
def layer(layer_idx, layer_cache, h):
# write
layer_cache.append(make_cache_row(layer_idx, h))
# read
h = attend(layer_idx, layer_cache, h)
# think
return mlp(layer_idx, h)
def process_token(cache, token):
h = embed_table[token]
for layer_idx in range(NUM_LAYERS):
h = layer(layer_idx, cache[layer_idx], h)
return cache, h
def generate_token(h) -> int:
probs = softmax(MODEL_TO_VOCAB(h))
return max(range(VOCAB_DIM), key=lambda i: probs[i])
def generate(ids, n):
cache = empty_cache()
# prefill
for token in ids:
cache, h = process_token(cache, token)
# decode
for _ in range(n):
next_id = generate_token(h)
cache, h = process_token(cache, next_id)
ids.append(next_id)
return ids
generate([1, 2, 5, 20, 2], n=10)
This covers end to end, how a generative LLM processes the input and produces text. Real models accumulate a lot of variations and optimizations on top, some of which are covered below.
Extensions
Each extension below builds on the code above.
Positional encoding
Attention ignores order
Looking at the attention step from read from cache, and compute attention - scores depend only on what the key vectors contain, not on which row of the cache they sit in. So if we shuffled the cache rows, the output would stay exactly the same.
def attend(layer_idx, layer_cache, h):
query = MODEL_TO_QUERY[layer_idx](h)
# same for any row order
weights = softmax([dot(query, key) for key, _ in layer_cache])
...The only signal for order in the model comes from causality - which tokens have made it into the cache so far.
Rotate vectors by position
The typical way to add positional information per token is called RoPE (rotary position embedding). Before taking the dot product, the query and key vectors are rotated by an angle that grows with their position.
- Split the vector into pairs of dimensions.
- Rotate each pair like a 2D point, by
pos * frequency. - Each pair gets its own frequency, so some pairs spin fast and others slow.
BASE = 10_000
def rotate(vec: list[float], pos: int) -> list[float]:
out = []
for i in range(0, MODEL_DIM, 2):
theta = pos * BASE ** (-i / MODEL_DIM)
c, s = math.cos(theta), math.sin(theta)
x, y = vec[i], vec[i + 1]
out += [x * c - y * s, x * s + y * c]
return outApply it at write and at read
The position of a token is the number of rows already in the layer cache, when it is processed. We rotate the key vector when writing it to the cache, and the query vector when reading from it.
def layer(layer_idx, layer_cache, h):
pos = len(layer_cache)
# write
layer_cache.append(make_cache_row(layer_idx, h))
layer_cache.append(make_cache_row(layer_idx, h, pos))
# read
h = attend(layer_idx, layer_cache, h)
h = attend(layer_idx, layer_cache, h, pos)
# think
h = mlp(layer_idx, h)
return h
def make_cache_row(layer_idx, h, pos):
return MODEL_TO_KEY[layer_idx](h), MODEL_TO_VALUE[layer_idx](h)
key = rotate(MODEL_TO_KEY[layer_idx](h), pos)
return key, MODEL_TO_VALUE[layer_idx](h)
def attend(layer_idx, layer_cache, h, pos):
query = MODEL_TO_QUERY[layer_idx](h)
query = rotate(MODEL_TO_QUERY[layer_idx](h), pos)
weights = softmax([dot(query, key) for key, _ in layer_cache])
...Note that when both vectors are rotated by their positions, the math works out so dot(query, key) depend only on the distance between the two tokens.
Attention heads
One set of weights per token
In read from cache, and compute attention, each token computes a single list of weights over the cache. So it can only follow one pattern at a time - for ex, look at the previous token, or at the matching open bracket, but not both.
Split attention into heads
Instead, we could run several smaller attention operations in parallel, called heads.
- Split
MODEL_DIMintoNUM_HEADSchunks of sizeHEAD_DIM. - Each head gets its own query, key and value maps.
- A cache row now holds one (key, value) pair per head.
NUM_HEADS = 4
HEAD_DIM = MODEL_DIM // NUM_HEADS
MODEL_TO_QUERY = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_KEY = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_VALUE = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_TO_QUERY = [[make_learned_fn(MODEL_DIM, HEAD_DIM) for _ in range(NUM_HEADS)] for _ in range(NUM_LAYERS)]
MODEL_TO_KEY = [[make_learned_fn(MODEL_DIM, HEAD_DIM) for _ in range(NUM_HEADS)] for _ in range(NUM_LAYERS)]
MODEL_TO_VALUE = [[make_learned_fn(MODEL_DIM, HEAD_DIM) for _ in range(NUM_HEADS)] for _ in range(NUM_LAYERS)]
CacheRow = tuple[list[float], list[float]]
# one (key, value) per head
CacheRow = list[tuple[list[float], list[float]]]
def make_cache_row(layer_idx, h):
return MODEL_TO_KEY[layer_idx](h), MODEL_TO_VALUE[layer_idx](h)
return [
(MODEL_TO_KEY[layer_idx][hd](h), MODEL_TO_VALUE[layer_idx][hd](h))
for hd in range(NUM_HEADS)
] Run each head, then combine
The attention step now runs once per head, each on its own slice of the cache row.
- Each head computes the same weighted sum as before, and outputs a
HEAD_DIMvector. - The head outputs are concatenated back into a single
MODEL_DIMvector. - Another learned map mixes the head outputs together, before adding the result to
h.
MODEL_ATTN_OUT = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
def attend(layer_idx, layer_cache, h):
query = MODEL_TO_QUERY[layer_idx](h)
weights = softmax([dot(query, key) for key, _ in layer_cache])
for wt, (_, value) in zip(weights, layer_cache):
h = add(h, [wt * v for v in value])
return h
head_outputs = []
for hd in range(NUM_HEADS):
query = MODEL_TO_QUERY[layer_idx][hd](h)
scores = [dot(query, row[hd][0]) / math.sqrt(HEAD_DIM) for row in layer_cache]
weights = softmax(scores)
out = [0.0] * HEAD_DIM
for wt, row in zip(weights, layer_cache):
out = add(out, [wt * v for v in row[hd][1]])
head_outputs.extend(out)
return add(h, MODEL_ATTN_OUT[layer_idx](head_outputs)) Note that the scores are also divided by sqrt(HEAD_DIM). Dot products grow with vector size, and large scores make softmax put nearly all the weight on a single token.
Task-specific heads
A different kind of head
This “head” is unrelated to attention heads. A task head is the final learned map that turns the last residual vector into whatever output the task needs.
We already have one in implement token generation - MODEL_TO_VOCAB is the language modelling head, while the transformer layers are called the backbone.
Swap in a classification head
To classify text instead of generating it, we replace the head and leave the backbone untouched.
NUM_CLASSES = 3
MODEL_TO_CLASS = make_learned_fn(MODEL_DIM, NUM_CLASSES)
def classify(h: list[float]) -> int:
probs = softmax(MODEL_TO_CLASS(h))
return max(range(NUM_CLASSES), key=lambda i: probs[i]) This works since the last residual vector already has to hold whatever is needed to predict the next token. So it is a decent summary of the whole input, and a single linear map on top of it can be trained with few examples.
Drop the causal mask
A classifier never generates tokens, so there is no reason for a token to only see the ones before it.
- First compute the key and value vectors for every token in the layer.
- Then let each token attend over all of them, including the ones after it.
def process_sequence(ids: list[int]) -> list[list[float]]:
hs = [embed_table[token] for token in ids]
for layer_idx in range(NUM_LAYERS):
# write: every key and value first
rows = [make_cache_row(layer_idx, h) for h in hs]
# read: each token sees all of them
hs = [attend(layer_idx, rows, h) for h in hs]
# think
hs = [mlp(layer_idx, h) for h in hs]
return hs
def classify_sequence(ids: list[int]) -> int:
hs = process_sequence(ids)
# average over all positions
pooled = [sum(col) / len(hs) for col in zip(*hs)]
probs = softmax(MODEL_TO_CLASS(pooled))
return max(range(NUM_CLASSES), key=lambda i: probs[i])Note that there is no KV cache anymore, since all the tokens are processed together. Earlier tokens now get context from later ones too, which is helpful when the task doesn’t require token generation.
Mixture of experts
Replace the MLP with experts
The attention computation and the KV cache in each layer stay as they are. Only the MLP step from add back the MLP op changes.
- The single
MODEL_MLPmap per layer is replaced by a list ofexperts, each with their own linear map. - A small learned map, called the
router, scores each expert for the current token.
NUM_EXPERTS = 8
TOP_K = 2
MODEL_MLP = [make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_LAYERS)]
MODEL_EXPERTS = [[make_learned_fn(MODEL_DIM, MODEL_DIM) for _ in range(NUM_EXPERTS)] for _ in range(NUM_LAYERS)]
MODEL_ROUTER = [make_learned_fn(MODEL_DIM, NUM_EXPERTS) for _ in range(NUM_LAYERS)] Route each token to a few experts
For each token, the MLP step now does the following.
- Score every expert with the router, and pick the top
TOP_K. - Run only the chosen experts on
h. - Take the weighted sum of their outputs, using the router scores as weights.
- Add the result to
h, same as before.
def mlp(layer_idx, h):
out = MODEL_MLP[layer_idx](h)
return add(h, [math.tanh(x) for x in out])
# one score per expert
gate = softmax(MODEL_ROUTER[layer_idx](h))
chosen = sorted(range(NUM_EXPERTS), key=lambda e: -gate[e])[:TOP_K]
total = sum(gate[e] for e in chosen)
out = [0.0] * MODEL_DIM
for e in chosen:
expert_out = [math.tanh(x) for x in MODEL_EXPERTS[layer_idx][e](h)]
out = add(out, [(gate[e] / total) * x for x in expert_out])
return add(h, out) Now the model has NUM_EXPERTS times more MLP parameters, but each token only runs TOP_K of them. This allows MoE models to store much more “knowledge” without each token costing a lot more compute.