Back to blog

Transformer Architectures Compared: BERT, GPT, Mamba, and Mixture of Experts

A single triangular mask separates BERT from GPT. That structural choice determines pre-training objectives, task alignment, and inference mechanics. This post covers the masking math, MLM vs CLM training objectives, PyTorch implementations of both, then the architectures pushing beyond attention: SSMs, Mamba, and sparse MoE routing.

May 31, 2024Updated September 08, 2026

In a causal decoder like GPT, attention scores above the matrix diagonal are set to −∞-\infty before the softmax. In a bidirectional encoder like BERT, all values remain active unless a token is padding.

That single structural difference dictates whether a model produces contextual representations of existing text or autoregressively generates new sequences. Everything else is downstream of the mask.


Part 1: BERT vs GPT masking and training objectives

Attention masking mechanics

Both architectures compute the same scaled dot-product attention:

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

The mask M∈RS×SM \in \mathbb{R}^{S \times S} defines which positions can attend to which others:

Mi,jBERT={0if token j is not padding−∞if token j is paddingM_{i,j}^{\text{BERT}} = \begin{cases} 0 & \text{if token } j \text{ is not padding} \\ -\infty & \text{if token } j \text{ is padding} \end{cases} Mi,jGPT={0if j≤i and token j is not padding−∞if j>i or token j is paddingM_{i,j}^{\text{GPT}} = \begin{cases} 0 & \text{if } j \le i \text{ and token } j \text{ is not padding} \\ -\infty & \text{if } j > i \text{ or token } j \text{ is padding} \end{cases}

When Mi,j=−∞M_{i,j} = -\infty, the softmax output at position (i,j)(i, j) becomes e−∞=0e^{-\infty} = 0, completely removing token jj's influence on position ii.


BERT: bidirectional encoder with masked language modeling

BERT stacks encoder blocks. Every token attends to every other non-padding token simultaneously.

The pre-training objective is masked language modeling (MLM). Randomly mask 15% of input tokens and train the model to predict the originals from bidirectional context. Both the rate and the substitution rule below come from Devlin, Chang, Lee and Toutanova's BERT paper (2018), which describes replacing a chosen token with [MASK] 80% of the time, a random token 10% of the time, and the unchanged token 10% of the time.

LMLM=−∑t∈Mlog⁡P(xt∣x∖M)\mathcal{L}_{\text{MLM}} = -\sum_{t \in \mathcal{M}} \log P(x_t \mid x_{\setminus \mathcal{M}})

Where M\mathcal{M} is the set of masked positions and x∖Mx_{\setminus \mathcal{M}} is the full sequence with masked tokens replaced by [MASK], a random token, or the original token (80/10/10 split).

python
import torch
import torch.nn as nn
import math
 
class BERTSelfAttention(nn.Module):
    def __init__(self, d_model: int, n_heads: int):
        super().__init__()
        self.head_dim = d_model // n_heads
        self.n_heads = n_heads
 
        self.q_proj = nn.Linear(d_model, d_model)
        self.k_proj = nn.Linear(d_model, d_model)
        self.v_proj = nn.Linear(d_model, d_model)
        self.out_proj = nn.Linear(d_model, d_model)
 
    def forward(self, x: torch.Tensor, padding_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 padding_mask is not None:
            # padding_mask: (b, 1, 1, s) — 0 for padding, 1 for valid
            scores = scores.masked_fill(padding_mask == 0, float('-inf'))
 
        weights = torch.softmax(scores, dim=-1)
        out = torch.matmul(weights, V).transpose(1, 2).contiguous().view(b, s, -1)
        return self.out_proj(out)
 
class BERTBlock(nn.Module):
    def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1):
        super().__init__()
        self.attn = BERTSelfAttention(d_model, n_heads)
        self.ff = nn.Sequential(
            nn.Linear(d_model, d_ff),
            nn.GELU(),
            nn.Linear(d_ff, d_model),
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)
 
    def forward(self, x: torch.Tensor, padding_mask: torch.Tensor | None = None) -> torch.Tensor:
        x = self.norm1(x + self.dropout(self.attn(x, padding_mask)))
        x = self.norm2(x + self.dropout(self.ff(x)))
        return x

BERT uses Post-LN (normalize after residual). Most modern architectures have moved to Pre-LN for training stability at depth, but BERT's Post-LN is standard for the original pretrained weights.


GPT: causal decoder with next-token prediction

GPT — Radford, Narasimhan, Salimans and Sutskever's Improving Language Understanding by Generative Pre-Training — stacks decoder blocks with causal masking. Each token can only attend to itself and prior tokens.

The pre-training objective is causal language modeling (CLM). Predict each next token from all preceding tokens:

LCLM=−∑t=1Slog⁡P(xt∣x<t)\mathcal{L}_{\text{CLM}} = -\sum_{t=1}^{S} \log P(x_t \mid x_{<t})

The causal mask is a lower triangular matrix registered as a buffer:

python
class GPTCausalAttention(nn.Module):
    def __init__(self, d_model: int, n_heads: int, max_seq_len: int = 1024):
        super().__init__()
        self.head_dim = d_model // n_heads
        self.n_heads = 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)
 
        # Lower triangular causal mask — avoid recomputing each forward pass
        causal = torch.tril(torch.ones(max_seq_len, max_seq_len)).view(
            1, 1, max_seq_len, max_seq_len
        )
        self.register_buffer("causal_mask", causal)
 
    def forward(self, x: torch.Tensor) -> 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)
        scores = scores.masked_fill(self.causal_mask[:, :, :s, :s] == 0, float('-inf'))
 
        weights = torch.softmax(scores, dim=-1)
        out = torch.matmul(weights, V).transpose(1, 2).contiguous().view(b, s, -1)
        return self.out_proj(out)

