Building a Transformer from Scratch: Attention, Architecture, and Memory
A ground-up implementation of the Transformer: scaled dot-product attention in NumPy and PyTorch, the full decoder block with RoPE, RMSNorm, and causal masking, then the memory optimizations that make large models practical: FlashAttention, KV caching, GQA, and LoRA.
An RNN folds the entire input sequence into a single fixed-size vector before making a prediction. Attention removes that constraint. The model computes similarity weights between tokens and retrieves information across the full sequence at every step.
Mathematically, scaled dot-product attention — the formulation Vaswani et al. introduced in Attention Is All You Need in 2017 — maps query vectors against key-value pairs:
This post builds the mechanism from raw tensor operations, assembles it into a full decoder block, then works through the optimizations that make it run at scale.
Part 1: Scaled dot-product attention
The query, key, value abstraction
Every token simultaneously acts as a query, a key, and a value. When the model processes token , its query vector is compared against every key vector via an inner product. Softmax normalizes those scores into a probability distribution, which weights a sum over all value vectors .
Every arrow is a batched matrix multiplication. The whole operation parallelizes across GPU tensor cores with no recurrent loop.
import numpy as np
import torch
import torch.nn as nn
import math
seq_len = 5
d_model = 8
d_k = d_model
input_embeddings = np.random.rand(1, seq_len, d_model)Step 1: Project into Q, K, V spaces
# Learned projection weights
W_Q = np.random.randn(d_model, d_k) * 0.01
W_K = np.random.randn(d_model, d_k) * 0.01
W_V = np.random.randn(d_model, d_k) * 0.01
X = input_embeddings[0] # (seq_len, d_model)
Q = X @ W_Q # (seq_len, d_k)
K = X @ W_K
V = X @ W_V
print(f"Q shape: {Q.shape}") # (5, 8)Step 2: Compute attention scores and apply causal mask
Without scaling, dot products grow with , pushing softmax into regions where gradients vanish.
scores = Q @ K.T / np.sqrt(d_k) # (seq_len, seq_len)
# Causal mask: token i cannot attend to j > i
mask = np.triu(np.ones((seq_len, seq_len)), k=1) * -1e9
scores += mask
def softmax(x):
x -= x.max(axis=-1, keepdims=True)
e = np.exp(x)
return e / e.sum(axis=-1, keepdims=True)
attention_weights = softmax(scores) # (seq_len, seq_len)
output = attention_weights @ V # (seq_len, d_k)
print(f"Attention output shape: {output.shape}") # (5, 8)Step 3: Multi-head attention in PyTorch
Multi-head attention runs attention heads in parallel, each projecting into a lower-dimensional subspace (), then concatenates and projects the results:
class MultiHeadAttention(nn.Module):
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.head_dim = 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, mask: torch.Tensor | None = None) -> torch.Tensor:
b, s, _ = x.shape
Q = self.q_proj(x).view(b, s, self.n_heads, self.head_dim).transpose(1, 2)
K = self.k_proj(x).view(b, s, self.n_heads, self.head_dim).transpose(1, 2)
V = self.v_proj(x).view(b, s, self.n_heads, self.head_dim).transpose(1, 2)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
weights = torch.softmax(scores, dim=-1)
out = torch.matmul(weights, V)
out = out.transpose(1, 2).contiguous().view(b, s, self.d_model)
return self.out_proj(out)Part 2: The full decoder block
RMSNorm and Pre-LN residuals
Modern LLMs use Pre-LN placement, normalizing before each sub-layer rather than after. This keeps gradient flow stable at large depths. RMSNorm, from Zhang and Sennrich's Root Mean Square Layer Normalization (2019), replaces LayerNorm's mean subtraction with root-mean-square normalization only — their hypothesis was that LayerNorm's re-centering invariance is dispensable and only the re-scaling matters:
class RMSNorm(nn.Module):
def __init__(self, d_model: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(d_model))
def forward(self, x: torch.Tensor) -> torch.Tensor:
rms = torch.sqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
return (x / rms) * self.weightRotary Position Embeddings (RoPE)
RoPE, introduced by Su et al. in RoFormer (2021), encodes position by rotating Q and K vectors in 2D subspaces. Unlike absolute positional embeddings, rotation preserves the relative distance between tokens and extends naturally to longer sequences:
def precompute_rope_freqs(head_dim: int, max_seq_len: int, base: float = 10000.0) -> torch.Tensor:
theta = 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim))
t = torch.arange(max_seq_len)
freqs = torch.outer(t, theta)
return torch.polar(torch.ones_like(freqs), freqs)
def apply_rope(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
freqs_cis = freqs_cis[:x.shape[-2], :].unsqueeze(0).unsqueeze(0)
x_rotated = torch.view_as_real(x_complex * freqs_cis).flatten(-2)
return x_rotated.type_as(x)Feed-forward with SwiGLU
The SwiGLU activation, defined in Noam Shazeer's GLU Variants Improve Transformer (2020), replaces ReLU in modern decoder FFNs. It uses a gating mechanism that allows the network to suppress activations selectively:
class FeedForward(nn.Module):
def __init__(self, d_model: int, d_ff: int):
super().__init__()
self.gate_proj = nn.Linear(d_model, d_ff, bias=False)
self.up_proj = nn.Linear(d_model, d_ff, bias=False)
self.down_proj = nn.Linear(d_ff, d_model, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.down_proj(
torch.nn.functional.silu(self.gate_proj(x)) * self.up_proj(x)
)Full decoder block
class TransformerDecoderBlock(nn.Module):
def __init__(self, d_model: int, n_heads: int, d_ff: int):
super().__init__()
self.attn = MultiHeadAttention(d_model, n_heads)
self.ff = FeedForward(d_model, d_ff)
self.norm1 = RMSNorm(d_model)
self.norm2 = RMSNorm(d_model)
def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None) -> torch.Tensor:
# Pre-LN: normalize before each sub-layer, add residual after
x = x + self.attn(self.norm1(x), mask)
x = x + self.ff(self.norm2(x))
return xPart 3: Memory and compute optimizations
KV caching
During autoregressive generation, each new token recomputes attention over the full sequence history. KV caching stores past key and value tensors and appends to them, reducing per-step compute from to :
class KVCache:
def __init__(self):
self.k_cache: list[torch.Tensor] = []
self.v_cache: list[torch.Tensor] = []
def update(self, k: torch.Tensor, v: torch.Tensor):
self.k_cache.append(k)
self.v_cache.append(v)
def get(self) -> tuple[torch.Tensor, torch.Tensor]:
return torch.cat(self.k_cache, dim=2), torch.cat(self.v_cache, dim=2)
def clear(self):
self.k_cache.clear()
self.v_cache.clear()Grouped Query Attention (GQA)
Multi-head attention with heads stores key and value tensors per layer. GQA — Ainslie et al., 2023 — reduces this by sharing one key-value head across query heads. Llama 2 adopts it for its two largest models; the released 70B config sets 8 KV heads for 64 query heads, an 8x reduction in KV cache memory at inference:
class GroupedQueryAttention(nn.Module):
def __init__(self, d_model: int, n_query_heads: int, n_kv_heads: int):
super().__init__()
assert n_query_heads % n_kv_heads == 0
self.n_query_heads = n_query_heads
self.n_kv_heads = n_kv_heads
self.n_rep = n_query_heads // n_kv_heads
self.head_dim = d_model // n_query_heads
self.q_proj = nn.Linear(d_model, d_model, bias=False)
self.k_proj = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False)
self.out_proj = nn.Linear(d_model, d_model, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
b, s, _ = x.shape
Q = self.q_proj(x).view(b, s, self.n_query_heads, self.head_dim).transpose(1, 2)
K = self.k_proj(x).view(b, s, self.n_kv_heads, self.head_dim).transpose(1, 2)
V = self.v_proj(x).view(b, s, self.n_kv_heads, self.head_dim).transpose(1, 2)
# Repeat KV heads to match query head count
K = K.repeat_interleave(self.n_rep, dim=1)
V = V.repeat_interleave(self.n_rep, dim=1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim)
weights = torch.softmax(scores, dim=-1)
out = torch.matmul(weights, V)
out = out.transpose(1, 2).contiguous().view(b, s, -1)
return self.out_proj(out)FlashAttention: fused tiled SRAM kernels
Standard attention materializes the full attention matrix in HBM (high-bandwidth memory). For a 4,096-token sequence with 32 heads, that's ~2GB of intermediate tensors per layer, with multiple round trips between HBM and SRAM.
FlashAttention (Dao, Fu, Ermon, Rudra and Ré, 2022) fuses the three attention operations (QKT, softmax, AV) into a single kernel that tiles Q, K, V into SRAM blocks, computes attention locally, and accumulates the running softmax normalization without writing the full attention matrix back to HBM.
The kernel implements a numerically stable online softmax:
In practice: torch.nn.functional.scaled_dot_product_attention uses FlashAttention kernels automatically when inputs are in the right format. As of PyTorch 2.14, the docs list three backends it dispatches across — FlashAttention-2, a memory-efficient kernel, and a C++ implementation matching the formula above — selected automatically from the inputs. No manual kernel writing needed for most applications.
Peak memory drops from to . On an A100 with a 2,048-token sequence, FlashAttention 2 achieves approximately 2.2x speedup and 5-10x memory reduction compared to standard attention.
Low-Rank Adaptation (LoRA)
LoRA, from Hu et al., 2021, freezes the pre-trained weight matrix and injects a low-rank decomposition where , , with :
Only and are trained. A rank-16 adapter on a 4,096-dimensional projection reduces trainable parameters from 16.7M to 131K, a 128x reduction.
class LoRALinear(nn.Module):
def __init__(self, in_features: int, out_features: int, rank: int = 16, alpha: float = 16.0):
super().__init__()
self.weight = nn.Parameter(torch.empty(out_features, in_features), requires_grad=False)
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
self.lora_A = nn.Parameter(torch.randn(rank, in_features) * 0.01)
self.lora_B = nn.Parameter(torch.zeros(out_features, rank))
self.scale = alpha / rank
def forward(self, x: torch.Tensor) -> torch.Tensor:
base_out = nn.functional.linear(x, self.weight)
lora_out = nn.functional.linear(nn.functional.linear(x, self.lora_A), self.lora_B)
return base_out + self.scale * lora_outScale with alpha / rank instead of a raw learning rate. It decouples adapter sensitivity from rank choice, so you can change rank without re-tuning the learning rate.
Memory and compute summary
| Optimization | Memory impact | Compute impact |
|---|---|---|
| KV caching | KV grows linearly with sequence length | Per-step compute drops from to |
| GQA | KV cache reduced by | Minimal compute reduction |
| FlashAttention | peak SRAM | 2x to 3x throughput on A100 for long sequences |
| LoRA | ~1% of full fine-tune parameters | Same forward pass cost; backward over only |
| INT8 quantization | ~50% model size reduction | Minor throughput gain; small accuracy drop on some tasks |
When this is the wrong choice
- You are shipping, not learning.
torch.nn.functional.scaled_dot_product_attentionalready dispatches to FlashAttention kernels when the inputs are in the right format. The NumPy and PyTorch blocks above exist to show the mechanism. Writing your own attention into a production path buys you a maintenance burden and the same math. - Your sequences are short. Everything in Part 3 buys headroom that only appears as grows. At a few hundred tokens the intermediate is not what is costing you, and KV caching, GQA, and tiled kernels are three more moving parts between you and a working model.
- You copied the
KVCacheabove into a serving loop. It appends to a list and concatenates the whole history on everyget(), so the copy grows with exactly the history the cache was supposed to stop recomputing. It is a teaching sketch. A serving cache preallocates. - The base weights are wrong about your task. LoRA trains 131K parameters against a frozen projection. That is the right trade when you are moving style or a narrow domain. When the pre-trained model does not represent the task at all, a low-rank update on top of it will not put the representation there.
Part 2 of this series covers encoder vs. decoder architecture differences: BERT's bidirectional masking against GPT's causal masking, their pre-training objectives, and the structural consequences for downstream tasks.