mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
feat(phase-07/01): why transformers — the problems with RNNs
This commit is contained in:
@@ -0,0 +1,100 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 900 460" font-family="Georgia, 'Times New Roman', serif">
|
||||
<defs>
|
||||
<marker id="arrow" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="6" markerHeight="6" orient="auto">
|
||||
<path d="M0,0 L10,5 L0,10 z" fill="#1a1a1a"/>
|
||||
</marker>
|
||||
<style>
|
||||
.box { fill: #faf6ef; stroke: #1a1a1a; stroke-width: 1.5; }
|
||||
.hot { fill: #fff1d6; stroke: #c0392b; stroke-width: 1.5; }
|
||||
.label { font-size: 14px; font-weight: 600; fill: #1a1a1a; }
|
||||
.content { font-size: 12px; fill: #333; font-family: 'Menlo', monospace; }
|
||||
.caption { font-size: 11px; fill: #555; font-style: italic; }
|
||||
.title { font-size: 16px; font-weight: 700; fill: #1a1a1a; }
|
||||
</style>
|
||||
</defs>
|
||||
|
||||
<text x="450" y="30" text-anchor="middle" class="title">serial depth: the only thing a GPU actually cares about</text>
|
||||
|
||||
<!-- RNN row -->
|
||||
<text x="40" y="80" class="label">RNN</text>
|
||||
<text x="40" y="98" class="caption">depth = N</text>
|
||||
|
||||
<rect x="100" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="140" y="94" text-anchor="middle" class="content">h_1</text>
|
||||
|
||||
<line x1="180" y1="90" x2="210" y2="90" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
|
||||
<rect x="210" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="250" y="94" text-anchor="middle" class="content">h_2</text>
|
||||
|
||||
<line x1="290" y1="90" x2="320" y2="90" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
|
||||
<rect x="320" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="360" y="94" text-anchor="middle" class="content">h_3</text>
|
||||
|
||||
<line x1="400" y1="90" x2="430" y2="90" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
|
||||
<rect x="430" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="470" y="94" text-anchor="middle" class="content">h_4</text>
|
||||
|
||||
<line x1="510" y1="90" x2="540" y2="90" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
|
||||
<rect x="540" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="580" y="94" text-anchor="middle" class="content">h_5</text>
|
||||
|
||||
<line x1="620" y1="90" x2="650" y2="90" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
|
||||
<rect x="650" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="690" y="94" text-anchor="middle" class="content">h_6</text>
|
||||
|
||||
<line x1="730" y1="90" x2="760" y2="90" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
|
||||
<rect x="760" y="70" width="80" height="40" class="hot"/>
|
||||
<text x="800" y="94" text-anchor="middle" class="content">h_7</text>
|
||||
|
||||
<text x="450" y="145" text-anchor="middle" class="caption">each box must wait for the one to its left — 7 steps of wall-clock</text>
|
||||
|
||||
<!-- Divider -->
|
||||
<line x1="40" y1="180" x2="860" y2="180" stroke="#999" stroke-width="0.5" stroke-dasharray="4,3"/>
|
||||
|
||||
<!-- Transformer row -->
|
||||
<text x="40" y="230" class="label">Transformer</text>
|
||||
<text x="40" y="248" class="caption">depth = 1</text>
|
||||
|
||||
<rect x="100" y="220" width="80" height="40" class="box"/>
|
||||
<text x="140" y="244" text-anchor="middle" class="content">x_1</text>
|
||||
|
||||
<rect x="210" y="220" width="80" height="40" class="box"/>
|
||||
<text x="250" y="244" text-anchor="middle" class="content">x_2</text>
|
||||
|
||||
<rect x="320" y="220" width="80" height="40" class="box"/>
|
||||
<text x="360" y="244" text-anchor="middle" class="content">x_3</text>
|
||||
|
||||
<rect x="430" y="220" width="80" height="40" class="box"/>
|
||||
<text x="470" y="244" text-anchor="middle" class="content">x_4</text>
|
||||
|
||||
<rect x="540" y="220" width="80" height="40" class="box"/>
|
||||
<text x="580" y="244" text-anchor="middle" class="content">x_5</text>
|
||||
|
||||
<rect x="650" y="220" width="80" height="40" class="box"/>
|
||||
<text x="690" y="244" text-anchor="middle" class="content">x_6</text>
|
||||
|
||||
<rect x="760" y="220" width="80" height="40" class="box"/>
|
||||
<text x="800" y="244" text-anchor="middle" class="content">x_7</text>
|
||||
|
||||
<!-- all-to-all attention arrows (simplified, from middle down) -->
|
||||
<line x1="140" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="0.8" opacity="0.4"/>
|
||||
<line x1="250" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="0.8" opacity="0.4"/>
|
||||
<line x1="360" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="0.8" opacity="0.4"/>
|
||||
<line x1="470" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="1.5" marker-end="url(#arrow)"/>
|
||||
<line x1="580" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="0.8" opacity="0.4"/>
|
||||
<line x1="690" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="0.8" opacity="0.4"/>
|
||||
<line x1="800" y1="260" x2="470" y2="305" stroke="#c0392b" stroke-width="0.8" opacity="0.4"/>
|
||||
|
||||
<rect x="380" y="310" width="180" height="38" class="hot"/>
|
||||
<text x="470" y="334" text-anchor="middle" class="content">softmax(QK^T / √d) V</text>
|
||||
|
||||
<text x="450" y="378" text-anchor="middle" class="caption">every x_i contributes to every output in a single matmul — 1 step of wall-clock</text>
|
||||
|
||||
<text x="450" y="440" text-anchor="middle" class="caption">same op count. different dependency graph. that is the entire thesis of Attention Is All You Need.</text>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 5.0 KiB |
@@ -0,0 +1,110 @@
|
||||
"""Why Transformers - demonstrate the serial-depth gap between RNN-style
|
||||
recurrence and attention-style parallel reduction.
|
||||
|
||||
Runs in pure stdlib. No numpy, no torch.
|
||||
"""
|
||||
|
||||
import math
|
||||
import time
|
||||
|
||||
|
||||
def rnn_style(xs, decay=0.9):
|
||||
"""Sequential recurrence: h_t depends on h_{t-1}. Cannot parallelize."""
|
||||
h = 0.0
|
||||
for x in xs:
|
||||
h = decay * h + x
|
||||
return h
|
||||
|
||||
|
||||
def attention_style(xs):
|
||||
"""Order-independent reduction: every element is independent."""
|
||||
return sum(xs) / len(xs)
|
||||
|
||||
|
||||
def serial_scan(xs):
|
||||
"""Prefix sum computed serially. Depth O(N)."""
|
||||
out = []
|
||||
acc = 0.0
|
||||
for x in xs:
|
||||
acc += x
|
||||
out.append(acc)
|
||||
return out
|
||||
|
||||
|
||||
def parallel_scan(xs):
|
||||
"""Hillis-Steele parallel prefix sum. Depth O(log N).
|
||||
|
||||
In pure Python each step is still serial, but the data-dependency
|
||||
graph has depth log2(N). On real hardware with N-wide SIMD this
|
||||
gets you a log-depth scan; on a CPU it's the same wall-clock but
|
||||
the graph shape is what matters for GPU kernels.
|
||||
"""
|
||||
out = list(xs)
|
||||
step = 1
|
||||
n = len(out)
|
||||
while step < n:
|
||||
new = list(out)
|
||||
for i in range(step, n):
|
||||
new[i] = out[i] + out[i - step]
|
||||
out = new
|
||||
step *= 2
|
||||
return out
|
||||
|
||||
|
||||
def benchmark(n, reps=3):
|
||||
xs = [0.001 * (i % 17) for i in range(n)]
|
||||
|
||||
best_rnn = math.inf
|
||||
for _ in range(reps):
|
||||
t0 = time.perf_counter()
|
||||
_ = rnn_style(xs)
|
||||
best_rnn = min(best_rnn, time.perf_counter() - t0)
|
||||
|
||||
best_attn = math.inf
|
||||
for _ in range(reps):
|
||||
t0 = time.perf_counter()
|
||||
_ = attention_style(xs)
|
||||
best_attn = min(best_attn, time.perf_counter() - t0)
|
||||
|
||||
return best_rnn, best_attn
|
||||
|
||||
|
||||
def depth(n):
|
||||
"""Serial-depth count for RNN vs attention-style reductions."""
|
||||
rnn_depth = n
|
||||
attn_depth = max(1, math.ceil(math.log2(n)))
|
||||
return rnn_depth, attn_depth
|
||||
|
||||
|
||||
def main():
|
||||
print("=== serial-depth comparison ===")
|
||||
print(f"{'N':>8} {'rnn depth':>12} {'attn depth':>12} {'speedup (ops)':>16}")
|
||||
for n in [64, 512, 4096, 32768, 262144]:
|
||||
rd, ad = depth(n)
|
||||
print(f"{n:>8} {rd:>12} {ad:>12} {rd / ad:>15.0f}x")
|
||||
|
||||
print()
|
||||
print("=== wall-clock on this machine (pure Python) ===")
|
||||
print(f"{'N':>8} {'rnn (ms)':>10} {'attn (ms)':>10} {'ratio':>8}")
|
||||
for n in [1_000, 10_000, 100_000, 1_000_000]:
|
||||
rnn_t, attn_t = benchmark(n)
|
||||
ratio = rnn_t / attn_t if attn_t > 0 else float("inf")
|
||||
print(f"{n:>8} {rnn_t * 1000:>10.2f} {attn_t * 1000:>10.2f} {ratio:>7.1f}x")
|
||||
|
||||
print()
|
||||
print("=== prefix-sum equivalence check ===")
|
||||
xs = [float(i) for i in range(16)]
|
||||
ser = serial_scan(xs)
|
||||
par = parallel_scan(xs)
|
||||
mismatches = sum(1 for a, b in zip(ser, par) if abs(a - b) > 1e-9)
|
||||
print(f"length: {len(xs)}, mismatches between serial and parallel scan: {mismatches}")
|
||||
print(f"last value (serial): {ser[-1]}")
|
||||
print(f"last value (parallel): {par[-1]}")
|
||||
|
||||
print()
|
||||
print("takeaway: attention wins on every dimension but memory.")
|
||||
print("memory cost is O(N^2) for full attention; Lesson 12 covers the fixes.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,108 @@
|
||||
# Why Transformers — The Problems with RNNs
|
||||
|
||||
> RNNs process tokens one at a time. Transformers process all tokens at once. That single architectural bet changed every scaling curve in deep learning after 2017.
|
||||
|
||||
**Type:** Learn
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 3 (Deep Learning Core), Phase 5 · 09 (Sequence-to-Sequence), Phase 5 · 10 (Attention Mechanism)
|
||||
**Time:** ~45 minutes
|
||||
|
||||
## The Problem
|
||||
|
||||
Before 2017, every state-of-the-art sequence model on the planet — language, translation, speech — was a recurrent neural network. LSTMs and GRUs won ImageNet-equivalent translation benchmarks for half a decade. They were the only tool anyone had.
|
||||
|
||||
They had three fatal weaknesses. Sequential computation meant you could not parallelize along the time axis: token `t+1` needs the hidden state from token `t`. A 1,024-token sequence meant 1,024 serial steps on a GPU that can do 1,000,000 floating-point ops per cycle. Training wall-clock time scaled linearly with sequence length on hardware designed for parallelism.
|
||||
|
||||
Vanishing gradients meant information 50 tokens back was already compressed through 50 non-linearities. Gated recurrent units (LSTM, GRU) softened the crush but never eliminated it. Long-range dependencies — "the book I read last summer on a plane to Kyoto was…" — routinely failed.
|
||||
|
||||
Fixed-width hidden states meant the encoder squeezed the entire source sequence into a single vector before the decoder saw anything. Doesn't matter if the source is 5 tokens or 500; the bottleneck is the same shape.
|
||||
|
||||
The 2017 paper "Attention Is All You Need" proposed something radical: drop recurrence entirely. Let every position attend to every other position in parallel. Train in one big matrix multiplication instead of 1,024 sequential ones.
|
||||
|
||||
The result dominates every modality by 2026. Language (GPT-5, Claude 4, Llama 4), vision (ViT, DINOv2, SAM 3), audio (Whisper), biology (AlphaFold 3), robotics (RT-2). Same block, different inputs.
|
||||
|
||||
## The Concept
|
||||
|
||||

|
||||
|
||||
**Recurrence as a bottleneck.** An RNN computes `h_t = f(h_{t-1}, x_t)`. Each step depends on the previous. You cannot compute `h_5` before `h_4`. On modern GPUs with 10,000+ parallel cores, this wastes 99% of the silicon on a long sequence.
|
||||
|
||||
**Attention as a broadcast.** Self-attention computes `output_i = sum_j(a_ij * v_j)` for every pair `(i, j)` simultaneously. The whole N×N attention matrix fills in one batched matmul. No step depends on another. GPUs love it.
|
||||
|
||||
**The speedup is not a constant.** It is the difference between `O(N)` serial depth and `O(1)` serial depth. In practice, transformers train 5–10× faster per epoch on matched hardware at N=512, and the gap widens with sequence length until you hit the `O(N²)` memory wall of attention (which Flash Attention later fixed — see Lesson 12).
|
||||
|
||||
**What transformers cost.** Attention memory scales as `O(N²)`. For 2K context, fine. For 128K context, you need sliding windows, RoPE extrapolation, Flash Attention tiling, or linear attention variants. Recurrence was `O(N)` in both time and memory; transformers trade time for memory and then win the time back through parallelism.
|
||||
|
||||
**The inductive bias shift.** RNNs assume locality and recency. Transformers assume nothing — every pair is a candidate for attention. That is why transformers need more data to train well but scale further once they have it. Chinchilla (2022) formalized this: given enough tokens, a transformer always beats an RNN of equal parameter count.
|
||||
|
||||
## Build It
|
||||
|
||||
No neural network here — we simulate the core bottleneck numerically so you feel the gap on your laptop.
|
||||
|
||||
### Step 1: measure serial depth
|
||||
|
||||
See `code/main.py`. We build two functions. One encodes a sequence as a chain of additions (serial, like an RNN). One encodes it as a parallel reduction (broadcast, like attention). Same math, different dependency graph.
|
||||
|
||||
```python
|
||||
def rnn_style(xs):
|
||||
h = 0.0
|
||||
for x in xs:
|
||||
h = 0.9 * h + x # can't parallelize: h depends on previous h
|
||||
return h
|
||||
|
||||
def attention_style(xs):
|
||||
return sum(xs) / len(xs) # every x is independent
|
||||
```
|
||||
|
||||
We time both on sequences up to 100,000 elements. The RNN version is O(N) and a single CPU pipeline. Even in pure Python, the attention-style reduction beats it at length ≥ 1,000 because Python's `sum()` is implemented in C and iterates without interpreter overhead per step.
|
||||
|
||||
### Step 2: count theoretical operations
|
||||
|
||||
Both algorithms do N adds. The difference is *dependency depth*: how many operations must happen sequentially before the next can start. RNN depth = N. Attention depth = log(N) with a tree reduction, or 1 with a parallel scan. Depth, not op count, decides GPU time.
|
||||
|
||||
### Step 3: empirical scaling on long sequences
|
||||
|
||||
We print a timing table that makes the O(N) gap visible. On a 2026 Mac laptop, sequences under 1,000 elements are too fast to measure. Sequences of 100,000 show a clean linear scan. Scale that to a 16,384-token transformer with a 12-layer LSTM equivalent and you see why training wall-clock was a blocker in 2016.
|
||||
|
||||
## Use It
|
||||
|
||||
When to still pick an RNN in 2026:
|
||||
|
||||
| Situation | Pick |
|
||||
|-----------|------|
|
||||
| Streaming inference, one token at a time, constant memory | RNN or state-space model (Mamba, RWKV) |
|
||||
| Very long sequences (>1M tokens) where attention memory explodes | Linear attention, Mamba 2, Hyena |
|
||||
| Edge device with no matmul accelerator | Depthwise-separable RNN still wins on FLOPs/watt |
|
||||
| Anything else (training, batched inference, context up to 128K) | Transformer |
|
||||
|
||||
State-space models (SSMs) like Mamba are essentially RNNs with structured parameterization that gives them the best of both: `O(N)` scan memory, parallel training via selective scan. They recover 90% of transformer quality with better long-context scaling. In 2026 most frontier labs train hybrid SSM+transformer models (e.g. Jamba, Samba) — recurrence is not dead, it is a component.
|
||||
|
||||
## Ship It
|
||||
|
||||
See `outputs/skill-architecture-picker.md`. The skill picks an architecture for a new sequence problem given length, throughput, and training-budget constraints. It should always refuse to recommend a pure RNN for training runs above 1B tokens without stating the trade-off.
|
||||
|
||||
## Exercises
|
||||
|
||||
1. **Easy.** Take `rnn_style` from `code/main.py` and replace the scalar hidden state with a length-64 vector of hidden states. Re-measure. How much does the serial overhead grow with hidden-state dimension?
|
||||
2. **Medium.** Implement a parallel prefix-sum (Hillis-Steele scan) in pure Python. Verify it produces the same numerical output as a serial scan on length 1024. Count the depth.
|
||||
3. **Hard.** Port the attention-style reduction to PyTorch on GPU. Time both as you sweep sequence length from 64 to 65,536. Plot and explain the curve shape.
|
||||
|
||||
## Key Terms
|
||||
|
||||
| Term | What people say | What it actually means |
|
||||
|------|-----------------|-----------------------|
|
||||
| Recurrence | "RNNs are sequential" | Computation where step `t` depends on step `t-1`, forcing serial execution along the time axis. |
|
||||
| Serial depth | "How deep the graph is" | Longest chain of dependent ops; bounds wall-clock even on infinite hardware. |
|
||||
| Attention | "Let tokens look at each other" | Weighted sum `sum_j a_ij v_j` where `a_ij` comes from a similarity score between positions i and j. |
|
||||
| Context window | "How much the model sees" | Number of positions an attention layer can take as input; quadratic memory cost scales here. |
|
||||
| Inductive bias | "Assumptions baked into the architecture" | Prior about what the data looks like; CNNs assume translation invariance, RNNs assume recency. |
|
||||
| State-space model | "RNN with algebra behind it" | Recurrence parameterized for parallel training via structured state-space matrices. |
|
||||
| Quadratic bottleneck | "Why context costs so much" | Attention memory = `O(N²)` in sequence length; Flash Attention hides the constants, not the scaling. |
|
||||
|
||||
## Further Reading
|
||||
|
||||
- [Vaswani et al. (2017). Attention Is All You Need](https://arxiv.org/abs/1706.03762) — the paper that killed recurrence in mainstream NLP.
|
||||
- [Bahdanau, Cho, Bengio (2014). Neural MT by Jointly Learning to Align and Translate](https://arxiv.org/abs/1409.0473) — where attention was born, bolted onto an RNN.
|
||||
- [Hochreiter, Schmidhuber (1997). Long Short-Term Memory](https://www.bioinf.jku.at/publications/older/2604.pdf) — the original LSTM paper, for the record.
|
||||
- [Gu, Dao (2023). Mamba: Linear-Time Sequence Modeling with Selective State Spaces](https://arxiv.org/abs/2312.00752) — modern recurrent answer to transformers.
|
||||
- [Karpathy's "Unreasonable Effectiveness of Recurrent Neural Networks" (2015)](https://karpathy.github.io/2015/05/21/rnn-effectiveness/) — best intuition for what RNNs were used for before transformers.
|
||||
+18
@@ -0,0 +1,18 @@
|
||||
---
|
||||
name: sequence-architecture-picker
|
||||
description: Pick sequence architecture (RNN, transformer, SSM, hybrid) given length, throughput, and training budget.
|
||||
version: 1.0.0
|
||||
phase: 7
|
||||
lesson: 1
|
||||
tags: [transformers, architecture, rnn, ssm]
|
||||
---
|
||||
|
||||
Given a sequence problem (max length, batch shape, training tokens budgeted, inference latency target, device class), output:
|
||||
|
||||
1. Primary architecture. One of: transformer, state-space model (Mamba/RWKV), hybrid SSM+attention, RNN. One-sentence reason tied to the dominant constraint.
|
||||
2. Context length strategy. If transformer: full attention cutoff, sliding window size, RoPE scaling factor. If SSM: scan chunk size. If RNN: hidden width.
|
||||
3. Training FLOP profile. Approximate FLOPs per token from architecture + context; note whether the spec fits the compute budget.
|
||||
4. Inference memory profile. KV cache for transformers, state size for SSMs, per-token memory for RNNs. Flag if the target device can hold a single batch of 1.
|
||||
5. Risk note. One specific failure mode that this choice is known to have at the scale of the spec (e.g. transformer OOM at 64K context on a 24GB GPU without Flash Attention).
|
||||
|
||||
Refuse to recommend a pure RNN for any training run above 1B tokens without explicitly stating the gradient-flow and parallelism penalties. Refuse to recommend a full-attention transformer for >64K context without stating the `O(N^2)` memory cost. Refuse to recommend a brand-new architecture (published <12 months ago) for production without a named fallback.
|
||||
Reference in New Issue
Block a user