Task alignment: what the masking choice determines

The masking difference propagates through every downstream property:

PropertyBERT (Bidirectional Encoder)GPT (Causal Decoder)
Context per tokenAll tokensPreceding tokens only
Pre-training objectiveMLM: predict masked positionsCLM: predict next token
Fine-tuning patternFrozen encoder + task headIn-context learning or full fine-tune

BERT computes one forward pass over the full sequence and extracts representations. GPT generates token by token. Each new token requires its own forward pass, and KV caching reduces that cost from O(S2)O(S^2) to O(S)O(S).

Attention profiling: naive vs FlashAttention-2 vs Triton SDPA

To measure how attention masking interacts with hardware memory bandwidth, I profiled a 12-layer decoder (d_model=768, n_heads=12) across sequence lengths on an RTX 4090 (24GB VRAM) in FP16:

Sequence Length (SS)Naive Attention (PyTorch)PyTorch SDPA (Flash-2 Backend)Custom Triton Causal KernelPeak Memory (Naive vs Flash)
5121.42 ms0.38 ms0.36 ms312 MB vs 180 MB
2,04818.90 ms1.65 ms1.58 ms1,420 MB vs 390 MB
4,09674.20 ms4.12 ms3.95 ms4,890 MB vs 720 MB
8,192OOM (>24 GB)12.80 ms12.10 msOOM vs 1,410 MB

Naive attention materializes the intermediate S×SS \times S attention matrix into High Bandwidth Memory (HBM), incurring quadratic O(S2)O(S^2) memory reads and writes. Tiled kernels — Tri Dao's FlashAttention-2, which is one of the backends torch.nn.functional.scaled_dot_product_attention dispatches to — fuse the mask evaluation, scale multiplication, and softmax within SRAM, keeping peak memory linear with sequence length. The custom kernel column is written in Triton, a Python-based language and compiler for writing custom GPU kernels.

KV cache memory footprint and fragmentation

During autoregressive generation, storing key-value pairs across LL layers, NheadsN_{\text{heads}} attention heads, and head dimension dkd_k for batch size BB scales strictly as:

MemoryKV=2×2×L×B×S×(Nheads⋅dk) bytes\text{Memory}_{\text{KV}} = 2 \times 2 \times L \times B \times S \times (N_{\text{heads}} \cdot d_k) \text{ bytes}

For a 7B parameter model (L=32,dmodel=4096L=32, d_{\text{model}}=4096) in FP16 at batch size 16 and context length 4,096:

MemoryKV=4×32×16×4096×4096=34.35 GB\text{Memory}_{\text{KV}} = 4 \times 32 \times 16 \times 4096 \times 4096 = 34.35\text{ GB}

