Module 12: Attention

Attention’s O(N²) compute cost is only half the story. Its O(N²) memory footprint for the scores matrix and O(N) KV-cache growth during decoding are what actually limit real inference workloads, turning attention from a FLOP problem into an HBM bandwidth problem. This module builds scaled dot-product attention, multi-head attention, and causal masking from scratch so you can see exactly where the quadratic wall comes from before FlashAttention and friends try to climb it.

NoteModule Info

ARCHITECTURE TIER | Difficulty: ●●●○ | Time: 5-7 hours | Prerequisites: 01-08, 10-11

Prerequisites: Modules 01-08 and 10-11 means you should understand:

  • Tensor operations and shape manipulation (Module 01)
  • Activations, particularly softmax (Module 02)
  • Linear layers and weight projections (Module 03)
  • Autograd for gradient computation (Module 06)
  • Tokenization and embeddings (Modules 10-11)

If you can explain why softmax(x).sum(axis=-1) equals 1.0 and how embeddings convert token IDs to dense vectors, you’re ready.

Audio overview (AI-generated)
Open in Binder → Runs in your browser with nothing to install; the session is discarded when you leave, so download your notebook to keep it.
Lecture slides AI-generated · opens an in-page viewer
🔥 Slide Deck · AI-generated
1 / -
Loading slides...

Overview

Attention is the mechanism behind GPT, BERT, and every modern LLM. In this module you build it from scratch — scaled dot-product attention and multi-head attention — the same math that runs in production transformers, written in NumPy so you can read every line.

The shift attention introduced is simple to state. RNNs squeeze a whole sequence into one fixed-size hidden state and hope nothing important falls out. Attention lets every position read directly from every other position, weighted by relevance computed on the fly. That’s it. The cost of that freedom is the equation you’ll implement: Attention(Q, K, V) = softmax(QK^T / √d_k) V. The QK^T term creates an n × n matrix — quadratic in sequence length. By the end of this module you’ll have written that matrix, watched it dominate memory at long context, and understood exactly why FlashAttention and friends exist.

Commands

# first time
tito module start 12

# later sessions
tito module resume 12

# when your tests pass
tito module complete 12

Your notebook is modules/12_attention/attention.ipynb.

Learning objectives

TipBy completing this module, you will:
  • Implement scaled dot-product attention with vectorized operations that reveal O(n²) memory complexity
  • Build multi-head attention for parallel processing of different relationship types across representation subspaces
  • Master attention weight computation, normalization, and the query-key-value paradigm
  • Understand quadratic memory scaling and why attention becomes the bottleneck in long-context transformers
  • Connect your implementation to production frameworks and understand why efficient attention research matters at scale

What you’ll build

Figure 1: Scaled dot-product attention. Per-head Q and K produce scaled scores. An optional mask blocks positions before softmax over keys; multiplying weights by V produces the per-head output. MultiHeadAttention first projects and splits heads, then merges head outputs and applies an output projection. For self-attention, the implementation materializes S × S scores and weights for each batch item and head.

The pattern you’ll enable:

# Multi-head attention for sequence processing
mha = MultiHeadAttention(embed_dim=512, num_heads=8)
output = mha(embeddings, mask)  # Learn different relationship types in parallel

What you’re not building yet

To keep this module focused, you will not implement:

  • Full transformer blocks (that’s Module 13: Transformers)
  • Positional encoding (you built this in Module 11: Embeddings)
  • Efficient attention variants like FlashAttention (production optimization beyond scope)
  • Cross-attention for encoder-decoder models (PyTorch does this with separate Q vs K/V inputs)

You are building the core attention mechanism. Complete transformer architectures come next.

What you write

The notebook arrives with the surrounding code already written and explained. You write 2 functions, each marked # YOUR CODE HERE and followed by a test cell:

scaled_dot_product_attention
Complete scaled dot-product attention.
MultiHeadAttention.forward
Forward pass through multi-head attention.

How you know it works

tito module complete stops at the first step that fails:

  1. the unit tests inside your notebook run;
  2. your code is exported into tinytorch.core.attention;
  3. the integration tests run against that exported package, together with the modules before it;
  4. the module is recorded as done, and tito module status shows it.

Unit tests in your notebook (7). Each prints a ✅ line when it passes.

  • Attention Scores
  • Score Scaling
  • Causal Masking
  • Scaled Dot-Product Attention
  • Split Heads
  • Merge Heads
  • Multi-Head Attention (End-to-End)

Integration tests after export (17).

  • tests/12_attention/test_12_attention_progressive.py

When it fails

A bare NotImplementedError with no message means a cell reached a function you have not written yet: the notebook ships each one as # YOUR CODE HERE followed by raise NotImplementedError(). The messages below are ones this module actually prints when an implementation is present but wrong.

Attention weights don't sum to 1
From test_unit_scaled_dot_product_attention. The softmax ran over queries. Normalize over keys, the last axis: dim=-1.
Future attention not masked at (0,1)
From test_unit_scaled_dot_product_attention. The mask was not applied. Before the softmax, pass the scores through _apply_mask(scores, mask) whenever a mask is given.

Finished? Read why

The reasoning behind this module (why it is built this way, what it costs, and how production frameworks differ) is the chapter Attention Mechanisms: Scaled Dot-Product and Causal Masking in the companion book, TinyTorch: From Tensors to Transformers (PDF). The book prints complete reference implementations, so read it after you finish the module, not while you are working on it.

What’s next

NoteUp next: Module 13, Transformers

You have built the engine. Module 13 builds the chassis around it. The question Module 13 answers is: what do you wrap attention in to actually train it? Bare attention has two problems — its outputs collapse to similar values across positions (no per-token computation), and stacking it deeply makes gradients vanish. The transformer block fixes both, with feed-forward networks for per-position transformation, layer normalization to keep activations well-scaled, and residual connections to keep gradients flowing through arbitrarily many layers. Stack the result and you have GPT.

Next: Module 13: Transformers

How later modules use this one

Table 1: How attention feeds into the transformer block assembly.
Module What It Does Your Attention In Action
13: Transformers Complete transformer blocks TransformerBlock(attention + MLP + LayerNorm)
13: Transformers Residual connections x + attention(x) keeps gradients flowing
13: Transformers Stacked layers attention → FFN → attention → FFN…
Back to top