Learn ML from Scratch: Beginner to Advanced

A rigorous, step-by-step interactive curriculum. Master foundational mathematical derivations, vector calculus, pure NumPy/PyTorch implementations from scratch, and modern 2026 reasoning LLM architectures.

🎓 Track Progress: 0 Completed
0%
💻 Modern ArchitecturesLEVEL 3 · ADVANCEDChapter 8

Modern Decoder-Only LLM Block (Llama / Mistral Style)

Grouped-Query Attention (GQA), SwiGLU FFN, Pre-RMSNorm, and KV Caching.

⏱ 22 min read🎯 Prerequisites: Multi-Head Attention & RMSNorm
🔍 Inspect Architecture: Decoder-only language models (GPT series)

💡 1. Core Intuition & Concepts

Modern frontier Large Language Models (Llama 3, Mistral, Gemma) share an optimized autoregressive decoder block combining Pre-RMSNorm, Grouped-Query Attention (GQA) for 8x KV-cache memory compression, SwiGLU gated activations, and residual streams.

📐 2. Mathematical Formulations & Derivations

SwiGLU Gated Feed-Forward Network
SwiGLU(x)=(SiLU(xWgate)⊙(xWup))WdownSiLU(z)=z⋅σ(z)\text{SwiGLU}(x) = \left( \text{SiLU}(x W_{\text{gate}}) \odot (x W_{\text{up}}) \right) W_{\text{down}} \qquad \text{SiLU}(z) = z \cdot \sigma(z)
Gated linear unit with continuous SiLU non-linearity for enhanced representation capacity.
Grouped-Query Attention (GQA) Memory Scaling
KV Memory Reduction=HQHKV(e.g. 32 query heads share 8 KV heads  ⟹  4× smaller KV cache)\text{KV Memory Reduction} = \frac{H_Q}{H_{KV}} \qquad (\text{e.g. } 32 \text{ query heads share } 8 \text{ KV heads} \implies 4\times \text{ smaller KV cache})
Multiple query heads share single key/value head pairs during inference.

⚙️ 3. Step-by-Step Computational Mechanism

1
Pre-RMSNorm Normalization
Apply scale normalization x_norm = RMSNorm(x) prior to attention.
2
Grouped-Query Attention
Compute GQA self-attention and add residual: x <- x + GQA(x_norm).
3
Pre-RMSNorm FFN
Normalize residual stream: x_ffn = RMSNorm(x).
4
SwiGLU Projection & Residual
Execute gated FFN and add residual: x <- x + SwiGLU(x_ffn).

💻 4. Code from Scratch (python)

llm_block.py
import torch
import torch.nn as nn
import torch.nn.functional as F

class RMSNorm(nn.Module):
    def __init__(self, dim: int, eps: float = 1e-6):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        rms = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
        return x * rms * self.weight

class SwiGLU(nn.Module):
    def __init__(self, dim: int, hidden_dim: int):
        super().__init__()
        self.w_gate = nn.Linear(dim, hidden_dim, bias=False)
        self.w_up = nn.Linear(dim, hidden_dim, bias=False)
        self.w_down = nn.Linear(hidden_dim, dim, bias=False)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))

class ModernLLMBlock(nn.Module):
    """Complete Llama 3 / Mistral style Transformer Decoder Block."""
    def __init__(self, dim: int, n_heads: int, hidden_dim: int):
        super().__init__()
        self.attn_norm = RMSNorm(dim)
        self.attn = nn.MultiheadAttention(dim, n_heads, batch_first=True)
        self.ffn_norm = RMSNorm(dim)
        self.ffn = SwiGLU(dim, hidden_dim)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        normed = self.attn_norm(x)
        attn_out, _ = self.attn(normed, normed, normed, need_weights=False)
        x = x + attn_out
        x = x + self.ffn(self.ffn_norm(x))
        return x

🧠 5. Comprehension Checkpoint

Answer all 1 questions correctly to complete the chapter · 0 / 1 done
Q1/1 What is the primary operational benefit of Grouped-Query Attention (GQA) over Multi-Head Attention (MHA)?

Finished this chapter?