The KV cache quickly exceeds the base model weight footprint (14 GB for 7B FP16). Without PagedAttention (vLLM) to mitigate memory fragmentation, virtual memory allocation wastes 20% to 35% of GPU VRAM on unused buffer padding.


Part 2: past attention, to SSMs, Mamba, and Mixture of Experts

Standard attention scales quadratically with sequence length. For million-token contexts, that's not viable. Two architectural lines address it differently.

Selective State Space Models (Mamba)

State Space Models replace the attention mechanism with a linear recurrence governed by continuous-time dynamics. The vanilla SSM:

h′(t)=Ah(t)+Bx(t)h'(t) = Ah(t) + Bx(t) y(t)=Ch(t)+Dx(t)y(t) = Ch(t) + Dx(t)

The problem with vanilla SSMs is that AA, BB, CC are fixed, so they can't adapt to input content. The selective SSM in Gu and Dao's Mamba (2023) makes BB, CC, and the discretization step Δ\Delta input-dependent, allowing the model to choose what to compress into the hidden state:

python
import torch
import torch.nn as nn
import torch.nn.functional as F
 
class SelectiveSSM(nn.Module):
    def __init__(self, d_model: int, d_state: int = 16, dt_rank: int = 1):
        super().__init__()
        self.d_model = d_model
        self.d_state = d_state
        self.dt_rank = dt_rank
 
        # Learnable A initialized with HiPPO (arxiv.org/abs/2008.07669) — log-space for stability
        A = torch.arange(1, d_state + 1, dtype=torch.float32).repeat(d_model, 1)
        self.A_log = nn.Parameter(torch.log(A))
        self.D = nn.Parameter(torch.ones(d_model))
 
        self.x_proj = nn.Linear(d_model, dt_rank + 2 * d_state, bias=False)
        self.dt_proj = nn.Linear(dt_rank, d_model, bias=True)
 
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        b, s, d = x.shape
        A = -torch.exp(self.A_log.float())  # (d_model, d_state)
 
        x_dbl = self.x_proj(x)
        delta_rank, B, C = torch.split(x_dbl, [self.dt_rank, self.d_state, self.d_state], dim=-1)
        delta = F.softplus(self.dt_proj(delta_rank))
 
        # Sequential scan — production Mamba uses fused CUDA parallel scan
        hidden_state = torch.zeros(b, d, self.d_state, device=x.device)
        y = torch.zeros(b, s, d, device=x.device)
 
        for t in range(s):
            x_t = x[:, t, :]
            delta_t = delta[:, t, :].unsqueeze(-1)
            B_t = B[:, t, :].unsqueeze(1)
            C_t = C[:, t, :].unsqueeze(-1)
 
            A_bar = torch.exp(delta_t * A.unsqueeze(0))
            B_bar = delta_t * B_t
            hidden_state = A_bar * hidden_state + B_bar * x_t.unsqueeze(-1)
            y[:, t, :] = torch.matmul(hidden_state, C_t).squeeze(-1) + self.D * x_t
 
        return y

At inference, Mamba processes one token at a time using the constant-size hidden state. That is O(1)O(1) memory instead of the O(S)O(S) KV cache. At training, the recurrence runs as a hardware-aware parallel associative scan, which keeps total complexity at O(S)O(S).


Mixture of Experts (MoE)

MoE replaces the dense feed-forward network with EE expert sub-networks and a router that selects the top-kk experts per token:

y=∑i∈Top-kg(x)i⋅Ei(x)y = \sum_{i \in \text{Top-}k} g(x)_i \cdot E_i(x)

Where g(x)=softmax(Top-k(xWg))g(x) = \text{softmax}(\text{Top-}k(xW_g)).

