8.6 Residual connections
Stack 20 layers and something sneaky happens: the gradient has to pass backwards through all of them, and each one can shrink it. By the bottom, nothing is left, so the early layers never learn. One plus sign fixes it.
1Watch
A residual connection means a block (one big step of the network: attention or the MLP of lesson 9.2) adds its result to what came in, instead of replacing it: x = x + block(x). A layer is one round of blocks; big GPTs stack dozens of them (9.3 builds microgpt’s), and this lab’s toy layer(x) stands in for one. microgpt does this around attention (line 134) and around the MLP (line 141).
Why it works, in the language of 7.2 and 7.3: the + x path has a local rate of exactly 1. So each layer has two roads back: straight through the +, at rate 1, and through the layer, at 0.1 in the lab. Roads add, so each layer’s rate is 1.1. In the lab’s stack of 20 layers, the plain version passes back 0.1²⁰ = 10⁻²⁰ of the gradient (Python prints that as 1e-20). The residual version passes back 1.1²⁰ ≈ 6.73.
What if a layer’s own rate were negative, say −0.9? Then 1 + (−0.9) = 0.1, and the total could shrink again. But follow only the straight roads: among all the roads back, one goes straight through every +, and its rate is 1 × 1 × … × 1 = 1. The layers can add to that road or pull on others, but that highway is always there to carry the gradient home.
The numbers do grow in the lab (6.73 at 20 layers, 45 at 40). In microgpt the rmsnorm before each block (8.5) resets the size of what the block receives, so the + path can’t blow up a block’s input.
That is also why line 112 is “not redundant”. Line 117 normalises only the copy that goes into attention; the highway (x_residual, line 116) carries x exactly as it arrived. Line 112 is the one rmsnorm on the highway itself: it sets the size of what the highway carries, and on the way back, the highway’s gradient reaches the embeddings through it.
Now the whole picture of one letter’s trip: embed (8.1) → normalise (8.5) → attention (8.2 to 8.4) → add back (this lesson) → MLP (9.2) → lm_head gives 27 next-letter scores. Residual connections are why networks can be hundreds of layers deep. microgpt has just one layer, but the same two lines are in every GPT.
2Explore
Gradient reaching the input (log scale), when each layer is “×0.1 then relu”. Without the shortcut it shrinks 10× per layer (10⁻²⁰ means a 1 that is 20 places after the decimal point). With it, every layer adds 1 to the rate, and among all the roads back, one goes straight through every + with rate 1 × 1 × … × 1 = 1, so the signal always gets through.
3Build
One blank in residual_stack: add the layer’s output to x instead of replacing it. Your Value class and backward() from Module 7 measure the gradient that reaches the input.
import math
# microgpt's Value class: you wrote it in lessons 7.4 and 7.5.
class Value:
__slots__ = ('data', 'grad', '_children', '_local_grads') # Python optimization for memory usage
def __init__(self, data, children=(), local_grads=()):
self.data = data # scalar value of this node calculated during forward pass
self.grad = 0 # derivative of the loss w.r.t. this node, calculated in backward pass
self._children = children # children of this node in the computation graph
self._local_grads = local_grads # local derivative of this node w.r.t. its children
def __add__(self, other):
other = other if isinstance(other, Value) else Value(other)
return Value(self.data + other.data, (self, other), (1, 1))
def __mul__(self, other):
other = other if isinstance(other, Value) else Value(other)
return Value(self.data * other.data, (self, other), (other.data, self.data))
def __pow__(self, other): return Value(self.data**other, (self,), (other * self.data**(other-1),))
def log(self): return Value(math.log(self.data), (self,), (1/self.data,))
def exp(self): return Value(math.exp(self.data), (self,), (math.exp(self.data),))
def relu(self): return Value(max(0, self.data), (self,), (float(self.data > 0),))
def __neg__(self): return self * -1
def __radd__(self, other): return self + other
def __sub__(self, other): return self + (-other)
def __rsub__(self, other): return other + (-self)
def __rmul__(self, other): return self * other
def __truediv__(self, other): return self * other**-1
def __rtruediv__(self, other): return other * self**-1
def backward(self):
topo = []
visited = set()
def build_topo(v):
if v not in visited:
visited.add(v)
for child in v._children:
build_topo(child)
topo.append(v)
build_topo(self)
self.grad = 1
for v in reversed(topo):
for child, local_grad in zip(v._children, v._local_grads):
child.grad += local_grad * v.grad
# A deep stack of 20 small layers, each "multiply by 0.1, then relu".
def layer(x):
return (x * 0.1).relu()
def plain_stack(x, depth=20):
for _ in range(depth):
x = layer(x)
return x
# The residual version: each layer ADDS its output to what came in
# (microgpt lines 134 and 141: x = [a + b for a, b in zip(x, x_residual)]).
def residual_stack(x, depth=20):
for _ in range(depth):
x = layer(x) # TODO: add the layer's output to x instead of replacing it
return x
for build in [plain_stack, residual_stack]:
x = Value(1.0)
out = build(x)
out.backward()
print(f'{build.__name__:15s} output {out.data:.3e} gradient reaching the input {x.grad:.3e}')
print('(1.000e-20 means 1.000 × 10 to the power -20: twenty places after the decimal point)')
4Check yourself
5Unlocked in microgpt
Lines 116 and 117 save the input and normalise; line 134 adds it back after attention. The MLP block repeats the pattern (lines 136 to 141, in 9.2). With this, every attention line is yours except the head split and the output matrix (lines 124 to 128 and 132 to 133, lesson 9.1) and the layer loop (lines 113 to 115, lesson 9.3).
Progress is saved in this browser. Sign in to keep it across devices.