π» Modern ArchitecturesLEVEL 3 Β· ADVANCEDChapter 7
Scaled Dot-Product & Multi-Head Attention
Query-Key-Value routing, scaling factor sqrt(d_k), causal autoregressive masks, and RoPE.
π Inspect Architecture: Attention mechanismπ‘ 1. Core Intuition & Concepts
Self-attention computes dynamic, data-dependent routing coefficients across all sequence positions. Each token projects into Query, Key, and Value vectors, computing similarity dot products scaled by 1/βd_k to prevent gradient saturation.
π 2. Mathematical Formulations & Derivations
Scaled Dot-Product Attention Formula
Q: queries, K: keys, V: values, M: causal triangular mask (-inf above diagonal).
Rotary Position Embedding (RoPE)
Encodes relative token distances via 2D rotation of Query/Key orthogonal coordinate pairs.
βοΈ 3. Step-by-Step Computational Mechanism
1
Q, K, V Linear Projections
Project input sequence X in R^(B x T x D) into Q, K, V and reshape into h parallel attention heads.
2
Apply RoPE Rotations
Rotate query and key head vectors with frequency angles ΞΈ_i = 10000^(-2(i-1)/d) to inject relative position.
3
Causal Scaled Dot-Product
Compute scores (Q K^T)/sqrt(d_k), apply upper-triangular -inf mask, and take Softmax.
4
Aggregate Values & Output Projection
Multiply attention weights by V and project concatenated heads back to hidden dimension D.
π» 4. Code from Scratch (python)
mha.py
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttentionScratch(nn.Module):
"""Multi-Head Self-Attention with Causal Masking in pure PyTorch."""
def __init__(self, d_model: int, n_heads: int):
super().__init__()
assert d_model % n_heads == 0
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
self.q_proj = nn.Linear(d_model, d_model, bias=False)
self.k_proj = nn.Linear(d_model, d_model, bias=False)
self.v_proj = nn.Linear(d_model, d_model, bias=False)
self.out_proj = nn.Linear(d_model, d_model, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, T, C = x.shape
q = self.q_proj(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
mask = torch.triu(torch.full((T, T), float("-inf"), device=x.device), diagonal=1)
scores = scores + mask
attn_weights = F.softmax(scores, dim=-1)
out = torch.matmul(attn_weights, v)
out = out.transpose(1, 2).contiguous().view(B, T, C)
return self.out_proj(out)π§ 5. Comprehension Checkpoint
Answer all 1 questions correctly to complete the chapter Β· 0 / 1 done
Q1/1 What is the computational complexity of standard self-attention with respect to sequence length T?
In the catalog
Finished this chapter?