Understanding Transformer Architecture
A visual guide to the transformer architecture that powers modern LLMs like GPT and LLaMA.
prerequisites
3- Linear algebra basics
- Neural network fundamentals
- Python/PyTorch experience
The Transformer architecture, introduced in the landmark paper "Attention Is All You Need" (2017), revolutionized natural language processing and laid the foundation for modern LLMs like GPT, LLaMA, and Claude. In this tutorial, we'll break down the architecture piece by piece.
The Big Picture
At its core, a Transformer processes sequences by:
- Converting tokens to embeddings
- Adding positional information
- Processing through attention and feedforward layers
- Producing output predictions
Input Tokens → Embeddings → [N × Transformer Blocks] → Output
↓
Each block contains:
- Multi-Head Attention
- Feed-Forward Network
- Layer Normalization
- Residual ConnectionsStep 1: Token Embeddings
First, we convert discrete tokens (words, subwords) into continuous vectors:
import torch
import torch.nn as nn
class TokenEmbedding(nn.Module):
def __init__(self, vocab_size: int, d_model: int):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.d_model = d_model
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [batch_size, seq_len] → [batch_size, seq_len, d_model]
# Scale by sqrt(d_model) as per original paper
return self.embedding(x) * (self.d_model ** 0.5)
# Example
vocab_size = 50000
d_model = 512
embedding = TokenEmbedding(vocab_size, d_model)
tokens = torch.tensor([[1, 42, 156, 7]]) # [1, 4]
embedded = embedding(tokens) # [1, 4, 512]Step 2: Positional Encoding
Unlike RNNs, Transformers process all tokens in parallel. To capture sequence order, we add positional information:
import math
class PositionalEncoding(nn.Module):
def __init__(self, d_model: int, max_seq_len: int = 5000, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
# Create position encodings
position = torch.arange(max_seq_len).unsqueeze(1) # [max_seq_len, 1]
div_term = torch.exp(
torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)
)
pe = torch.zeros(max_seq_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term) # Even indices
pe[:, 1::2] = torch.cos(position * div_term) # Odd indices
# Register as buffer (not a parameter)
self.register_buffer('pe', pe.unsqueeze(0)) # [1, max_seq_len, d_model]
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [batch_size, seq_len, d_model]
seq_len = x.size(1)
x = x + self.pe[:, :seq_len, :]
return self.dropout(x)Why sinusoidal? The sine/cosine functions allow the model to:
- Learn relative positions (PE[pos+k] can be represented as a function of PE[pos])
- Generalize to longer sequences than seen during training
Step 3: Self-Attention (The Core Innovation)
Self-attention computes relationships between all pairs of tokens:
class ScaledDotProductAttention(nn.Module):
def __init__(self, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
def forward(
self,
query: torch.Tensor, # [batch, heads, seq_len, d_k]
key: torch.Tensor, # [batch, heads, seq_len, d_k]
value: torch.Tensor, # [batch, heads, seq_len, d_v]
mask: torch.Tensor = None
) -> tuple[torch.Tensor, torch.Tensor]:
d_k = query.size(-1)
# Compute attention scores
# [batch, heads, seq_len, seq_len]
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k)
# Apply mask (for causal attention in decoders)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
# Softmax to get attention weights
attention_weights = torch.softmax(scores, dim=-1)
attention_weights = self.dropout(attention_weights)
# Apply attention to values
output = torch.matmul(attention_weights, value)
return output, attention_weightsIntuition: Each token "queries" for relevant information from all other tokens. The dot product between query and key determines relevance, and values carry the actual information.
Step 4: Multi-Head Attention
Instead of single attention, we use multiple "heads" to capture different types of relationships:
class MultiHeadAttention(nn.Module):
def __init__(self, d_model: int, num_heads: int, dropout: float = 0.1):
super().__init__()
assert d_model % num_heads == 0
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads
# Linear projections
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
self.W_o = nn.Linear(d_model, d_model)
self.attention = ScaledDotProductAttention(dropout)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
mask: torch.Tensor = None
) -> torch.Tensor:
batch_size = query.size(0)
# Project and reshape for multi-head: [batch, seq, d_model] → [batch, heads, seq, d_k]
Q = self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
K = self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
V = self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
# Apply attention
attn_output, _ = self.attention(Q, K, V, mask)
# Concatenate heads: [batch, heads, seq, d_k] → [batch, seq, d_model]
attn_output = attn_output.transpose(1, 2).contiguous().view(
batch_size, -1, self.d_model
)
# Final projection
return self.W_o(attn_output)Why multiple heads? Different heads can focus on:
- Syntactic relationships (subject-verb agreement)
- Semantic relationships (word meanings)
- Positional patterns (nearby words)
Step 5: Feed-Forward Network
After attention, each position passes through a feedforward network:
class FeedForward(nn.Module):
def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1):
super().__init__()
self.linear1 = nn.Linear(d_model, d_ff)
self.linear2 = nn.Linear(d_ff, d_model)
self.dropout = nn.Dropout(dropout)
self.activation = nn.GELU() # Modern transformers use GELU
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [batch, seq, d_model]
x = self.linear1(x) # [batch, seq, d_ff]
x = self.activation(x)
x = self.dropout(x)
x = self.linear2(x) # [batch, seq, d_model]
return xTypically, d_ff = 4 * d_model. This expansion allows the model to process information in a higher-dimensional space.
Step 6: Transformer Block
Combining everything with residual connections and layer normalization:
class TransformerBlock(nn.Module):
def __init__(
self,
d_model: int,
num_heads: int,
d_ff: int,
dropout: float = 0.1,
pre_norm: bool = True # Modern transformers use pre-norm
):
super().__init__()
self.pre_norm = pre_norm
self.attention = MultiHeadAttention(d_model, num_heads, dropout)
self.ff = FeedForward(d_model, d_ff, dropout)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor, mask: torch.Tensor = None) -> torch.Tensor:
if self.pre_norm:
# Pre-norm (GPT-style)
attn_out = self.attention(
self.norm1(x), self.norm1(x), self.norm1(x), mask
)
x = x + self.dropout(attn_out)
ff_out = self.ff(self.norm2(x))
x = x + self.dropout(ff_out)
else:
# Post-norm (original transformer)
attn_out = self.attention(x, x, x, mask)
x = self.norm1(x + self.dropout(attn_out))
ff_out = self.ff(x)
x = self.norm2(x + self.dropout(ff_out))
return xStep 7: Complete Decoder (GPT-style)
Putting it all together for a decoder-only model (like GPT):
class GPTModel(nn.Module):
def __init__(
self,
vocab_size: int,
d_model: int = 512,
num_heads: int = 8,
num_layers: int = 6,
d_ff: int = 2048,
max_seq_len: int = 1024,
dropout: float = 0.1,
):
super().__init__()
self.token_embedding = TokenEmbedding(vocab_size, d_model)
self.pos_encoding = PositionalEncoding(d_model, max_seq_len, dropout)
self.layers = nn.ModuleList([
TransformerBlock(d_model, num_heads, d_ff, dropout)
for _ in range(num_layers)
])
self.norm = nn.LayerNorm(d_model)
self.output = nn.Linear(d_model, vocab_size)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [batch, seq_len] token indices
seq_len = x.size(1)
# Create causal mask
mask = torch.triu(
torch.ones(seq_len, seq_len, device=x.device), diagonal=1
).bool()
mask = ~mask # Invert: True = attend, False = mask
# Embedding + positional encoding
x = self.token_embedding(x)
x = self.pos_encoding(x)
# Transformer blocks
for layer in self.layers:
x = layer(x, mask)
# Output projection
x = self.norm(x)
logits = self.output(x) # [batch, seq, vocab_size]
return logitsKey Concepts Summary
| Component | Purpose |
|---|---|
| Token Embedding | Convert discrete tokens to vectors |
| Positional Encoding | Add sequence order information |
| Self-Attention | Model relationships between all tokens |
| Multi-Head | Capture different types of relationships |
| Feed-Forward | Process each position independently |
| Layer Norm | Stabilize training |
| Residual Connections | Enable gradient flow in deep networks |
Modern Improvements
Since the original paper, several improvements have been made:
- Pre-normalization: Apply LayerNorm before (not after) attention and FFN
- Rotary Position Embeddings (RoPE): Better handling of relative positions
- Grouped Query Attention (GQA): More efficient multi-head attention
- SwiGLU Activation: Improved feedforward networks
Next Steps
Now that you understand the architecture:
- Implement a small GPT from scratch
- Experiment with different hyperparameters
- Explore pre-trained models on Hugging Face
- Study specific improvements like FlashAttention
This tutorial is part of the ML Fundamentals series. Understanding transformers is essential for working with modern LLMs.