Blogsintermediate

Understanding Transformer Architecture

A visual guide to the transformer architecture that powers modern LLMs like GPT and LLaMA.

25 minDeep Learning, NLP, AttentionSeries: ML Fundamentals

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:

  1. Converting tokens to embeddings
  2. Adding positional information
  3. Processing through attention and feedforward layers
  4. Producing output predictions
Text
Input Tokens → Embeddings → [N × Transformer Blocks] → Output
                    ↓
            Each block contains:
            - Multi-Head Attention
            - Feed-Forward Network
            - Layer Normalization
            - Residual Connections

Step 1: Token Embeddings

First, we convert discrete tokens (words, subwords) into continuous vectors:

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

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

Python
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_weights

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

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

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

Typically, 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:

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

Step 7: Complete Decoder (GPT-style)

Putting it all together for a decoder-only model (like GPT):

Python
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 logits

Key Concepts Summary

ComponentPurpose
Token EmbeddingConvert discrete tokens to vectors
Positional EncodingAdd sequence order information
Self-AttentionModel relationships between all tokens
Multi-HeadCapture different types of relationships
Feed-ForwardProcess each position independently
Layer NormStabilize training
Residual ConnectionsEnable gradient flow in deep networks

Modern Improvements

Since the original paper, several improvements have been made:

  1. Pre-normalization: Apply LayerNorm before (not after) attention and FFN
  2. Rotary Position Embeddings (RoPE): Better handling of relative positions
  3. Grouped Query Attention (GQA): More efficient multi-head attention
  4. 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.

Zizhao Huloading加载中