Tokenwerk · The LLM textbook

Chapter 19 · III · Building the transformer · 9 minutes

The complete transformer block

Attention alone is not enough. The residual path, normalization and the feed-forward network work together.

Putting the familiar parts together

A transformer block is a repeatable unit of computation. It receives a list of numbers for each text position and outputs a list of numbers of the same length. It combines two ideas you already know: attention gathers information from the context; a small neural network processes the numbers further at each position.

We are not building a mysterious new mechanism. We arrange familiar computation steps and add two aids for training: direct addition paths and normalization. Normalization here means adjusting the scale of the numbers according to a fixed calculation.

Adding to a result instead of replacing everything

Picture a position representation [2, 5]. An operation computes the change [0.3, -0.2]. We add component by component and get [2.3, 4.8]. The original representation is kept as a direct term in the sum. This is a residual connection. Residual here means: the operation contributes an addition on top.

In short form we write y=x+f(x)y=x+f(x). x is the input, f the operation and y the result. f(x) must have the same shape as x so the addition is possible. In the backward pass the path splits: one part goes straight back through the addition, the other through f. The contributions are then added together. Remember the branching computation paths in the backpropagation chapter.

This often makes deep networks easier to train. But it does not guarantee that any network, however deep, will learn stably at any learning rate.

Normalizing with two numbers

Take the list [1, 3]. Its mean is (1 + 3) / 2 = 2. Subtract it from both values: you are left with [-1, 1]. The average squared deviation is (1 + 1) / 2 = 1. This number is called the variance. Its square root is called the standard deviation; here it is also 1.

If we divide the deviations by the standard deviation, we get [-1, 1] again. Without a safety term, [10, 30] would give the same normalized list. The original level and the original scale are removed.

What happens with [2, 2]? Both deviations and the variance are zero. We must not divide by zero. That's why we add a very small positive number, epsilon, before taking the square root. It serves as a safety term. In addition, the model learns a multiplier and a shift for each component. With these, it may rescale the standardized values as needed.

LayerNorm: the exact notation

Normalizing over the components of a single position is called LayerNorm. Our implementation computes it separately for each position:

LN⁡(x)i=γixi−μσ2+ϵ+βi.\operatorname{LN}(x)_i=\gamma_i\frac{x_i-\mu}{\sqrt{\sigma^2+\epsilon}}+\beta_i.
Symbol Meaning in this formula
x List of numbers for one position
i Which component of the list we are computing right now
μ\mu (mu) Mean of all components of this list
σ2\sigma^2 (sigma squared) Variance: mean of the squared deviations
ϵ\epsilon (epsilon) Small positive safety constant
γi\gamma_i (gamma) Trainable multiplier for component i
βi\beta_i (beta) Trainable shift for component i
LN⁡(x)i\operatorname{LN}(x)_i The finished normalized i-th component

Read from the inside out: subtract the mean, divide by the protected standard deviation, multiply by gamma, add beta. With gamma 1 and beta 0, the standardized list stays as it is. Epsilon causes a small deviation from the ideal calculation above.

nn.LayerNorm(D) processes the last D components each time. The different texts in a batch and the different positions are not mixed into one shared statistic. LayerNorm also does not limit the output to the interval from zero to one.

The small network inside the block

After attention, each position goes through a feed-forward network, also called an MLP. MLP stands for multilayer perceptron: a network of several consecutive layers. You know the basic principle from our first small network.

A first linear layer turns D components into a wider intermediate list, here with 4D components. An activation function changes the values nonlinearly. A second linear layer brings the width back to D. With D = 8 that means: 8 → 32 → 8.

We use GELU, a smooth activation function. Unlike ReLU, it does not cut negative values off hard at zero; it weights them softly. You don't need its exact special form to assemble this block. What matters is the same property as in the activation chapter: the whole operation can express more than a single linear layer.

