ishan's notebook
← all writing

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 j can be computed from the input h vector to the layer at j.
  • Since position j never sees anything after it, that h doesn’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 h from 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 query using the current h vector
  • 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_DIM to the VOCAB_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 out

Apply 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_DIM into NUM_HEADS chunks of size HEAD_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_DIM vector.
  • The head outputs are concatenated back into a single MODEL_DIM vector.
  • 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_MLP map per layer is replaced by a list of experts, 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.