Back to blog

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.

May 10, 2024Updated September 08, 2026

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:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

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 ii, its query vector qiq_i is compared against every key vector kjk_j via an inner product. Softmax normalizes those scores into a probability distribution, which weights a sum over all value vectors vjv_j.

Every arrow is a batched matrix multiplication. The whole operation parallelizes across GPU tensor cores with no recurrent loop.

python
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

python
# 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

scores=QKTdk\text{scores} = \frac{QK^T}{\sqrt{d_k}}

Without scaling, dot products grow with dkd_k, pushing softmax into regions where gradients vanish.

python
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 hh attention heads in parallel, each projecting into a lower-dimensional subspace (dk=dmodel/hd_k = d_{\text{model}} / h), then concatenates and projects the results:

MultiHead(Q,K,V)=Concat(head1,…,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h)W^O
python
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:

RMSNorm(x)=xRMS(x)⋅γ,RMS(x)=1d∑i=1dxi2\text{RMSNorm}(x) = \frac{x}{\text{RMS}(x)} \cdot \gamma, \quad \text{RMS}(x) = \sqrt{\frac{1}{d} \sum_{i=1}^{d} x_i^2}
python
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.weight

Rotary 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:

qm⋅kn=Re[∑jqm,jkn,j∗ei(m−n)θj]q_m \cdot k_n = \text{Re}\left[\sum_{j} q_{m,j} k_{n,j}^* e^{i(m-n)\theta_j}\right]
python
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:

SwiGLU(x)=SiLU(W1x)⊗W3x\text{SwiGLU}(x) = \text{SiLU}(W_1 x) \otimes W_3 x
python
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

python
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 x

Part 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 O(S2)O(S^2) to O(S)O(S):

python
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 hh heads stores hh key and value tensors per layer. GQA — Ainslie et al., 2023 — reduces this by sharing one key-value head across GG 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:

python
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 S×SS \times S 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:

mi(j)=max⁡(mi(j−1),max⁡lSil(j))m_i^{(j)} = \max(m_i^{(j-1)}, \max_l S_{il}^{(j)}) ℓi(j)=emi(j−1)−mi(j)ℓi(j−1)+∑leSil(j)−mi(j)\ell_i^{(j)} = e^{m_i^{(j-1)} - m_i^{(j)}} \ell_i^{(j-1)} + \sum_l e^{S_{il}^{(j)} - m_i^{(j)}}

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 O(S2)O(S^2) to O(S)O(S). 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 W0∈Rd×kW_0 \in \mathbb{R}^{d \times k} and injects a low-rank decomposition ΔW=BA\Delta W = BA where B∈Rd×rB \in \mathbb{R}^{d \times r}, A∈Rr×kA \in \mathbb{R}^{r \times k}, with r≪min⁡(d,k)r \ll \min(d, k):

h=W0x+ΔWx=W0x+BAxh = W_0 x + \Delta W x = W_0 x + B A x

Only AA and BB are trained. A rank-16 adapter on a 4,096-dimensional projection reduces trainable parameters from 16.7M to 131K, a 128x reduction.

python
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_out

Scale 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

OptimizationMemory impactCompute impact
KV cachingKV grows linearly with sequence lengthPer-step compute drops from O(S2)O(S^2) to O(S)O(S)
GQAKV cache reduced by n_query_heads/n_kv_headsn\_query\_heads / n\_kv\_headsMinimal compute reduction
FlashAttentionO(S2)→O(S)O(S^2) \to O(S) peak SRAM2x to 3x throughput on A100 for long sequences
LoRA~1% of full fine-tune parametersSame forward pass cost; backward over A,BA, B only
INT8 quantization~50% model size reductionMinor 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_attention already 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 SS grows. At a few hundred tokens the S×SS \times S 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 KVCache above into a serving loop. It appends to a list and concatenates the whole history on every get(), 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.