MLP⁡(x)=W2GELU⁡(W1x+b1)+b2.\operatorname{MLP}(x)=W_2\operatorname{GELU}(W_1x+b_1)+b_2.
Symbol Meaning
W1W_1, b1b_1 Weight table and bias of the first layer
W1x+b1W_1x+b_1 Wider intermediate list, here with 4D components
GELU Activation function, applied to each component individually
W2W_2, b2b_2 Weight table and bias of the second layer
MLP(x) Output, again with D components

Here we write mathematically with column vectors. PyTorch stores and multiplies its batches according to its own shape convention; nn.Linear takes care of the right arrangement.

python
self.mlp = nn.Sequential(
    nn.Linear(D, 4 * D),
    nn.GELU(),
    nn.Linear(4 * D, D),
)

The same MLP with the same parameters is applied to every position. It does not mix different positions itself. That is attention's job.

Two additions, two normalizations

Now the block can be read in words:

  1. Normalize the input and compute attention from it.
  2. Add this result to the original input. Call the new intermediate list u.
  3. Normalize u and process it in the MLP.
  4. Add its result to u. That is the output y.
u=x+Attention⁡(Norm⁡1(x)),u=x+\operatorname{Attention}(\operatorname{Norm}_1(x)),
y=u+MLP⁡(Norm⁡2(u)).y=u+\operatorname{MLP}(\operatorname{Norm}_2(u)).

x is the block input, u the result after the first addition, and y the block output. Norm₁ and Norm₂ are two separate LayerNorm modules with their own trainable numbers. Pre-norm means: normalization comes before attention or the MLP. The original transformer from 2017 normalized only after the addition (post-norm); pre-norm later turned out to be more stable to train (Xiong et al., 2020). On the second addition path, u is added, not the old x again.

python
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x

Where does the position information come from?

Our token embedding table initially gives the same token the same list everywhere. For language, though, the place in the sequence also has to count. In the classic GPT-2-style decoder, we add a second learned list, the position embedding. The original transformer from 2017 used fixed sine-cosine encodings instead.

For a token list [0.2, 0.5] and a position list [0.1, -0.2], the first block receives [0.3, 0.3]. Both tables are trained. The index t of a position selects the corresponding row from the position table.

python
pos = torch.arange(T, device=ids.device)
x = self.token_embedding(ids) + self.position_embedding(pos)

T is the length of the current token sequence. arange(T) creates the position numbers 0 to T−1. device makes sure both arrays live on the same compute device.

The learned position table has a fixed number of rows. With context length 256, there is no row for position 10,000. That's why our generator later keeps at most the last context tokens (for example 256) during generation and renumbers this window. Old tokens that have been dropped are no longer available to the model in that computation step.

The complete decoder

The flow is: add token and position embeddings, apply several transformer blocks, normalize once more, and finally run a linear output layer. This last layer is called the output head. For each position, it produces one logit for every possible next token. With V = 100 vocabulary entries, that is 100 numbers.

The head does not pick a token yet. Softmax and the sampling strategy that comes later turn this into a continuation. The same learned parameters are used for all positions and for every new generation step.

Going deeper: a rough parameter count

A D-by-D weight table holds D² numbers. Four such tables for Q, K, V and the attention output hold about 4D² weights. The two MLP tables have D times 4D and 4D times D entries, together 8D². So one block holds about 12D² weights. Biases and norm parameters come on top.

For L blocks, V vocabulary entries and T learned positions, with a shared input/output table, the approximation is:

N≈VD+12LD2+TD.N\approx VD+12LD^2+TD.

N is the total number of parameters; D the width of a position list; L the number of blocks; V the vocabulary size; T the maximum number of positions. The three terms count the token embedding, the blocks and the position table. With V = 100, D = 8, L = 2 and T = 16, that's 800 + 1536 + 128 = 2464 weights, plus the small parameter groups left out.

This estimate applies to the classic block described here. Later architecture changes alter the calculation. The code can count the parameters that actually exist directly: sum(p.numel() for p in model.parameters()). numel() returns the number of values in a parameter array.