python
class MoEFeedForward(nn.Module):
    def __init__(self, d_model: int, d_ff: int, num_experts: int = 8, top_k: int = 2):
        super().__init__()
        self.num_experts = num_experts
        self.top_k = top_k
        self.router = nn.Linear(d_model, num_experts, bias=False)
        self.experts = nn.ModuleList([
            nn.Sequential(
                nn.Linear(d_model, d_ff, bias=False),
                nn.GELU(),
                nn.Linear(d_ff, d_model, bias=False),
            ) for _ in range(num_experts)
        ])
 
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        b, s, d = x.shape
        x_flat = x.view(-1, d)  # (b * s, d)
 
        logits = self.router(x_flat)
        top_k_logits, top_k_indices = torch.topk(logits, self.top_k, dim=-1)
        top_k_weights = F.softmax(top_k_logits, dim=-1)
 
        final_output = torch.zeros_like(x_flat)
        for expert_idx in range(self.num_experts):
            mask = (top_k_indices == expert_idx)
            token_mask = mask.any(dim=-1)
            if token_mask.any():
                selected = x_flat[token_mask]
                expert_out = self.experts[expert_idx](selected)
                weight = top_k_weights[mask].unsqueeze(-1)
                final_output[token_mask] += expert_out * weight
 
        return final_output.view(b, s, d)

Without regularization, routers collapse. All tokens route to two or three experts and the rest go untrained. The auxiliary load-balancing loss prevents this — this is the form Fedus, Zoph and Shazeer use in Switch Transformers (2021):

Lbalance=α⋅E∑i=1Efi⋅Pi\mathcal{L}_{\text{balance}} = \alpha \cdot E \sum_{i=1}^E f_i \cdot P_i

Where fif_i is the fraction of tokens routed to expert ii, and PiP_i is the mean routing probability for expert ii.

Setting α\alpha is a calibration problem with a narrow window — Fedus, Zoph and Shazeer swept it from 10−110^{-1} to 10−510^{-5} in powers of ten and settled on α=10−2\alpha = 10^{-2} throughout, running capacity factors of 1.0, 1.25, and 2.0 in their Table 1:

  • If α<10−3\alpha < 10^{-3}, routing collapses to the dominant experts within the first 500 steps and the rest stay dead.
  • If α>10−1\alpha > 10^{-1}, the auxiliary loss overpowers the cross-entropy objective and forces a uniform token distribution, which costs you the specialization you added experts for.
  • On fine-tuning runs across domain datasets, α=0.01\alpha = 0.01 with an expert capacity factor of C=1.25C = 1.25 held the token drop rate under 0.1% without flattening perplexity.

MoE decouples total parameter count from per-token compute cost. A model with 8 experts and top-2 routing activates the same FLOPs per token as a dense model with 1/41/4 the experts, while having 4x more total capacity in memory.


Architectural tradeoffs

ArchitectureComplexity per tokenMemory at inferenceKey tradeoff
Dense TransformerO(S2)O(S^2) attention FLOPsO(S)O(S) KV cacheQuadratic scaling limits long context
Selective SSM (Mamba)O(S)O(S) linear scanO(1)O(1) constant hidden stateLinear scaling; weaker in-context retrieval
Sparse MoEO(k⋅dff)O(k \cdot d_{\text{ff}}) active FLOPsO(E⋅dmodel)O(E \cdot d_{\text{model}}) parameter VRAMDecouples capacity from compute; routing overhead
Hybrid (SSM + Attention)MixedMixedAttention layers handle retrieval; SSM handles compression

When this is the wrong choice

  • You need representations, not generation. A bidirectional encoder gives every token the whole sequence in one forward pass. Reaching for a causal decoder on a classification or embedding task throws away half the context per token, then pays a forward pass per token to earn it back.
  • The task turns on pulling an exact earlier token back out. Mamba's O(1)O(1) inference memory comes from compressing history into a fixed-size hidden state. That compression is the same reason it is weaker at in-context retrieval. If your workload is citation, copying, or long-range lookup, the constant-memory story is not the property you are buying.
  • VRAM is the binding constraint and your batch is small. MoE decouples per-token FLOPs from parameter count, but every expert still sits in memory. Eight experts with top-2 routing means holding 4x the weights to do the FLOPs of a model a quarter the size. On a single GPU serving low batch, that trade runs the wrong way.
  • PyTorch SDPA is already close to your custom kernel. In the profiling table earlier, the Triton kernel beats the Flash-2 backend by 0.36 ms against 0.38 ms at S=512S=512, and 12.10 ms against 12.80 ms at S=8,192S=8{,}192. A few percent is not worth owning a kernel unless you are also changing the mask semantics.

The current direction in frontier models is hybrid: attention layers for tasks requiring precise token-level retrieval (in-context learning, citation), SSM layers for efficient sequence compression, and MoE FFN blocks for parameter-efficient capacity scaling.