PyTorch: The Complete Reference Guide
From Absolute Beginner to Advanced Practitioner
A structured, self-contained course organized into 50 curriculum modules, 2 operational guides, and 50 interactive notebooks. Each curriculum module contains detailed explanations, theory, formulas, runnable Python scripts, and a Jupyter playbook.
Module 01: Foundations — PyTorch and the Mathematics of Deep Learning
Table of Contents
- What is PyTorch?
- Installation
- Core Philosophy
- Architecture Stack
- Mathematical Prerequisites
- Linear Algebra
- Calculus for Deep Learning
- Probability and Statistics
- Optimization Theory
- Information Theory
What is PyTorch?
PyTorch is an open-source machine learning framework developed primarily by Meta AI (formerly Facebook AI Research). It provides two core capabilities:
- N-dimensional tensor computation — similar to NumPy but with GPU acceleration
- Automatic differentiation — computes gradients of arbitrary computational graphs
A Brief History
- 2002: Torch was created in Lua at NYU by Ronan Collobert and others.
- 2016: PyTorch 0.1 was released by Facebook AI Research (FAIR), bringing the
Torch tensor library to Python with automatic differentiation built in.
- 2018: PyTorch 1.0 merged the research-focused PyTorch with the production-focused
Caffe2, adding TorchScript for model export.
- 2022: PyTorch 2.0 introduced
torch.compile(), a compiler-based approach that
can dramatically speed up models with a single line of code.
- 2023: PyTorch moved to the Linux Foundation, becoming a truly community-governed project.
- 2024-2025: Continued evolution with FlexAttention, torch.export improvements,
and expanded hardware support (Intel XPU, Apple MPS, AMD ROCm).
PyTorch vs TensorFlow
| Aspect | PyTorch | TensorFlow |
|---|---|---|
| Execution | Eager by default (define-by-run) | Historically graph-based (define-then-run), now eager via tf.function |
| Debugging | Standard Python debugger works | Harder to debug graph mode |
| Research adoption | Dominant in academia (~80%+ of papers) | Strong in industry/production |
| Deployment | TorchServe, ONNX, torch.export | TF Serving, TFLite, TF.js |
| API Style | Pythonic, object-oriented | Keras-based high-level API |
| Compilation | torch.compile (TorchDynamo + Inductor) | XLA compiler |
| Community | Massive open-source ecosystem (HuggingFace, etc.) | Google-driven ecosystem |
Why PyTorch won in research: The key reason is debuggability. When your model produces NaN values or wrong outputs, you can insert a breakpoint() anywhere in your forward pass, inspect tensors, and understand what's happening. In graph-based frameworks, you're debugging a compiled representation, not your original code.
The PyTorch Ecosystem
PyTorch is not just one library — it's a constellation:
- torchvision: Computer vision (datasets, models, transforms)
- torchaudio: Audio processing
- torchtext: NLP utilities
- PyTorch Lightning / Fabric: Training framework that reduces boilerplate
- HuggingFace Transformers: Built on PyTorch, the dominant NLP/LLM library
- TorchServe: Model serving for production
- ONNX Runtime: Export PyTorch models to a portable format
- torch.export / ExecuTorch: On-device deployment (mobile, embedded)
Installation
CPU-Only Installation (Recommended for Learning)
# Using pip (simplest)
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
# Using conda
conda install pytorch torchvision torchaudio cpuonly -c pytorch
With CUDA (for NVIDIA GPUs)
# CUDA 11.8
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118
# CUDA 12.1
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121
# CUDA 12.4
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu124
Verify Installation
import torch
print(f"PyTorch version: {torch.__version__}")
print(f"CUDA available: {torch.cuda.is_available()}")
print(f"Number of GPUs: {torch.cuda.device_count()}")
# Quick test
x = torch.rand(3, 3)
print(f"Random tensor:\n{x}")
Core Philosophy
1. Eager Execution (Define-by-Run)
In PyTorch, operations execute immediately as Python runs them. There is no separate "compilation" step before you can see results:
import torch
x = torch.tensor([1.0, 2.0, 3.0])
y = x * 2 # This executes RIGHT NOW — not later in a session
print(y) # tensor([2., 4., 6.])
Why this matters: You can use standard Python control flow (if, for, while) inside your models. The computation graph is built dynamically as your code runs, which means different inputs can follow different code paths.
2. Dynamic Computation Graphs
Unlike static graph frameworks, PyTorch rebuilds the computation graph every time you run a forward pass. This means:
def dynamic_model(x):
if x.sum() > 0: # This condition is evaluated at runtime
return x * 2
else:
return x * 3
Each call to dynamic_model can follow a different path through the code. The autograd graph is constructed on the fly and torn down after .backward().
3. Python-First Design
PyTorch is designed to feel like an extension of Python and NumPy, not a separate language embedded in Python. You think in Python, you debug in Python, you profile in Python. The C++ backend handles performance; you rarely need to think about it.
4. torch.compile — The Best of Both Worlds
Starting with PyTorch 2.0, you can optionally compile your models for speed while keeping the eager-mode development experience:
model = MyModel()
compiled_model = torch.compile(model) # One line for significant speedup
output = compiled_model(input_data)
This uses TorchDynamo to trace your Python code, TorchInductor to generate optimized kernels, and falls back to eager mode for unsupported patterns.
Architecture Stack
Understanding PyTorch's layered architecture helps you know where to look when debugging or optimizing. From top to bottom:
┌──────────────────────────────────────────────────┐
│ Python Frontend (torch, torch.nn, etc.) │ ← You write code here
├──────────────────────────────────────────────────┤
│ torch.compile (TorchDynamo + TorchInductor) │ ← Optional compilation
├──────────────────────────────────────────────────┤
│ Autograd Engine │ ← Automatic differentiation
├──────────────────────────────────────────────────┤
│ ATen (A Tensor Library) │ ← Core tensor operations
├──────────────────────────────────────────────────┤
│ C10 (Caffe2 + PyTorch Core) │ ← Dispatcher, memory, dtypes
├──────────────────────────────────────────────────┤
│ Hardware Backends (CPU/CUDA/MPS/XPU/ROCm) │ ← Actual computation
└──────────────────────────────────────────────────┘
Python Frontend: The torch module, torch.nn, torch.optim, etc. — the high-level API you interact with daily.
torch.compile: TorchDynamo captures your Python code as an FX graph. TorchInductor generates optimized C++/CUDA/Triton kernels from that graph.
Autograd: The engine that tracks operations on tensors with requires_grad=True and computes gradients via reverse-mode automatic differentiation.
ATen: "A Tensor Library" — over 2,000 operators written in C++ that implement the actual math (add, matmul, conv2d, etc.). When you call torch.add(a, b), this is where the computation happens.
C10: The core library providing the dispatcher (routes operations to the right backend), memory allocators, dtype system, and device abstraction.
Hardware Backends: BLAS libraries (MKL for CPU, cuBLAS for GPU), cuDNN for convolutions, and other hardware-specific optimized libraries.
Mathematical Prerequisites
Deep learning sits at the intersection of linear algebra, calculus, probability, and optimization. You don't need a PhD in mathematics, but you need working familiarity with these concepts. This section teaches them through PyTorch code.
Linear Algebra
Linear algebra is the language of deep learning. Neural networks are fundamentally sequences of matrix multiplications interspersed with non-linear functions.
Scalars, Vectors, Matrices, and Tensors
import torch
scalar = torch.tensor(3.14) # 0-D tensor (scalar)
vector = torch.tensor([1.0, 2.0, 3.0]) # 1-D tensor (vector)
matrix = torch.tensor([[1, 2], [3, 4]]) # 2-D tensor (matrix)
tensor_3d = torch.randn(2, 3, 4) # 3-D tensor
print(f"Scalar shape: {scalar.shape}") # torch.Size([])
print(f"Vector shape: {vector.shape}") # torch.Size([3])
print(f"Matrix shape: {matrix.shape}") # torch.Size([2, 2])
print(f"3D shape: {tensor_3d.shape}") # torch.Size([2, 3, 4])
Why tensors? In deep learning, data naturally has multiple dimensions. An image is a 3D tensor (channels × height × width). A batch of images is 4D (batch × channels × height × width). A batch of sequences of word embeddings is 3D (batch × sequence_length × embedding_dim).
Norms — Measuring Vector Magnitude
A norm measures the "size" of a vector. Different norms emphasize different properties:
- L1 norm (Manhattan): Sum of absolute values. Encourages sparsity in optimization.
\( \|x\|_1 = \sum_i |x_i| \)
- L2 norm (Euclidean): Square root of sum of squares. The "ordinary" distance.
\( \|x\|_2 = \sqrt{\sum_i x_i^2} \)
- L∞ norm (Max): Largest absolute value. Used in adversarial robustness.
\( \|x\|_\infty = \max_i |x_i| \)
v = torch.tensor([3.0, -4.0])
print(f"L1 norm: {torch.norm(v, p=1)}") # 7.0
print(f"L2 norm: {torch.norm(v, p=2)}") # 5.0
print(f"L∞ norm: {torch.norm(v, p=float('inf'))}") # 4.0
Why norms matter in deep learning: L2 regularization (weight decay) penalizes large L2 norms of weight vectors, preventing overfitting. L1 regularization drives weights to exactly zero, performing feature selection. Gradient clipping uses norms to prevent exploding gradients.
Dot Product — Measuring Similarity
The dot product of two vectors measures their alignment:
\( a \cdot b = \sum_i a_i b_i = \|a\| \|b\| \cos\theta \)
a = torch.tensor([1.0, 0.0])
b = torch.tensor([0.0, 1.0])
c = torch.tensor([1.0, 1.0])
print(f"a·b (perpendicular): {torch.dot(a, b)}") # 0.0
print(f"a·c (45 degrees): {torch.dot(a, c)}") # 1.0
print(f"a·a (parallel): {torch.dot(a, a)}") # 1.0
Why dot products matter: Attention mechanisms in Transformers compute dot products between query and key vectors to measure relevance. The output of a linear layer y = Wx + b is a batch of dot products between weight rows and the input vector.
Matrix Multiplication — The Core Operation
Matrix multiplication is the single most important operation in deep learning. Every linear layer, every attention head, every convolution can be expressed as matrix multiplications.
For matrices A (m×n) and B (n×p), the result C = AB is (m×p):
\( C_{ij} = \sum_k A_{ik} B_{kj} \)
A = torch.tensor([[1., 2.], [3., 4.]]) # 2×2
B = torch.tensor([[5., 6.], [7., 8.]]) # 2×2
C = A @ B # or torch.matmul(A, B) or torch.mm(A, B)
print(f"A @ B =\n{C}")
# tensor([[19., 22.],
# [43., 50.]])
The n dimensions are "consumed" by the multiplication.
Eigendecomposition
A square matrix A can be decomposed as A = QΛQ⁻¹, where Q contains eigenvectors and Λ is a diagonal matrix of eigenvalues. An eigenvector v satisfies Av = λv — the matrix only scales it, doesn't change its direction.
A = torch.tensor([[2., 1.], [1., 2.]], dtype=torch.float)
eigenvalues, eigenvectors = torch.linalg.eig(A)
print(f"Eigenvalues: {eigenvalues}")
print(f"Eigenvectors:\n{eigenvectors}")
Why eigendecomposition matters: Principal Component Analysis (PCA) uses eigendecomposition to find the directions of maximum variance in data. The condition number (ratio of largest to smallest eigenvalue) tells you how numerically stable a problem is.
Singular Value Decomposition (SVD)
SVD generalizes eigendecomposition to non-square matrices: A = UΣVᵀ.
- U: left singular vectors (column space basis)
- Σ: singular values (diagonal, non-negative, sorted)
- Vᵀ: right singular vectors (row space basis)
A = torch.tensor([[1., 2., 3.], [4., 5., 6.]], dtype=torch.float)
U, S, Vh = torch.linalg.svd(A)
print(f"U shape: {U.shape}, S shape: {S.shape}, Vh shape: {Vh.shape}")
Why SVD matters: Low-rank approximation via SVD is used in LoRA (Low-Rank Adaptation) for efficient fine-tuning of large language models. It's also the mathematical foundation of matrix factorization in recommender systems.
Calculus for Deep Learning
Derivatives and the Chain Rule
A derivative measures how a function's output changes when its input changes:
\( f'(x) = \lim_{h \to 0} \frac{f(x+h) - f(x)}{h} \)
The chain rule is the single most important calculus concept for deep learning. If y = f(g(x)), then:
\( \frac{dy}{dx} = \frac{dy}{dg} \cdot \frac{dg}{dx} \)
Neural networks are compositions of functions: output = f_n(f_{n-1}(...f_1(x)...)). Backpropagation is just the chain rule applied repeatedly.
x = torch.tensor(2.0, requires_grad=True)
y = x**3 + 2*x**2 + x # y = x³ + 2x² + x
y.backward() # dy/dx = 3x² + 4x + 1
print(f"dy/dx at x=2: {x.grad}") # 3(4) + 4(2) + 1 = 21
Gradients — Multivariable Derivatives
When a function has multiple inputs, the gradient is the vector of all partial derivatives:
\( \nabla f = \left[\frac{\partial f}{\partial x_1}, \frac{\partial f}{\partial x_2}, \ldots\right] \)
The gradient points in the direction of steepest ascent. To minimize a loss, we move in the opposite direction: gradient descent.
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
f = (x**2).sum() # f = x₁² + x₂² + x₃²
f.backward()
print(f"Gradient: {x.grad}") # [2, 4, 6] — the gradient ∇f = 2x
Jacobians and Hessians
The Jacobian generalizes the gradient to vector-valued functions. If f: ℝⁿ → ℝᵐ, the Jacobian J is an m×n matrix where J_ij = ∂f_i/∂x_j.
The Hessian is the matrix of second derivatives: H_ij = ∂²f/∂x_i∂x_j. It tells you about the curvature of the loss surface — whether you're at a minimum, maximum, or saddle point.
from torch.autograd.functional import jacobian, hessian
def f(x):
return torch.stack([x[0]**2 + x[1], x[0] * x[1]**2])
x = torch.tensor([1.0, 2.0])
J = jacobian(f, x)
print(f"Jacobian:\n{J}")
# [[2*x0, 1 ], = [[2, 1],
# [x1^2, 2*x0*x1]] [4, 4]]
Probability and Statistics
Probability Distributions
PyTorch's torch.distributions module provides a rich set of probability distributions. These are essential for:
- Initializing weights (normal, uniform)
- Variational autoencoders (reparameterization trick)
- Reinforcement learning (sampling actions)
- Bayesian neural networks
from torch.distributions import Normal, Bernoulli, Categorical
normal = Normal(loc=0.0, scale=1.0) # mean=0, std=1
sample = normal.sample((5,))
log_prob = normal.log_prob(torch.tensor(0.0))
print(f"Samples: {sample}")
print(f"Log prob of 0: {log_prob}") # log(1/√(2π)) ≈ -0.9189
Expectation and Variance
- Expectation E[X]: The average value of a random variable.
- Variance Var(X) = E[(X - E[X])²]: How spread out the values are.
samples = torch.randn(100000) # Standard normal
print(f"Mean ≈ {samples.mean():.4f}") # ≈ 0
print(f"Var ≈ {samples.var():.4f}") # ≈ 1
Cross-Entropy — The Loss Function of Classification
Cross-entropy measures the difference between two probability distributions p (true) and q (predicted):
\( H(p, q) = -\sum_i p_i \log(q_i) \)
When p is a one-hot vector (classification), this simplifies to: \( H(p, q) = -\log(q_{\text{true class}}) \)
This is why we use log-softmax + negative log likelihood, which PyTorch combines into nn.CrossEntropyLoss:
import torch.nn as nn
logits = torch.tensor([[2.0, 1.0, 0.1]]) # Raw model output
target = torch.tensor([0]) # True class is 0
loss_fn = nn.CrossEntropyLoss()
loss = loss_fn(logits, target)
print(f"Cross-entropy loss: {loss.item():.4f}")
KL Divergence
KL divergence measures how one probability distribution differs from a reference:
\( D_{KL}(P \| Q) = \sum_i P(i) \log\frac{P(i)}{Q(i)} \)
It's asymmetric: D_KL(P||Q) ≠ D_KL(Q||P). Used in VAEs (variational autoencoders) to keep the learned latent distribution close to a prior (usually standard normal).
import torch.nn.functional as F
p = torch.tensor([0.4, 0.3, 0.3]) # True distribution
q = torch.tensor([0.33, 0.33, 0.34]) # Predicted distribution
kl = F.kl_div(q.log(), p, reduction='sum')
print(f"KL divergence: {kl.item():.4f}")
Optimization Theory
Gradient Descent — The Foundation
Gradient descent minimizes a function by repeatedly taking steps opposite to the gradient. The update rule:
\( \theta_{t+1} = \theta_t - \alpha \nabla L(\theta_t) \)
where α is the learning rate — the most important hyperparameter in deep learning.
- Too large: Overshoots the minimum, loss oscillates or diverges
- Too small: Converges extremely slowly, may get stuck in local minima
- Just right: Smooth convergence to a good minimum
# Minimizing f(x) = (x - 3)² from scratch
x = torch.tensor(0.0, requires_grad=True)
lr = 0.1
for step in range(50):
loss = (x - 3) ** 2
loss.backward()
with torch.no_grad():
x -= lr * x.grad
x.grad.zero_()
print(f"Final x: {x.item():.6f}") # ≈ 3.0
Stochastic Gradient Descent (SGD)
In practice, computing the gradient over the entire dataset is expensive. SGD estimates the gradient using a random mini-batch:
\( \theta_{t+1} = \theta_t - \alpha \nabla L_{\text{batch}}(\theta_t) \)
The noise from mini-batch sampling actually helps — it can escape local minima and leads to solutions that generalize better.
SGD with Momentum
Plain SGD can oscillate in narrow valleys. Momentum adds a "velocity" term that accumulates past gradients, smoothing the trajectory:
\( v_t = \beta v_{t-1} + \nabla L(\theta_t) \) \( \theta_{t+1} = \theta_t - \alpha v_t \)
β (typically 0.9) controls how much history to keep. Think of it like a ball rolling downhill — it builds up speed in consistent directions.
Adam — Adaptive Moment Estimation
Adam combines momentum with per-parameter adaptive learning rates. It maintains two running averages:
- m (first moment): exponential moving average of gradients (like momentum)
- v (second moment): exponential moving average of squared gradients
\( m_t = \beta_1 m_{t-1} + (1 - \beta_1) g_t \) \( v_t = \beta_2 v_{t-1} + (1 - \beta_2) g_t^2 \) \( \hat{m}_t = m_t / (1 - \beta_1^t) \) (bias correction) \( \hat{v}_t = v_t / (1 - \beta_2^t) \) (bias correction) \( \theta_{t+1} = \theta_t - \alpha \hat{m}_t / (\sqrt{\hat{v}_t} + \epsilon) \)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001, betas=(0.9, 0.999))
Why Adam works well: Parameters that receive sparse, infrequent gradients get larger effective learning rates (because v is small). Parameters with frequent, large gradients get smaller effective learning rates (because v is large). This adaptive behavior is crucial for training models where different parameters need different learning rates.
When to use what:
- SGD + Momentum: Often achieves better final performance but requires careful
learning rate tuning and scheduling. Preferred for vision models (ResNet, etc.).
- Adam/AdamW: Faster convergence, less sensitive to learning rate. Preferred
for Transformers, LLMs, and when you want quick results.
- AdamW: Adam with decoupled weight decay — almost always preferred over plain Adam.
Information Theory
Entropy — Measuring Uncertainty
Entropy quantifies the uncertainty in a probability distribution:
\( H(p) = -\sum_i p_i \log p_i \)
- Maximum entropy: uniform distribution (maximum uncertainty)
- Minimum entropy (0): all probability mass on one outcome (certainty)
uniform = torch.tensor([0.25, 0.25, 0.25, 0.25])
certain = torch.tensor([1.0, 0.0, 0.0, 0.0])
H_uniform = -(uniform * uniform.log()).sum()
H_certain = -(certain * (certain + 1e-8).log()).sum()
print(f"Entropy of uniform: {H_uniform:.4f}") # 1.3863 (= ln(4))
print(f"Entropy of certain: {H_certain:.4f}") # ≈ 0.0
Cross-Entropy Loss Derivation
Why do we use cross-entropy as a loss function for classification?
- Maximum likelihood: We want to find model parameters θ that maximize the
probability of the observed data: argmax_θ P(data|θ).
- Log transformation: Taking the log converts the product of probabilities
into a sum (numerically stable, easier to optimize): argmax_θ Σ log P(y_i | x_i, θ)
- Negation: Maximizing log-likelihood = minimizing negative log-likelihood:
argmin_θ -Σ log P(y_i | x_i, θ)
- Softmax output: For classification, P(y=k|x) = softmax(logits)_k, so:
Loss = -log(softmax(logits)_{true_class})
This is exactly cross-entropy between the one-hot true distribution and the softmax predicted distribution. The connection is not coincidental — cross-entropy loss is the information-theoretically optimal loss for classification.
logits = torch.tensor([2.0, 1.0, 0.1])
probs = torch.softmax(logits, dim=0)
true_class = 0
nll_loss = -torch.log(probs[true_class])
print(f"NLL loss: {nll_loss.item():.4f}")
ce_loss = nn.CrossEntropyLoss()(logits.unsqueeze(0), torch.tensor([true_class]))
print(f"CE loss: {ce_loss.item():.4f}") # Same value
What's Next?
With these mathematical foundations, you're ready to dive into PyTorch's tensor system in Module 02. The linear algebra you learned here will appear every time you work with layers (matrix multiplication), losses (norms, cross-entropy), and optimization (gradients, Adam). The key insight is that deep learning is applied linear algebra + calculus, automated by PyTorch's autograd system.
Run math_with_pytorch.py in this directory to see all these concepts in action with runnable code.
📓 Open Notebook — Interactive version of this module
Source Files
math_with_pytorch.py— Mathematical foundations with PyTorch
Module 02: Tensors — The Complete Guide
Table of Contents
- What is a Tensor?
- Tensor Creation
- Data Types
- Device Management
- Tensor Properties
- Element-wise Operations
- Reduction Operations
- Matrix Operations
- Tensor Manipulation
- Indexing Deep Dive
- Broadcasting
- Views vs Copies
- In-place Operations
- Strides Explained
- NumPy Interop
What is a Tensor?
A tensor is a multi-dimensional array — the fundamental data structure of PyTorch. If you know NumPy, a tensor is like np.ndarray but with two superpowers: GPU acceleration and automatic differentiation.
The word "tensor" comes from mathematics, where it refers to objects that obey certain transformation rules. In deep learning, we use the word more loosely to mean "an n-dimensional array of numbers."
Dimensionality taxonomy:
| Dimensions | Math Name | PyTorch Shape | Example |
|---|---|---|---|
| 0 | Scalar | torch.Size([]) | A single loss value: 0.543 |
| 1 | Vector | torch.Size([n]) | A word embedding: [0.2, -0.1, ...] |
| 2 | Matrix | torch.Size([m, n]) | A linear layer's weights |
| 3 | 3-tensor | torch.Size([a, b, c]) | A batch of sequences (batch, seq, embed) |
| 4 | 4-tensor | torch.Size([a, b, c, d]) | A batch of images (batch, channels, H, W) |
| N | N-tensor | torch.Size([...]) | Anything higher-dimensional |
import torch
scalar = torch.tensor(3.14) # 0-D
vector = torch.tensor([1, 2, 3]) # 1-D
matrix = torch.tensor([[1, 2], [3, 4]]) # 2-D
cube = torch.randn(2, 3, 4) # 3-D
print(scalar.ndim, vector.ndim, matrix.ndim, cube.ndim)
# 0, 1, 2, 3
Why "tensor" instead of "array"? The name emphasizes that these objects carry metadata (dtype, device, gradient tracking) and participate in PyTorch's autograd system. A NumPy array can't run on a GPU or compute gradients.
Tensor Creation
PyTorch offers many ways to create tensors. Each serves a specific purpose:
From Python data
# From a list
t = torch.tensor([1, 2, 3]) # int64 by default
t = torch.tensor([1.0, 2.0, 3.0]) # float32 by default
# From nested lists (matrix)
t = torch.tensor([[1, 2], [3, 4]])
# With explicit dtype
t = torch.tensor([1, 2, 3], dtype=torch.float32)
Constant-fill tensors
# Zeros and ones
torch.zeros(3, 4) # 3x4 matrix of zeros
torch.ones(2, 3, 5) # 2x3x5 tensor of ones
torch.full((2, 3), 7.0) # 2x3 matrix filled with 7.0
# Identity matrix
torch.eye(4) # 4x4 identity matrix
Why these exist: torch.zeros is used to initialize bias terms, accumulation buffers, and masks. torch.ones creates multiplicative identities. torch.eye creates identity matrices for residual connections and initialization schemes.
Random tensors
# Uniform [0, 1)
torch.rand(3, 4)
# Standard normal (mean=0, std=1)
torch.randn(3, 4)
# Random integers
torch.randint(low=0, high=10, size=(3, 4))
# Random permutation
torch.randperm(10) # A random shuffling of [0, 1, ..., 9]
Why random tensors matter: Weight initialization is crucial for training. torch.randn is the basis of Gaussian initialization. torch.randperm is used for shuffling dataset indices. Setting torch.manual_seed(42) makes random operations reproducible.
Sequences
# Integer sequence [0, 1, 2, ..., 9]
torch.arange(10)
torch.arange(2, 10) # [2, 3, ..., 9]
torch.arange(0, 1, 0.1) # [0.0, 0.1, ..., 0.9]
# Evenly spaced (specify count, not step)
torch.linspace(0, 1, steps=5) # [0.0, 0.25, 0.5, 0.75, 1.0]
torch.logspace(0, 2, steps=3) # [10^0, 10^1, 10^2] = [1, 10, 100]
Why these exist: torch.arange creates position indices. torch.linspace is used for evaluating functions over a range (plotting, interpolation). torch.logspace creates logarithmic learning rate schedules.
Uninitialized tensors
# Uninitialized — contains whatever was in memory!
torch.empty(3, 4)
Warning: torch.empty does NOT fill with zeros. It allocates memory and returns whatever garbage was there. Use it only when you'll immediately overwrite all values, as it avoids the cost of zero-filling. This matters in performance-critical code.
Like-functions (matching shape/dtype/device of an existing tensor)
x = torch.randn(3, 4, dtype=torch.float32)
torch.zeros_like(x) # Same shape, dtype, device — filled with zeros
torch.ones_like(x) # Same but with ones
torch.randn_like(x) # Same but with random normal values
torch.empty_like(x) # Same but uninitialized
torch.full_like(x, 5.0) # Same but filled with 5.0
Why these exist: When writing generic code (custom layers, loss functions), you often need to create a tensor with the same properties as an input. These functions handle dtype and device automatically, avoiding common bugs.
Data Types
Every tensor has a dtype (data type). Choosing the right one affects memory, speed, and numerical precision:
| dtype | Bits | Range/Precision | Use Case |
|---|---|---|---|
torch.float32 (default) | 32 | ~7 decimal digits | Standard training |
torch.float64 | 64 | ~15 decimal digits | Numerical verification, scientific computing |
torch.float16 | 16 | ~3 decimal digits | Mixed-precision training (with loss scaling) |
torch.bfloat16 | 16 | ~3 decimal digits, wider range | LLM training (Transformer-preferred) |
torch.int8 | 8 | [-128, 127] | Quantized inference |
torch.int16 | 16 | [-32768, 32767] | Rarely used |
torch.int32 | 32 | [-2^31, 2^31-1] | Indices, counts |
torch.int64 (default for ints) | 64 | [-2^63, 2^63-1] | Default integer type |
torch.bool | 8 | True/False | Masks, conditions |
torch.complex64 | 64 | Two float32 | Signal processing, FFT |
torch.complex128 | 128 | Two float64 | High-precision complex math |
torch.float8_e4m3fn | 8 | ~2 digits, narrow | Transformer engine inference |
torch.float8_e5m2 | 8 | ~1 digit, wider range | Transformer engine training |
float32 vs float16 vs bfloat16
This choice is one of the most important practical decisions in training:
- float32: Full precision. Use for the optimizer state and any numerically
sensitive operations (loss computation, normalization).
- float16: Half precision. 2x memory savings, faster on GPUs with tensor cores.
BUT has a limited range (max ~65504), which means gradients can overflow. Requires loss scaling for stable training.
- bfloat16: Same memory as float16 but with the same exponent range as float32.
This means it almost never overflows. Preferred for LLM training because you get memory savings without the numerical headaches of float16.
x = torch.randn(1000, 1000)
print(f"float32: {x.element_size()} bytes per element, {x.nelement() * x.element_size() / 1e6:.1f} MB")
x16 = x.to(torch.float16)
print(f"float16: {x16.element_size()} bytes per element, {x16.nelement() * x16.element_size() / 1e6:.1f} MB")
xbf = x.to(torch.bfloat16)
print(f"bfloat16: {xbf.element_size()} bytes per element")
Casting
x = torch.tensor([1, 2, 3]) # int64
x_float = x.float() # → float32
x_half = x.half() # → float16
x_double = x.double() # → float64
x_int = x_float.int() # → int32
x_long = x_float.long() # → int64
x_bool = x.bool() # → bool (0=False, nonzero=True)
x_cast = x.to(torch.float32) # General casting
Device Management
Tensors live on a specific device. Operations between tensors require them to be on the same device — this is a common source of runtime errors.
# CPU (default)
x_cpu = torch.randn(3, 4)
print(x_cpu.device) # cpu
# Check for GPU availability
if torch.cuda.is_available():
x_gpu = torch.randn(3, 4, device='cuda')
x_gpu = x_cpu.to('cuda') # Move to GPU
x_gpu = x_cpu.cuda() # Shorthand
x_back = x_gpu.cpu() # Move back to CPU
# Apple Silicon
if torch.backends.mps.is_available():
x_mps = x_cpu.to('mps')
# Context manager for default device
with torch.device('cpu'):
x = torch.randn(3, 4) # Created on CPU
Common error: RuntimeError: Expected all tensors to be on the same device. This happens when you mix CPU and GPU tensors in an operation. Always check .device when debugging.
Tensor Properties
Every tensor carries metadata that you can inspect:
x = torch.randn(2, 3, 4, dtype=torch.float32)
print(f"Shape: {x.shape}") # torch.Size([2, 3, 4])
print(f"Size (same): {x.size()}") # torch.Size([2, 3, 4])
print(f"Dimensions: {x.ndim}") # 3
print(f"Data type: {x.dtype}") # torch.float32
print(f"Device: {x.device}") # cpu
print(f"Total elements: {x.numel()}") # 24 (= 2*3*4)
print(f"Element size: {x.element_size()} bytes") # 4 (float32)
print(f"Total memory: {x.nelement() * x.element_size()} bytes") # 96
print(f"Strides: {x.stride()}") # (12, 4, 1)
print(f"Is contiguous: {x.is_contiguous()}") # True
print(f"Requires grad: {x.requires_grad}") # False
Understanding shape vs size: They're identical. x.shape is a property, x.size() is a method. Most people use .shape (following NumPy convention). You can index into shape: x.shape[0] gives the first dimension.
Element-wise Operations
Element-wise operations apply a function independently to each element. The shapes of input tensors must be compatible (same shape or broadcastable).
Arithmetic
a = torch.tensor([1.0, 2.0, 3.0])
b = torch.tensor([4.0, 5.0, 6.0])
a + b # tensor([5., 7., 9.]) — or torch.add(a, b)
a - b # tensor([-3., -3., -3.]) — or torch.sub(a, b)
a * b # tensor([4., 10., 18.]) — or torch.mul(a, b)
a / b # tensor([0.25, 0.4, 0.5]) — or torch.div(a, b)
a ** 2 # tensor([1., 4., 9.]) — or torch.pow(a, 2)
a // b # Floor division
a % b # Modulo
Mathematical functions
x = torch.tensor([0.0, 1.0, 2.0])
torch.exp(x) # e^x: [1.0, 2.718, 7.389]
torch.log(x + 1) # ln(x+1) — add 1 to avoid log(0)
torch.sqrt(x) # [0.0, 1.0, 1.414]
torch.abs(x - 1) # |x-1|: [1.0, 0.0, 1.0]
torch.sin(x) # Sine
torch.cos(x) # Cosine
torch.tanh(x) # Hyperbolic tangent (activation function)
torch.sigmoid(x) # 1 / (1 + e^(-x)) (activation function)
Clamping
x = torch.tensor([-3.0, -1.0, 0.5, 2.0, 5.0])
torch.clamp(x, min=0.0) # ReLU! [0, 0, 0.5, 2, 5]
torch.clamp(x, min=-1.0, max=1.0) # Clip to [-1, 1]
Why clamping matters: torch.clamp(x, min=0) is literally the ReLU activation function. Gradient clipping uses torch.clamp on gradient norms. Value clipping prevents numerical instability.
Reduction Operations
Reductions collapse one or more dimensions by aggregating values. Understanding the dim parameter is critical.
The dim parameter
When you specify dim=k, that dimension is "collapsed" (removed from the output):
x = torch.tensor([[1., 2., 3.],
[4., 5., 6.]]) # Shape: (2, 3)
x.sum() # 21.0 — all elements (scalar output)
x.sum(dim=0) # [5., 7., 9.] — collapse rows → shape (3,)
x.sum(dim=1) # [6., 15.] — collapse columns → shape (2,)
x.sum(dim=1, keepdim=True) # [[6.], [15.]] — shape (2, 1) — keeps the dim
Mental model for dim: Think "I'm reducing ALONG this axis." dim=0 means "go down the rows" (collapse them). dim=1 means "go across the columns." keepdim=True keeps the reduced dimension as size 1, which is essential for broadcasting the result back.
Common reductions
x = torch.tensor([[1., 2., 3.],
[4., 5., 6.]])
x.mean(dim=1) # [2., 5.] — row means
x.std(dim=1) # Standard deviation per row
x.var(dim=1) # Variance per row
x.max(dim=1) # Returns (values, indices) — both the max and where
x.min(dim=0) # Returns (values, indices) along dim 0
x.argmax(dim=1) # Index of max in each row: [2, 2]
x.argmin(dim=0) # Index of min in each column: [0, 0, 0]
x.prod(dim=1) # Product: [6., 120.]
torch.norm(x, dim=1) # L2 norm per row
The max/min return type
max and min return a named tuple with .values and .indices:
vals, idxs = x.max(dim=1)
print(f"Max values: {vals}") # [3., 6.]
print(f"Max indices: {idxs}") # [2, 2]
Matrix Operations
Matrix multiplication varieties
# 2D @ 2D: standard matrix multiply
A = torch.randn(3, 4)
B = torch.randn(4, 5)
C = A @ B # Shape: (3, 5)
C = torch.mm(A, B) # Equivalent, only for 2D
C = torch.matmul(A, B) # Most general
# Batched matrix multiply: each batch is independent
A = torch.randn(10, 3, 4) # 10 matrices of shape 3x4
B = torch.randn(10, 4, 5) # 10 matrices of shape 4x5
C = A @ B # Shape: (10, 3, 5) — 10 independent matmuls
C = torch.bmm(A, B) # Equivalent, only for 3D
# Vector-matrix products
v = torch.randn(4)
M = torch.randn(4, 5)
result = v @ M # Shape: (5,) — vector × matrix
result = M.T @ v # Shape: (5,) — equivalent
Einstein summation (einsum)
einsum is the Swiss army knife of tensor operations. It uses subscript notation to specify arbitrary contractions:
# Matrix multiply: "ik,kj->ij" means sum over k
A = torch.randn(3, 4)
B = torch.randn(4, 5)
C = torch.einsum('ik,kj->ij', A, B) # Same as A @ B
# Batch matrix multiply
A = torch.randn(10, 3, 4)
B = torch.randn(10, 4, 5)
C = torch.einsum('bik,bkj->bij', A, B) # Same as torch.bmm
# Dot product
a = torch.randn(5)
b = torch.randn(5)
d = torch.einsum('i,i->', a, b) # Same as torch.dot
# Outer product
outer = torch.einsum('i,j->ij', a, b) # Shape: (5, 5)
# Trace
M = torch.randn(4, 4)
tr = torch.einsum('ii->', M) # Same as torch.trace(M)
# Batch diagonal
B_diag = torch.einsum('bii->bi', torch.randn(3, 4, 4)) # Diagonal of each batch
Why einsum? It makes complex tensor operations readable. Instead of chains of transpose, reshape, and matmul, one einsum string describes the operation declaratively. Attention mechanisms are often clearest in einsum notation.
Tensor Manipulation
Reshaping
x = torch.arange(12) # [0, 1, 2, ..., 11]
# view: returns a view (shares memory) — requires contiguous input
x.view(3, 4) # Shape: (3, 4)
x.view(2, 2, 3) # Shape: (2, 2, 3)
x.view(-1, 4) # -1 is inferred: (3, 4)
x.view(-1) # Flatten: (12,)
# reshape: like view but works on non-contiguous tensors (may copy)
x.reshape(3, 4)
x.reshape(-1, 6) # (2, 6)
# flatten: collapse dimensions
y = torch.randn(2, 3, 4)
y.flatten() # Shape: (24,) — all dims
y.flatten(1) # Shape: (2, 12) — flatten from dim 1 onward
y.flatten(start_dim=1, end_dim=2) # Shape: (2, 12)
view vs reshape: view always returns a view (no memory copy), but requires the tensor to be contiguous in memory. reshape returns a view if possible, but will copy data if necessary. Rule of thumb: use reshape unless you specifically need to guarantee no copy (then use view and handle the contiguity yourself).
Transposing and permuting
x = torch.randn(2, 3, 4)
# transpose: swap exactly two dimensions
x.transpose(0, 1) # Shape: (3, 2, 4)
x.transpose(1, 2) # Shape: (2, 4, 3)
# For 2D matrices, .T is shorthand
m = torch.randn(3, 4)
m.T # Shape: (4, 3)
# permute: reorder ALL dimensions at once
x.permute(2, 0, 1) # Shape: (4, 2, 3) — moved dim 2 to front
Common use case: Images come as (batch, H, W, C) from some libraries but PyTorch expects (batch, C, H, W). Fix with: img.permute(0, 3, 1, 2).
Squeezing and unsqueezing
x = torch.randn(1, 3, 1, 4)
x.squeeze() # Remove ALL size-1 dims → shape (3, 4)
x.squeeze(0) # Remove dim 0 if size 1 → shape (3, 1, 4)
x.squeeze(2) # Remove dim 2 if size 1 → shape (1, 3, 4)
y = torch.randn(3, 4)
y.unsqueeze(0) # Add dim at position 0 → shape (1, 3, 4)
y.unsqueeze(1) # Add dim at position 1 → shape (3, 1, 4)
y.unsqueeze(-1) # Add dim at end → shape (3, 4, 1)
Why these matter: Many PyTorch operations expect specific numbers of dimensions. A single image (3, H, W) needs unsqueeze(0) to become a batch of 1 (1, 3, H, W) for a model. After processing, squeeze(0) removes the batch dimension.
Concatenation and stacking
a = torch.randn(2, 3)
b = torch.randn(2, 3)
# cat: join along EXISTING dimension
torch.cat([a, b], dim=0) # Shape: (4, 3) — stack vertically
torch.cat([a, b], dim=1) # Shape: (2, 6) — join horizontally
# stack: join along NEW dimension
torch.stack([a, b], dim=0) # Shape: (2, 2, 3) — new dim 0 indexes a vs b
torch.stack([a, b], dim=1) # Shape: (2, 2, 3) — new dim inserted at 1
Key difference: cat glues tensors along an existing dimension (dimensions must match elsewhere). stack creates a new dimension and places tensors along it (all tensors must have exactly the same shape).
Splitting and chunking
x = torch.arange(12).reshape(4, 3)
# split: split into pieces of given size
pieces = torch.split(x, 2, dim=0) # Two pieces of size 2 each
# pieces[0] shape: (2, 3), pieces[1] shape: (2, 3)
# chunk: split into N roughly equal pieces
chunks = torch.chunk(x, 3, dim=0) # Three chunks (sizes 2, 1, 1)
Expanding and repeating
x = torch.tensor([[1], [2], [3]]) # Shape: (3, 1)
# expand: view-based (no memory copy), only expands size-1 dims
x.expand(3, 4) # Shape: (3, 4) — [1,1,1,1], [2,2,2,2], [3,3,3,3]
x.expand(-1, 4) # -1 means "keep this dim's size"
# repeat: actually copies data
x.repeat(1, 4) # Shape: (3, 4) — same result but data is copied
x.repeat(2, 3) # Shape: (6, 3) — repeat 2x along dim 0, 3x along dim 1
expand vs repeat: expand is free (just changes strides, no memory allocation). repeat copies data. Always prefer expand when possible. But be careful: modifying an expanded tensor affects all "copies" since they share memory.
Indexing Deep Dive
Basic indexing (returns views)
x = torch.arange(20).reshape(4, 5)
# tensor([[ 0, 1, 2, 3, 4],
# [ 5, 6, 7, 8, 9],
# [10, 11, 12, 13, 14],
# [15, 16, 17, 18, 19]])
x[0] # Row 0: [0, 1, 2, 3, 4]
x[0, 2] # Element at row 0, col 2: 2
x[-1] # Last row: [15, 16, 17, 18, 19]
x[-1, -1] # Last element: 19
Slicing (returns views)
x[1:3] # Rows 1 and 2
x[:, 2:4] # All rows, columns 2 and 3
x[::2] # Every other row (stride 2)
x[:, ::-1] # Reverse columns (or use torch.flip)
x[1:3, 2:5] # Submatrix: rows 1-2, cols 2-4
Boolean (mask) indexing (returns copies)
x = torch.randn(4, 4)
mask = x > 0
positives = x[mask] # 1-D tensor of all positive values
x[mask] = 0 # Set all positive values to 0
x[x < -1] = -1 # Clamp from below
# Common pattern: conditional replacement
scores = torch.randn(5)
scores[scores < 0] = 0 # ReLU by hand
Why boolean indexing returns copies: The selected elements aren't contiguous in memory, so they can't form a view. Any modification to the result won't affect the original tensor (but assigning back with x[mask] = val does work).
Fancy (advanced) indexing
x = torch.arange(20).reshape(4, 5)
rows = torch.tensor([0, 2, 3])
cols = torch.tensor([1, 3, 4])
x[rows] # Select rows 0, 2, 3 → shape (3, 5)
x[rows, cols] # Select (0,1), (2,3), (3,4) → [1, 13, 19]
x[:, [0, 2, 4]] # Select columns 0, 2, 4 → shape (4, 3)
gather and scatter
gather collects values from a tensor according to an index tensor. Think of it as "for each position in the index, fetch the value at that index from the source."
# gather: collect elements
src = torch.tensor([[1, 2, 3],
[4, 5, 6]])
index = torch.tensor([[0, 2],
[1, 0]])
result = torch.gather(src, dim=1, index=index)
# result[0,0] = src[0, index[0,0]] = src[0, 0] = 1
# result[0,1] = src[0, index[0,1]] = src[0, 2] = 3
# result = tensor([[1, 3], [5, 4]])
# scatter: the inverse of gather — place values at index positions
dst = torch.zeros(2, 3, dtype=torch.long)
dst.scatter_(1, index, src=torch.tensor([[10, 20], [30, 40]]))
Where gather is used: Gathering token embeddings by index, selecting log probabilities at target positions (for computing cross-entropy manually), implementing top-k selection.
index_select and masked_select
x = torch.randn(4, 5)
torch.index_select(x, dim=0, index=torch.tensor([0, 3])) # Select rows 0 and 3
torch.index_select(x, dim=1, index=torch.tensor([1, 4])) # Select cols 1 and 4
mask = x > 0
selected = torch.masked_select(x, mask) # 1-D tensor of all True-mask elements
Broadcasting
Broadcasting lets you operate on tensors with different shapes without explicit copying. PyTorch follows NumPy's broadcasting rules.
The Rules (applied right-to-left)
- If tensors have different numbers of dimensions, prepend 1s to the smaller
tensor's shape until they match.
- For each dimension, sizes must either be equal or one of them must be 1.
- A dimension of size 1 is "stretched" to match the other.
# Example: (3, 4) + (4,) → (3, 4) + (1, 4) → (3, 4)
A = torch.randn(3, 4)
b = torch.randn(4) # Automatically becomes (1, 4), then broadcast to (3, 4)
C = A + b # Works! Each row of A gets b added
# Example: (3, 1) + (1, 4) → (3, 4)
col = torch.tensor([[1.], [2.], [3.]]) # Shape: (3, 1)
row = torch.tensor([[10., 20., 30., 40.]]) # Shape: (1, 4)
result = col + row # Shape: (3, 4) — addition table!
# [[11, 21, 31, 41],
# [12, 22, 32, 42],
# [13, 23, 33, 43]]
Common broadcasting patterns
# Subtract the mean from each row (feature normalization)
x = torch.randn(100, 10) # 100 samples, 10 features
mean = x.mean(dim=0) # Shape: (10,)
x_centered = x - mean # Broadcasting: (100, 10) - (10,) → (100, 10)
# Subtract mean from each column
col_mean = x.mean(dim=1, keepdim=True) # Shape: (100, 1)
x_col_centered = x - col_mean # (100, 10) - (100, 1) → (100, 10)
Why keepdim=True matters: Without it, x.mean(dim=1) has shape (100,). You can't subtract shape (100,) from shape (100, 10) — it's ambiguous. With keepdim=True, shape is (100, 1), which broadcasts unambiguously.
Broadcasting failure
# This FAILS:
a = torch.randn(3, 4)
b = torch.randn(5)
# a + b → Error! 4 ≠ 5 and neither is 1
Views vs Copies
Understanding when PyTorch shares memory vs copies it is crucial for both performance and correctness.
Operations that return views (share memory)
x = torch.arange(12).reshape(3, 4)
# These all share memory with x:
y = x.view(4, 3) # Reshape (requires contiguous)
y = x.reshape(4, 3) # Reshape (may return view)
y = x.T # Transpose
y = x.transpose(0, 1) # Transpose
y = x[0] # Basic indexing
y = x[:2] # Slicing
y = x.unsqueeze(0) # Add dimension
y = x.squeeze() # Remove size-1 dims
y = x.expand(2, 3, 4) # Expand size-1 dims
y = x.narrow(0, 0, 2) # Narrowing
# Modifying y also modifies x!
y = x.view(4, 3)
y[0, 0] = 999
print(x[0, 0]) # 999 — shared memory!
Operations that return copies
x = torch.arange(12).reshape(3, 4)
# These create new tensors (independent memory):
y = x.clone() # Explicit copy
y = x.contiguous() # Copy if not already contiguous
y = x[x > 5] # Boolean indexing
y = x[[0, 2]] # Fancy indexing
y = x.repeat(2, 1) # Repeat (not expand)
How to check
x = torch.arange(6).reshape(2, 3)
y = x.view(3, 2)
z = x.clone()
print(x.data_ptr() == y.data_ptr()) # True — same memory
print(x.data_ptr() == z.data_ptr()) # False — different memory
print(x.storage().data_ptr() == y.storage().data_ptr()) # True
In-place Operations
In-place operations modify a tensor's data directly without allocating new memory. They are indicated by a trailing underscore:
x = torch.tensor([1.0, 2.0, 3.0])
x.add_(1) # x is now [2, 3, 4] — no new tensor created
x.mul_(2) # x is now [4, 6, 8]
x.zero_() # x is now [0, 0, 0]
x.fill_(5) # x is now [5, 5, 5]
x.clamp_(min=0) # ReLU in-place
x.uniform_() # Fill with uniform random numbers
x.normal_() # Fill with normal random numbers
Autograd implications
In-place operations can break autograd because they modify data that the computation graph may need for backward:
x = torch.tensor([1.0, 2.0], requires_grad=True)
y = x * 2
y.add_(1) # This MAY cause an error during backward!
# PyTorch tracks in-place modifications via a version counter.
# If backward() needs y's original value, this will fail.
Rule of thumb: Avoid in-place operations on tensors that require gradients, unless you're explicitly operating on gradient buffers (like zeroing them) or on detached tensors. The one exception is param.data manipulation, which bypasses autograd entirely.
Strides Explained
Strides are the mechanism that makes views, transposes, and slices work without copying data. A stride tells PyTorch how many elements to skip in physical memory to advance one position along each dimension.
How strides work
x = torch.arange(12).reshape(3, 4)
# Memory layout: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
# Shape: (3, 4), Strides: (4, 1)
#
# To go from x[0,0] to x[1,0]: skip 4 elements (stride of dim 0)
# To go from x[0,0] to x[0,1]: skip 1 element (stride of dim 1)
print(f"Strides: {x.stride()}") # (4, 1)
Transpose changes strides, not data
x = torch.arange(6).reshape(2, 3)
print(f"x strides: {x.stride()}") # (3, 1) — row-major
xt = x.T
print(f"x.T strides: {xt.stride()}") # (1, 3) — column-major
print(f"x.T is contiguous: {xt.is_contiguous()}") # False!
# The data in memory hasn't changed! Only the interpretation (strides) changed.
# x: [0, 1, 2, 3, 4, 5] read as 2x3 with strides (3, 1)
# x.T: [0, 1, 2, 3, 4, 5] read as 3x2 with strides (1, 3)
Why contiguity matters: Some operations (like view) require contiguous memory. If a tensor isn't contiguous, call .contiguous() first (which copies data to a contiguous layout) or use .reshape() which handles this automatically.
Strides enable zero-copy slicing
x = torch.arange(20).reshape(4, 5)
y = x[::2, ::2] # Every other row, every other column
print(f"y strides: {y.stride()}") # (10, 2)
# Stride of 10 means skip 10 elements (2 rows of 5) to get next row
# Stride of 2 means skip 2 elements to get next column
# No data was copied!
NumPy Interop
PyTorch and NumPy can share memory, enabling seamless interoperation.
Conversion
import numpy as np
# NumPy → PyTorch (shares memory by default!)
np_array = np.array([1.0, 2.0, 3.0])
tensor = torch.from_numpy(np_array)
np_array[0] = 999
print(tensor[0]) # 999.0 — shared memory!
# PyTorch → NumPy (shares memory by default!)
tensor = torch.tensor([1.0, 2.0, 3.0])
np_array = tensor.numpy()
tensor[0] = 999
print(np_array[0]) # 999.0 — shared memory!
The shared memory warning
This is the single most common source of subtle bugs when mixing PyTorch and NumPy. If you modify one, the other changes too. To get an independent copy:
# Safe conversion (independent copy)
np_array = tensor.clone().numpy() # PyTorch → NumPy (safe)
tensor = torch.from_numpy(np_array.copy()) # NumPy → PyTorch (safe)
GPU tensors and NumPy
NumPy only works on CPU. You must move GPU tensors to CPU first:
# If x is on GPU:
# x.numpy() # ERROR! Can't convert CUDA tensor to numpy
# x.cpu().numpy() # OK — move to CPU first
# x.detach().cpu().numpy() # OK — also detaches from autograd
dtype compatibility
# NumPy float64 → PyTorch float64 (NOT the default float32!)
np_f64 = np.array([1.0, 2.0]) # numpy default is float64
t = torch.from_numpy(np_f64)
print(t.dtype) # torch.float64
# If you want float32, cast explicitly
t32 = torch.from_numpy(np_f64).float()
What's Next?
With tensors mastered, Module 03 covers autograd — PyTorch's automatic differentiation engine that makes all of deep learning possible. You'll learn how PyTorch tracks operations on tensors and computes gradients automatically.
Run the example files in this directory to practice:
creation_and_properties.py— tensor creation and inspectionoperations.py— element-wise and reduction operationsindexing_and_slicing.py— all forms of indexingbroadcasting.py— broadcasting rules with examplesviews_strides_memory.py— views, strides, and memory layout
📓 Open Notebook — Interactive version of this module
Source Files
creation_and_properties.py— Tensor creation and propertiesoperations.py— Tensor operationsindexing_and_slicing.py— Indexing and slicingbroadcasting.py— Broadcastingviews_strides_memory.py— Views, strides, and memory layout
Module 03: Autograd — Automatic Differentiation
Table of Contents
- What is Automatic Differentiation?
- Forward Mode vs Reverse Mode
- The Computation Graph
- Leaf Tensors vs Intermediate Tensors
- requires_grad, grad_fn, grad
- The backward() Function
- Gradient Accumulation
- torch.no_grad() vs torch.inference_mode()
- detach()
- Custom Autograd Functions
- gradcheck and gradgradcheck
- Higher-Order Gradients
- torch.autograd.grad()
- Jacobian and Hessian Computation
- Common Pitfalls
- Autograd Hooks
- Compiled Autograd
What is Automatic Differentiation?
Automatic differentiation (AD) is a technique for computing exact derivatives of functions expressed as computer programs. It is NOT:
- Symbolic differentiation (like Mathematica/Sympy): These manipulate
mathematical expressions symbolically, which can lead to expression swell (the derivative expression becomes exponentially larger than the original).
- Numerical differentiation (finite differences): Computing
(f(x+h) - f(x)) / h is simple but introduces truncation and rounding errors, and scales poorly with the number of parameters (requires one function evaluation per parameter).
Instead, AD exploits the fact that every computer program, no matter how complex, is composed of elementary operations (+, *, sin, exp, etc.) whose derivatives are known. By applying the chain rule systematically, AD computes exact derivatives at machine precision.
import torch
x = torch.tensor(3.0, requires_grad=True)
y = torch.sin(x) * torch.exp(x)
# PyTorch knows the derivative of sin, exp, and *.
# It chains them together to get dy/dx exactly.
y.backward()
print(f"dy/dx at x=3: {x.grad.item():.6f}")
# Verify: d/dx[sin(x)*exp(x)] = cos(x)*exp(x) + sin(x)*exp(x)
manual = (torch.cos(torch.tensor(3.0)) + torch.sin(torch.tensor(3.0))) * torch.exp(torch.tensor(3.0))
print(f"Manual: {manual.item():.6f}")
Forward Mode vs Reverse Mode
There are two ways to apply the chain rule through a computation:
Forward Mode (Tangent Mode)
Propagates derivatives forward through the computation, alongside the primal (original) computation. For f: R^n → R^m:
- Computes one column of the Jacobian per forward pass
- Cost: O(n) forward passes for full Jacobian
- Efficient when n << m (few inputs, many outputs)
Reverse Mode (Adjoint Mode) = Backpropagation
- Computes one row of the Jacobian per backward pass
- Cost: O(m) backward passes for full Jacobian
- Efficient when m << n (many inputs, few outputs)
Why neural networks use reverse mode: A neural network has millions of parameters (inputs to the loss function) but produces a single scalar loss (one output). Reverse mode computes ALL gradients (∂loss/∂param for every param) in a single backward pass — regardless of how many parameters there are. Forward mode would require one pass per parameter, which is millions of times slower.
This asymmetry is why:
- Training uses reverse mode (backpropagation): 1 scalar output, millions of inputs
- Forward-mode AD is used for Jacobian-vector products in some optimization methods
- PyTorch supports both:
backward()for reverse mode,torch.autograd.forward_ad
for forward mode
The Computation Graph
When you perform operations on tensors with requires_grad=True, PyTorch builds a directed acyclic graph (DAG) that records every operation. This graph is essential for computing gradients.
How the graph is built
x = torch.tensor(2.0, requires_grad=True)
y = torch.tensor(3.0, requires_grad=True)
# Each operation creates a node in the graph
a = x * y # MulBackward node
b = a + x # AddBackward node
c = b.sin() # SinBackward node
print(f"c.grad_fn: {c.grad_fn}") # SinBackward
print(f"b.grad_fn: {b.grad_fn}") # AddBackward
print(f"a.grad_fn: {a.grad_fn}") # MulBackward
print(f"x.grad_fn: {x.grad_fn}") # None (leaf tensor)
The graph looks like:
x ──→ MulBackward(a=x*y) ──→ AddBackward(b=a+x) ──→ SinBackward(c=sin(b))
y ──↗ x ──↗
Graph lifecycle
The graph is built dynamically during the forward pass and consumed (destroyed) during backward. This is the "dynamic graph" feature of PyTorch:
- Forward pass: Operations are recorded as nodes in the graph.
- backward(): The graph is traversed in reverse order, computing gradients.
- Graph destroyed: After backward, the graph is freed (by default).
- Next forward: A new graph is built from scratch.
This means every iteration can follow a different code path — if statements, for loops, and variable-length sequences all just work.
Leaf Tensors vs Intermediate Tensors
Leaf tensors
A leaf tensor is one that was created directly by the user (not by an operation). Leaf tensors are the "starting points" of the computation graph.
x = torch.tensor(1.0, requires_grad=True) # Leaf
w = torch.randn(3, 4, requires_grad=True) # Leaf
b = torch.zeros(4, requires_grad=True) # Leaf
print(f"x.is_leaf: {x.is_leaf}") # True
print(f"w.is_leaf: {w.is_leaf}") # True
Intermediate tensors
Tensors created by operations on other tensors are intermediate. They have a grad_fn that records how they were created.
y = x * 2 # Intermediate
z = w @ b # Intermediate
print(f"y.is_leaf: {y.is_leaf}") # False
print(f"y.grad_fn: {y.grad_fn}") # MulBackward0
Key difference: where gradients are stored
By default, PyTorch only stores gradients for leaf tensors. This is a memory optimization — intermediate gradients are computed during backward but discarded immediately after use.
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2
z = y * 3
z.backward()
print(f"x.grad: {x.grad}") # 12.0 — stored because x is a leaf
print(f"y.grad: {y.grad}") # None — not stored because y is intermediate
If you need gradients for intermediate tensors, use retain_grad():
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2
y.retain_grad() # Tell PyTorch to keep y's gradient
z = y * 3
z.backward()
print(f"y.grad: {y.grad}") # 3.0 — now it's stored
requires_grad, grad_fn, grad
These three attributes form the autograd trinity:
requires_grad
A boolean flag indicating whether this tensor participates in gradient computation.
# Tensors that need gradients
w = torch.randn(3, 3, requires_grad=True) # Set at creation
b = torch.zeros(3)
b.requires_grad_(True) # Set after creation (in-place)
# Tensors that don't
data = torch.randn(100, 3) # Data doesn't need gradients
labels = torch.randint(0, 2, (100,)) # Labels don't need gradients
Rule: If ANY input to an operation has requires_grad=True, the output will also have requires_grad=True (gradient tracking propagates forward).
grad_fn
A reference to the backward function (the operation that created this tensor). Leaf tensors have grad_fn=None.
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2 # y.grad_fn = PowBackward0
z = y.sum() # z.grad_fn = SumBackward0
You can traverse the graph by following grad_fn links:
print(z.grad_fn) # SumBackward0
print(z.grad_fn.next_functions) # Links to y's grad_fn
grad
The accumulated gradient tensor. Only populated after backward() is called. Only leaf tensors have .grad by default.
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
loss = (x ** 2).sum()
loss.backward()
print(f"x.grad: {x.grad}") # [2., 4., 6.] — the gradient ∂loss/∂x = 2x
The backward() Function
backward() is the workhorse of training. It computes gradients by traversing the computation graph in reverse order (topological sort).
How it works
x = torch.tensor(3.0, requires_grad=True)
y = x ** 3 + 2 * x ** 2 - 5 * x
y.backward()
# Computes dy/dx = 3x² + 4x - 5 = 27 + 12 - 5 = 34
print(f"dy/dx at x=3: {x.grad.item()}") # 34.0
The gradient argument
When the output is not a scalar, you must provide a gradient argument (also called the "upstream gradient" or "grad_output"):
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x ** 2 # y is a vector, not a scalar
# Must provide gradient (Jacobian-vector product)
y.backward(gradient=torch.tensor([1.0, 1.0, 1.0]))
print(f"x.grad: {x.grad}") # [2., 4., 6.]
# The gradient argument acts as weights in a weighted sum:
# effectively computing d(sum(gradient * y))/dx
Why? PyTorch's backward pass always computes vector-Jacobian products (VJPs). For a scalar output, the "vector" is implicitly 1.0. For vector outputs, you must supply it. In practice, this gradient comes from the downstream loss.
Graph destruction
By default, backward() destroys the computation graph after use:
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2
y.backward()
# y.backward() # ERROR: graph already freed
# To keep the graph, use retain_graph=True
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2
y.backward(retain_graph=True)
y.backward() # OK — graph was retained
Gradient Accumulation
Gradients accumulate by default. This is the single most common source of bugs for PyTorch beginners.
x = torch.tensor(1.0, requires_grad=True)
# First backward
y = x * 2
y.backward()
print(f"After first backward: x.grad = {x.grad}") # 2.0
# Second backward WITHOUT zeroing
y = x * 3
y.backward()
print(f"After second backward: x.grad = {x.grad}") # 5.0 (= 2.0 + 3.0)
# The gradient ACCUMULATED!
Why accumulation exists
Gradient accumulation is actually a feature, not a bug. It's useful for:
- Large effective batch sizes: Process mini-batches one at a time, accumulate
gradients, then take one optimizer step. This lets you simulate large batches on limited GPU memory.
- Multiple losses: If your model has multiple loss terms, you can backward
each one separately and the gradients add up correctly.
How to zero gradients
# Method 1: Manual zeroing
x.grad.zero_()
# Method 2: Optimizer (most common)
optimizer.zero_grad() # Zeros all parameter gradients
# Method 3: Set to None (slightly faster)
optimizer.zero_grad(set_to_none=True)
# or
x.grad = None
The standard training loop pattern
# for batch in dataloader:
# optimizer.zero_grad() # 1. Zero gradients
# output = model(batch) # 2. Forward pass
# loss = criterion(output) # 3. Compute loss
# loss.backward() # 4. Backward pass (compute gradients)
# optimizer.step() # 5. Update parameters
torch.no_grad() vs torch.inference_mode()
Both disable gradient tracking, but for different purposes:
torch.no_grad()
Disables gradient computation. Tensors created inside still have requires_grad based on their inputs, but no graph is built.
x = torch.tensor(1.0, requires_grad=True)
with torch.no_grad():
y = x * 2
print(f"y.requires_grad: {y.requires_grad}") # False
print(f"y.grad_fn: {y.grad_fn}") # None
Use for: Validation loops, parameter updates, evaluation metrics.
torch.inference_mode()
A stricter, more optimized version of no_grad(). Disables autograd entirely and enables additional optimizations.
x = torch.tensor(1.0, requires_grad=True)
with torch.inference_mode():
y = x * 2
# Even stricter: y is an InferenceTensor, cannot be used with autograd at all
Use for: Production inference, deployment, any time you're certain you won't need gradients.
When to use which
| Situation | Use |
|---|---|
| Validation loop during training | torch.no_grad() |
| Updating parameters manually | torch.no_grad() |
| Production inference | torch.inference_mode() |
| Need to use result in autograd later | Neither (or no_grad carefully) |
detach()
detach() creates a new tensor that shares data but is disconnected from the computation graph.
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2
z = y.detach() # z shares data with y but has no grad_fn
print(f"y.requires_grad: {y.requires_grad}") # True
print(f"z.requires_grad: {z.requires_grad}") # False
print(f"z.grad_fn: {z.grad_fn}") # None
print(f"Same data: {y.data_ptr() == z.data_ptr()}") # True
Common uses
- Stopping gradient flow: In GANs, you detach the generator output when
training the discriminator, so gradients don't flow back to the generator.
- Converting to NumPy:
tensor.detach().numpy()— you must detach first
if the tensor requires gradients.
- Target networks: In reinforcement learning, the target network's output
is detached so it's treated as a constant.
# GAN-style gradient stopping
# fake = generator(noise)
# d_loss = discriminator(fake.detach()) # Stops gradient to generator
Custom Autograd Functions
When PyTorch's built-in operations aren't sufficient, you can define custom forward and backward passes.
class MyReLU(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
return input.clamp(min=0)
@staticmethod
def backward(ctx, grad_output):
input, = ctx.saved_tensors
grad_input = grad_output.clone()
grad_input[input < 0] = 0
return grad_input
# Usage
x = torch.randn(5, requires_grad=True)
y = MyReLU.apply(x)
y.sum().backward()
print(f"x: {x}")
print(f"x.grad: {x.grad}")
The ctx object
ctx (context) is used to pass information from forward to backward:
ctx.save_for_backward(tensor1, tensor2, ...): Save tensors needed for backward.
These are stored efficiently and checked for version consistency.
ctx.saved_tensors: Retrieve saved tensors in backward.ctx.needs_input_grad: Tuple of booleans indicating which inputs need gradients.ctx.mark_dirty(tensor): Mark tensors modified in-place.ctx.mark_non_differentiable(tensor): Mark outputs that don't need gradients.
Rules for custom functions
forwardreceives the actual tensor values and returns output tensors.backwardreceives the upstream gradient (grad_output) for each output
and must return one gradient per input (or None if that input doesn't need gradients).
- The number of tensors returned by
backwardmust match the number of
inputs to forward (excluding ctx).
gradcheck and gradgradcheck
These utilities verify that your custom autograd function computes correct gradients by comparing against numerical finite differences.
from torch.autograd import gradcheck, gradgradcheck
class MySigmoid(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
result = 1 / (1 + torch.exp(-x))
ctx.save_for_backward(result)
return result
@staticmethod
def backward(ctx, grad_output):
result, = ctx.saved_tensors
return grad_output * result * (1 - result)
# gradcheck requires float64 for numerical precision
x = torch.randn(5, dtype=torch.float64, requires_grad=True)
test = gradcheck(MySigmoid.apply, (x,), eps=1e-6, atol=1e-4)
print(f"Gradient check passed: {test}")
# gradgradcheck checks second derivatives
test2 = gradgradcheck(MySigmoid.apply, (x,), eps=1e-6, atol=1e-4)
print(f"Double gradient check passed: {test2}")
Higher-Order Gradients
By default, backward() only computes first-order gradients. To compute higher-order gradients (gradients of gradients), use create_graph=True.
x = torch.tensor(2.0, requires_grad=True)
y = x ** 4 # y = x^4
# First derivative: dy/dx = 4x³
grad1 = torch.autograd.grad(y, x, create_graph=True)[0]
print(f"dy/dx = {grad1.item()}") # 32.0
# Second derivative: d²y/dx² = 12x²
grad2 = torch.autograd.grad(grad1, x, create_graph=True)[0]
print(f"d²y/dx² = {grad2.item()}") # 48.0
# Third derivative: d³y/dx³ = 24x
grad3 = torch.autograd.grad(grad2, x)[0]
print(f"d³y/dx³ = {grad3.item()}") # 48.0
Why create_graph=True
Without create_graph=True, the backward pass doesn't build a graph for itself. This means you can't differentiate through the backward pass. With it enabled, the gradient computation is itself differentiable.
Use cases:
- Gradient penalty (WGAN-GP): Penalize the norm of gradients, which requires
computing the gradient of the gradient norm.
- MAML (meta-learning): Differentiate through the inner optimization loop.
- Physics-informed neural networks: Loss functions involving derivatives of
the network output.
torch.autograd.grad()
torch.autograd.grad() computes gradients without storing them in .grad attributes. This is useful for:
- Computing gradients of specific outputs w.r.t. specific inputs
- Higher-order gradients
- Functional-style gradient computation
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = (x ** 2).sum()
# Compute gradient without storing in x.grad
grad = torch.autograd.grad(y, x)[0]
print(f"Gradient: {grad}") # [2., 4., 6.]
print(f"x.grad: {x.grad}") # None — not stored
Multiple outputs and inputs
x = torch.tensor(1.0, requires_grad=True)
y = torch.tensor(2.0, requires_grad=True)
out1 = x * y
out2 = x ** 2 + y ** 2
# Gradients of both outputs w.r.t. both inputs
grads = torch.autograd.grad(
outputs=[out1, out2],
inputs=[x, y],
grad_outputs=[torch.tensor(1.0), torch.tensor(1.0)]
)
print(f"d(out1+out2)/dx: {grads[0].item()}") # y + 2x = 2 + 2 = 4
print(f"d(out1+out2)/dy: {grads[1].item()}") # x + 2y = 1 + 4 = 5
Jacobian and Hessian Computation
Jacobian
The Jacobian matrix contains all partial derivatives of a vector-valued function. For f: R^n → R^m, J is m × n where J_ij = ∂f_i/∂x_j.
from torch.autograd.functional import jacobian
def f(x):
return torch.stack([x[0]**2 + x[1], x[0] * x[1]**2])
x = torch.tensor([1.0, 2.0])
J = jacobian(f, x)
print(f"Jacobian:\n{J}")
# [[2*x0, 1 ], = [[2, 1],
# [x1^2, 2*x0*x1]] [4, 4]]
Hessian
H is n × n where H_ij = ∂²f/∂x_i∂x_j.
from torch.autograd.functional import hessian
def g(x):
return x[0]**3 + x[0]*x[1]**2 + x[1]**3
x = torch.tensor([1.0, 2.0])
H = hessian(g, x)
print(f"Hessian:\n{H}")
# [[6*x0, 2*x1 ], = [[6, 4],
# [2*x1, 2*x0+6*x1]] [4, 14]]
Efficient Jacobian-vector products (JVPs) and vector-Jacobian products (VJPs)
Computing the full Jacobian is expensive (O(n) backward passes). Often you only need the product of the Jacobian with a specific vector:
from torch.autograd.functional import jvp, vjp
def f(x):
return torch.stack([x[0]**2, x[0]*x[1]])
x = torch.tensor([1.0, 2.0])
v = torch.tensor([1.0, 0.0]) # Direction vector
# JVP: J @ v (forward mode — one forward pass)
_, jvp_result = jvp(f, (x,), (v,))
print(f"JVP (J @ v): {jvp_result}")
# VJP: v^T @ J (reverse mode — one backward pass)
_, vjp_fn = vjp(f, x)
vjp_result = vjp_fn(torch.tensor([1.0, 0.0]))
print(f"VJP (v^T @ J): {vjp_result}")
Common Pitfalls
1. Forgetting to zero gradients
x = torch.tensor(1.0, requires_grad=True)
for i in range(3):
loss = x * (i + 1)
loss.backward()
print(f"Step {i}: x.grad = {x.grad.item()}")
# Without zeroing: 1.0, 3.0, 6.0 (accumulated!)
# x.grad.zero_() # Uncomment to fix
2. In-place operations on tensors that need gradients
x = torch.tensor([1.0, 2.0], requires_grad=True)
y = x * 2
# This WILL cause problems:
# y.add_(1) # In-place modification of a tensor needed for backward
# y.backward(torch.tensor([1.0, 1.0])) # RuntimeError!
3. Gradient not flowing through integer operations
x = torch.tensor(3.0, requires_grad=True)
y = x.int() # Casting to int breaks gradient flow
# y has no grad_fn — gradient is lost!
4. NaN gradients
Common causes:
log(0): Uselog(x + epsilon)ortorch.clamp(x, min=1e-8)sqrt(0): derivative of sqrt at 0 is infinity. Usesqrt(x + epsilon)0/0: Can occur in normalization layers when variance is zero- Division by a very small number: use
torch.clampon denominators
x = torch.tensor(0.0, requires_grad=True)
# y = torch.log(x) # Will give -inf, grad is inf
y = torch.log(x + 1e-8) # Safe
y.backward()
print(f"Safe log grad: {x.grad}")
5. Modifying parameters without torch.no_grad()
x = torch.tensor(1.0, requires_grad=True)
# x = x - 0.1 * x.grad # This creates a NEW tensor, breaking the leaf status
# Instead:
# with torch.no_grad():
# x -= 0.1 * x.grad
Autograd Hooks
Hooks let you inspect or modify gradients during the backward pass without changing the model code.
Tensor hooks
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
def print_grad(grad):
print(f" Hook received gradient: {grad}")
x.register_hook(print_grad)
y = (x ** 2).sum()
y.backward() # Hook fires during backward
Gradient modification hooks
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
def clip_grad(grad):
return torch.clamp(grad, -1.0, 1.0)
x.register_hook(clip_grad)
y = (x ** 3).sum()
y.backward()
print(f"Clipped gradient: {x.grad}") # Gradients are clipped to [-1, 1]
Module hooks
For nn.Module, you can register hooks on the entire module:
import torch.nn as nn
model = nn.Linear(5, 3)
def forward_hook(module, input, output):
print(f"Forward: input shape={input[0].shape}, output shape={output.shape}")
def backward_hook(module, grad_input, grad_output):
print(f"Backward: grad_output shape={grad_output[0].shape}")
model.register_forward_hook(forward_hook)
model.register_full_backward_hook(backward_hook)
x = torch.randn(2, 5)
y = model(x)
y.sum().backward()
Compiled Autograd
PyTorch 2.x introduced compiled autograd, which applies torch.compile to the backward pass. This can significantly speed up training by:
- Fusing backward operations into optimized kernels
- Eliminating Python overhead in the backward pass
- Enabling whole-graph optimization across forward and backward
# Basic usage (requires torch >= 2.0)
model = torch.nn.Linear(100, 10)
compiled_model = torch.compile(model)
# Both forward AND backward are compiled
x = torch.randn(32, 100)
y = compiled_model(x)
loss = y.sum()
loss.backward() # This backward is also compiled
The compilation happens lazily — the first iteration traces the graph, subsequent iterations use the compiled version. This means the first iteration is slower (compilation overhead) but all subsequent iterations are faster.
Summary
Autograd is the engine that makes deep learning practical. Key takeaways:
- Reverse-mode AD is efficient for neural networks (one backward pass for
all parameters).
- The computation graph is built dynamically during forward and consumed
during backward.
- Always zero gradients before each optimization step.
- Use
torch.no_grad()for validation andinference_mode()for deployment. - Custom Functions let you define operations with hand-written gradients.
- Hooks let you inspect and modify gradients without changing model code.
Run the example files to see these concepts in action:
gradient_basics.py— basic gradient computation and the training loopcomputation_graph.py— visualizing and understanding the computation graphcustom_functions.py— writing custom autograd functionshigher_order_gradients.py— second derivatives, Jacobians, and Hessians
📓 Open Notebook — Interactive version of this module
Source Files
gradient_basics.py— Gradient basicscomputation_graph.py— Computation graphcustom_functions.py— Custom autograd functionshigher_order_gradients.py— Higher-order gradients, Jacobians, and Hessians
Module 04: Neural Networks in PyTorch
Table of Contents
- What is nn.Module?
- Module Lifecycle
- Parameters vs Buffers
- Container Modules
- Linear Layers
- Convolution Layers
- Pooling Layers
- Normalization Layers
- Activation Functions
- Dropout
- Recurrent Layers
- Transformer Layers
- Embedding Layers
- Loss Functions
- Functional API
- Weight Initialization
- Hooks
- State Dict
What is nn.Module?
torch.nn.Module is the base class for ALL neural network components in PyTorch. Every layer, every model, every building block inherits from this class. When you build a neural network in PyTorch, you create a class that inherits from nn.Module.
Think of nn.Module as a container that:
- Holds learnable parameters (weights and biases)
- Defines how data flows through the network (the
forwardmethod) - Provides utilities for moving to GPU, saving/loading, switching train/eval modes
- Builds a tree of sub-modules (layers inside your model)
import torch
import torch.nn as nn
class MyNetwork(nn.Module):
def __init__(self):
super().__init__() # MUST call parent __init__
self.linear1 = nn.Linear(784, 256)
self.linear2 = nn.Linear(256, 10)
def forward(self, x):
x = torch.relu(self.linear1(x))
x = self.linear2(x)
return x
Key insight: When you assign an nn.Module or nn.Parameter as an attribute of your module (in __init__), PyTorch automatically registers it. This means it will show up in .parameters(), be moved when you call .to(device), and be saved in .state_dict().
Module Lifecycle
__init__: Construction
In __init__, you define all the layers and parameters your network needs. You MUST call super().__init__() first. Any nn.Module or nn.Parameter assigned as an attribute gets automatically registered.
class Model(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super().__init__()
# These are automatically registered as sub-modules
self.layer1 = nn.Linear(input_dim, hidden_dim)
self.layer2 = nn.Linear(hidden_dim, output_dim)
# This is NOT registered (plain Python attribute)
self.activation_name = "relu"
forward: The Forward Pass
The forward method defines how input data flows through your network. You NEVER call forward() directly — instead, you call the module as a function: model(x). This is because __call__ does extra work (hooks, checks) before calling forward.
def forward(self, x):
x = torch.relu(self.layer1(x))
x = self.layer2(x)
return x
# Correct usage:
output = model(input_tensor) # calls __call__ which calls forward
# WRONG — never do this:
# output = model.forward(input_tensor)
Train vs Eval Mode
Modules have two modes that affect behavior of certain layers (Dropout, BatchNorm):
model.train() # Sets training mode (dropout active, batchnorm uses batch stats)
model.eval() # Sets evaluation mode (dropout disabled, batchnorm uses running stats)
# Check current mode
print(model.training) # True or False
This is critical: forgetting model.eval() during inference causes incorrect results because Dropout still drops neurons and BatchNorm uses batch statistics instead of learned running statistics.
Parameters vs Buffers
Parameters
Parameters are tensors that require gradients and are updated by the optimizer. They represent the learnable weights of your model.
class CustomLayer(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
# Manual parameter creation
self.weight = nn.Parameter(torch.randn(out_features, in_features))
self.bias = nn.Parameter(torch.zeros(out_features))
def forward(self, x):
return x @ self.weight.t() + self.bias
Buffers
Buffers are tensors that are part of the module's state but do NOT require gradients. They are saved in state_dict() and moved with .to(device), but the optimizer ignores them.
Common use cases:
- Running mean/variance in BatchNorm
- Fixed positional encodings
- Binary masks that don't change during training
class MyModule(nn.Module):
def __init__(self):
super().__init__()
# Buffer: saved in state_dict, moved with .to(), but NOT optimized
self.register_buffer('running_mean', torch.zeros(10))
# Non-persistent buffer: moved with .to() but NOT saved
self.register_buffer('temp_mask', torch.ones(10), persistent=False)
register_parameter and register_buffer
class ExplicitRegistration(nn.Module):
def __init__(self):
super().__init__()
# Explicit parameter registration (equivalent to self.weight = nn.Parameter(...))
self.register_parameter('weight', nn.Parameter(torch.randn(5, 3)))
# Can register None — useful for optional parameters
self.register_parameter('optional_bias', None)
# Buffer registration
self.register_buffer('counter', torch.tensor(0))
Traversing the Module Tree
model = MyNetwork()
# All parameters (recursively)
for name, param in model.named_parameters():
print(f"{name}: shape={param.shape}, requires_grad={param.requires_grad}")
# All sub-modules (recursively)
for name, module in model.named_modules():
print(f"{name}: {type(module).__name__}")
# Direct children only
for name, module in model.named_children():
print(f"{name}: {type(module).__name__}")
# All buffers
for name, buf in model.named_buffers():
print(f"{name}: shape={buf.shape}")
Container Modules
nn.Sequential
Chains modules in order. Input flows through each module sequentially.
model = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 128),
nn.ReLU(),
nn.Linear(128, 10)
)
# Equivalent to calling each in order: output = layer3(relu(layer2(relu(layer1(x)))))
output = model(input_tensor)
Use when: Your network is a simple chain of operations with no branching.
nn.ModuleList
A list of modules. Does NOT define a forward pass — you iterate manually.
class MultiHeadModel(nn.Module):
def __init__(self, num_heads):
super().__init__()
self.heads = nn.ModuleList([nn.Linear(256, 10) for _ in range(num_heads)])
def forward(self, x):
return [head(x) for head in self.heads]
Use when: You need a variable number of layers that you'll iterate over yourself. WARNING: A plain Python list [] will NOT register the modules!
nn.ModuleDict
A dictionary of modules, accessed by string keys.
class MultiTaskModel(nn.Module):
def __init__(self):
super().__init__()
self.backbone = nn.Linear(784, 256)
self.heads = nn.ModuleDict({
'classification': nn.Linear(256, 10),
'regression': nn.Linear(256, 1),
})
def forward(self, x, task):
features = torch.relu(self.backbone(x))
return self.heads[task](features)
Use when: You need named access to different sub-modules (multi-task, configurable architectures).
nn.ParameterList and nn.ParameterDict
Same idea but for raw parameters instead of modules:
class CustomModel(nn.Module):
def __init__(self, num_experts):
super().__init__()
self.expert_weights = nn.ParameterList(
[nn.Parameter(torch.randn(256, 256)) for _ in range(num_experts)]
)
self.config_params = nn.ParameterDict({
'scale': nn.Parameter(torch.ones(1)),
'shift': nn.Parameter(torch.zeros(1)),
})
Linear Layers
nn.Linear
Applies a linear transformation: y = xW^T + b
# Input: (batch_size, in_features) e.g., (32, 784)
# Output: (batch_size, out_features) e.g., (32, 256)
linear = nn.Linear(in_features=784, out_features=256, bias=True)
# Weight shape: (out_features, in_features) = (256, 784)
# Bias shape: (out_features,) = (256,)
The math: For input x of shape (, in_features), output y of shape (, out_features): y_i = sum_j(x_j * W_ij) + b_i
nn.Bilinear
Applies a bilinear transformation: y = x1^T A x2 + b
# Two inputs of potentially different sizes
bilinear = nn.Bilinear(in1_features=20, in2_features=30, out_features=40)
input1 = torch.randn(128, 20)
input2 = torch.randn(128, 30)
output = bilinear(input1, input2) # shape: (128, 40)
nn.LazyLinear
Infers in_features from the first input — useful for prototyping:
lazy = nn.LazyLinear(out_features=256)
# in_features is determined on first forward pass
output = lazy(torch.randn(32, 784)) # Now it knows in_features=784
Convolution Layers
Core Concepts
Convolutions slide a kernel (filter) across the input, computing dot products at each position.
Key parameters:
kernel_size: Size of the sliding window (e.g., 3 means 3x3 for Conv2d)stride: How far the kernel moves each step (default=1)padding: Zero-padding added to input bordersdilation: Spacing between kernel elements (dilated/atrous convolution)groups: Split input channels into groups for grouped convolution
Output size formula (for each spatial dimension):
output_size = floor((input_size + 2*padding - dilation*(kernel_size-1) - 1) / stride + 1)
nn.Conv1d
For sequential/temporal data (text, audio, time series).
# Input: (batch, in_channels, length) e.g., (32, 1, 100)
# Output: (batch, out_channels, new_length) e.g., (32, 16, 98)
conv1d = nn.Conv1d(in_channels=1, out_channels=16, kernel_size=3)
nn.Conv2d
For image data. The most commonly used convolution.
# Input: (batch, in_channels, height, width) e.g., (32, 3, 224, 224)
# Output: (batch, out_channels, new_height, new_width) e.g., (32, 64, 112, 112)
conv2d = nn.Conv2d(
in_channels=3, # RGB input
out_channels=64, # 64 filters
kernel_size=3, # 3x3 kernel
stride=2, # Downsample by 2
padding=1 # Same padding for stride=1
)
# Weight shape: (out_channels, in_channels/groups, kernel_h, kernel_w)
# = (64, 3, 3, 3)
nn.Conv3d
For volumetric data (video, 3D medical images).
# Input: (batch, channels, depth, height, width)
conv3d = nn.Conv3d(in_channels=3, out_channels=64, kernel_size=3, padding=1)
ConvTranspose2d (Transposed Convolution)
Used for upsampling — goes from smaller to larger spatial dimensions. Often called "deconvolution" (technically incorrect name).
# Input: (batch, in_channels, H, W)
# Output: (batch, out_channels, H*2, W*2) with stride=2
upsample = nn.ConvTranspose2d(
in_channels=64, out_channels=32,
kernel_size=4, stride=2, padding=1
)
Depthwise Separable Convolution
A two-step convolution that's much more efficient:
- Depthwise: Apply one filter per input channel (groups=in_channels)
- Pointwise: 1x1 convolution to mix channels
class DepthwiseSeparable(nn.Module):
def __init__(self, in_ch, out_ch, kernel_size=3, padding=1):
super().__init__()
# Depthwise: each input channel gets its own filter
self.depthwise = nn.Conv2d(in_ch, in_ch, kernel_size,
padding=padding, groups=in_ch)
# Pointwise: 1x1 conv to combine channels
self.pointwise = nn.Conv2d(in_ch, out_ch, kernel_size=1)
def forward(self, x):
x = self.depthwise(x)
x = self.pointwise(x)
return x
Pooling Layers
Pooling reduces spatial dimensions while retaining important information.
MaxPool2d
Takes the maximum value in each pooling window:
# Input: (batch, channels, 224, 224)
# Output: (batch, channels, 112, 112) — halves spatial dims
pool = nn.MaxPool2d(kernel_size=2, stride=2)
AvgPool2d
Takes the average value in each pooling window:
pool = nn.AvgPool2d(kernel_size=2, stride=2)
AdaptiveAvgPool2d — Global Average Pooling
Outputs a fixed spatial size regardless of input size. Setting output to (1,1) gives "global average pooling" — commonly used before the final classifier.
# No matter what spatial size comes in, output is (batch, channels, 1, 1)
gap = nn.AdaptiveAvgPool2d(output_size=(1, 1))
# Then flatten: (batch, channels, 1, 1) -> (batch, channels)
This is the modern replacement for large fully-connected layers at the end of CNNs.
Normalization Layers
BatchNorm (nn.BatchNorm1d, BatchNorm2d)
Normalizes across the batch dimension. For each feature/channel:
Formula:
y = (x - E[x]) / sqrt(Var[x] + eps) * gamma + beta
Where gamma (weight) and beta (bias) are learnable parameters.
Training behavior: Uses batch mean and variance, updates running statistics. Eval behavior: Uses stored running mean and variance (fixed).
bn = nn.BatchNorm2d(num_features=64) # 64 channels
# Maintains: running_mean, running_var (buffers), weight, bias (parameters)
When to use: CNNs with large batch sizes. Not suitable for batch_size=1 or variable batch sizes (use LayerNorm or GroupNorm instead).
LayerNorm
Normalizes across the feature dimensions (not the batch). Each sample is normalized independently.
# For a transformer with hidden_size=512
ln = nn.LayerNorm(normalized_shape=512)
# For image data: normalize over (C, H, W)
ln_image = nn.LayerNorm([64, 32, 32])
When to use: Transformers, RNNs, any case where batch stats are unreliable.
GroupNorm
Splits channels into groups and normalizes within each group. A middle ground between BatchNorm (all channels) and InstanceNorm (each channel separately).
gn = nn.GroupNorm(num_groups=32, num_channels=256)
When to use: When batch size is small, or when you want BatchNorm-like behavior without batch dependency.
InstanceNorm
Normalizes each channel of each sample independently. Equivalent to GroupNorm with num_groups = num_channels.
inst_norm = nn.InstanceNorm2d(num_features=64)
When to use: Style transfer, image generation tasks.
RMSNorm
Root Mean Square Layer Normalization — simpler than LayerNorm (no mean subtraction):
rms_norm = nn.RMSNorm(normalized_shape=512)
Formula: y = x / RMS(x) * gamma where RMS(x) = sqrt(mean(x^2) + eps)
When to use: Modern LLMs (LLaMA, etc.) — slightly faster than LayerNorm.
Activation Functions
Activations introduce non-linearity. Without them, stacking linear layers is equivalent to a single linear layer.
ReLU: f(x) = max(0, x)
The default choice. Simple, fast, but can "die" (output 0 for all inputs).
nn.ReLU(inplace=False) # inplace=True saves memory but can cause issues
LeakyReLU: f(x) = max(alpha*x, x) (default alpha=0.01)
Prevents dying ReLU by allowing small negative slope.
nn.LeakyReLU(negative_slope=0.01)
PReLU: f(x) = max(alpha*x, x) where alpha is LEARNED
nn.PReLU(num_parameters=1) # One alpha per channel if num_parameters=num_channels
GELU: f(x) = x * Phi(x) where Phi is the CDF of standard normal
Used in Transformers (BERT, GPT). Smooth approximation of ReLU.
nn.GELU(approximate='none') # 'tanh' for faster approximation
SiLU/Swish: f(x) = x * sigmoid(x)
Smooth, non-monotonic. Used in EfficientNet, many modern architectures.
nn.SiLU()
Mish: f(x) = x * tanh(softplus(x))
Similar to SiLU but slightly different properties.
nn.Mish()
Sigmoid: f(x) = 1 / (1 + exp(-x))
Squashes to [0, 1]. Used for binary classification output, gates.
nn.Sigmoid()
Tanh: f(x) = (exp(x) - exp(-x)) / (exp(x) + exp(-x))
Squashes to [-1, 1]. Used in RNN gates.
nn.Tanh()
Softmax: f(x_i) = exp(x_i) / sum(exp(x_j))
Outputs a probability distribution (sums to 1). Used for multi-class classification.
nn.Softmax(dim=-1) # Usually along the last dimension
Dropout
Dropout randomly zeros elements during training to prevent overfitting. During evaluation, dropout is disabled and outputs are unchanged.
Key insight: During training, remaining elements are scaled by 1/(1-p) so that expected values remain the same at test time (inverted dropout).
nn.Dropout
dropout = nn.Dropout(p=0.5) # 50% of elements zeroed during training
nn.Dropout2d
Drops entire channels (feature maps) for Conv2d outputs:
dropout2d = nn.Dropout2d(p=0.1) # Drops entire channels
nn.AlphaDropout
For use with SELU activation — maintains self-normalizing property:
alpha_dropout = nn.AlphaDropout(p=0.1)
Recurrent Layers
nn.RNN
Basic recurrent layer: h_t = tanh(x_t W_ih^T + h_{t-1} W_hh^T + b)
rnn = nn.RNN(input_size=128, hidden_size=256, num_layers=2,
batch_first=True, bidirectional=False, dropout=0.1)
# Input: (batch, seq_len, input_size)
# Output: (batch, seq_len, hidden_size * num_directions)
# Hidden: (num_layers * num_directions, batch, hidden_size)
output, h_n = rnn(input_seq, h_0)
nn.LSTM
Long Short-Term Memory — solves vanishing gradients with gates:
Gate equations:
f_t = sigmoid(W_f [h_{t-1}, x_t] + b_f) # Forget gate
i_t = sigmoid(W_i [h_{t-1}, x_t] + b_i) # Input gate
g_t = tanh(W_g [h_{t-1}, x_t] + b_g) # Cell candidate
o_t = sigmoid(W_o [h_{t-1}, x_t] + b_o) # Output gate
c_t = f_t * c_{t-1} + i_t * g_t # Cell state
h_t = o_t * tanh(c_t) # Hidden state
lstm = nn.LSTM(input_size=128, hidden_size=256, num_layers=2,
batch_first=True, bidirectional=True, dropout=0.1)
output, (h_n, c_n) = lstm(input_seq)
# output shape: (batch, seq_len, hidden_size * 2) for bidirectional
nn.GRU
Gated Recurrent Unit — simpler than LSTM with fewer parameters:
Gate equations:
r_t = sigmoid(W_r [h_{t-1}, x_t]) # Reset gate
z_t = sigmoid(W_z [h_{t-1}, x_t]) # Update gate
n_t = tanh(W_n [r_t * h_{t-1}, x_t]) # New gate
h_t = (1 - z_t) * n_t + z_t * h_{t-1} # Hidden state
gru = nn.GRU(input_size=128, hidden_size=256, num_layers=2,
batch_first=True, bidirectional=True)
output, h_n = gru(input_seq)
Transformer Layers
nn.MultiheadAttention
Computes scaled dot-product attention across multiple heads:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V
mha = nn.MultiheadAttention(embed_dim=512, num_heads=8, dropout=0.1,
batch_first=True)
# Self-attention: query = key = value
attn_output, attn_weights = mha(query, key, value, key_padding_mask=mask)
nn.TransformerEncoderLayer
One layer of a Transformer encoder (self-attention + feedforward):
encoder_layer = nn.TransformerEncoderLayer(
d_model=512, nhead=8, dim_feedforward=2048,
dropout=0.1, activation='gelu', batch_first=True,
norm_first=True # Pre-norm (more stable training)
)
nn.TransformerEncoder
Stack of TransformerEncoderLayers:
encoder = nn.TransformerEncoder(encoder_layer, num_layers=6)
output = encoder(src, src_key_padding_mask=padding_mask)
Embedding Layers
nn.Embedding
Lookup table that maps integer indices to dense vectors:
# vocab_size=10000, embedding_dim=256
embed = nn.Embedding(num_embeddings=10000, embedding_dim=256, padding_idx=0)
# Input: (batch, seq_len) of integer indices
# Output: (batch, seq_len, embedding_dim)
token_ids = torch.tensor([[1, 45, 234, 0, 0]]) # 0 = padding
embeddings = embed(token_ids) # shape: (1, 5, 256)
padding_idx: The embedding at this index is always zero and is not updated.
nn.EmbeddingBag
More efficient when you need to sum/mean/max embeddings (e.g., bag-of-words):
embed_bag = nn.EmbeddingBag(num_embeddings=10000, embedding_dim=256, mode='mean')
# Returns one vector per "bag" — no need to manually average
Loss Functions
nn.CrossEntropyLoss
For multi-class classification. Combines LogSoftmax + NLLLoss.
Formula: loss = -log(exp(x_y) / sum(exp(x_j))) where y is the true class.
criterion = nn.CrossEntropyLoss()
# logits: (batch, num_classes) — RAW scores, NOT softmax
# target: (batch,) — class indices (integers)
loss = criterion(logits, targets)
nn.BCEWithLogitsLoss
For binary or multi-label classification. Combines Sigmoid + BCELoss.
Formula: loss = -[ylog(sigmoid(x)) + (1-y)log(1-sigmoid(x))]
criterion = nn.BCEWithLogitsLoss()
# logits: (batch, num_labels) — RAW scores
# target: (batch, num_labels) — 0.0 or 1.0
nn.MSELoss
Mean Squared Error for regression.
Formula: loss = mean((y_pred - y_true)^2)
criterion = nn.MSELoss()
nn.L1Loss
Mean Absolute Error for regression.
Formula: loss = mean(|y_pred - y_true|)
nn.HuberLoss (Smooth L1)
Combination of L1 and L2 — less sensitive to outliers than MSE.
Formula: L2 for |error| < delta, L1 for |error| >= delta.
criterion = nn.HuberLoss(delta=1.0)
nn.KLDivLoss
Kullback-Leibler Divergence — measures how one distribution differs from another.
Formula: loss = y_true * (log(y_true) - x)
criterion = nn.KLDivLoss(reduction='batchmean', log_target=False)
# Input must be log-probabilities!
nn.TripletMarginLoss
For metric learning with (anchor, positive, negative) triplets.
Formula: loss = max(d(anchor, positive) - d(anchor, negative) + margin, 0)
criterion = nn.TripletMarginLoss(margin=1.0)
loss = criterion(anchor, positive, negative)
nn.CosineEmbeddingLoss
Measures cosine similarity between pairs.
criterion = nn.CosineEmbeddingLoss(margin=0.0)
# target: +1 (similar) or -1 (dissimilar)
loss = criterion(x1, x2, target)
Functional API
torch.nn.functional (commonly imported as F) provides the same operations as nn.Module layers but as pure functions without stored state.
import torch.nn.functional as F
# Module version (has stored parameters):
relu_module = nn.ReLU()
output = relu_module(x)
# Functional version (stateless):
output = F.relu(x)
When to use Module vs Functional:
- Use Module when the operation has learnable parameters (Linear, Conv, BatchNorm)
- Use Module when the operation has different train/eval behavior (Dropout, BatchNorm)
- Use Functional for stateless operations (relu, softmax in forward pass)
- Use Functional when you need the operation in a custom forward pass without
wanting to register it as a sub-module
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.conv = nn.Conv2d(3, 64, 3) # Module: has parameters
self.bn = nn.BatchNorm2d(64) # Module: has state (running stats)
def forward(self, x):
x = self.conv(x)
x = self.bn(x)
x = F.relu(x) # Functional: no state needed
x = F.dropout(x, p=0.5, training=self.training) # Must pass training flag!
return x
Weight Initialization
Proper initialization prevents vanishing/exploding gradients at the start of training.
Xavier (Glorot) Initialization
Designed for sigmoid/tanh activations. Keeps variance constant across layers.
Formula: W ~ Uniform(-sqrt(6/(fan_in+fan_out)), sqrt(6/(fan_in+fan_out)))
nn.init.xavier_uniform_(layer.weight)
nn.init.xavier_normal_(layer.weight)
Kaiming (He) Initialization
Designed for ReLU activations. Accounts for the fact that ReLU zeros out half the inputs.
Formula: W ~ Normal(0, sqrt(2/fan_in))
nn.init.kaiming_uniform_(layer.weight, mode='fan_in', nonlinearity='relu')
nn.init.kaiming_normal_(layer.weight, mode='fan_in', nonlinearity='relu')
When to use each:
- Xavier: sigmoid, tanh activations
- Kaiming: ReLU, LeakyReLU activations
- Normal/Uniform: When you want simple random initialization
- Zeros: For biases (common default)
- Ones: For normalization layer weights
def init_weights(module):
if isinstance(module, nn.Linear):
nn.init.kaiming_normal_(module.weight, nonlinearity='relu')
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Conv2d):
nn.init.kaiming_normal_(module.weight, mode='fan_out', nonlinearity='relu')
model.apply(init_weights) # Recursively apply to all modules
Hooks
Hooks let you inspect or modify intermediate values during forward/backward passes without changing the model code.
Forward Hook
Called after forward() completes. Receives (module, input, output).
def print_output_shape(module, input, output):
print(f"{module.__class__.__name__}: output shape = {output.shape}")
hook_handle = model.layer1.register_forward_hook(print_output_shape)
# Later: hook_handle.remove()
Forward Pre-Hook
Called before forward(). Receives (module, input). Can modify the input.
def modify_input(module, args):
# args is a tuple of inputs
return (args[0] * 2,) # Double the input
handle = model.layer1.register_forward_pre_hook(modify_input)
Backward Hook
Called during backward pass. Can inspect or modify gradients.
def print_grad(module, grad_input, grad_output):
print(f"Grad output norm: {grad_output[0].norm()}")
handle = model.layer1.register_full_backward_hook(print_grad)
Common Use Cases:
- Feature extraction from intermediate layers
- Gradient visualization/debugging
- Gradient clipping per layer
- Activation statistics for debugging training
State Dict
The state dict is an OrderedDict mapping parameter/buffer names to tensors. It's the standard way to save and load models.
Saving and Loading
# Save
torch.save(model.state_dict(), 'model_weights.pth')
# Load
model = MyModel() # Create model with same architecture
model.load_state_dict(torch.load('model_weights.pth', weights_only=True))
Partial Loading (strict=False)
When architectures don't match exactly:
# Load only matching keys, ignore missing/unexpected keys
state_dict = torch.load('pretrained.pth', weights_only=True)
model.load_state_dict(state_dict, strict=False)
Inspecting State Dict
state_dict = model.state_dict()
for key, tensor in state_dict.items():
print(f"{key}: shape={tensor.shape}, dtype={tensor.dtype}")
Modifying Before Loading
# Remove prefix from keys (e.g., from DataParallel)
state_dict = torch.load('model.pth', weights_only=True)
new_state_dict = {}
for k, v in state_dict.items():
new_key = k.replace('module.', '') # Remove DataParallel prefix
new_state_dict[new_key] = v
model.load_state_dict(new_state_dict)
Summary
| Concept | Key Takeaway |
|---|---|
| nn.Module | Base class; register layers in __init__, compute in forward |
| Parameters | Learnable tensors (weights), updated by optimizer |
| Buffers | State tensors without gradients (running stats, masks) |
| Sequential | Simple chain of layers |
| Conv2d | Spatial feature extraction with kernel sliding |
| BatchNorm | Normalize across batch; different train/eval behavior |
| LayerNorm | Normalize across features; batch-independent |
| Dropout | Random zeroing during training only |
| LSTM | Recurrent with forget/input/output gates |
| Transformer | Self-attention + feedforward |
| CrossEntropyLoss | Multi-class classification standard |
| Kaiming init | Default for ReLU networks |
| Hooks | Inspect/modify without changing model code |
| state_dict | Standard save/load mechanism |
📓 Open Notebook — Interactive version of this module
Source Files
module_basics.py— Neural network basics — complete nn.Module tutorialcommon_layers.py— Common neural network layersloss_functions.py— Loss functions — complete guideweight_initialization.py— Weight initialization strategieshooks_and_state_dict.py— Hooks and state dict — advanced module features
Module 05: Optimizers and Learning Rate Schedulers
Table of Contents
- Optimizer Fundamentals
- SGD — Stochastic Gradient Descent
- Adam — Adaptive Moment Estimation
- AdamW — Decoupled Weight Decay
- Other Optimizers
- Learning Rate Schedulers
- Practical Advice
- Gradient Clipping
- Compiled Optimizers
Optimizer Fundamentals
An optimizer updates model parameters to minimize the loss function. In PyTorch, all optimizers inherit from torch.optim.Optimizer and share this interface:
import torch.optim as optim
# Create optimizer — pass it the parameters to optimize
optimizer = optim.Adam(model.parameters(), lr=0.001)
# Training loop
for batch in dataloader:
optimizer.zero_grad() # Clear old gradients
output = model(batch) # Forward pass
loss = criterion(output, target)
loss.backward() # Compute gradients
optimizer.step() # Update parameters
Parameter Groups
Optimizers support different settings for different parameter groups:
optimizer = optim.SGD([
{'params': model.backbone.parameters(), 'lr': 0.001}, # Lower LR for backbone
{'params': model.head.parameters(), 'lr': 0.01}, # Higher LR for head
], momentum=0.9, weight_decay=1e-4)
Optimizer State Dict
# Save optimizer state (for training resumption)
torch.save(optimizer.state_dict(), 'optimizer.pth')
# Load optimizer state
optimizer.load_state_dict(torch.load('optimizer.pth'))
The state dict contains:
state: per-parameter state (momentum buffers, Adam moments, step counts)param_groups: hyperparameters (lr, momentum, weight_decay, etc.)
SGD
Stochastic Gradient Descent is the simplest optimizer but with momentum is still competitive.
Vanilla SGD
Update rule: theta = theta - lr * gradient
SGD with Momentum
Momentum accumulates past gradients to smooth updates and escape shallow local minima.
Classical momentum:
v_t = momentum * v_{t-1} + gradient_t
theta_t = theta_{t-1} - lr * v_t
Nesterov momentum (look-ahead):
v_t = momentum * v_{t-1} + gradient(theta - lr * momentum * v_{t-1})
theta_t = theta_{t-1} - lr * v_t
Nesterov is generally better — it evaluates the gradient at the "look-ahead" position, which provides better correction.
SGD with Weight Decay (L2 Regularization)
gradient_with_wd = gradient + weight_decay * theta
theta = theta - lr * gradient_with_wd
Weight decay adds a penalty proportional to parameter magnitude, pushing weights toward zero.
optimizer = optim.SGD(
model.parameters(),
lr=0.1,
momentum=0.9,
weight_decay=1e-4,
nesterov=True
)
When to use SGD:
- Computer vision tasks (often gives better generalization than Adam)
- When you have a good learning rate schedule
- Large batch training
- When you want maximum control
Adam
Adam (Adaptive Moment Estimation) adapts the learning rate for each parameter based on first and second moments of the gradients.
Algorithm Step-by-Step
Initialize: m_0 = 0, v_0 = 0, t = 0
For each step:
t = t + 1
g_t = gradient at step t
# Update biased first moment estimate (mean of gradients)
m_t = beta1 * m_{t-1} + (1 - beta1) * g_t
# Update biased second moment estimate (mean of squared gradients)
v_t = beta2 * v_{t-1} + (1 - beta2) * g_t^2
# Bias correction (crucial in early steps)
m_hat_t = m_t / (1 - beta1^t)
v_hat_t = v_t / (1 - beta2^t)
# Update parameters
theta_t = theta_{t-1} - lr * m_hat_t / (sqrt(v_hat_t) + eps)
Why bias correction? Since m_0 = 0 and v_0 = 0, the estimates are biased toward zero in early training. Dividing by (1 - beta^t) corrects this — as t grows large, the correction approaches 1 and has no effect.
Default hyperparameters:
lr = 0.001beta1 = 0.9(momentum for first moment)beta2 = 0.999(momentum for second moment)eps = 1e-8(numerical stability)
optimizer = optim.Adam(
model.parameters(),
lr=0.001,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=0 # This is L2 regularization, NOT decoupled
)
When to use Adam:
- Default choice for most tasks
- NLP, transformers
- When you don't want to tune the learning rate carefully
- Faster convergence early in training
AdamW
AdamW fixes a subtle but important problem with Adam's weight decay implementation.
The Problem with Adam + Weight Decay
In standard Adam with weight_decay, the weight decay is applied to the gradient:
g_t = gradient + weight_decay * theta # L2 regularization added to gradient
But then Adam's adaptive learning rate scales this differently per parameter, effectively applying DIFFERENT regularization strengths to different parameters. This breaks the intended uniform regularization.
Decoupled Weight Decay (AdamW)
AdamW applies weight decay directly to the parameters AFTER the Adam update:
theta_t = theta_{t-1} - lr * adam_update - lr * weight_decay * theta_{t-1}
This means every parameter gets the same relative shrinkage regardless of its gradient history.
Impact: AdamW generalizes better than Adam+L2, especially for transformers and large models. It's now the default for most modern training.
optimizer = optim.AdamW(
model.parameters(),
lr=0.001,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=0.01 # Decoupled! Common value: 0.01 to 0.1
)
When to use AdamW:
- Transformers (GPT, BERT, ViT, etc.)
- When using weight decay (almost always)
- Default recommendation for most tasks today
Other Optimizers
Adagrad
Adapts learning rate based on accumulated squared gradients. Good for sparse data but learning rate decays to zero over time (problematic for long training).
optimizer = optim.Adagrad(model.parameters(), lr=0.01)
RMSprop
Fixes Adagrad's decaying learning rate by using exponential moving average of squared gradients. Predecessor to Adam (Adam = RMSprop + momentum).
optimizer = optim.RMSprop(model.parameters(), lr=0.01, alpha=0.99)
Adadelta
Similar to RMSprop but eliminates the need to set an initial learning rate.
optimizer = optim.Adadelta(model.parameters(), lr=1.0, rho=0.9)
LBFGS
Limited-memory BFGS — a quasi-Newton method. Uses second-order information (Hessian approximation). Much more expensive per step but converges in fewer steps. Requires a closure.
optimizer = optim.LBFGS(model.parameters(), lr=1.0, max_iter=20)
def closure():
optimizer.zero_grad()
output = model(input)
loss = criterion(output, target)
loss.backward()
return loss
optimizer.step(closure)
RAdam (Rectified Adam)
Adam with a variance-rectification term that provides an automatic warmup effect. Addresses the high variance of Adam in early training.
Muon
A newer optimizer designed for large language models. Uses momentum and sign-based updates for more stable training at scale.
Adafactor
Memory-efficient alternative to Adam. Factorizes the second moment matrix to reduce memory from O(mn) to O(m+n). Popular for training very large models.
Learning Rate Schedulers
Schedulers adjust the learning rate during training. The general pattern:
optimizer = optim.Adam(model.parameters(), lr=0.001)
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
for epoch in range(100):
train(...)
scheduler.step() # Update LR after each epoch
StepLR
Decay by gamma every step_size epochs.
# LR: 0.1 -> 0.01 (at epoch 30) -> 0.001 (at epoch 60)
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
MultiStepLR
Decay at specific milestones.
# LR: 0.1 -> 0.01 (at epoch 30) -> 0.001 (at epoch 80)
scheduler = optim.lr_scheduler.MultiStepLR(optimizer, milestones=[30, 80], gamma=0.1)
ExponentialLR
Multiply LR by gamma every epoch.
scheduler = optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.95)
CosineAnnealingLR
Smoothly decays LR following a cosine curve from initial LR to eta_min.
# Decays from lr to eta_min over T_max epochs
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-6)
OneCycleLR
Implements the 1-cycle policy: ramp up LR, then decay. Often achieves best results.
scheduler = optim.lr_scheduler.OneCycleLR(
optimizer, max_lr=0.01, total_steps=1000,
pct_start=0.3, # 30% warmup, 70% decay
anneal_strategy='cos'
)
# Note: step() after each BATCH, not each epoch!
ReduceLROnPlateau
Reduce LR when a metric plateaus (most practical for validation loss).
scheduler = optim.lr_scheduler.ReduceLROnPlateau(
optimizer, mode='min', factor=0.5, patience=10
)
# Must pass the metric value:
scheduler.step(val_loss)
LinearLR
Linearly scale LR from start_factor to end_factor over total_iters.
# Warmup: LR goes from 0.001*0.1 to 0.001 over 10 epochs
scheduler = optim.lr_scheduler.LinearLR(
optimizer, start_factor=0.1, end_factor=1.0, total_iters=10
)
SequentialLR (Warmup + Cosine)
Chain multiple schedulers together:
warmup = optim.lr_scheduler.LinearLR(optimizer, start_factor=0.1, total_iters=10)
cosine = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=90)
scheduler = optim.lr_scheduler.SequentialLR(
optimizer, schedulers=[warmup, cosine], milestones=[10]
)
CosineAnnealingWarmRestarts
Cosine annealing with periodic restarts (warm restarts increase exploration).
scheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(
optimizer, T_0=10, T_mult=2 # First restart at 10, then 20, 40, ...
)
Practical Advice
Which Optimizer for Which Task?
| Task | Recommended | LR Range | Notes |
|---|---|---|---|
| Vision (CNNs) | SGD+momentum or AdamW | 0.01-0.1 | SGD often generalizes better |
| NLP/Transformers | AdamW | 1e-5 to 5e-4 | With cosine schedule + warmup |
| Fine-tuning | AdamW | 1e-5 to 3e-5 | Lower LR for pretrained weights |
| GANs | Adam (beta1=0.0) | 1e-4 to 2e-4 | Two separate optimizers |
| RL | Adam | 3e-4 | Simpler schedules |
| Small datasets | SGD+momentum | 0.01 | Better generalization |
Learning Rate Selection
- Learning Rate Finder: Start very small, increase exponentially, plot loss vs LR.
Choose LR where loss is decreasing steepest (typically 1/10 of the minimum).
- Rule of thumb: If training is unstable, reduce LR by 3-10x.
- Linear scaling rule: When increasing batch size by N, multiply LR by N (with warmup).
Warmup Strategies
Warmup is crucial for:
- Large learning rates
- Large batch sizes
- Transformer training
Common approach: Linear warmup for 5-10% of total training, then cosine decay.
Gradient Clipping
Prevents exploding gradients by limiting gradient magnitude.
clip_grad_norm_ (recommended)
Clips the total norm of all gradients:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
This scales all gradients uniformly to ensure the total L2 norm doesn't exceed max_norm. It preserves gradient direction.
clip_grad_value_
Clips each gradient element independently:
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)
Each element is clamped to [-clip_value, clip_value]. Can change gradient direction.
Usage in training loop:
for batch in dataloader:
optimizer.zero_grad()
loss = model(batch).sum()
loss.backward()
# Clip AFTER backward(), BEFORE step()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
Compiled Optimizers
PyTorch 2.0+ can compile optimizers with torch.compile for significant speedups:
optimizer = optim.AdamW(model.parameters(), lr=0.001)
# The optimizer step is fused and optimized
# This happens automatically when the model is compiled
@torch.compile
def train_step(model, x, y):
output = model(x)
loss = F.cross_entropy(output, y)
loss.backward()
optimizer.step()
optimizer.zero_grad()
return loss
Benefits:
- Fused kernels: multiple optimizer ops become one kernel launch
- Reduced memory traffic
- Horizontal fusion across parameter groups
- Can give 10-20% speedup on optimizer step
Summary
| Optimizer | Key Feature | Best For |
|---|---|---|
| SGD+momentum | Simple, good generalization | Vision, large-scale |
| Adam | Adaptive per-param LR | General, fast convergence |
| AdamW | Proper weight decay | Transformers, modern default |
| Adagrad | Adapts to sparse features | NLP with sparse embeddings |
| RMSprop | Fixes Adagrad decay | RNNs (historically) |
| LBFGS | Second-order | Small problems, fine-tuning |
| Scheduler | Pattern | Best For |
|---|---|---|
| CosineAnnealing | Smooth decay to 0 | Most tasks |
| OneCycleLR | Warmup + decay | Fastest convergence |
| ReduceLROnPlateau | Adaptive decay | When you have val metric |
| Sequential(Linear+Cosine) | Warmup + cosine | Transformers |
| CosineWarmRestarts | Periodic resets | Long training, exploration |
📓 Open Notebook — Interactive version of this module
Source Files
optimizer_basics.py— Optimizer basicslr_schedulers.py— Learning rate schedulersoptimizer_comparison.py— Optimizer comparison
Module 06: Data Loading in PyTorch
Table of Contents
- Dataset Class
- IterableDataset
- Built-in Datasets
- DataLoader
- Worker Processes
- Custom Collate Functions
- Samplers
- Data Augmentation
- Memory-Mapped Datasets
- Best Practices
Dataset Class
The torch.utils.data.Dataset is the base class for all map-style datasets. You implement two methods:
__len__(): Returns the total number of samples__getitem__(index): Returns one sample given an index
from torch.utils.data import Dataset
class MyDataset(Dataset):
def __init__(self, data, labels):
self.data = data
self.labels = labels
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx], self.labels[idx]
Key properties:
- Supports random access by index
- Length is known ahead of time
- DataLoader can shuffle efficiently
- Can be split into train/val subsets
When data doesn't fit in memory:
class LargeFileDataset(Dataset):
def __init__(self, file_paths):
self.file_paths = file_paths
def __len__(self):
return len(self.file_paths)
def __getitem__(self, idx):
# Load only one file at a time
data = load_file(self.file_paths[idx])
return process(data)
IterableDataset
For streaming data or when random access is not possible/efficient:
from torch.utils.data import IterableDataset
class StreamDataset(IterableDataset):
def __init__(self, url):
self.url = url
def __iter__(self):
# Yield samples one at a time
for line in open_stream(self.url):
yield process(line)
Use cases:
- Reading from network streams
- Very large files that can only be read sequentially
- Data generated on-the-fly
- Databases with cursor-based access
Key differences from map-style Dataset:
- No
__len__()— size may be unknown - No
__getitem__()— no random access - DataLoader cannot shuffle (must shuffle upstream)
- Multi-worker requires careful work splitting
Multi-worker IterableDataset:
class ShardedDataset(IterableDataset):
def __init__(self, file_list):
self.file_list = file_list
def __iter__(self):
worker_info = torch.utils.data.get_worker_info()
if worker_info is None:
# Single-process loading
files = self.file_list
else:
# Split files among workers
per_worker = len(self.file_list) // worker_info.num_workers
start = worker_info.id * per_worker
end = start + per_worker
files = self.file_list[start:end]
for f in files:
for sample in read_file(f):
yield sample
Built-in Datasets
TensorDataset
Wraps tensors — each sample is a tuple indexed along the first dimension:
from torch.utils.data import TensorDataset
X = torch.randn(1000, 20)
y = torch.randint(0, 5, (1000,))
dataset = TensorDataset(X, y)
# dataset[0] returns (X[0], y[0])
ConcatDataset
Concatenates multiple datasets end-to-end:
from torch.utils.data import ConcatDataset
combined = ConcatDataset([dataset_train, dataset_extra])
# len(combined) == len(dataset_train) + len(dataset_extra)
Subset
Selects a subset of a dataset by indices:
from torch.utils.data import Subset
# Manual train/val split
indices = torch.randperm(len(dataset))
train_set = Subset(dataset, indices[:800])
val_set = Subset(dataset, indices[800:])
random_split
Convenience function for splitting:
from torch.utils.data import random_split
train_set, val_set, test_set = random_split(
dataset, [0.7, 0.15, 0.15], # Fractions
generator=torch.Generator().manual_seed(42)
)
DataLoader
The DataLoader is the workhorse that turns a Dataset into an iterable of batches:
from torch.utils.data import DataLoader
loader = DataLoader(
dataset,
batch_size=32, # Samples per batch
shuffle=True, # Randomize order each epoch
num_workers=4, # Parallel data loading processes
pin_memory=True, # Speed up CPU->GPU transfer
drop_last=False, # Drop incomplete final batch?
persistent_workers=True,# Keep workers alive between epochs
prefetch_factor=2, # Batches to prefetch per worker
)
for batch_X, batch_y in loader:
# batch_X: (32, ...), batch_y: (32, ...)
output = model(batch_X)
loss = criterion(output, batch_y)
...
Key Parameters:
batch_size: Number of samples per batch. Larger = faster training but more memory. Typical values: 16, 32, 64, 128, 256.
shuffle: Randomize sample order each epoch. Always True for training, False for validation/test (for reproducible evaluation).
num_workers: Number of subprocess workers for data loading. Set to 0 for debugging (main process only). Typical: 2-8 depending on CPU cores and data complexity.
pin_memory: Pre-allocates batch tensors in pinned (page-locked) memory. Makes CPU->GPU transfers faster. Always True when training on GPU.
drop_last: If True, drops the last batch if it's smaller than batch_size. Important for BatchNorm (needs consistent batch size) and distributed training.
persistent_workers: Keeps worker processes alive across epochs instead of respawning. Saves the cost of worker initialization. Use with num_workers > 0.
prefetch_factor: Number of batches each worker pre-loads. Higher values use more memory but can hide I/O latency better. Default is 2.
Worker Processes
How Workers Work
When num_workers > 0, DataLoader spawns separate processes (not threads!) that:
- Receive batch indices from the main process
- Call
dataset.__getitem__()for each index - Apply the collate function to form a batch
- Send the batch back to the main process via shared memory
Why num_workers > 0 is Faster
- Data loading (disk I/O, decompression, augmentation) happens in parallel
- While GPU trains on batch N, workers prepare batch N+1, N+2, ...
- Python's GIL doesn't affect separate processes
Common Pitfalls
- Too many workers: More workers != always faster. Each worker duplicates the
dataset object in memory. Start with num_workers=4 and tune.
- Fork vs Spawn: On Linux, workers are forked (fast, shares memory). On macOS/Windows,
workers are spawned (slow startup, separate memory). Set: ``python torch.multiprocessing.set_start_method('spawn') # If needed ``
- Random state in workers: Each worker gets the same random seed by default.
Use worker_init_fn to set different seeds: ```python def worker_init_fn(worker_id): seed = torch.initial_seed() % 2**32 numpy.random.seed(seed + worker_id)
loader = DataLoader(..., worker_init_fn=worker_init_fn) ```
- File handles: If your dataset opens files, each worker opens its own copy.
With many workers, you can run out of file descriptors.
- Shared memory limits: Workers send data via shared memory. Very large batches
can exhaust /dev/shm. Increase shared memory or reduce batch size.
Custom Collate Functions
The collate function converts a list of samples into a batch. The default collator stacks tensors, but custom collation handles variable-length sequences:
def custom_collate(batch):
"""Pad variable-length sequences to same length."""
sequences, labels = zip(*batch)
# Pad sequences to max length in this batch
lengths = [len(s) for s in sequences]
max_len = max(lengths)
padded = torch.zeros(len(sequences), max_len)
for i, (seq, length) in enumerate(zip(sequences, lengths)):
padded[i, :length] = seq
labels = torch.tensor(labels)
lengths = torch.tensor(lengths)
return padded, labels, lengths
loader = DataLoader(dataset, batch_size=32, collate_fn=custom_collate)
Common custom collate patterns:
- Padding variable-length sequences
- Creating attention masks
- Handling nested data structures (dicts, lists)
- Filtering out None values (corrupt samples)
Samplers
Samplers control the order in which indices are provided to the Dataset.
SequentialSampler
Iterates indices 0, 1, 2, ..., N-1. Used when shuffle=False.
RandomSampler
Random permutation of indices. Used when shuffle=True.
from torch.utils.data import RandomSampler
# With replacement (for bootstrap):
sampler = RandomSampler(dataset, replacement=True, num_samples=10000)
WeightedRandomSampler
For handling class imbalance — oversamples rare classes:
from torch.utils.data import WeightedRandomSampler
# Assign weight to each sample (higher weight = sampled more often)
class_counts = [1000, 100, 50] # Imbalanced classes
class_weights = 1.0 / torch.tensor(class_counts, dtype=torch.float)
sample_weights = class_weights[labels] # Weight per sample
sampler = WeightedRandomSampler(
weights=sample_weights,
num_samples=len(dataset),
replacement=True
)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
# Note: cannot use shuffle=True with a sampler
SubsetRandomSampler
Random sampling from a fixed set of indices (for train/val splits):
from torch.utils.data import SubsetRandomSampler
indices = list(range(len(dataset)))
train_indices = indices[:800]
val_indices = indices[800:]
train_loader = DataLoader(dataset, batch_size=32,
sampler=SubsetRandomSampler(train_indices))
val_loader = DataLoader(dataset, batch_size=32,
sampler=SubsetRandomSampler(val_indices))
BatchSampler
Groups indices into batches (useful for custom batching logic):
from torch.utils.data import BatchSampler, SequentialSampler
# Create batches of similar-length sequences (for efficient padding)
sampler = BatchSampler(
SequentialSampler(dataset),
batch_size=32,
drop_last=False
)
DistributedSampler
For multi-GPU training — ensures each GPU sees different data:
from torch.utils.data import DistributedSampler
sampler = DistributedSampler(
dataset,
num_replicas=world_size, # Number of GPUs
rank=rank, # This GPU's index
shuffle=True
)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
# IMPORTANT: Must call set_epoch each epoch for proper shuffling
for epoch in range(num_epochs):
sampler.set_epoch(epoch)
for batch in loader:
...
Data Augmentation
Data augmentation applies random transformations during training to increase data diversity and reduce overfitting.
torchvision.transforms.v2
The modern transforms API for images:
from torchvision.transforms import v2
train_transform = v2.Compose([
v2.RandomResizedCrop(224),
v2.RandomHorizontalFlip(p=0.5),
v2.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
v2.RandomRotation(15),
v2.ToImage(),
v2.ToDtype(torch.float32, scale=True),
v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
val_transform = v2.Compose([
v2.Resize(256),
v2.CenterCrop(224),
v2.ToImage(),
v2.ToDtype(torch.float32, scale=True),
v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
MixUp
Linearly interpolates between pairs of samples:
mixed_input = lambda * input_i + (1-lambda) * input_j
mixed_target = lambda * target_i + (1-lambda) * target_j
CutMix
Cuts a patch from one image and pastes it onto another:
mixed_input = mask * input_i + (1-mask) * input_j
mixed_target = lambda * target_i + (1-lambda) * target_j
Where lambda is the area ratio.
Custom Augmentations
class CustomAugmentation:
def __init__(self, p=0.5):
self.p = p
def __call__(self, image):
if torch.rand(1) < self.p:
return apply_augmentation(image)
return image
Memory-Mapped Datasets
For datasets too large to fit in RAM, memory-mapping lets the OS page data in/out:
import numpy as np
class MemmapDataset(Dataset):
def __init__(self, data_path, shape, dtype='float32'):
# Memory-map the file (doesn't load into RAM)
self.data = np.memmap(data_path, dtype=dtype, mode='r', shape=shape)
def __len__(self):
return self.data.shape[0]
def __getitem__(self, idx):
# Only this sample is loaded into RAM
return torch.from_numpy(self.data[idx].copy())
Benefits:
- Works with datasets much larger than RAM
- OS handles caching transparently
- Sequential access patterns are very efficient
- Multiple workers can share the same memory map
Best Practices
Performance Checklist
- pin_memory=True when using GPU training
- persistent_workers=True with num_workers > 0 to avoid respawn cost
- prefetch_factor=2 (default) is usually good; increase if I/O-bound
- num_workers: Start with 4, increase until CPU is saturated or you run out of RAM
- drop_last=True for training with BatchNorm
Data Loading Bottleneck Detection
If your GPU utilization is low, data loading might be the bottleneck:
import time
for batch in loader:
start = time.time()
output = model(batch)
loss.backward()
optimizer.step()
gpu_time = time.time() - start
# If data loading time >> gpu_time, you're data-bound
Memory Efficiency
- Use
uint8for images until the augmentation step - Convert to float32 only in the transform/collate
- Use memory-mapped files for very large datasets
- Consider chunked reading for CSV/parquet files
Reproducibility
# Seed everything for reproducible data loading
def seed_worker(worker_id):
worker_seed = torch.initial_seed() % 2**32
import numpy as np
np.random.seed(worker_seed)
import random
random.seed(worker_seed)
g = torch.Generator()
g.manual_seed(42)
loader = DataLoader(
dataset,
batch_size=32,
shuffle=True,
num_workers=4,
worker_init_fn=seed_worker,
generator=g,
)
Common Patterns
Training with validation:
train_loader = DataLoader(train_set, batch_size=64, shuffle=True,
num_workers=4, pin_memory=True)
val_loader = DataLoader(val_set, batch_size=128, shuffle=False,
num_workers=4, pin_memory=True)
Infinite data loader (for step-based training):
def infinite_loader(loader):
while True:
for batch in loader:
yield batch
for step, batch in enumerate(infinite_loader(train_loader)):
if step >= max_steps:
break
train_step(batch)
Progress bar with tqdm:
from tqdm import tqdm
for batch in tqdm(loader, desc="Training"):
...
📓 Open Notebook — Interactive version of this module
Source Files
dataset_basics.py— Dataset basicscustom_datasets.py— Custom datasetsdataloader_advanced.py— Advanced DataLoader featuresaugmentation.py— Data augmentation patterns
Module 07: The Complete Training Guide
Overview
Training a neural network is where theory meets practice. This module covers everything from the basic training loop to advanced techniques used by state-of-the-art models. By the end, you'll understand not just how to train models, but why each step exists and how to diagnose problems.
1. Basic Training Loop Anatomy
Every training loop in PyTorch follows the same fundamental pattern:
forward pass → compute loss → backward pass → optimizer step
Let's break each step down:
Step 1: Forward Pass
predictions = model(inputs)
Data flows through the model's layers. Each layer applies its transformation (matrix multiply, activation, normalization, etc.) and PyTorch records the operations in a computational graph for later backpropagation.
Step 2: Compute Loss
loss = loss_fn(predictions, targets)
The loss function measures how far the model's predictions are from the true targets. Common losses:
nn.CrossEntropyLoss()— classification (combines LogSoftmax + NLLLoss)nn.MSELoss()— regression (mean squared error)nn.BCEWithLogitsLoss()— binary classification (numerically stable)
Step 3: Backward Pass
loss.backward()
PyTorch walks backward through the computational graph, computing the gradient of the loss with respect to every parameter that has requires_grad=True. These gradients accumulate in each parameter's .grad attribute.
Step 4: Optimizer Step
optimizer.step()
The optimizer uses the computed gradients to update the model's parameters. Different optimizers (SGD, Adam, AdamW) use different update rules, but all read from .grad and modify .data.
The Missing Step: zero_grad()
optimizer.zero_grad()
This clears old gradients before computing new ones. Without it, gradients accumulate across iterations (which is sometimes intentional — see gradient accumulation below).
Complete Minimal Loop
model.train()
for epoch in range(num_epochs):
for batch_inputs, batch_targets in dataloader:
optimizer.zero_grad() # Clear old gradients
predictions = model(batch_inputs) # Forward pass
loss = loss_fn(predictions, batch_targets) # Compute loss
loss.backward() # Compute gradients
optimizer.step() # Update parameters
2. train() vs eval() Mode
What model.train() Does
Sets the model to training mode. This affects layers that behave differently during training vs inference:
- Dropout: Randomly zeros elements during training. During eval, all
elements pass through (scaled appropriately).
- BatchNorm: During training, uses batch statistics (mean/var of current
batch) and updates running statistics. During eval, uses the accumulated running statistics.
What model.eval() Does
model.eval()
Sets the model to evaluation mode. Dropout is disabled, BatchNorm uses running statistics instead of batch statistics.
Common Mistake
# WRONG: Forgetting to switch modes
def evaluate(model, test_loader):
# model is still in train() mode!
# Dropout is randomly zeroing activations
# BatchNorm is using (and updating!) batch statistics
total_correct = 0
for inputs, targets in test_loader:
outputs = model(inputs)
...
# CORRECT:
def evaluate(model, test_loader):
model.eval()
with torch.no_grad(): # Also disable gradient computation
total_correct = 0
for inputs, targets in test_loader:
outputs = model(inputs)
...
model.train() # Switch back after evaluation
torch.no_grad() vs model.eval()
These are different things:
model.eval()— changes layer behavior (dropout, batchnorm)torch.no_grad()— disables gradient computation (saves memory/compute)
For evaluation, you typically want BOTH.
3. zero_grad() — Why and How
Why Gradients Accumulate
PyTorch accumulates gradients by default. After loss.backward(), the .grad attribute of each parameter ADDS to whatever was already there:
param.grad += new_gradient # This is what happens internally
This design choice enables gradient accumulation (discussed later), but it means you must manually clear gradients each iteration.
set_to_none=True Optimization
optimizer.zero_grad(set_to_none=True)
Instead of setting gradients to zero tensors, this sets them to None. Benefits:
- Slightly less memory (no zero tensor allocated)
- Can be marginally faster
- The gradient will be lazily created on the next backward pass
This is now the default in modern PyTorch (>= 2.0). The only reason to use set_to_none=False is if your code explicitly checks if param.grad is not None.
4. Mixed Precision Training (AMP)
What It Is
Mixed precision training uses lower-precision floating point numbers (float16 or bfloat16) for most operations, while keeping critical operations in float32. This is faster because:
- Lower precision operations use less memory bandwidth
- Hardware (GPUs, modern CPUs) has specialized units for half-precision math
- Smaller tensors mean more data fits in cache
float16 vs float32
| Property | float32 | float16 | bfloat16 |
|---|---|---|---|
| Total bits | 32 | 16 | 16 |
| Exponent bits | 8 | 5 | 8 |
| Mantissa bits | 23 | 10 | 7 |
| Max value | ~3.4 × 10³⁸ | 65504 | ~3.4 × 10³⁸ |
| Min positive | ~1.2 × 10⁻³⁸ | ~6.0 × 10⁻⁸ | ~1.2 × 10⁻³⁸ |
The torch.amp API (Modern, Device-Agnostic)
# The modern way (PyTorch 2.0+)
with torch.amp.autocast(device_type='cpu', dtype=torch.bfloat16):
output = model(input)
loss = loss_fn(output, target)
The autocast context manager automatically casts operations to the specified lower precision where safe, and keeps float32 where needed (e.g., loss computation, softmax, layer norm).
GradScaler (for float16 only)
float16 has a limited range. Small gradients can underflow to zero. GradScaler solves this by:
- Scaling the loss UP before backward (so gradients are larger)
- Unscaling gradients before optimizer step
- Skipping steps where gradients contain inf/nan (and reducing the scale)
scaler = torch.amp.GradScaler()
for inputs, targets in dataloader:
optimizer.zero_grad()
with torch.amp.autocast(device_type='cuda', dtype=torch.float16):
output = model(inputs)
loss = loss_fn(output, targets)
scaler.scale(loss).backward() # Scaled backward
scaler.step(optimizer) # Unscale + step (or skip)
scaler.update() # Adjust scale factor
BFloat16 Doesn't Need GradScaler
bfloat16 has the same exponent range as float32, so gradients don't underflow. You can use it without GradScaler:
with torch.amp.autocast(device_type='cpu', dtype=torch.bfloat16):
output = model(input)
loss = loss_fn(output, target)
loss.backward() # No scaler needed
optimizer.step()
5. BFloat16 vs Float16 — When to Use Each
Float16 (FP16)
- Pros: Widely supported, maximum speed on older GPUs (V100)
- Cons: Limited range (max 65504), requires GradScaler, can overflow/underflow
- Use when: Training on older NVIDIA GPUs, inference where range is known
BFloat16 (BF16)
- Pros: Same range as float32, no GradScaler needed, more stable training
- Cons: Less precision (7 mantissa bits vs 10), requires Ampere+ GPU or modern CPU
- Use when: Training large models (LLMs), when numerical stability matters
Practical Recommendation
- For training: prefer bfloat16 if your hardware supports it
- For inference: either works, float16 slightly more precise per value
- For CPU: bfloat16 is supported on modern x86 (AMX) and ARM
6. Gradient Accumulation
The Problem
You want a batch size of 256 but your memory only fits 32 samples.
The Solution
Accumulate gradients over multiple mini-batches before stepping:
accumulation_steps = 8 # Effective batch = 32 * 8 = 256
for i, (inputs, targets) in enumerate(dataloader):
# Forward + backward (gradients accumulate)
output = model(inputs)
loss = loss_fn(output, targets)
loss = loss / accumulation_steps # Scale loss!
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
Why Divide the Loss?
Without division, accumulated gradients are accumulation_steps times larger than a true large-batch gradient. Dividing the loss by accumulation_steps makes the accumulated gradient equivalent to computing it over one large batch.
Mathematically: grad(L/N) summed N times = grad(sum(L_i)/N) which equals the gradient of the mean loss over all N mini-batches.
7. Gradient Checkpointing (Activation Checkpointing)
The Problem
During the forward pass, PyTorch saves all intermediate activations for use in the backward pass. For deep models, this uses enormous amounts of memory.
The Solution
Don't save activations — recompute them during the backward pass. This trades compute time (~33% more) for memory (~60-80% savings).
Usage
from torch.utils.checkpoint import checkpoint
class DeepModel(nn.Module):
def __init__(self):
super().__init__()
self.block1 = HeavyBlock()
self.block2 = HeavyBlock()
self.block3 = HeavyBlock()
def forward(self, x):
x = checkpoint(self.block1, x, use_reentrant=False)
x = checkpoint(self.block2, x, use_reentrant=False)
x = checkpoint(self.block3, x, use_reentrant=False)
return x
use_reentrant Parameter
use_reentrant=False(recommended): Uses a newer, more robust implementation.
Supports all autograd features correctly.
use_reentrant=True(legacy): The old implementation. Has subtle bugs with
certain autograd features. Being phased out.
Always use use_reentrant=False for new code.
8. Gradient Clipping
The Problem
Exploding gradients: when gradients become extremely large, the optimizer takes huge steps that destabilize training. Common in RNNs and deep networks.
clip_grad_norm_
Scales all gradients so their combined L2 norm doesn't exceed a threshold:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
This preserves the direction of gradients but limits their magnitude. The most common approach.
clip_grad_value_
Clamps each gradient element independently to [-value, +value]:
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)
More aggressive — changes gradient direction. Rarely used in practice.
Where to Place Clipping
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step() # Step with clipped gradients
Always clip AFTER backward, BEFORE step.
9. Transfer Learning
The Concept
Take a model pre-trained on a large dataset (e.g., ImageNet with 1M images) and adapt it to your smaller dataset. The pre-trained model already knows useful features (edges, textures, shapes).
Strategy 1: Feature Extraction (Freeze Everything)
model = torchvision.models.resnet18(weights='IMAGENET1K_V1')
# Freeze all parameters
for param in model.parameters():
param.requires_grad = False
# Replace the final classifier
model.fc = nn.Linear(512, num_classes)
# Only model.fc parameters will be trained
Strategy 2: Fine-tuning with Different Learning Rates
# Lower LR for pretrained backbone, higher LR for new head
optimizer = torch.optim.Adam([
{'params': model.features.parameters(), 'lr': 1e-5},
{'params': model.classifier.parameters(), 'lr': 1e-3},
])
Strategy 3: Progressive Unfreezing
Start with everything frozen except the head. Gradually unfreeze layers from top to bottom:
# Epoch 1-5: Only train the head
# Epoch 5-10: Unfreeze last block + head
# Epoch 10+: Unfreeze everything with small LR
This prevents the pre-trained features from being destroyed by large early gradients.
10. Fine-Tuning Strategies
Full Fine-Tuning
Unfreeze everything, train with a small learning rate:
for param in model.parameters():
param.requires_grad = True
optimizer = Adam(model.parameters(), lr=1e-5)
Best when you have enough data and the domain differs from pre-training.
Linear Probing
Only train a linear layer on top of frozen features:
for param in model.parameters():
param.requires_grad = False
probe = nn.Linear(feature_dim, num_classes)
Good baseline to check how useful the features are.
Progressive Unfreezing
Unfreeze one layer group at a time, starting from the output end:
# Phase 1: Just the head
# Phase 2: Head + last block
# Phase 3: Head + last 2 blocks
# Phase N: Everything
Which to Choose?
| Strategy | Data Amount | Domain Similarity | Risk |
|---|---|---|---|
| Linear probe | Very small | Any | Low |
| Freeze + new head | Small | Similar | Low |
| Progressive unfreeze | Medium | Different | Medium |
| Full fine-tune | Large | Different | Higher |
11. Knowledge Distillation
The Concept
Train a small "student" model to mimic a large "teacher" model. The student learns from the teacher's soft predictions (probability distributions) which contain more information than hard labels.
Why Soft Targets Help
Hard label: [0, 0, 1, 0] — "this is a cat" Soft prediction: [0.01, 0.05, 0.85, 0.09] — "this is mostly cat, slightly dog"
The soft predictions encode relationships between classes that hard labels miss.
Temperature Scaling
Higher temperature makes the distribution softer (more informative):
soft_teacher = F.softmax(teacher_logits / temperature, dim=-1)
soft_student = F.log_softmax(student_logits / temperature, dim=-1)
distill_loss = F.kl_div(soft_student, soft_teacher, reduction='batchmean')
distill_loss = distill_loss * (temperature ** 2) # Scale back
The temperature ** 2 factor compensates for the reduced gradient magnitude at higher temperatures.
Combined Loss
total_loss = alpha * distill_loss + (1 - alpha) * hard_loss
Typically alpha = 0.5-0.9 (emphasize soft targets).
12. EMA (Exponential Moving Average)
What It Is
Maintain a running average of model parameters that smooths out training noise:
ema_param = decay * ema_param + (1 - decay) * current_param
Typical decay: 0.999 or 0.9999 (very slow moving average).
Why It Helps
- Reduces variance in the final model
- Often achieves better generalization than the final checkpoint
- Used in many SOTA models (diffusion models, GANs, etc.)
Implementation
class EMA:
def __init__(self, model, decay=0.999):
self.decay = decay
self.shadow = {name: p.clone().detach()
for name, p in model.named_parameters()}
@torch.no_grad()
def update(self, model):
for name, param in model.named_parameters():
self.shadow[name].mul_(self.decay).add_(
param.data, alpha=1 - self.decay
)
def apply(self, model):
for name, param in model.named_parameters():
param.data.copy_(self.shadow[name])
13. SWA (Stochastic Weight Averaging)
The Concept
Average model weights from multiple points in training (typically from a cyclical or high constant LR schedule). This tends to find flatter minima which generalize better.
PyTorch Built-in Support
from torch.optim.swa_utils import AveragedModel, SWALR
swa_model = AveragedModel(model)
swa_scheduler = SWALR(optimizer, swa_lr=0.05)
for epoch in range(swa_start, total_epochs):
train_one_epoch(model)
swa_model.update_parameters(model)
swa_scheduler.step()
# Update batch normalization statistics
torch.optim.swa_utils.update_bn(train_loader, swa_model)
SWA vs EMA
- EMA: Continuous exponential average, gives more weight to recent params
- SWA: Equal-weight average of checkpoints, typically from later training
14. Label Smoothing
What It Does
Instead of training against hard targets [0, 0, 1, 0], use soft targets [0.033, 0.033, 0.9, 0.033]. This prevents the model from becoming overconfident.
Formula
smooth_target = (1 - smoothing) * one_hot + smoothing / num_classes
With smoothing=0.1 and 4 classes:
- True class: 0.9 + 0.1/4 = 0.925
- Other classes: 0.1/4 = 0.025
PyTorch Implementation
# Built-in support in CrossEntropyLoss
loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)
When to Use
- Large models prone to overconfidence
- When calibrated probabilities matter (not just accuracy)
- Generally helps with 0.05-0.1 smoothing; higher can hurt
15. Early Stopping
The Pattern
Stop training when validation loss stops improving to prevent overfitting:
class EarlyStopping:
def __init__(self, patience=10, min_delta=0.001):
self.patience = patience
self.min_delta = min_delta
self.counter = 0
self.best_loss = float('inf')
def __call__(self, val_loss):
if val_loss < self.best_loss - self.min_delta:
self.best_loss = val_loss
self.counter = 0
return False # Continue training
self.counter += 1
return self.counter >= self.patience # Stop if patience exceeded
Best Practice
Always save the model at the best validation loss, not the final epoch:
if val_loss < best_val_loss:
best_val_loss = val_loss
torch.save(model.state_dict(), 'best_model.pt')
16. Logging and Monitoring
What to Track
- Training loss (per batch and per epoch average)
- Validation loss (per epoch)
- Learning rate (especially with schedulers)
- Gradient norm (detect exploding/vanishing gradients)
- Parameter statistics (weight magnitudes per layer)
Simple Logging Pattern
for epoch in range(num_epochs):
running_loss = 0.0
for i, (inputs, targets) in enumerate(train_loader):
loss = train_step(inputs, targets)
running_loss += loss.item()
if (i + 1) % log_every == 0:
avg_loss = running_loss / log_every
print(f"Epoch {epoch}, Step {i+1}, Loss: {avg_loss:.4f}")
running_loss = 0.0
val_loss = evaluate(model, val_loader)
print(f"Epoch {epoch}, Val Loss: {val_loss:.4f}")
When to Save Checkpoints
- After each epoch (for resume capability)
- When validation metric improves (for best model)
- At fixed intervals for long training runs
17. Reproducibility
The Full Reproducibility Recipe
import torch
import numpy as np
import random
def set_seed(seed=42):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
set_seed(42)
# For fully deterministic operations
torch.use_deterministic_algorithms(True)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
Caveats
torch.use_deterministic_algorithms(True)may raise errors for ops without
deterministic implementations
- Setting
benchmark = Falsecan slow down training - DataLoader with
num_workers > 0needs worker seeding:
def seed_worker(worker_id):
worker_seed = torch.initial_seed() % 2**32
np.random.seed(worker_seed)
random.seed(worker_seed)
dataloader = DataLoader(
dataset,
worker_init_fn=seed_worker,
generator=torch.Generator().manual_seed(42),
)
18. Common Training Debugging
Loss Not Decreasing
- Learning rate too high: Loss oscillates wildly. Try 10x smaller.
- Learning rate too low: Loss decreases extremely slowly. Try 10x larger.
- Bug in data pipeline: Verify labels match inputs.
- Wrong loss function: Ensure loss matches the task (e.g., CrossEntropy for
classification needs raw logits, not softmax outputs).
- Model too small: May lack capacity to fit even training data.
NaN Loss
- Learning rate too high: Gradients explode. Lower LR or add clipping.
- Division by zero: Check for zero denominators in custom loss.
- Log of zero/negative: Ensure inputs to
log()are positive. - Overflow in float16: Use GradScaler or switch to bfloat16.
Overfitting (Train Loss Low, Val Loss High)
- Add regularization (dropout, weight decay)
- Reduce model size
- Add data augmentation
- Use early stopping
- Get more data
Underfitting (Both Losses High)
- Increase model capacity (more layers, wider layers)
- Train longer
- Reduce regularization
- Check for bugs in the model architecture
- Verify the task is learnable with this architecture
Gradient Debugging
# Check for vanishing/exploding gradients
for name, param in model.named_parameters():
if param.grad is not None:
grad_norm = param.grad.norm()
if grad_norm == 0:
print(f"WARNING: Zero gradient in {name}")
elif grad_norm > 100:
print(f"WARNING: Large gradient in {name}: {grad_norm}")
19. Putting It All Together
A production-ready training loop combines many of these techniques:
def train(model, train_loader, val_loader, config):
optimizer = AdamW(model.parameters(), lr=config.lr, weight_decay=0.01)
scheduler = CosineAnnealingLR(optimizer, T_max=config.epochs)
early_stop = EarlyStopping(patience=10)
ema = EMA(model, decay=0.999)
for epoch in range(config.epochs):
model.train()
for i, (inputs, targets) in enumerate(train_loader):
with torch.amp.autocast('cpu', dtype=torch.bfloat16):
output = model(inputs)
loss = loss_fn(output, targets)
loss = loss / config.accumulation_steps
loss.backward()
if (i + 1) % config.accumulation_steps == 0:
clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
optimizer.zero_grad()
ema.update(model)
scheduler.step()
val_loss = evaluate(model, val_loader)
if early_stop(val_loss):
break
Summary
| Technique | Purpose | Memory Impact | Speed Impact |
|---|---|---|---|
| Mixed precision | Faster math, less memory | Reduces 50% | 2-3x faster |
| Gradient accumulation | Simulate large batches | No change | Slight slow |
| Gradient checkpointing | Reduce activation memory | Saves 60-80% | ~33% slower |
| Gradient clipping | Prevent exploding gradients | None | Negligible |
| EMA | Smoother final model | 2x params | Negligible |
| Label smoothing | Prevent overconfidence | None | None |
| Early stopping | Prevent overfitting | None | Saves time |
📓 Open Notebook — Interactive version of this module
Source Files
basic_training_loop.py— Basic training loop — complete annotated examplemixed_precision.py— Mixed precision training — AMP with float16 and bfloat16gradient_techniques.py— Gradient techniques — accumulation, checkpointing, and clippingtransfer_learning.py— Transfer learning — freeze/unfreeze and differential learning ratesregularization.py— Regularization techniques — EMA, label smoothing, weight decay
Module 08: torch.compile — The Complete Guide
Overview
torch.compile is PyTorch's JIT compiler that makes your models run faster without changing your code. Introduced in PyTorch 2.0, it captures your Python model into an optimized graph and generates efficient machine code.
1. What is torch.compile?
The Problem
Standard PyTorch executes operations one-by-one ("eager mode"). Each operation:
- Dispatches from Python to C++
- Launches a kernel
- Reads/writes memory
- Returns to Python
This means:
- Python overhead on every operation
- No opportunity to fuse adjacent operations
- Suboptimal memory access patterns
The Solution
torch.compile captures a sequence of operations, optimizes them as a group, and generates fused kernels that minimize memory traffic and Python overhead.
Expected Speedup
- Typical: 20-50% faster on GPU, 10-30% on CPU
- Best case: 2-3x for memory-bound models (transformers)
- Worst case: No speedup if the model has many graph breaks
2. How It Works — The 3 Stages
Stage 1: TorchDynamo (Graph Capture)
TorchDynamo intercepts Python bytecode execution and traces the operations into an FX graph (a simple intermediate representation).
How it works (simplified):
- Your function runs normally
- Dynamo watches which torch operations are called
- It records them into a graph
- For subsequent calls with the same "shape signature," it replays the graph
Guards: Dynamo records assumptions about inputs (dtype, device, shape). If an assumption is violated on a later call, the graph is invalidated and Dynamo retraces (recompiles) the function.
Graph breaks: When Dynamo encounters something it can't trace (like print() with a tensor, or data-dependent control flow), it "breaks" the graph — splits execution into: compiled-graph-1 → Python → compiled-graph-2.
Stage 2: AOTAutograd (Ahead-of-Time Autograd)
After capturing the forward graph, AOTAutograd:
- Traces the backward pass as well (ahead of time)
- Partitions into a forward graph and backward graph
- Both are passed to the backend for optimization
This is important because the backward pass can also be optimized and fused.
Stage 3: TorchInductor (Code Generation)
Inductor takes the graph and generates optimized code:
- For GPU: Generates Triton kernels (a Python-based GPU programming language)
- For CPU: Generates C++/OpenMP code
Key optimizations:
- Operator fusion: Combine multiple ops into one kernel (e.g., linear + relu)
- Memory planning: Reuse memory buffers, minimize allocations
- Layout optimization: Choose best memory layout for the hardware
3. Basic Usage
Three ways to use torch.compile:
# Method 1: Compile a model
compiled_model = torch.compile(model)
output = compiled_model(input)
# Method 2: Decorator on a function
@torch.compile
def my_function(x, y):
return torch.matmul(x, y) + x
# Method 3: Compile a specific function
compiled_fn = torch.compile(my_function)
First Call is Slow
The first call triggers compilation (tracing + code generation). Subsequent calls with the same input shapes are fast:
compiled_model = torch.compile(model)
# First call: SLOW (compiles)
output = compiled_model(input_batch_1)
# Second call: FAST (uses compiled code)
output = compiled_model(input_batch_2)
4. Compilation Modes
torch.compile(model, mode="default") # Balanced
torch.compile(model, mode="reduce-overhead") # Minimize framework overhead
torch.compile(model, mode="max-autotune") # Maximum optimization effort
"default"
- Balanced between compilation time and runtime performance
- Uses a reasonable set of optimizations
- Good starting point
"reduce-overhead"
- Uses CUDA graphs to eliminate kernel launch overhead
- Best for models with many small kernels
- May increase memory usage
- GPU only
"max-autotune"
- Tries many kernel implementations and picks the fastest
- Much longer compilation time
- Best runtime performance
- Good for production deployment after development
Comparison
| Mode | Compile Time | Runtime Speed | Memory | Best For |
|---|---|---|---|---|
| default | Fast | Good | Normal | Development |
| reduce-overhead | Medium | Better | Higher | Small ops (GPU) |
| max-autotune | Slow | Best | Normal | Production |
5. Graph Breaks
What Causes Graph Breaks
A graph break splits the compiled region into multiple segments, with Python execution in between. This reduces optimization opportunities.
Common causes:
- print() with tensor values — requires Python execution
- Unsupported Python operations — certain builtins
- Data-dependent control flow —
if tensor.item() > 0: - Calling non-compilable functions — some third-party code
- In-place operations on views (in some cases)
- Python side-effects — logging, writing to files
How to Find Graph Breaks
# Method 1: explain()
explanation = torch._dynamo.explain(model)(sample_input)
print(explanation)
# Method 2: fullgraph=True raises an error on any break
compiled = torch.compile(model, fullgraph=True)
try:
compiled(input)
except Exception as e:
print(f"Graph break: {e}")
# Method 3: Logging
import logging
torch._logging.set_logs(graph_breaks=True)
How to Fix Graph Breaks
# BAD: print causes graph break
def forward(self, x):
x = self.linear(x)
print(f"Shape: {x.shape}") # GRAPH BREAK!
return self.relu(x)
# GOOD: remove print or use torch._dynamo.config.suppress_errors
def forward(self, x):
x = self.linear(x)
return self.relu(x)
# BAD: data-dependent control flow
def forward(self, x):
if x.sum() > 0: # GRAPH BREAK! (value depends on data)
return x * 2
return x
# GOOD: use torch.where for data-dependent logic
def forward(self, x):
return torch.where(x.sum() > 0, x * 2, x)
6. fullgraph=True
Forces compilation of the entire function as a single graph. If there would be any graph break, compilation fails with an error instead of silently degrading.
@torch.compile(fullgraph=True)
def my_fn(x):
# This MUST be fully traceable — no graph breaks allowed
return x.sin() + x.cos()
Use this when:
- You want maximum performance (no graph breaks = fully optimized)
- You want to catch non-compilable code early
- In production code where you've already fixed all breaks
7. Dynamic Shapes
The Problem: Recompilation
By default, torch.compile captures the exact shapes of inputs. If shapes change, it must recompile:
compiled_fn = torch.compile(fn)
compiled_fn(torch.randn(32, 64)) # Compiles for shape [32, 64]
compiled_fn(torch.randn(16, 64)) # Recompiles for shape [16, 64]!
compiled_fn(torch.randn(8, 64)) # Recompiles again!
The Solution: dynamic=True
compiled_fn = torch.compile(fn, dynamic=True)
compiled_fn(torch.randn(32, 64)) # Compiles with symbolic shapes
compiled_fn(torch.randn(16, 64)) # Reuses compiled code!
compiled_fn(torch.randn(8, 64)) # Reuses again!
With dynamic=True, Dynamo uses symbolic shapes (e.g., s0 instead of 32), and the generated code handles any batch size.
Automatic Dynamic Shapes
PyTorch can also automatically detect that a dimension varies and switch to dynamic shapes after seeing multiple sizes:
# After 2 recompilations on the same dimension, Dynamo marks it dynamic
compiled_fn(torch.randn(32, 64)) # Compile for [32, 64]
compiled_fn(torch.randn(16, 64)) # Recompile, mark dim 0 as dynamic
compiled_fn(torch.randn(8, 64)) # Uses dynamic code — no recompile!
mark_dynamic
For fine-grained control:
x = torch.randn(32, 64)
torch._dynamo.mark_dynamic(x, 0) # Mark dimension 0 as dynamic
compiled_fn(x) # Compiles with dynamic dim 0, static dim 1 (64)
8. Compilation Cache
How Caching Works
Compiled code is cached so you don't recompile every time you restart:
- In-memory cache: Within a process, same function + same guards = reuse
- Persistent cache: Across process restarts (PyTorch 2.1+), compiled
artifacts are stored on disk
Persistent Cache
# Enable persistent cache (stored in ~/.cache/torch/inductor/)
import torch._inductor.config
torch._inductor.config.fx_graph_cache = True
This means:
- First run: full compilation (slow)
- Second run: loads from cache (fast startup)
9. Compiler Stances
Stances control how the compiler handles recompilation:
# Don't compile at all — run in eager mode
torch.compiler.set_stance("force_eager")
# Warn (log) on recompilation instead of silently recompiling
torch.compiler.set_stance("eager_on_recompile")
# Error if recompilation would occur
torch.compiler.set_stance("fail_on_recompile")
# Default behavior
torch.compiler.set_stance("default")
Use fail_on_recompile in production to catch unexpected dynamic behavior that would hurt performance.
10. Debugging torch.compile
torch._dynamo.explain()
Shows what happened during compilation:
explanation = torch._dynamo.explain(compiled_fn)(input)
print(explanation)
# Shows: number of graphs, graph breaks, break reasons
TORCH_LOGS Environment Variable
# See graph breaks
TORCH_LOGS="graph_breaks" python script.py
# See what Dynamo captured
TORCH_LOGS="dynamo" python script.py
# See generated code
TORCH_LOGS="output_code" python script.py
# See recompilation reasons
TORCH_LOGS="recompiles" python script.py
Common Errors
- "torch._dynamo.exc.Unsupported" — An operation can't be traced.
Fix: rewrite using supported operations.
- Recompilation spam — Model keeps recompiling.
Fix: Use dynamic=True or ensure consistent input shapes.
- Incorrect results — Compiled code gives different outputs.
Fix: Report as a bug. Workaround: mark the function with torch._dynamo.disable().
11. torch._dynamo.reset()
Clears all compiled graphs and cached state:
torch._dynamo.reset()
Useful when:
- Testing different compilation settings
- Debugging compilation issues
- Benchmarking (to force recompilation)
12. Custom Backends
You can write your own backend that receives the FX graph:
def my_backend(gm: torch.fx.GraphModule, example_inputs):
"""
Custom backend that receives the graph and returns a callable.
Args:
gm: The captured FX graph module
example_inputs: Example inputs used during tracing
Returns:
A callable that takes the same inputs and produces outputs
"""
# Inspect the graph
print(f"Graph has {len(list(gm.graph.nodes))} nodes")
# You can transform the graph here, or just return it as-is
# (gm is already callable)
return gm
compiled_fn = torch.compile(fn, backend=my_backend)
This is useful for:
- Profiling what operations are captured
- Custom optimizations
- Research on graph transformations
13. Compiled Autograd
By default, torch.compile only compiles the forward pass. The backward pass still runs in eager mode. Compiled Autograd compiles the backward too:
with torch._dynamo.compiled_autograd.enable(torch.compile(backend="inductor")):
loss.backward()
Benefits:
- Backward pass also gets operator fusion
- End-to-end compilation of training step
14. Performance Tips
When torch.compile Helps Most
- Transformer models (lots of small ops to fuse)
- Memory-bound operations (fusion reduces memory traffic)
- Models with many element-wise operations in sequence
- Standard architectures using torch.nn modules
When It Doesn't Help Much
- Already compute-bound (large matmuls with batch dim)
- Heavy graph breaks (too much falls back to Python)
- Very small models (compilation overhead > runtime savings)
- Highly dynamic models (frequent recompilation)
Best Practices
- Profile first: Know where time is spent before compiling
- Start simple:
torch.compile(model)with defaults - Check for graph breaks: Use
explain()orfullgraph=True - Use dynamic=True if batch sizes vary
- Use max-autotune for production deployments
- Cache compiled code for fast startup
15. FX Graph Basics
The intermediate representation (IR) used by torch.compile is an FX graph:
import torch.fx
def fn(x, y):
z = x + y
return z.relu()
# Trace into FX graph
traced = torch.fx.symbolic_trace(fn)
print(traced.graph)
Output:
graph():
%x : [num_users=1] = placeholder[target=x]
%y : [num_users=1] = placeholder[target=y]
%add : [num_users=1] = call_function[target=operator.add](args = (%x, %y))
%relu : [num_users=1] = call_method[target=relu](args = (%add,))
return relu
Node types:
placeholder— function inputscall_function— calls to functions (like torch.add)call_method— method calls on tensors (.relu(), .view(), etc.)call_module— calls to nn.Module submodulesoutput— return value
Understanding FX graphs helps when:
- Writing custom backends
- Debugging compilation issues
- Understanding what optimizations are applied
Summary
| Feature | What It Does | When to Use |
|---|---|---|
| torch.compile | Compiles model for speed | Always (production) |
| fullgraph=True | Errors on graph breaks | Ensuring no breaks |
| dynamic=True | Handles varying shapes | Variable batch sizes |
| max-autotune | Maximum optimization | Deployment |
| reduce-overhead | Minimizes launch overhead | Many small GPU ops |
| explain() | Shows compilation info | Debugging |
| custom backend | Custom graph processing | Research/profiling |
Quick Start
# Step 1: Just compile it
model = torch.compile(model)
# Step 2: Check for issues
explanation = torch._dynamo.explain(model)(sample_input)
# Step 3: Optimize
model = torch.compile(model, mode="max-autotune", fullgraph=True)
📓 Open Notebook — Interactive version of this module
Source Files
compile_basics.py— torch.compile basics — fundamental usage and getting startedcompilation_modes.py— Compilation modes — default, reduce-overhead, max-autotunegraph_breaks.py— Graph breaks — examples, causes, and fixesdynamic_shapes.py— Dynamic shapes — handling varying input sizescustom_backend.py— Custom backends — writing your own torch.compile backend
Module 09: Attention Mechanisms — From Scratch to FlexAttention
Overview
Attention is the core mechanism behind transformers — the architecture powering GPT, BERT, LLaMA, and virtually all modern AI models. This module builds attention from first principles and works up to PyTorch's most advanced APIs.
1. What is Attention? (Intuitive Explanation)
The Analogy
Imagine you're reading a sentence: "The cat sat on the mat because it was tired."
What does "it" refer to? You (unconsciously) attend to earlier words and determine "it" = "cat". Attention is this mechanism: for each position in a sequence, it looks at ALL other positions and decides which ones are relevant.
The Core Idea
Given a query ("what am I looking for?"), attention:
- Compares the query against all keys ("what information is available?")
- Produces attention weights (how relevant each key is)
- Uses those weights to aggregate values ("what to return?")
2. Scaled Dot-Product Attention
The Formula
Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
Where:
- Q (Query): [batch, seq_len, d_k] — what each position is looking for
- K (Key): [batch, seq_len, d_k] — what each position offers for matching
- V (Value): [batch, seq_len, d_v] — what each position actually provides
- d_k: dimension of keys/queries
Why Scale by sqrt(d_k)?
Without scaling, dot products grow with d_k. Large dot products push softmax into regions where it has extremely small gradients (saturation). Dividing by sqrt(d_k) keeps variance ~1 regardless of dimension.
Example: if Q and K entries are independent with mean 0, variance 1:
Q @ K^Thas variance = d_k- After scaling: variance = d_k / d_k = 1
Step-by-Step with Shapes
# Input: Q, K, V each with shape [batch, seq_len, d_model]
# For simplicity, assume d_k = d_v = d_model
# Step 1: Compute attention scores
scores = Q @ K.transpose(-2, -1) # [batch, seq_len, seq_len]
# Step 2: Scale
scores = scores / math.sqrt(d_k) # [batch, seq_len, seq_len]
# Step 3: (Optional) Apply mask
scores = scores.masked_fill(mask == 0, float('-inf'))
# Step 4: Softmax to get attention weights
weights = softmax(scores, dim=-1) # [batch, seq_len, seq_len]
# Step 5: Weighted sum of values
output = weights @ V # [batch, seq_len, d_v]
3. Causal (Autoregressive) Attention
What It Is
In language models that generate text left-to-right, position i should only attend to positions 0, 1, ..., i (not future positions). This is enforced with a causal mask:
mask = [[1, 0, 0, 0],
[1, 1, 0, 0],
[1, 1, 1, 0],
[1, 1, 1, 1]]
Implementation
# Create causal mask (lower triangular)
seq_len = Q.shape[1]
causal_mask = torch.tril(torch.ones(seq_len, seq_len))
# Apply: set future positions to -inf before softmax
scores = scores.masked_fill(causal_mask == 0, float('-inf'))
# After softmax, -inf becomes 0 (no attention to future)
4. Multi-Head Attention
Why Multiple Heads?
A single attention head can only focus on one pattern (e.g., subject-verb agreement). Multiple heads let the model attend to different patterns simultaneously:
- Head 1 might attend to syntactic relationships
- Head 2 might attend to semantic similarity
- Head 3 might attend to positional proximity
How It Works
Instead of one attention with d_model dimensions, use h heads each with d_k = d_model / h dimensions:
# Input: x with shape [batch, seq_len, d_model]
# Project to Q, K, V for each head
Q = W_q(x) # [batch, seq_len, d_model]
K = W_k(x) # [batch, seq_len, d_model]
V = W_v(x) # [batch, seq_len, d_model]
# Reshape to [batch, num_heads, seq_len, head_dim]
Q = Q.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)
K = K.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)
V = V.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)
# Apply attention independently per head
# output: [batch, num_heads, seq_len, head_dim]
output = scaled_dot_product_attention(Q, K, V)
# Concatenate heads and project
output = output.transpose(1, 2).reshape(batch, seq_len, d_model)
output = W_o(output) # Final linear projection
5. F.scaled_dot_product_attention (SDPA)
PyTorch provides an optimized implementation:
import torch.nn.functional as F
output = F.scaled_dot_product_attention(
query, # [batch, heads, seq_q, dim]
key, # [batch, heads, seq_kv, dim]
value, # [batch, heads, seq_kv, dim_v]
attn_mask=None, # Optional mask
dropout_p=0.0, # Dropout probability
is_causal=False, # If True, applies causal mask automatically
scale=None, # Custom scale factor (default: 1/sqrt(dim))
)
Benefits
- Automatically selects the fastest available backend
- Memory-efficient (doesn't materialize the full attention matrix)
- Fused implementation (fewer memory reads/writes)
6. SDPA Backends
PyTorch has multiple backends for SDPA, each optimized for different cases:
Flash Attention
- Memory: O(N) instead of O(N^2)
- Speed: Fastest for long sequences
- How: Tiles the computation, never materializes full attention matrix
- Requires: GPU with compute capability >= 8.0 (A100+)
Memory-Efficient Attention
- Memory: O(N) via chunked computation
- Speed: Good general-purpose
- How: Processes attention in chunks
- Requires: Any GPU
cuDNN Attention
- Speed: Can be fastest for standard shapes
- Requires: cuDNN 8.9+
Math (Fallback)
- Memory: O(N^2) — materializes full attention matrix
- Speed: Slowest for long sequences
- Works: Everywhere (CPU and GPU)
- Use: When other backends aren't available or for debugging
Backend Selection
from torch.nn.attention import sdpa_kernel, SDPBackend
# Force a specific backend
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
output = F.scaled_dot_product_attention(q, k, v)
# Use math backend for debugging (can inspect attention weights)
with sdpa_kernel(SDPBackend.MATH):
output = F.scaled_dot_product_attention(q, k, v)
7. Flash Attention Explained Simply
The Problem
Standard attention computes:
S = Q @ K^T # [N, N] — this is HUGE for long sequences
P = softmax(S) # [N, N] — stored in memory
O = P @ V # output
For N = 16384 (16K tokens): S alone is 16384^2 * 4 bytes = 1 GB!
The Flash Attention Trick
Instead of computing the full N×N matrix:
- Tile Q, K, V into small blocks that fit in GPU fast memory (SRAM)
- Compute attention block by block
- Use the online softmax trick to accumulate correct softmax results
across blocks without needing the full row
Result: O(N) memory instead of O(N^2), and faster due to fewer memory accesses.
Online Softmax
The key insight: you can compute softmax incrementally. As you process each block of keys, maintain a running maximum and running sum, then correct at the end. This avoids needing all scores simultaneously.
8. FlexAttention
What It Is
FlexAttention (torch.nn.attention.flex_attention) lets you define custom attention patterns using simple Python functions, while still getting the performance benefits of Flash Attention-like fused kernels.
Why It Exists
Previously, custom attention patterns (sliding window, ALiBi, document masking) required:
- Writing custom CUDA kernels (hard)
- Materializing masks (memory expensive)
- Giving up on fused implementations (slow)
FlexAttention compiles your pattern into an efficient kernel.
Core API
from torch.nn.attention.flex_attention import (
flex_attention,
create_block_mask,
)
# score_mod: modifies attention scores before softmax
def causal_score_mod(score, b, h, q_idx, kv_idx):
return torch.where(q_idx >= kv_idx, score, float('-inf'))
# mask_mod: defines which positions can attend to which (for BlockMask)
def causal_mask_mod(b, h, q_idx, kv_idx):
return q_idx >= kv_idx
# Create a BlockMask for efficiency
block_mask = create_block_mask(causal_mask_mod, B=1, H=1, Q_LEN=seq_len, KV_LEN=seq_len)
# Apply flex attention
output = flex_attention(query, key, value, block_mask=block_mask)
score_mod vs mask_mod
- score_mod(score, b, h, q_idx, kv_idx): Modifies the attention score
for a specific (query, key) pair. Can add biases, apply causal masking, etc. Returns the modified score.
- mask_mod(b, h, q_idx, kv_idx): Returns True/False for whether this
(query, key) pair should be allowed. Used by BlockMask to skip entire blocks of computation.
9. FlexAttention Patterns
Causal Attention
def causal(b, h, q_idx, kv_idx):
return q_idx >= kv_idx
Sliding Window
def sliding_window(b, h, q_idx, kv_idx):
return (q_idx - kv_idx).abs() <= window_size
Causal + Sliding Window
def causal_sliding(b, h, q_idx, kv_idx):
return (q_idx >= kv_idx) & (q_idx - kv_idx <= window_size)
ALiBi (Attention with Linear Biases)
def alibi_score_mod(score, b, h, q_idx, kv_idx):
slope = 2 ** (-(h + 1) * 8 / num_heads)
bias = -slope * (q_idx - kv_idx).abs()
return score + bias
Document Masking (Multiple Documents in One Sequence)
# document_id[i] tells which document position i belongs to
def document_mask(b, h, q_idx, kv_idx):
return document_id[q_idx] == document_id[kv_idx]
Prefix LM (Bidirectional prefix + Causal suffix)
def prefix_lm(b, h, q_idx, kv_idx):
# Allow bidirectional attention within prefix
# Causal attention after prefix
return (kv_idx < prefix_length) | (q_idx >= kv_idx)
10. Building a Transformer Block
Pre-Norm vs Post-Norm
# Post-norm (original Transformer paper):
x = x + attention(x)
x = layer_norm(x)
# Pre-norm (GPT-2, modern models — more stable training):
x = x + attention(layer_norm(x))
RMSNorm vs LayerNorm
# LayerNorm: normalize, then scale + shift
# Subtracts mean AND divides by std
y = (x - mean(x)) / std(x) * gamma + beta
# RMSNorm: only divides by RMS (no mean subtraction)
# Faster, often works just as well
y = x / rms(x) * gamma
Activation Functions
- GELU: Smooth approximation of ReLU, used in BERT/GPT
- SiLU (Swish): x * sigmoid(x), used in LLaMA, PaLM
- SwiGLU: Gated variant, used in modern LLMs
Complete Block
class TransformerBlock(nn.Module):
def __init__(self, dim, num_heads):
self.norm1 = RMSNorm(dim)
self.attn = MultiHeadAttention(dim, num_heads)
self.norm2 = RMSNorm(dim)
self.ffn = SwiGLU(dim)
def forward(self, x):
x = x + self.attn(self.norm1(x))
x = x + self.ffn(self.norm2(x))
return x
11. Positional Encoding
Attention is permutation-invariant — it doesn't know token order. We add positional information explicitly.
Sinusoidal (Original Transformer)
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
Each dimension oscillates at a different frequency. The model can learn to attend to relative positions via linear combinations.
Learned Positional Embeddings
self.pos_embed = nn.Embedding(max_seq_len, d_model)
x = x + self.pos_embed(positions)
Simple but limited to max_seq_len seen during training.
RoPE (Rotary Position Embeddings)
Used in LLaMA, Mistral, and most modern LLMs.
Key idea: encode position by rotating Q and K vectors in 2D subspaces. The dot product Q·K then naturally depends on relative position.
# For each pair of dimensions (2i, 2i+1):
q_rotated[2i] = q[2i] * cos(theta) - q[2i+1] * sin(theta)
q_rotated[2i+1] = q[2i] * sin(theta) + q[2i+1] * cos(theta)
# Where theta = position * base_freq^(-2i/d)
Benefits:
- Relative position awareness
- Can extrapolate to longer sequences than training
- No additional parameters
12. KV Cache for Inference
The Problem
During autoregressive generation, each new token needs to attend to ALL previous tokens. Without caching, you recompute K and V for all previous tokens at every step.
The Solution: KV Cache
Cache the K and V projections from previous steps:
# Step 1: Process prompt
K_cache = key_projection(prompt) # [batch, num_heads, prompt_len, head_dim]
V_cache = value_projection(prompt) # [batch, num_heads, prompt_len, head_dim]
# Step 2+: Generate each new token
for _ in range(max_new_tokens):
# Only project the NEW token
new_k = key_projection(new_token) # [batch, num_heads, 1, head_dim]
new_v = value_projection(new_token) # [batch, num_heads, 1, head_dim]
# Append to cache
K_cache = torch.cat([K_cache, new_k], dim=2)
V_cache = torch.cat([V_cache, new_v], dim=2)
# Query only needs the new token, but attends to full cache
new_q = query_projection(new_token) # [batch, num_heads, 1, head_dim]
output = attention(new_q, K_cache, V_cache)
Memory vs Speed Tradeoff
- Without cache: O(n) compute per step, O(1) extra memory
- With cache: O(1) amortized compute per step, O(n) extra memory
The KV cache makes generation ~10x faster for long sequences.
Summary
| Concept | Purpose | Key Shape |
|---|---|---|
| Scaled dot-product | Core attention computation | [B, N, N] scores |
| Causal mask | Prevent attending to future | Lower triangular |
| Multi-head | Multiple attention patterns | h heads, d_k = d/h |
| SDPA | Optimized PyTorch attention | Same as manual |
| Flash Attention | O(N) memory attention | Tiled computation |
| FlexAttention | Custom patterns, fused | score_mod/mask_mod |
| RoPE | Relative position encoding | 2D rotations |
| KV Cache | Fast autoregressive generation | Cached K, V tensors |
📓 Open Notebook — Interactive version of this module
Source Files
manual_attention.py— Manual attention — implementing attention from scratch with shape annotationssdpa_and_backends.py— SDPA and backend control — PyTorch's optimized attentionmultihead_attention.py— Multi-head attention — complete implementation with shape trackingtransformer_block.py— Transformer block — full implementation with RMSNorm/LayerNormflex_attention_patterns.py— FlexAttention patterns — custom attention with compiled kernels
Module 10: Distributed Training in PyTorch
Table of Contents
- Why Distributed Training?
- Key Concepts
- Launching with torchrun
- Collective Operations
- DistributedDataParallel (DDP)
- DeviceMesh
- DTensor (Distributed Tensor)
- FSDP1 vs FSDP2
- FSDP2 (fully_shard)
- Tensor Parallelism
- Pipeline Parallelism
- Combining Strategies: 3D Parallelism
- Distributed Checkpointing (DCP)
- SymmetricMemory
- Context Parallel
- Practical Advice
Why Distributed Training?
As models grow larger and datasets expand, a single GPU becomes insufficient. Distributed training addresses three fundamental bottlenecks:
1. Model too large for one GPU (Memory) A model like Llama 70B requires ~140 GB just for parameters in FP16. Even the largest GPUs (H100 with 80 GB) cannot hold this model, let alone the optimizer states and activations needed for training. Distributed strategies split the model across GPUs.
2. Training too slow (Compute) Even when a model fits on one GPU, training can take weeks or months. By distributing data across N GPUs, each GPU processes 1/N of the data per step, achieving near-linear speedup. Training that takes 30 days on 1 GPU takes ~4 days on 8 GPUs.
3. Data too large (I/O and throughput) With terabytes of training data, increasing throughput by processing more samples in parallel reduces wall-clock training time proportionally.
The Parallelism Taxonomy
| Strategy | What is split? | When to use |
|---|---|---|
| Data Parallel (DDP) | Data (each GPU has full model copy) | Model fits on 1 GPU |
| Fully Sharded Data Parallel (FSDP) | Data + model parameters | Model barely fits or doesn't fit on 1 GPU |
| Tensor Parallel (TP) | Individual layers/tensors | Very large layers (e.g., huge linear layers) |
| Pipeline Parallel (PP) | Model stages (groups of layers) | Very deep models, many GPUs |
| Context Parallel (CP) | Sequence dimension | Very long sequences |
| 3D Parallelism | Combination of DP + TP + PP | Large-scale training (100s-1000s of GPUs) |
Key Concepts
World Size, Rank, and Local Rank
When you launch distributed training, you create multiple processes, each driving one GPU. These processes form a process group and coordinate via collective communication.
Node 0 (Machine 0) Node 1 (Machine 1)
┌──────────────────┐ ┌──────────────────┐
│ GPU0 GPU1 │ │ GPU0 GPU1 │
│ rank=0 rank=1 │ │ rank=2 rank=3 │
│ local_rank=0 │ │ local_rank=0 │
│ local_rank=1 │ local_rank=1
└──────────────────┘ └──────────────────┘
world_size = 4
- World size: Total number of processes across all machines
- Rank: Unique global identifier for each process (0 to world_size-1)
- Local rank: Identifier within a single machine (0 to num_local_gpus-1)
import torch.distributed as dist
dist.init_process_group(backend="nccl") # or "gloo" for CPU
rank = dist.get_rank()
world_size = dist.get_world_size()
local_rank = int(os.environ["LOCAL_RANK"])
Process Groups
A process group is a subset of all processes that can communicate. The default process group includes all processes and is created by init_process_group. You can create sub-groups for specialized communication:
# Create a group with only ranks 0 and 1
subgroup = dist.new_group(ranks=[0, 1])
# Only ranks in the group participate in collectives on this group
if rank in [0, 1]:
dist.all_reduce(tensor, group=subgroup)
Backends
| Backend | Devices | Use Case |
|---|---|---|
| NCCL | GPU (NVIDIA) | Default for GPU training. Highly optimized for NVIDIA hardware |
| Gloo | CPU, GPU | CPU training, or as a fallback. Also used for CPU collectives in GPU training |
| UCC | GPU | Alternative to NCCL, supports additional hardware |
NCCL (pronounced "nickel") is the standard for GPU training and provides the best performance. Gloo is useful for CPU-based experiments and prototyping.
# GPU training (most common)
dist.init_process_group(backend="nccl")
# CPU training or prototyping
dist.init_process_group(backend="gloo")
# Use NCCL for GPU ops, Gloo for CPU ops (advanced)
dist.init_process_group(backend="nccl")
cpu_group = dist.new_group(backend="gloo")
Launching with torchrun
torchrun is PyTorch's built-in launcher for distributed training. It replaces the older torch.distributed.launch. It sets up environment variables and spawns processes for you.
Single-Node Launch
# 4 GPUs on one machine
torchrun --nproc_per_node=4 train.py --arg1 val1
# Can also use for CPU (with gloo backend)
torchrun --nproc_per_node=2 train.py
Multi-Node Launch
On each machine, run torchrun with the same --master_addr and --master_port:
# Machine 0 (master)
torchrun \
--nproc_per_node=8 \
--nnodes=2 \
--node_rank=0 \
--master_addr=192.168.1.100 \
--master_port=29500 \
train.py
# Machine 1
torchrun \
--nproc_per_node=8 \
--nnodes=2 \
--node_rank=1 \
--master_addr=192.168.1.100 \
--master_port=29500 \
train.py
Environment Variables Set by torchrun
| Variable | Description |
|---|---|
RANK | Global rank of this process |
LOCAL_RANK | Local rank on this node |
WORLD_SIZE | Total number of processes |
MASTER_ADDR | Address of the master node |
MASTER_PORT | Port for master node communication |
LOCAL_WORLD_SIZE | Number of processes on this node |
Your training script reads these:
import os
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
Elastic Launch
torchrun supports elastic training where nodes can join or leave:
torchrun \
--nproc_per_node=4 \
--nnodes=2:8 \ # min 2, max 8 nodes
--rdzv_backend=c10d \
--rdzv_endpoint=master:29500 \
train.py
Collective Operations
Collective operations are communication primitives where all processes in a group participate. Understanding these is essential because all distributed strategies are built on top of them.
All-Reduce
Every process starts with a tensor. After all-reduce, every process has the element-wise sum (or other reduction) of all tensors.
Before: After all_reduce(SUM):
Rank 0: [1, 2] Rank 0: [10, 20]
Rank 1: [3, 4] → Rank 1: [10, 20]
Rank 2: [6, 14] Rank 2: [10, 20]
This is the core operation in DDP: after each backward pass, gradients are all-reduced so every replica has the same averaged gradients.
tensor = torch.tensor([rank * 2.0, rank * 3.0])
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
# Now tensor is the same on all ranks
All-Gather
Each process contributes a tensor, and every process receives the concatenation of all tensors.
Before: After all_gather:
Rank 0: [A] Rank 0: [A, B, C]
Rank 1: [B] → Rank 1: [A, B, C]
Rank 2: [C] Rank 2: [A, B, C]
FSDP uses all-gather to reconstruct full parameter tensors before the forward pass.
local_tensor = torch.tensor([rank])
gathered = [torch.zeros(1) for _ in range(world_size)]
dist.all_gather(gathered, local_tensor)
# gathered = [tensor([0]), tensor([1]), tensor([2])]
Reduce-Scatter
The inverse of all-gather. First reduces (sums) all tensors element-wise, then scatters the result so each rank gets a different chunk.
Before: After reduce_scatter:
Rank 0: [1, 2, 3] Rank 0: [6] (sum of position 0: 1+2+3)
Rank 1: [2, 3, 4] → Rank 1: [9] (sum of position 1: 2+3+4+... wait)
Rank 2: [3, 4, 5] Rank 2: [12] (sum of position 2: 3+4+5)
More precisely, with 3 ranks each holding a 3-element tensor:
- Element-wise sum: [1+2+3, 2+3+4, 3+4+5] = [6, 9, 12]
- Scatter: Rank 0 gets [6], Rank 1 gets [9], Rank 2 gets [12]
FSDP uses reduce-scatter after backward to reduce gradients and distribute shards back to their owners.
output = torch.zeros(2)
input_tensor = torch.arange(world_size * 2, dtype=torch.float) + rank
dist.reduce_scatter_tensor(output, input_tensor, op=dist.ReduceOp.SUM)
Broadcast
One process sends a tensor to all other processes.
Before: After broadcast(src=0):
Rank 0: [42, 7] Rank 0: [42, 7]
Rank 1: [0, 0] → Rank 1: [42, 7]
Rank 2: [0, 0] Rank 2: [42, 7]
tensor = torch.tensor([42.0, 7.0]) if rank == 0 else torch.zeros(2)
dist.broadcast(tensor, src=0)
Barrier
Synchronizes all processes. Every process blocks until all processes have reached the barrier. No data is exchanged.
dist.barrier() # All ranks wait here until everyone arrives
Use barriers sparingly; they are expensive and often unnecessary when collectives already imply synchronization.
Reduce
Like all-reduce, but the result only goes to one destination rank.
Before: After reduce(dst=0, SUM):
Rank 0: [1, 2] Rank 0: [6, 9]
Rank 1: [2, 3] → Rank 1: [2, 3] (unchanged)
Rank 2: [3, 4] Rank 2: [3, 4] (unchanged)
Scatter
One process distributes different chunks to each process.
Before (rank 0 has all data): After scatter(src=0):
Rank 0: [[A], [B], [C]] Rank 0: [A]
Rank 1: [] → Rank 1: [B]
Rank 2: [] Rank 2: [C]
Gather
The inverse of scatter. All processes send data to one destination.
Before: After gather(dst=0):
Rank 0: [A] Rank 0: [[A], [B], [C]]
Rank 1: [B] → Rank 1: [B] (unchanged)
Rank 2: [C] Rank 2: [C] (unchanged)
DistributedDataParallel (DDP)
DDP is the simplest and most commonly used distributed training strategy. Each GPU holds a complete copy of the model. Training data is split across GPUs, and gradients are synchronized via all-reduce after each backward pass.
How DDP Works Internally
- Initialization: The model is replicated on each GPU. Parameters are
broadcast from rank 0 to ensure all replicas start identically.
- Forward pass: Each rank processes its own mini-batch independently.
No communication occurs during forward.
- Backward pass: As gradients are computed, DDP groups them into
buckets (default ~25 MB each). When a bucket is full, all-reduce starts immediately -- overlapping communication with computation for the remaining layers. This is called gradient bucketing.
- Optimizer step: After all-reduce completes, every rank has identical
averaged gradients. Each rank runs the optimizer independently, producing identical updated parameters.
Forward (independent)
┌─────────┐
Rank 0: Data₀ → │ Model₀ │ → Loss₀
└─────────┘
┌─────────┐
Rank 1: Data₁ → │ Model₁ │ → Loss₁
└─────────┘
Backward (all-reduce gradients)
┌────────────────────────┐
│ All-Reduce Gradients │
│ (bucketed, overlapped │
│ with backward comp) │
└────────────────────────┘
Optimizer Step (independent, identical)
Complete DDP Setup
import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
def setup():
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
return local_rank
def cleanup():
dist.destroy_process_group()
def main():
local_rank = setup()
device = torch.device(f"cuda:{local_rank}")
# Create model and move to GPU
model = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 10),
).to(device)
# Wrap with DDP
model = DDP(model, device_ids=[local_rank])
# DistributedSampler ensures each rank sees different data
dataset = MyDataset()
sampler = DistributedSampler(dataset)
dataloader = DataLoader(dataset, batch_size=32, sampler=sampler)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
for epoch in range(10):
sampler.set_epoch(epoch) # Shuffle differently each epoch
for batch_x, batch_y in dataloader:
batch_x, batch_y = batch_x.to(device), batch_y.to(device)
optimizer.zero_grad()
output = model(batch_x)
loss = loss_fn(output, batch_y)
loss.backward()
optimizer.step()
cleanup()
if __name__ == "__main__":
main()
DistributedSampler
The DistributedSampler partitions the dataset indices so each rank gets a different subset. It pads the dataset to make it evenly divisible by world_size.
Key: call sampler.set_epoch(epoch) each epoch to get different shuffling. Without this, every epoch uses the same data order per rank.
DDP Tips
- Access the underlying model via
model.module(DDP wraps it). - Save checkpoints only on rank 0 to avoid file conflicts.
- Use
torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)before wrapping
with DDP if your model uses BatchNorm.
- The
find_unused_parameters=Trueflag is needed if some parameters don't
receive gradients every iteration (e.g., conditional branches).
DeviceMesh
DeviceMesh is the foundation of modern distributed training in PyTorch. It provides a multi-dimensional abstraction over a group of devices, replacing the manual management of process groups.
What is a DeviceMesh?
A DeviceMesh represents a logical grid of devices. Each dimension of the mesh corresponds to a parallelism strategy:
from torch.distributed.device_mesh import init_device_mesh
# 1D mesh: all 8 GPUs in one dimension (simple data parallelism)
mesh_1d = init_device_mesh("cuda", (8,), mesh_dim_names=("dp",))
# 2D mesh: 4 data-parallel groups × 2 tensor-parallel groups
# GPUs arranged as:
# TP=0 TP=1
# ──────────
# GPU0 GPU1 ← DP group 0
# GPU2 GPU3 ← DP group 1
# GPU4 GPU5 ← DP group 2
# GPU6 GPU7 ← DP group 3
mesh_2d = init_device_mesh("cuda", (4, 2), mesh_dim_names=("dp", "tp"))
# 3D mesh: DP × TP × PP
mesh_3d = init_device_mesh(
"cuda", (2, 2, 2), mesh_dim_names=("dp", "tp", "pp")
)
Why DeviceMesh?
Before DeviceMesh, you had to manually create process groups for each parallelism dimension and carefully track which ranks belonged to which groups. DeviceMesh automates this:
# Old way: manually creating groups
dp_groups = []
for i in range(0, 8, 2):
dp_groups.append(dist.new_group([i, i+1]))
# New way: DeviceMesh handles it
mesh = init_device_mesh("cuda", (4, 2), mesh_dim_names=("dp", "tp"))
dp_mesh = mesh["dp"] # Automatically creates the right groups
tp_mesh = mesh["tp"]
Accessing Sub-Meshes
You can slice a DeviceMesh to get a sub-mesh for a specific dimension:
mesh = init_device_mesh("cuda", (4, 2), mesh_dim_names=("dp", "tp"))
# Get the sub-mesh for data parallelism
dp_mesh = mesh["dp"] # 1D mesh with 4 devices (for this rank's DP group)
# Get the sub-mesh for tensor parallelism
tp_mesh = mesh["tp"] # 1D mesh with 2 devices (for this rank's TP group)
# These sub-meshes carry the correct process groups
# so you can pass them directly to FSDP, TP, etc.
DeviceMesh for 3D Parallelism
# 16 GPUs: 2 DP × 4 TP × 2 PP
mesh = init_device_mesh(
"cuda", (2, 4, 2), mesh_dim_names=("dp", "tp", "pp")
)
# Each parallelism strategy gets its own sub-mesh
dp_mesh = mesh["dp"]
tp_mesh = mesh["tp"]
pp_mesh = mesh["pp"]
# Apply each strategy using its sub-mesh
# TP on tp_mesh, FSDP on dp_mesh, PP on pp_mesh
DTensor (Distributed Tensor)
DTensor is a tensor abstraction that represents a tensor distributed across multiple devices. It knows how the tensor is distributed and automatically handles the communication needed for operations.
Placement Types
DTensor uses placements to describe how a tensor's data is distributed across the devices in a DeviceMesh dimension:
| Placement | Description | Example |
|---|---|---|
Shard(dim) | Tensor is sharded along dimension dim | A [4, 8] tensor Shard(1) across 2 GPUs → each gets [4, 4] |
Replicate() | Tensor is fully replicated on each device | A [4, 8] tensor Replicate() → each GPU has [4, 8] |
Partial() | Each device has a partial result; needs reduction | Intermediate matmul results before all-reduce |
Creating DTensors
from torch.distributed.tensor import DTensor, Shard, Replicate, distribute_tensor
mesh = init_device_mesh("cuda", (4,))
# Create a regular tensor and distribute it
big_tensor = torch.randn(16, 32)
# Shard along dim 0: each of 4 GPUs gets a [4, 32] chunk
sharded = distribute_tensor(big_tensor, mesh, placements=[Shard(0)])
# Replicate: each GPU gets the full [16, 32] tensor
replicated = distribute_tensor(big_tensor, mesh, placements=[Replicate()])
distribute_module
Instead of manually distributing each parameter, distribute_module distributes an entire module's parameters and handles input/output:
from torch.distributed.tensor import distribute_module, Shard, Replicate
def input_fn(mod, inputs, mesh):
# How to distribute inputs
return (distribute_tensor(inputs[0], mesh, [Shard(0)]),)
def output_fn(mod, outputs, mesh):
# How to gather outputs
return outputs.full_tensor()
model = nn.Linear(1024, 512)
distribute_module(
model,
device_mesh=mesh,
input_fn=input_fn,
output_fn=output_fn,
)
DTensor and Automatic Communication
The key insight: when you do operations on DTensors, PyTorch automatically inserts the right collectives. If you multiply a Shard(1) tensor by a Replicate() tensor, PyTorch knows it needs an all-reduce to get the correct result.
A: [Shard(1)] × B: [Replicate()] → C: [Partial()] → all_reduce → C: [Replicate()]
This is how Tensor Parallelism works under the hood.
FSDP1 vs FSDP2
FSDP1 (FullyShardedDataParallel - Legacy)
FSDP1 was PyTorch's first fully sharded data parallelism implementation. It wraps the entire module:
# FSDP1 (legacy)
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
model = FSDP(model, auto_wrap_policy=size_based_auto_wrap_policy)
Problems with FSDP1:
- Module wrapping: FSDP1 wraps modules, creating a new
FSDPmodule
hierarchy. This breaks model.layer1.weight access patterns.
- Non-composable: Hard to combine with other parallelism strategies.
- Complex API: Many constructor arguments, hard to reason about behavior.
- Flattened parameters: Parameters are flattened into a single 1D tensor,
making debugging and checkpointing harder.
FSDP2 (fully_shard - Current Standard)
FSDP2 is a ground-up rewrite with a composable, per-parameter design:
# FSDP2 (current)
from torch.distributed.fsdp import fully_shard
# Apply to submodules first, then root
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
Advantages of FSDP2:
- Composable: Works seamlessly with TP, PP, and other strategies.
- Per-parameter sharding: Each parameter is sharded independently (via
DTensor), preserving the original module structure.
- Simpler API: A single function call per module.
- Better debugging: Parameters remain accessible with their original names.
- DTensor-based: Built on DTensor, providing a clean abstraction.
Use FSDP2 for all new code. FSDP1 is maintained but not actively developed.
FSDP2 (fully_shard)
How Sharding Works
FSDP2 shards model parameters across data-parallel ranks. During training:
- Idle state: Each rank holds only its shard (1/N) of each parameter.
Memory usage is reduced by ~N× for parameters.
- Before forward:
all-gatherreconstructs the full parameters from all
shards. Each rank now temporarily has the full parameter for computation.
- Forward computation: Runs normally with the full parameters.
- After forward: Full parameters are freed (unless needed for backward).
Memory drops back to 1/N.
- Before backward:
all-gatheragain to get full parameters for gradient
computation.
- After backward:
reduce-scattersynchronizes gradients AND distributes
gradient shards. Each rank ends up with the gradient shard corresponding to its parameter shard.
- Optimizer step: Each rank updates only its parameter shard using its
gradient shard. No communication needed.
┌─────────────────────────────────────────────────┐
│ FSDP2 Training Loop │
│ │
│ Idle: Each rank holds 1/N of params │
│ ↓ │
│ all-gather → Full params → Forward │
│ ↓ │
│ Free full params (keep shards) │
│ ↓ │
│ all-gather → Full params → Backward │
│ ↓ │
│ reduce-scatter → Gradient shards │
│ ↓ │
│ Optimizer step (on shards only) │
│ ↓ │
│ Back to idle (1/N params + 1/N grads) │
└─────────────────────────────────────────────────┘
Basic FSDP2 Setup
import torch
import torch.nn as nn
from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
class TransformerBlock(nn.Module):
def __init__(self, d_model):
super().__init__()
self.attn = nn.MultiheadAttention(d_model, 8)
self.ffn = nn.Sequential(
nn.Linear(d_model, 4 * d_model),
nn.GELU(),
nn.Linear(4 * d_model, d_model),
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
def forward(self, x):
x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]
x = x + self.ffn(self.norm2(x))
return x
class TransformerModel(nn.Module):
def __init__(self, d_model=512, n_layers=6):
super().__init__()
self.embed = nn.Embedding(10000, d_model)
self.layers = nn.ModuleList(
[TransformerBlock(d_model) for _ in range(n_layers)]
)
self.head = nn.Linear(d_model, 10000)
def forward(self, x):
x = self.embed(x)
for layer in self.layers:
x = layer(x)
return self.head(x)
# Apply FSDP2: submodules first, then root
model = TransformerModel().cuda()
for layer in model.layers:
fully_shard(layer)
fully_shard(model) # Root module last
MixedPrecisionPolicy
Mixed precision reduces memory and increases throughput by using lower precision for computation while maintaining accuracy:
from torch.distributed.fsdp import MixedPrecisionPolicy
# BFloat16 for compute, FP32 for parameter storage
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16, # Cast params to bf16 for forward/backward
reduce_dtype=torch.float32, # Reduce gradients in fp32 for accuracy
)
for layer in model.layers:
fully_shard(layer, mp_policy=mp_policy)
fully_shard(model, mp_policy=mp_policy)
CPUOffloadPolicy
For extremely large models, offload parameters to CPU when not in use:
from torch.distributed.fsdp import CPUOffloadPolicy
offload_policy = CPUOffloadPolicy(pin_memory=True)
for layer in model.layers:
fully_shard(layer, offload_policy=offload_policy)
fully_shard(model, offload_policy=offload_policy)
This trades compute speed for memory: parameters are moved to CPU after forward/backward, freeing GPU memory. pin_memory=True uses pinned (page- locked) CPU memory for faster CPU-GPU transfers.
FSDP2 + DeviceMesh
When combining FSDP with other parallelism, use DeviceMesh to specify which dimension is for data parallelism:
mesh = init_device_mesh("cuda", (4, 2), mesh_dim_names=("dp", "tp"))
dp_mesh = mesh["dp"]
# fully_shard on the DP dimension
for layer in model.layers:
fully_shard(layer, mesh=dp_mesh)
fully_shard(model, mesh=dp_mesh)
Tensor Parallelism
Tensor Parallelism (TP) splits individual layers across GPUs. Unlike FSDP which shards parameters and reconstructs them for computation, TP keeps parameters split and performs distributed computation.
Why Tensor Parallelism?
- Reduces memory per GPU for very large individual layers
- Reduces latency per layer (each GPU does less work)
- Essential for models where single layers are too large for one GPU
- Works within a single node (requires fast interconnect like NVLink)
Column-wise and Row-wise Parallelism
For a linear layer Y = XW + b, there are two ways to split:
ColwiseParallel: Split W along columns (output dimension)
GPU 0 GPU 1
W = [w₀ | w₁] → W₀ = w₀ W₁ = w₁
Y = X @ W → Y₀ = X @ W₀ Y₁ = X @ W₁
Y = [Y₀ | Y₁] (gathered)
Each GPU computes a portion of the output. The input X is replicated.
RowwiseParallel: Split W along rows (input dimension)
GPU 0 GPU 1
W = [w₀] W₀ = w₀ W₁ = w₁
[w₁]
X = [x₀ | x₁] X₀ = x₀ X₁ = x₁
Y = X @ W → Y = X₀@W₀ + X₁@W₁ (all-reduce)
Input is split and each GPU computes a partial result that must be summed.
Using TP in PyTorch
from torch.distributed.tensor.parallel import (
parallelize_module,
ColwiseParallel,
RowwiseParallel,
SequenceParallel,
)
mesh = init_device_mesh("cuda", (world_size,))
# Parallelize specific layers
parallelize_plan = {
"attn.qkv_proj": ColwiseParallel(),
"attn.out_proj": RowwiseParallel(),
"ffn.up_proj": ColwiseParallel(),
"ffn.down_proj": RowwiseParallel(),
}
parallelize_module(model, mesh, parallelize_plan)
SequenceParallel
SequenceParallel splits the sequence dimension instead of replicating activations. Between ColwiseParallel and RowwiseParallel layers, activations can be kept split along the sequence dimension, reducing activation memory:
parallelize_plan = {
"norm1": SequenceParallel(),
"attn.qkv_proj": ColwiseParallel(input_layouts=Shard(0)),
"attn.out_proj": RowwiseParallel(output_layouts=Shard(0)),
"norm2": SequenceParallel(),
"ffn.up_proj": ColwiseParallel(input_layouts=Shard(0)),
"ffn.down_proj": RowwiseParallel(output_layouts=Shard(0)),
}
When to Use TP
- Large models with huge linear layers (e.g., LLMs with 8192+ hidden dim)
- When you have fast intra-node interconnect (NVLink)
- Typically TP degree of 2, 4, or 8 within a single node
- Beyond 8-way TP, the communication overhead usually outweighs benefits
Pipeline Parallelism
Pipeline Parallelism (PP) splits a model into sequential stages, each running on a different GPU (or set of GPUs). Data flows through stages sequentially, like an assembly line.
The Bubble Problem
Naive pipeline parallelism has a severe inefficiency: while stage 1 processes micro-batch 1, stages 2-N are idle. This idle time is called the pipeline bubble.
Naive (Sequential):
Stage 0: [F1][F2][F3][F4]
Stage 1: [F1][F2][F3][F4]
Stage 2: [F1][F2][F3][F4]
↑ huge bubble, most GPUs idle most of the time
Micro-batching
The solution: split each mini-batch into multiple micro-batches and pipeline them:
With micro-batches:
Stage 0: [F1][F2][F3][F4][B4][B3][B2][B1]
Stage 1: [F1][F2][F3][F4][B4][B3][B2][B1]
Stage 2: [F1][F2][F3][F4][B4][B3][B2][B1]
F = forward, B = backward, number = micro-batch id
Pipeline Schedules
Different schedules offer different trade-offs between memory, bubble ratio, and implementation complexity:
| Schedule | Bubble Ratio | Memory | Description |
|---|---|---|---|
| GPipe | (p-1)/m | High (all activations) | All forwards, then all backwards |
| 1F1B | (p-1)/m | Low (1 activation) | Alternating forward-backward in steady state |
| Interleaved 1F1B | (p-1)/(m×v) | Low | Multiple virtual stages per rank, smaller bubbles |
| Zero Bubble | ~0 | Moderate | Overlaps weight gradient with next forward |
| DualPipeV | ~0 | Moderate | Bidirectional pipeline for near-zero bubble |
Where p = number of pipeline stages, m = number of micro-batches, v = number of virtual stages (chunks).
GPipe: Simple but memory-hungry. All micro-batches do forward, then all do backward. Must store activations for all micro-batches simultaneously.
1F1B (One Forward One Backward): After a warm-up phase, alternates one forward and one backward. Steady-state memory is constant (only one micro- batch's activations at a time).
Interleaved 1F1B: Each rank handles multiple non-contiguous stages (e.g., rank 0 handles stages 0 and 4). This reduces the bubble because micro-batches cycle through stages faster.
Zero Bubble: Splits backward into two parts (input gradient and weight gradient) and overlaps weight gradient computation with the next forward pass. Nearly eliminates the bubble.
DualPipeV: A bidirectional schedule where micro-batches flow both forward and backward through the pipeline simultaneously, achieving near-zero bubble with better memory efficiency.
PP in PyTorch
from torch.distributed.pipelining import (
pipeline,
SplitPoint,
ScheduleGPipe,
Schedule1F1B,
ScheduleInterleaved1F1B,
)
# Split model into stages
pipe = pipeline(
model,
mb_args=(torch.randn(batch_size, seq_len, d_model),),
split_spec={
"layers.3": SplitPoint.BEGINNING, # Split before layer 3
"layers.6": SplitPoint.BEGINNING, # Split before layer 6
},
)
# Get this rank's stage
stage = pipe.get_stage(rank, device)
# Create schedule
schedule = Schedule1F1B(stage, n_microbatches=8)
# Run
if rank == 0:
schedule.step(input_data)
elif rank == num_stages - 1:
losses = schedule.step()
else:
schedule.step()
Combining Strategies: 3D Parallelism
For training the largest models (100B+ parameters), you combine multiple parallelism strategies. The standard combination is 3D parallelism: Data Parallel (FSDP) × Tensor Parallel × Pipeline Parallel.
DeviceMesh for 3D Parallelism
# 64 GPUs across 8 nodes, 8 GPUs per node
# 8 DP × 4 TP × 2 PP
mesh = init_device_mesh(
"cuda", (8, 4, 2), mesh_dim_names=("dp", "tp", "pp")
)
Typical Assignment
- TP within a node: TP requires fast interconnect, so TP groups are
within a single node (connected by NVLink).
- PP across nodes: PP has less communication (only activations at stage
boundaries), so it can span nodes.
- FSDP across remaining GPUs: FSDP handles the data parallelism dimension.
Node 0: [GPU0, GPU1, GPU2, GPU3] ← TP group, PP stage 0
Node 1: [GPU4, GPU5, GPU6, GPU7] ← TP group, PP stage 1
...
Across nodes: FSDP groups
Code Pattern for 3D Parallelism
mesh = init_device_mesh(
"cuda", (dp_size, tp_size, pp_size),
mesh_dim_names=("dp", "tp", "pp"),
)
# 1. Apply Tensor Parallelism first (innermost)
tp_mesh = mesh["tp"]
for layer in model.layers:
parallelize_module(layer, tp_mesh, tp_plan)
# 2. Apply FSDP (middle)
dp_mesh = mesh["dp"]
for layer in model.layers:
fully_shard(layer, mesh=dp_mesh)
fully_shard(model, mesh=dp_mesh)
# 3. Apply Pipeline Parallelism (outermost)
pp_mesh = mesh["pp"]
# Split model into stages along pp_mesh
Distributed Checkpointing (DCP)
When training with FSDP, TP, or other distributed strategies, each rank holds only a shard of the model. Distributed Checkpointing (DCP) handles saving and loading these sharded states correctly.
Why Not torch.save?
torch.save requires gathering the full model to one rank, which:
- May not fit in memory for very large models
- Creates a bottleneck (one rank doing all the work)
- Produces a format tied to the original parallelism configuration
DCP saves each rank's shard independently and can reshard when loading with a different parallelism configuration.
Basic Save and Load
import torch.distributed.checkpoint as dcp
# Save
state_dict = {"model": model.state_dict(), "optimizer": optimizer.state_dict()}
dcp.save(state_dict, checkpoint_id="checkpoints/step_1000")
# Load
state_dict = {"model": model.state_dict(), "optimizer": optimizer.state_dict()}
dcp.load(state_dict, checkpoint_id="checkpoints/step_1000")
model.load_state_dict(state_dict["model"])
optimizer.load_state_dict(state_dict["optimizer"])
Async Save
For large models, checkpointing can take minutes. Async save moves the checkpoint writing to a background thread so training can continue:
# Async save returns a Future
future = dcp.async_save(state_dict, checkpoint_id="checkpoints/step_1000")
# Training continues immediately...
# Optionally wait for completion before the next checkpoint
future.result()
get_model_state_dict / set_model_state_dict
These utilities handle the complexity of getting/setting state dicts for models wrapped with FSDP, TP, etc.:
from torch.distributed.checkpoint.state_dict import (
get_model_state_dict,
set_model_state_dict,
get_optimizer_state_dict,
set_optimizer_state_dict,
StateDictOptions,
)
# Get a "clean" state dict (handles FSDP/TP unwrapping)
model_state = get_model_state_dict(model)
optim_state = get_optimizer_state_dict(model, optimizer)
# Save
dcp.save({"model": model_state, "optim": optim_state}, checkpoint_id=path)
# Load
state = {"model": model_state, "optim": optim_state}
dcp.load(state, checkpoint_id=path)
set_model_state_dict(model, state["model"])
set_optimizer_state_dict(model, optimizer, state["optim"])
HuggingFace Format
DCP can save in HuggingFace-compatible format for interoperability:
from torch.distributed.checkpoint import HuggingFaceLoadPlanner
# Load a HuggingFace checkpoint
dcp.load(
state_dict,
checkpoint_id="path/to/hf_checkpoint",
planner=HuggingFaceLoadPlanner(),
)
SymmetricMemory
SymmetricMemory is an intra-node optimization that provides direct GPU-to-GPU memory access using NVLink, bypassing the traditional collective communication libraries.
What is SymmetricMemory?
On a multi-GPU node with NVLink, GPUs can directly read/write each other's memory. SymmetricMemory allocates a shared memory region accessible by all GPUs in a group, enabling custom, low-latency communication patterns.
import torch.distributed.symmetric_memory as sm
# Allocate symmetric memory (same virtual address on all GPUs)
t = sm.empty_strided_p2p(
size=(1024, 1024),
stride=(1024, 1),
dtype=torch.float32,
device=torch.device(f"cuda:{local_rank}"),
)
# Direct GPU-to-GPU operations
sm.memcpy_p2p(dst=t_on_gpu1, src=t_on_gpu0)
When to Use
- Custom all-reduce implementations that exploit NVLink topology
- Fine-grained producer-consumer patterns between GPUs
- Overlapping communication with computation at a granularity finer than
what standard collectives offer
- Typically used in advanced performance optimization, not in everyday training
Context Parallel
Context Parallelism (CP) addresses the challenge of training with very long sequences. When the sequence length is so large that a single GPU cannot hold the activations for one sequence, CP splits the sequence across GPUs.
How It Works
CP distributes the sequence dimension across GPUs in a process group. For attention computation, this requires specialized communication because each position needs to attend to all other positions:
Sequence: [token_0, token_1, ..., token_8191]
GPU 0: [token_0 ... token_2047]
GPU 1: [token_2048 ... token_4095]
GPU 2: [token_4096 ... token_6143]
GPU 3: [token_6144 ... token_8191]
For attention: GPU 0 needs KV from all GPUs → ring attention
Ring Attention
Ring attention is a common CP implementation where KV pairs are passed around in a ring. Each GPU computes attention with its local Q and the received KV, then passes KV to the next GPU:
Step 1: GPU0 attends to KV₀, GPU1 to KV₁, ...
Step 2: GPUs pass KV to neighbor: GPU0 gets KV₃, GPU1 gets KV₀, ...
Step 3: Repeat until all KV seen by all GPUs
When to Use
- Sequence lengths > 8K-16K tokens (depends on model size and GPU memory)
- Long-document training, video models, genomics
- Often combined with TP and FSDP
Practical Advice
Choosing a Parallelism Strategy
Start: Does the model fit on 1 GPU with your batch size?
│
├── YES → Use DDP
│ Still too slow? → Increase DDP world size
│
└── NO → Does the model fit on 1 GPU (batch_size=1)?
│
├── YES → Use FSDP2
│ Want even faster? → FSDP2 + TP
│
└── NO → Use FSDP2 + TP
Still doesn't fit? → Add PP
Very long sequences? → Add CP
Rules of Thumb
- Start with DDP if your model fits on one GPU. It's the simplest and
most efficient.
- Move to FSDP2 when memory is the bottleneck. FSDP reduces per-GPU
memory at the cost of extra communication.
- Add TP when individual layers are very large (hidden_dim > 4096) or
when you need to reduce per-GPU memory further. Keep TP within a node.
- Add PP when you have many GPUs across nodes and want to reduce
cross-node communication. PP only communicates activations at stage boundaries.
- Use CP specifically for long-sequence training where sequence
activations dominate memory.
- Match TP to NVLink topology: Use TP degree 2, 4, or 8 matching the
NVLink connectivity within your node.
- More micro-batches reduce PP bubble: For pipeline parallelism, use
at least 2-4× as many micro-batches as pipeline stages.
Common Pitfalls
- Forgetting
sampler.set_epoch(epoch): Causes same data order every
epoch with DDP/FSDP.
- Not sharding submodules before root in FSDP2: Always apply
fully_shard
to submodules first, then the root module.
- TP across nodes: TP requires fast interconnect. Putting a TP group
across nodes with only InfiniBand (instead of NVLink) kills performance.
- Saving checkpoints on all ranks: Use rank 0 for simple saves, or DCP
for distributed saves.
- Not using
no_sync()for gradient accumulation: When accumulating
gradients across multiple steps, wrap forward/backward in model.no_sync() to skip all-reduce on intermediate steps.
Gradient Accumulation with DDP/FSDP
accumulation_steps = 4
for i, (data, target) in enumerate(dataloader):
# Use no_sync for intermediate steps to avoid wasteful all-reduce
context = model.no_sync() if (i + 1) % accumulation_steps != 0 else nullcontext()
with context:
output = model(data)
loss = loss_fn(output, target) / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
Files in This Module
| File | Description | Run Command |
|---|---|---|
concepts_and_collectives.py | Collective operations on CPU with Gloo | torchrun --nproc_per_node=3 concepts_and_collectives.py |
ddp_example.py | Complete DDP training with synthetic data | torchrun --nproc_per_node=2 ddp_example.py |
fsdp2_example.py | FSDP2 API patterns and setup | torchrun --nproc_per_node=2 fsdp2_example.py |
device_mesh_example.py | DeviceMesh creation and sub-mesh access | torchrun --nproc_per_node=4 device_mesh_example.py |
parallelism_overview.py | API patterns for TP and PP | Reference script showing code patterns |
📓 Open Notebook — Interactive version of this module
Source Files
concepts_and_collectives.py— Distributed concepts and collective operations (CPU/Gloo)ddp_example.py— DistributedDataParallel (DDP) complete training examplefsdp2_example.py— FSDP2 (fully_shard) API patternsdevice_mesh_example.py— DeviceMesh creation and usageparallelism_overview.py— Parallelism overview: TP, PP, and combined strategies
Module 11: Export and Deployment
Table of Contents
- Why Export?
- The PyTorch Deployment Landscape
- torch.export
- Static vs Dynamic Shapes
- draft_export
- Saving and Loading
- Graph Inspection
- AOTInductor
- NativeRT
- ONNX Export
- TorchServe
- ExecuTorch
- Quantization for Deployment
- Practical Workflow
Why Export?
During research, you run PyTorch in eager mode: Python executes operations one at a time, giving you full flexibility (print statements, breakpoints, dynamic control flow). But production deployment has different requirements:
- No Python dependency: Servers, mobile devices, and embedded systems may
not have (or want) a Python runtime.
- Performance: Compiler optimizations (operator fusion, memory planning,
kernel selection) require seeing the full computation graph ahead of time.
- Predictability: Production systems need deterministic latency and memory
usage. Eager mode's dynamism makes this hard to guarantee.
- Portability: The same model may run on server GPUs, edge TPUs, mobile
phones, or web browsers.
Export bridges the gap: it captures your Python model as a self-contained computation graph that can be optimized and deployed without Python.
Research (Python, eager) → Export (capture graph) → Deploy (optimized, no Python)
The PyTorch Deployment Landscape
PyTorch provides multiple paths from research to production:
┌──────────────┐
│ Your Model │
│ (nn.Module) │
└──────┬───────┘
│
┌──────┴───────┐
│ torch.export │
│ (capture │
│ full graph) │
└──────┬───────┘
│
┌──────────┬───────┼────────┬──────────┐
▼ ▼ ▼ ▼ ▼
┌──────────┐ ┌───────┐ ┌────┐ ┌───────┐ ┌─────────┐
│AOTInductor│ │NativeRT│ │ONNX│ │Torch- │ │ExecuTorch│
│(.so lib) │ │(C++ │ │ │ │Serve │ │(mobile/ │
│ │ │engine)│ │ │ │ │ │edge) │
└──────────┘ └───────┘ └────┘ └───────┘ └─────────┘
Server Server Any Server Mobile/
(C++/Py) (C++) runtime (Py) Embedded
| Path | Output | Target | Key Advantage |
|---|---|---|---|
| AOTInductor | .so shared library | Server (C++ or Python) | Maximum performance, no Python needed |
| NativeRT | Serialized model | Server (C++) | C++ inference engine, easy deployment |
| ONNX | .onnx file | Any ONNX Runtime | Cross-framework interop |
| TorchServe | Model archive | Server (Python) | Full serving stack (batching, scaling) |
| ExecuTorch | .pte file | Mobile/edge devices | Small footprint, on-device inference |
torch.export
torch.export is the primary tool for capturing a PyTorch model as a complete, self-contained graph. It traces through your model and produces an ExportedProgram containing the full computation graph with no Python dependency.
Basic Export
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(10, 5)
self.relu = nn.ReLU()
def forward(self, x):
return self.relu(self.linear(x))
model = MyModel()
example_input = (torch.randn(3, 10),)
# Export the model
exported = torch.export.export(model, example_input)
# The exported program can be called like the original
result = exported.module()(torch.randn(3, 10))
What ExportedProgram Contains
An ExportedProgram captures:
- Graph Module: An
fx.GraphModulecontaining the computation graph as
a series of operations (ATen operators).
- Graph Signature: Maps graph inputs/outputs to parameters, buffers,
and user inputs.
- State Dict: The model's parameters and buffers.
- Range Constraints: Valid ranges for dynamic dimensions.
- Module Call graph: Preserves the module hierarchy for debugging.
exported = torch.export.export(model, example_input)
# Access the graph
print(exported.graph_module.graph)
# Access parameters
print(exported.state_dict.keys())
# The graph shows ATen-level operations
for node in exported.graph_module.graph.nodes:
print(f" {node.op}: {node.target}")
What Gets Captured
torch.export traces through your model's forward method and captures:
- All tensor operations as ATen ops
- Control flow (with restrictions — see below)
- Shape computations
It does NOT capture:
- Print statements or side effects
- Operations on non-tensor values that aren't related to shapes
- Dynamic Python control flow based on tensor values (use
torch.condinstead)
Handling Control Flow
Standard if/else based on tensor values won't work:
# This will fail with torch.export:
def forward(self, x):
if x.sum() > 0: # Can't branch on tensor value!
return x * 2
return x * 3
# Use torch.cond instead:
def forward(self, x):
return torch.cond(
x.sum() > 0,
lambda x: x * 2,
lambda x: x * 3,
(x,),
)
Control flow based on tensor shapes (not values) is fine because shapes are known at export time (or constrained for dynamic shapes):
# This is fine:
def forward(self, x):
if x.shape[0] > 5: # Shape is known at export time
return x[:5]
return x
Static vs Dynamic Shapes
By default, torch.export traces with static shapes: the exported model only accepts inputs with the exact same shapes as the example inputs. For production, you usually want dynamic shapes so the model handles variable batch sizes, sequence lengths, etc.
The Dim API
Use torch.export.Dim to declare which dimensions are dynamic:
from torch.export import Dim, export
# Declare a dynamic dimension named "batch"
batch = Dim("batch", min=1, max=128)
# Export with dynamic first dimension
exported = export(
model,
(torch.randn(4, 10),), # Example with batch=4
dynamic_shapes={"x": {0: batch}}, # dim 0 is dynamic
)
# Now works with any batch size from 1 to 128:
exported.module()(torch.randn(1, 10)) # batch=1
exported.module()(torch.randn(64, 10)) # batch=64
exported.module()(torch.randn(128, 10)) # batch=128
Multiple Dynamic Dimensions
batch = Dim("batch", min=1, max=256)
seq_len = Dim("seq_len", min=1, max=2048)
exported = export(
model,
(torch.randn(4, 128, 512),),
dynamic_shapes={"x": {0: batch, 1: seq_len}},
# dim 0 = batch (dynamic), dim 1 = seq_len (dynamic), dim 2 = 512 (static)
)
Constraints Between Dimensions
When multiple inputs share a dimension (e.g., same batch size), use the same Dim object:
batch = Dim("batch", min=1, max=128)
def forward(self, x, y):
return x + y # x and y must have same batch size
exported = export(
model,
(torch.randn(4, 10), torch.randn(4, 20)),
dynamic_shapes={
"x": {0: batch},
"y": {0: batch}, # Same Dim → enforces same batch size
},
)
Automatic Dynamic Shapes
For convenience, you can let PyTorch infer dynamic shapes:
from torch.export import export
exported = export(
model,
(torch.randn(4, 10),),
dynamic_shapes={"x": {0: Dim.AUTO}},
)
draft_export
When torch.export fails (due to unsupported Python constructs, dynamic control flow, etc.), draft_export helps you debug by producing a best- effort export with detailed error information.
from torch.export import draft_export
# If regular export fails:
try:
exported = torch.export.export(model, example_input)
except Exception as e:
print(f"Export failed: {e}")
# Use draft_export to get a partial graph + diagnostics
ep, report = draft_export(model, example_input)
# The report shows what went wrong and how to fix it
print(report)
draft_export returns:
- An
ExportedProgram(possibly with graph breaks or approximations) - A report detailing what couldn't be captured and suggesting fixes
Common fixes suggested by draft_export:
- Replace
if tensor_val > 0withtorch.cond - Replace
for i in range(tensor.shape[0])with bounded loop - Mark data-dependent shapes with
torch.export.Dim
Saving and Loading
PT2 Archive Format
The recommended format for saving exported programs is the PT2 Archive:
import torch
# Export
model = MyModel()
exported = torch.export.export(model, (torch.randn(1, 10),))
# Save as PT2 archive
torch.export.save(exported, "model.pt2")
# Load (no need for original model code!)
loaded = torch.export.load("model.pt2")
result = loaded.module()(torch.randn(1, 10))
The PT2 archive contains:
- The computation graph (serialized FX graph)
- Model weights (state dict)
- Metadata (dynamic shape constraints, etc.)
This is a self-contained format. You do not need the original Python model definition to load and run the model.
Saving for Different Backends
# Save for later AOTInductor compilation
torch.export.save(exported, "model.pt2")
# Save for ONNX (different path)
torch.onnx.export(model, example_input, "model.onnx", dynamo=True)
Versioning and Compatibility
PT2 archives are versioned. PyTorch maintains backward compatibility: a model saved with an older PyTorch version can be loaded with a newer version (within the same major version).
Graph Inspection
After exporting, you can inspect the computation graph to understand what was captured and verify correctness.
Viewing the Graph
exported = torch.export.export(model, example_input)
# Print the graph (human-readable)
print(exported.graph_module.graph)
# Print generated code (more readable)
print(exported.graph_module.code)
Walking the Graph Nodes
Each node in the graph represents an operation:
for node in exported.graph_module.graph.nodes:
print(f"Op: {node.op:15s} | Target: {node.target} | Args: {node.args}")
Node types:
placeholder: Input to the graph (parameters, buffers, user inputs)call_function: A function call (ATen operator)output: The graph's return valueget_attr: Access a stored attribute
Listing All Operations
# Get the set of all ATen ops used in the graph
ops = set()
for node in exported.graph_module.graph.nodes:
if node.op == "call_function":
ops.add(str(node.target))
print(f"Operations used ({len(ops)}):")
for op in sorted(ops):
print(f" {op}")
Understanding the Graph Signature
sig = exported.graph_signature
# What are the inputs?
print("User inputs:", sig.input_specs)
# What are the outputs?
print("Outputs:", sig.output_specs)
# Parameters and buffers
print("Parameters:", [s for s in sig.input_specs if s.kind.name == "PARAMETER"])
AOTInductor
AOTInductor (Ahead-Of-Time Inductor) compiles an exported model into a native shared library (.so on Linux, .dylib on macOS). The result runs without Python, with maximum performance.
Compilation Flow
ExportedProgram → AOTInductor → .so shared library → C++ or Python inference
Python API
import torch
model = MyModel().eval()
example_input = (torch.randn(1, 3, 224, 224),)
# Export
exported = torch.export.export(model, example_input)
# Compile to .so
compiled_path = torch._inductor.aot_compile(
exported.module(),
example_input,
options={"aot_inductor.output_path": "model.so"},
)
# Load and run in Python (for testing)
compiled_model = torch._inductor.aot_load(compiled_path)
result = compiled_model(torch.randn(1, 3, 224, 224))
Package API (Recommended)
The package API bundles the model, weights, and metadata:
# Compile and package
torch._inductor.aoti_compile_and_package(
exported,
package_path="model_package.pt2",
)
# Load the package
loaded = torch._inductor.aoti_load_package("model_package.pt2")
result = loaded(torch.randn(1, 3, 224, 224))
C++ Deployment
The primary use case for AOTInductor is deploying without Python:
#include <torch/csrc/inductor/aoti_runner/model_container_runner.h>
int main() {
// Load the compiled model
auto runner = torch::inductor::AOTIModelContainerRunner("model.so");
// Create input tensor
auto input = torch::randn({1, 3, 224, 224});
std::vector<torch::Tensor> inputs = {input};
// Run inference
auto outputs = runner.run(inputs);
auto result = outputs[0];
return 0;
}
When to Use AOTInductor
- Maximum inference performance (compiled, optimized kernels)
- Deployment without Python
- Server-side GPU inference
- When you can compile ahead of time (shapes known in advance)
NativeRT
NativeRT is a C++ inference engine for running PyTorch models. While AOTInductor compiles to native code, NativeRT interprets the exported graph using PyTorch's C++ runtime.
NativeRT vs AOTInductor
| Feature | AOTInductor | NativeRT |
|---|---|---|
| Compilation | Ahead of time (slow) | No compilation |
| Inference speed | Fastest (native code) | Fast (C++ runtime) |
| Startup time | Fast (pre-compiled) | Fast (load + interpret) |
| Flexibility | Fixed graph | More flexible |
| Python needed | No | No |
When to Use NativeRT
- When AOTInductor compilation is too slow or complex
- When you need C++ inference without ahead-of-time compilation
- Quick prototyping of C++ deployment
Usage Pattern
# Export and save
exported = torch.export.export(model, example_input)
torch.export.save(exported, "model.pt2")
# In C++ with NativeRT:
# auto runner = torch::nativert::ModelRunner("model.pt2");
# auto output = runner.run(inputs);
ONNX Export
ONNX (Open Neural Network Exchange) is an open format for representing ML models. PyTorch can export to ONNX for running on ONNX Runtime, which supports multiple hardware backends.
Modern ONNX Export (Dynamo-based)
import torch
model = MyModel().eval()
example_input = (torch.randn(1, 3, 224, 224),)
# Export to ONNX using the new dynamo-based exporter
onnx_program = torch.onnx.export(
model,
example_input,
dynamo=True,
)
# Save to file
onnx_program.save("model.onnx")
Running with ONNX Runtime
import onnxruntime as ort
import numpy as np
session = ort.InferenceSession("model.onnx")
# Get input name
input_name = session.get_inputs()[0].name
# Run inference
input_data = np.random.randn(1, 3, 224, 224).astype(np.float32)
result = session.run(None, {input_name: input_data})
Dynamic Shapes with ONNX
from torch.export import Dim
batch = Dim("batch", min=1, max=128)
onnx_program = torch.onnx.export(
model,
(torch.randn(4, 3, 224, 224),),
dynamo=True,
dynamic_shapes={"x": {0: batch}},
)
onnx_program.save("model_dynamic.onnx")
When to Use ONNX
- Cross-framework deployment (model trained in PyTorch, deployed with
TensorFlow Serving, etc.)
- Hardware with ONNX Runtime support but not PyTorch C++ runtime
- When inference hardware vendor provides an ONNX Runtime backend
TorchServe
TorchServe is PyTorch's model serving framework. It handles the operational concerns of serving models in production: batching, scaling, monitoring, and A/B testing.
Overview
Client → TorchServe ┬→ Worker 1 → Model
(HTTP/gRPC) ├→ Worker 2 → Model
└→ Worker N → Model
Key Features
- Dynamic batching: Aggregates individual requests into batches for
efficient GPU utilization
- Multi-model serving: Serve multiple models from one instance
- Model versioning: A/B testing, canary deployments
- Monitoring: Prometheus metrics, logging
- Auto-scaling: Scale workers based on load
Basic Workflow
# 1. Archive the model
torch-model-archiver \
--model-name my_model \
--version 1.0 \
--serialized-file model.pt \
--handler image_classifier \
--export-path model_store
# 2. Start TorchServe
torchserve --start \
--model-store model_store \
--models my_model=my_model.mar
# 3. Send requests
curl http://localhost:8080/predictions/my_model -T input.jpg
ExecuTorch
ExecuTorch is PyTorch's solution for on-device inference on mobile phones, wearables, and embedded systems. It produces small, efficient models that run without a full PyTorch runtime.
Key Features
- Small footprint: Runtime is ~100s of KB (vs PyTorch's ~100s of MB)
- Delegate system: Hardware-specific optimizations (Core ML, XNNPACK,
Qualcomm QNN, etc.)
- No Python: Runs on C++ runtime
Workflow
import torch
from executorch.exir import to_edge_transform_and_lower
model = MyModel().eval()
example_input = (torch.randn(1, 3, 224, 224),)
# Export
exported = torch.export.export(model, example_input)
# Lower to edge
edge_program = to_edge_transform_and_lower(exported)
# Save
edge_program.save("model.pte")
Target Platforms
| Platform | Delegate |
|---|---|
| iOS | Core ML, Metal |
| Android | XNNPACK, Qualcomm QNN |
| Microcontrollers | Custom delegates |
| Web | WebAssembly |
Quantization for Deployment
Quantization reduces model size and increases inference speed by using lower precision arithmetic (e.g., INT8 instead of FP32).
PT2E Quantization Flow
The modern quantization flow is built on torch.export:
import torch
from torch.ao.quantization.quantize_pt2e import (
prepare_pt2e,
convert_pt2e,
)
from torch.ao.quantization.quantizer.xnnpack_quantizer import (
XNNPACKQuantizer,
get_symmetric_quantization_config,
)
model = MyModel().eval()
example_input = (torch.randn(1, 3, 224, 224),)
# Step 1: Export
exported = torch.export.export(model, example_input)
# Step 2: Prepare for quantization (inserts observers)
quantizer = XNNPACKQuantizer().set_global(
get_symmetric_quantization_config()
)
prepared = prepare_pt2e(exported, quantizer)
# Step 3: Calibrate with representative data
with torch.no_grad():
for data in calibration_dataloader:
prepared(data)
# Step 4: Convert to quantized model
quantized = convert_pt2e(prepared)
# Step 5: Deploy (e.g., with AOTInductor or ExecuTorch)
Quantization Types
| Type | Precision | Speed | Accuracy | Use Case |
|---|---|---|---|---|
| FP32 | 32-bit | Baseline | Best | Training |
| FP16 | 16-bit | ~2× | Negligible loss | GPU inference |
| BF16 | 16-bit | ~2× | Negligible loss | GPU inference |
| INT8 | 8-bit | ~2-4× | Small loss | Server inference |
| INT4 | 4-bit | ~4-8× | Moderate loss | Edge/mobile, LLMs |
Dynamic vs Static Quantization
- Static quantization: Calibrate ranges with representative data. Best
accuracy. Requires calibration dataset.
- Dynamic quantization: Compute ranges at runtime. Simpler setup. Slightly
worse performance on some workloads.
Practical Workflow
The End-to-End Journey
1. RESEARCH & TRAINING
├── Train in eager mode (nn.Module, autograd)
├── Validate accuracy
└── Save checkpoint
2. EXPORT
├── torch.export.export(model, example_inputs)
├── Add dynamic shapes for variable inputs
├── Fix export issues (torch.cond for control flow, etc.)
├── Use draft_export to debug failures
└── Save: torch.export.save(exported, "model.pt2")
3. OPTIMIZE
├── Quantize (PT2E flow) for smaller/faster model
├── Profile to identify bottlenecks
└── Choose deployment target
4. DEPLOY
├── Server GPU → AOTInductor (.so) or NativeRT
├── Server CPU → ONNX Runtime or AOTInductor
├── Model serving → TorchServe
├── Mobile/Edge → ExecuTorch (.pte)
└── Cross-framework → ONNX
Common Patterns
Pattern 1: Quick Server Deployment
model = load_trained_model()
exported = torch.export.export(model, example_input)
torch.export.save(exported, "model.pt2")
# Load in production with torch.export.load("model.pt2")
Pattern 2: Maximum Performance Server
model = load_trained_model()
exported = torch.export.export(model, example_input)
torch._inductor.aoti_compile_and_package(exported, "model_pkg.pt2")
# Deploy .so or package without Python
Pattern 3: Mobile Deployment
model = load_trained_model()
exported = torch.export.export(model, example_input)
# Quantize → Lower to edge → Save .pte
Pattern 4: Cross-Platform
model = load_trained_model()
onnx_program = torch.onnx.export(model, example_input, dynamo=True)
onnx_program.save("model.onnx")
# Deploy with ONNX Runtime on any platform
Debugging Export Failures
- Start with draft_export: Get a partial graph and diagnostic report.
- Check control flow: Replace data-dependent branches with
torch.cond. - Check dynamic shapes: Use
Dimfor variable dimensions. - Simplify: Export a smaller part of the model first to isolate issues.
- Check operators: Some custom ops may need registration for export.
Performance Comparison (Rough Guidelines)
| Deployment Path | Relative Latency | Setup Effort |
|---|---|---|
| Eager (Python) | 1.0× (baseline) | None |
| torch.compile | 0.5-0.8× | One line |
| AOTInductor | 0.3-0.6× | Moderate |
| Quantized (INT8) | 0.2-0.4× | Significant |
| ExecuTorch (mobile) | Varies by hardware | Significant |
Files in This Module
| File | Description | Run Command |
|---|---|---|
export_basics.py | Basic export, run exported model | python export_basics.py |
dynamic_shapes.py | Dim API, multiple dynamic dims | python dynamic_shapes.py |
export_inspection.py | Inspect the graph, list ops | python export_inspection.py |
save_and_load.py | Save/load PT2 archives | python save_and_load.py |
📓 Open Notebook — Interactive version of this module
Source Files
export_basics.py— torch.export basics — fundamentals of exporting modelsdynamic_shapes.py— Dynamic shapes with torch.export — the Dim API for variable input sizesexport_inspection.py— Exported graph inspection — examining computation graphs after exportsave_and_load.py— Saving and loading exported models (PT2 archives)
Module 12: Model Architectures — From Paper to PyTorch
Building real neural network architectures is the bridge between understanding PyTorch basics and doing real deep learning research or engineering. This module walks through several landmark architectures, explaining the why behind each design choice and showing you how to translate ideas from papers into working PyTorch code.
How to Read an Architecture Paper and Translate It to Code
Most deep learning papers follow a predictable structure. Learning to read them systematically will save you enormous amounts of time.
Step 1: Understand the Problem Statement
Before looking at the architecture, understand what problem the paper solves. For ResNet, it's the "degradation problem" — deeper networks were performing worse than shallower ones, even on training data. For Transformers, it was the sequential bottleneck of RNNs that prevented parallelization.
Step 2: Identify the Core Innovation
Every architecture paper has one or two key ideas. Everything else is engineering around those ideas:
| Paper | Core Innovation |
|---|---|
| ResNet | Skip (residual) connections |
| Transformer | Scaled dot-product self-attention |
| GPT | Decoder-only Transformer + autoregressive pretraining |
| ViT | Treat image patches as token sequences |
| VAE | Reparameterization trick for differentiable sampling |
| U-Net | Encoder-decoder with skip connections for dense prediction |
Step 3: Map the Architecture Diagram to nn.Modules
Papers usually have an architecture diagram. Each box becomes either:
- An
nn.Modulesubclass (if it's a reusable block) - A line of code inside
forward()(if it's a simple operation)
Step 4: Match Dimensions
Papers describe tensor shapes. Track them through the network. A common approach: write comments with shapes at each step in your forward() method during development, then remove them once the code is tested.
Step 5: Implement, Test with Random Data, Then Train
Always verify your architecture with random tensors before training:
model = MyModel()
x = torch.randn(2, 3, 224, 224) # batch=2, channels=3, 224x224
out = model(x)
print(out.shape) # should be (2, num_classes)
ResNet (Residual Networks)
Paper: "Deep Residual Learning for Image Recognition" (He et al., 2015)
The Degradation Problem
Before ResNet, researchers observed a paradox: adding more layers to a neural network made it worse, even on the training set. This wasn't overfitting — it was a fundamental optimization problem. Deeper networks were harder to optimize because gradients had to flow through many layers.
The Key Insight: Skip Connections
Instead of learning H(x) directly, learn the residual F(x) = H(x) - x, then compute H(x) = F(x) + x. If the optimal transformation is close to identity, it's easier to learn a small residual than to learn identity from scratch.
Input x ──────────────────────┐
│ │
▼ │
┌──────────┐ │
│ Conv-BN │ │
│ ReLU │ │
│ Conv-BN │ │
└──────────┘ │
│ │
▼ │
F(x) ─────── + ◄──────────┘
│
▼
ReLU
│
▼
Output = ReLU(F(x) + x)
BasicBlock vs Bottleneck
ResNet uses two block types:
BasicBlock (for ResNet-18 and ResNet-34):
- Two 3x3 convolutions
- Each followed by BatchNorm
- A skip connection that adds the input to the output
- If dimensions don't match, a 1x1 conv "projection" shortcut is used
class BasicBlock(nn.Module):
expansion = 1
def __init__(self, in_channels, out_channels, stride=1, downsample=None):
super().__init__()
self.conv1 = nn.Conv2d(in_channels, out_channels, 3,
stride=stride, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, 3,
stride=1, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(out_channels)
self.downsample = downsample
def forward(self, x):
identity = x
out = F.relu(self.bn1(self.conv1(x)))
out = self.bn2(self.conv2(out))
if self.downsample is not None:
identity = self.downsample(x)
out += identity
return F.relu(out)
Bottleneck (for ResNet-50, 101, 152):
- Three convolutions: 1x1 (reduce), 3x3 (process), 1x1 (expand)
- The 1x1 convolutions reduce and restore channel dimensions
- This "bottleneck" design is more parameter-efficient for deep networks
The expansion factor controls how much the Bottleneck expands channels:
- BasicBlock: expansion = 1 (output channels == internal channels)
- Bottleneck: expansion = 4 (output channels == 4 * internal channels)
Pre-Activation ResNet
The original ResNet applies BN and ReLU after each convolution. A later paper ("Identity Mappings in Deep Residual Networks", He et al., 2016) showed that applying BN and ReLU before the convolution (pre-activation) improves both optimization and generalization:
# Post-activation (original): Conv -> BN -> ReLU
# Pre-activation (improved): BN -> ReLU -> Conv
The intuition: in the pre-activation design, the skip connection is a true identity mapping (no BN or ReLU on the shortcut path), allowing gradients to flow unimpeded.
ResNet Family Configurations
| Model | Block | Layers per stage | Total layers | Parameters |
|---|---|---|---|---|
| ResNet-18 | BasicBlock | [2, 2, 2, 2] | 18 | ~11M |
| ResNet-34 | BasicBlock | [3, 4, 6, 3] | 34 | ~21M |
| ResNet-50 | Bottleneck | [3, 4, 6, 3] | 50 | ~25M |
| ResNet-101 | Bottleneck | [3, 4, 23, 3] | 101 | ~44M |
See resnet.py for the complete implementation.
Transformer
Paper: "Attention Is All You Need" (Vaswani et al., 2017)
Motivation
RNNs process sequences one token at a time — you can't compute step t until step t-1 is done. This sequential bottleneck limits parallelism and makes it hard to learn long-range dependencies. The Transformer replaces recurrence entirely with attention mechanisms.
Scaled Dot-Product Attention
The core operation. Given queries Q, keys K, and values V:
Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
- Q @ K^T computes similarity scores between every query and every key
- Division by sqrt(d_k) prevents dot products from becoming too large (which
would push softmax into saturation, giving near-zero gradients)
- Softmax normalizes scores to a probability distribution
- Multiplication with V produces a weighted combination of values
Multi-Head Attention
Instead of one big attention computation, split Q, K, V into h heads:
# Instead of d_model-dimensional attention:
# Split into h heads, each of dimension d_k = d_model // h
# Compute attention independently per head
# Concatenate and project back to d_model
Why multiple heads? Each head can attend to different aspects of the input — one head might focus on syntactic relationships, another on semantic ones.
Self-Attention vs Cross-Attention
- Self-attention: Q, K, V all come from the same sequence. Each token
attends to all other tokens in the same sequence.
- Cross-attention: Q comes from one sequence (decoder), K and V come from
another (encoder output). This is how the decoder "reads" the encoder.
Positional Encoding
Attention is permutation-invariant — it has no notion of order. Positional encodings inject position information. The original paper uses sinusoidal encodings:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
These have a nice property: the encoding of position pos+k can be expressed as a linear function of the encoding of position pos, making it possible for the model to learn relative positions.
Encoder-Decoder Architecture
Encoder: Stack of N identical layers, each containing:
- Multi-head self-attention (with residual + LayerNorm)
- Position-wise feed-forward network (with residual + LayerNorm)
Decoder: Stack of N identical layers, each containing:
- Masked multi-head self-attention (causal mask prevents attending to future)
- Multi-head cross-attention (attends to encoder output)
- Position-wise feed-forward network
Each sub-layer uses the pattern: LayerNorm(x + Sublayer(x)) (post-norm) or x + Sublayer(LayerNorm(x)) (pre-norm, more common now).
See transformer.py for the complete implementation.
GPT (Decoder-Only Transformer)
Paper: "Language Models are Unsupervised Multitask Learners" (Radford et al., 2019)
Design Philosophy
GPT simplifies the Transformer by keeping only the decoder (with causal attention), removing the encoder and cross-attention entirely. The insight: a powerful enough language model trained to predict the next token can learn to perform many tasks without explicit task-specific architecture.
Causal (Autoregressive) Attention
In a decoder-only model, each token can only attend to itself and previous tokens. This is enforced by a triangular mask:
# For sequence length 4:
mask = [[1, 0, 0, 0], # token 0 sees only itself
[1, 1, 0, 0], # token 1 sees tokens 0-1
[1, 1, 1, 0], # token 2 sees tokens 0-2
[1, 1, 1, 1]] # token 3 sees tokens 0-3
Positions with 0 are set to -infinity before softmax, effectively zeroing those attention weights.
Weight Tying
GPT ties the token embedding matrix with the output projection (language model head). If embedding maps token IDs to vectors, the LM head maps vectors back to logits over the vocabulary — these are inverse operations, so sharing weights makes sense and reduces parameters significantly:
self.token_embedding = nn.Embedding(vocab_size, d_model)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
self.lm_head.weight = self.token_embedding.weight # weight tying
Pre-Norm vs Post-Norm
Original Transformer: x + Sublayer(LayerNorm(x)) — "post-norm" GPT-2 and later: x + Sublayer(LayerNorm(x)) — "pre-norm"
Wait, the formulas look the same? The difference is subtle:
- Post-norm:
LayerNorm(x + Sublayer(x))— norm is outside the residual - Pre-norm:
x + Sublayer(LayerNorm(x))— norm is inside the residual
Pre-norm is more stable for training deep networks because the residual connection carries un-normalized values, preserving gradient magnitude.
Generation Strategies
Given a trained model, how do you generate text?
Greedy: Always pick the most probable next token. Fast but repetitive.
Temperature: Divide logits by temperature T before softmax:
- T < 1: sharper distribution, more confident (less random)
- T = 1: original distribution
- T > 1: flatter distribution, more random
Top-k sampling: Keep only the top k most probable tokens, zero out the rest, renormalize, then sample. Prevents sampling very unlikely tokens.
Top-p (nucleus) sampling: Keep the smallest set of tokens whose cumulative probability exceeds p. Adapts the number of candidates dynamically — when the model is confident, few tokens are kept; when uncertain, more are kept.
def top_p_sample(logits, p=0.9):
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
# Remove tokens with cumulative probability above the threshold
mask = cumulative_probs - F.softmax(sorted_logits, dim=-1) >= p
sorted_logits[mask] = float('-inf')
# Sample from the filtered distribution
probs = F.softmax(sorted_logits, dim=-1)
idx = torch.multinomial(probs, 1)
return sorted_indices.gather(-1, idx)
KV Cache
During autoregressive generation, each new token only needs to attend to all previous tokens. Without caching, you'd recompute K and V for the entire sequence at every step. The KV cache stores previously computed K and V tensors and only computes the new token's K and V, then concatenates:
# Step 1: compute K, V for all tokens
# Step 2: only compute K_new, V_new for the new token
# K = cat(K_cached, K_new)
# V = cat(V_cached, V_new)
This reduces generation from O(n^2) to O(n) per token.
See gpt.py for the complete implementation.
Vision Transformer (ViT)
Paper: "An Image Is Worth 16x16 Words" (Dosovitskiy et al., 2020)
Core Idea
Treat an image as a sequence of patches and process them with a standard Transformer encoder. No convolutions needed.
Patch Embedding
Split the image into non-overlapping patches (e.g., 16x16 pixels), flatten each patch into a vector, then project to the model dimension:
# Image: (B, 3, 224, 224)
# Patches: 224/16 = 14 patches per side, 14*14 = 196 patches
# Each patch: 16*16*3 = 768 pixels
# Efficient implementation using Conv2d:
self.patch_embed = nn.Conv2d(3, d_model, kernel_size=16, stride=16)
# Output: (B, d_model, 14, 14) -> reshape to (B, 196, d_model)
CLS Token
A special learnable token prepended to the sequence. After processing through the Transformer, the CLS token's representation is used for classification:
self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))
# Prepend to patch sequence: (B, 196, d_model) -> (B, 197, d_model)
Why not just average-pool all patch representations? The CLS token provides a single, fixed-position summary token that the model can learn to aggregate global information into.
Position Embedding
Since patches have spatial relationships, learnable position embeddings are added (ViT uses learned embeddings, not sinusoidal):
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, d_model))
# +1 for the CLS token
Classification Head
After the Transformer encoder, take the CLS token's output and pass it through a simple MLP head:
cls_output = transformer_output[:, 0] # CLS token at position 0
logits = self.mlp_head(cls_output)
See vit.py for the complete implementation.
VAE (Variational Autoencoder)
Paper: "Auto-Encoding Variational Bayes" (Kingma & Welling, 2013)
Autoencoders vs Variational Autoencoders
A regular autoencoder learns to compress data to a latent code and reconstruct it. The latent space, however, has no structure — nearby points in latent space may decode to very different outputs, making it useless for generation.
A VAE forces the latent space to be structured (approximately Gaussian) by:
- Encoding to a distribution (mean and variance) instead of a point
- Sampling from that distribution (the reparameterization trick)
- Adding a KL divergence loss that pushes the distribution toward N(0, I)
The Reparameterization Trick
We can't backpropagate through random sampling. The trick: instead of sampling z ~ N(mu, sigma^2), compute z = mu + sigma * epsilon where epsilon ~ N(0, 1). Now the randomness is in epsilon (which doesn't need gradients), and z is a deterministic function of mu and sigma.
def reparameterize(self, mu, log_var):
std = torch.exp(0.5 * log_var)
eps = torch.randn_like(std)
return mu + eps * std
The ELBO Loss
The VAE loss has two terms:
Loss = Reconstruction Loss + KL Divergence
= E[||x - x_hat||^2] + KL(q(z|x) || p(z))
- Reconstruction loss: How well the decoder reconstructs the input.
Binary cross-entropy for binary data, MSE for continuous data.
- KL divergence: How far the encoder's distribution is from the prior
N(0, I). For Gaussians, this has a closed-form solution:
kl_loss = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp())
The KL term acts as a regularizer, preventing the encoder from collapsing the latent space to a single point (which would defeat the purpose of having a generative model).
See vae.py for the complete implementation.
U-Net
Paper: "U-Net: Convolutional Networks for Biomedical Image Segmentation" (Ronneberger et al., 2015)
Architecture Overview
U-Net has a symmetric encoder-decoder structure with skip connections:
Encoder Decoder
(downsampling) (upsampling)
[Input] ──────────────────────► [Output]
│ ▲
▼ │
[Down1] ──── skip connection ──► [Up1]
│ ▲
▼ │
[Down2] ──── skip connection ──► [Up2]
│ ▲
▼ │
[Down3] ──── skip connection ──► [Up3]
│ ▲
▼ │
[Down4] ──── skip connection ──► [Up4]
│ ▲
▼ │
[Bottleneck] ──────────────┘
Why Skip Connections?
The encoder captures "what" is in the image (semantic features) but loses spatial precision. The decoder recovers spatial resolution but lacks semantic context. Skip connections combine both: the decoder gets high-resolution features from the encoder concatenated with upsampled semantic features.
Encoder Path
Each encoder block: two 3x3 convolutions + ReLU, then 2x2 max pooling. Channels double at each level: 64 -> 128 -> 256 -> 512 -> 1024.
Decoder Path
Each decoder block: 2x2 transposed convolution (upsample), concatenate with the corresponding encoder features, then two 3x3 convolutions + ReLU.
Key Difference from ResNet Skip Connections
- ResNet: skip connections add the input (element-wise addition)
- U-Net: skip connections concatenate encoder features with decoder features
This is because U-Net needs to preserve both the high-resolution spatial information from the encoder and the semantic information from the decoder, while ResNet just needs to facilitate gradient flow.
Common Architectural Patterns
Residual Connections
Found in almost every modern architecture. The core pattern:
output = x + f(x) # residual connection
Benefits: easier optimization, better gradient flow, ability to train very deep networks.
Pre-Norm vs Post-Norm
Post-norm (original Transformer):
x = layer_norm(x + sublayer(x))
Pre-norm (GPT-2, modern practice):
x = x + sublayer(layer_norm(x))
Pre-norm is more stable for training. Post-norm can achieve slightly better performance but requires careful learning rate warmup.
Weight Tying
Sharing parameters between the input embedding and the output projection:
self.embed = nn.Embedding(vocab_size, d_model)
self.output_proj = nn.Linear(d_model, vocab_size, bias=False)
self.output_proj.weight = self.embed.weight # shared!
Reduces parameters by vocab_size * d_model and acts as regularization. Used in GPT, BERT, T5, and most modern language models.
Layer Scaling
Introduced in CaiT (Going Deeper with Image Transformers). Scale the output of each residual block by a learnable scalar, initialized to a small value (e.g., 0.1):
self.gamma = nn.Parameter(torch.ones(d_model) * 0.1)
# In forward:
x = x + self.gamma * sublayer(x)
This helps stabilize training of very deep Transformers by starting with near-identity blocks.
GELU Activation
Most modern Transformers use GELU (Gaussian Error Linear Unit) instead of ReLU:
F.gelu(x) # smooth approximation: x * Phi(x)
GELU is smoother than ReLU and has been empirically shown to work better for Transformers. It's the default in GPT, BERT, ViT, and most modern architectures.
Dropout Patterns in Transformers
Dropout is typically applied in three places:
- After attention weights (before multiplying with V)
- After the feed-forward network's output projection
- After adding positional embeddings (in some architectures)
attn_weights = F.softmax(scores, dim=-1)
attn_weights = F.dropout(attn_weights, p=0.1, training=self.training)
Weight Initialization
Different architectures use different initialization strategies:
# Xavier/Glorot (good for tanh/sigmoid): used in Transformer
nn.init.xavier_uniform_(self.weight)
# Kaiming/He (good for ReLU): used in ResNet
nn.init.kaiming_normal_(self.weight, mode='fan_out', nonlinearity='relu')
# Normal with small std: used in GPT
nn.init.normal_(self.weight, mean=0.0, std=0.02)
The right initialization prevents vanishing/exploding gradients at the start of training and can significantly affect convergence speed.
Summary
| Architecture | Year | Innovation | Key Pattern |
|---|---|---|---|
| ResNet | 2015 | Residual connections | Identity shortcut |
| U-Net | 2015 | Encoder-decoder + skip | Concatenation skip |
| VAE | 2013 | Reparameterization trick | Stochastic latent |
| Transformer | 2017 | Self-attention | Q/K/V attention |
| GPT | 2018 | Decoder-only + pretrain | Causal masking |
| ViT | 2020 | Patch tokenization | Image as sequence |
Files in This Module
resnet.py— Complete ResNet with BasicBlock, Bottleneck, configs for 18/34/50/101transformer.py— Full encoder-decoder Transformer from scratchgpt.py— GPT with generation: greedy, temperature, top-k, top-p samplingvit.py— Vision Transformer for image classificationvae.py— Variational Autoencoder with reparameterization trick and ELBO loss
📓 Open Notebook — Interactive version of this module
Source Files
resnet.py— ResNet (Residual Network) — complete implementation with BasicBlock and Bottlenecktransformer.py— Transformer — complete encoder-decoder implementationgpt.py— GPT (Generative Pre-trained Transformer) — decoder-only implementationvit.py— Vision Transformer (ViT) — complete implementation for image classificationvae.py— Variational Autoencoder (VAE) — complete implementation with reparameterization trick
Module 13: Advanced PyTorch Features
This module covers advanced PyTorch capabilities that go beyond standard model building and training. These are the tools that separate "I can train a model" from "I can build production-quality ML systems."
Functorch (torch.func): Functional Transformations
torch.func (formerly the standalone functorch library) provides composable function transforms inspired by JAX. The core idea: transform a plain Python function into a new function with different behavior.
vmap — Vectorized Map
vmap automatically vectorizes a function over a batch dimension. Instead of writing explicit batch loops or reshaping tensors, you write the function for a single example and vmap handles batching:
import torch
from torch.func import vmap
def compute_norm(x):
"""Compute L2 norm of a single vector."""
return torch.sqrt(torch.sum(x ** 2))
# Without vmap: need to handle batch dimension explicitly
batch = torch.randn(32, 10)
norms_loop = torch.stack([compute_norm(batch[i]) for i in range(32)])
# With vmap: automatic batching
norms_vmap = vmap(compute_norm)(batch)
Why use vmap instead of just writing batched code? Three reasons:
- Clarity: Write single-example logic, get batched execution
- Correctness: No batch dimension bugs
- Composition: Combine with
grad,jacrev, etc.
grad — Functional Gradient
torch.func.grad computes gradients functionally, without modifying tensors in-place or using .backward():
from torch.func import grad
def f(x):
return torch.sin(x).sum()
# grad returns a function that computes the gradient
grad_f = grad(f)
x = torch.tensor([1.0, 2.0, 3.0])
print(grad_f(x)) # cos(x) = [0.5403, -0.4161, -0.9900]
grad is particularly useful when composed with other transforms.
jacrev / jacfwd — Jacobian Computation
The Jacobian matrix contains all partial derivatives of a vector-valued function:
from torch.func import jacrev, jacfwd
def f(x):
return torch.stack([x[0]**2 + x[1], x[0] * x[1]**2])
x = torch.tensor([1.0, 2.0])
J_rev = jacrev(f)(x) # Reverse-mode: efficient when output dim < input dim
J_fwd = jacfwd(f)(x) # Forward-mode: efficient when input dim < output dim
Rule of thumb:
- Use
jacrevwhen output dimension is smaller than input dimension - Use
jacfwdwhen input dimension is smaller than output dimension
hessian — Second Derivatives
The Hessian matrix is the Jacobian of the gradient. torch.func.hessian is syntactic sugar for jacrev(jacrev(f)) or jacfwd(jacrev(f)):
from torch.func import hessian
def f(x):
return (x ** 3).sum()
H = hessian(f)(torch.tensor([1.0, 2.0, 3.0]))
# H[i,j] = d^2f / dx_i dx_j
Composing Transforms
The real power is in composition:
from torch.func import vmap, grad, jacrev
# Per-sample gradients: grad of loss for each sample in a batch
def loss_fn(params, x, y):
pred = model_fn(params, x)
return ((pred - y) ** 2).sum()
# Batched Jacobian: Jacobian for each sample in a batch
batched_jacobian = vmap(jacrev(model_fn), in_dims=(None, 0))
See functorch_transforms.py for complete examples.
Per-Sample Gradients
The classic vmap + grad use case. Standard training computes the average gradient across a batch. But sometimes you need the gradient for each sample individually — for example, in differential privacy (DP-SGD), where you need to clip per-sample gradients before averaging.
Without vmap, you'd need to loop over samples or use inefficient tricks. With vmap:
from torch.func import vmap, grad
from torch import nn
model = nn.Linear(10, 1)
params = dict(model.named_parameters())
def compute_loss(params, x, y):
# Stateless function: takes params explicitly
pred = torch.func.functional_call(model, params, (x,))
return ((pred - y) ** 2).squeeze()
# grad w.r.t. params for a single sample
grad_fn = grad(compute_loss)
# vmap over the batch dimension of x and y
per_sample_grads = vmap(grad_fn, in_dims=(None, 0, 0))(params, X_batch, Y_batch)
See per_sample_gradients.py for a complete walkthrough.
Sparse Tensors
Sparse tensors store only nonzero elements, saving memory and computation for data that is mostly zeros (e.g., adjacency matrices, bag-of-words features).
COO (Coordinate) Format
Stores row and column indices alongside values. Good for construction and conversion, less efficient for arithmetic:
indices = torch.tensor([[0, 1, 2], [1, 0, 2]]) # (2, nnz)
values = torch.tensor([3.0, 4.0, 5.0])
sparse_coo = torch.sparse_coo_tensor(indices, values, size=(3, 3))
CSR (Compressed Sparse Row) Format
Stores row pointers, column indices, and values. Efficient for row-slicing and matrix-vector products:
crow_indices = torch.tensor([0, 1, 2, 3]) # row pointers
col_indices = torch.tensor([1, 0, 2]) # column indices
values = torch.tensor([3.0, 4.0, 5.0])
sparse_csr = torch.sparse_csr_tensor(crow_indices, col_indices, values, size=(3, 3))
BSR (Block Sparse Row) Format
Like CSR but stores dense blocks instead of individual elements. Useful when sparsity has block structure (common in structured pruning):
crow_indices = torch.tensor([0, 1, 2])
col_indices = torch.tensor([0, 1])
values = torch.randn(2, 2, 2) # two 2x2 blocks
sparse_bsr = torch.sparse_bsr_tensor(crow_indices, col_indices, values, size=(4, 4))
When to use each:
- COO: Building sparse tensors, format conversion, unstructured updates
- CSR: Sparse matrix-vector products, row-based access patterns
- BSR: Block-structured sparsity, GPU-friendly operations
Sparse Operations
# Matrix multiply (sparse @ dense)
result = torch.sparse.mm(sparse_csr, dense_matrix)
# Element-wise operations
sparse_sum = sparse_coo + sparse_coo
sparse_scaled = sparse_coo * 2.0
# Convert between formats
dense = sparse_coo.to_dense()
csr = sparse_coo.to_sparse_csr()
Complex Numbers
PyTorch natively supports complex tensors, essential for signal processing, quantum computing simulations, and Fourier analysis.
# Creating complex tensors
z = torch.complex(torch.tensor([1.0, 2.0]), torch.tensor([3.0, 4.0]))
z = torch.tensor([1+3j, 2+4j]) # Python complex literals
# Operations
z.real # real part
z.imag # imaginary part
z.abs() # magnitude
z.angle() # phase angle
z.conj() # complex conjugate
# FFT
signal = torch.randn(1000)
spectrum = torch.fft.fft(signal)
freqs = torch.fft.fftfreq(1000)
# Inverse FFT
reconstructed = torch.fft.ifft(spectrum).real
See sparse_and_complex.py for complete examples.
Quantization Overview
Quantization reduces model size and increases inference speed by using lower precision (e.g., INT8 instead of FP32).
Why Quantize?
- Memory: INT8 uses 4x less memory than FP32
- Speed: Integer arithmetic is faster, especially on mobile/edge devices
- Accuracy: Modern quantization preserves most of the model's accuracy
Quantization Approaches
Post-Training Quantization (PTQ): Quantize a pre-trained model without retraining. Fast but may lose more accuracy.
Quantization-Aware Training (QAT): Simulate quantization during training so the model learns to be robust to quantization noise. Better accuracy but requires retraining.
PT2E Quantization Flow: The modern PyTorch 2 Export quantization approach. Uses torch.export to capture the model as a graph, applies quantization annotations, and lowers to optimized backends:
# Conceptual PT2E flow (simplified)
import torch
from torch.ao.quantization.quantize_pt2e import prepare_pt2e, convert_pt2e
model = MyModel()
exported = torch.export.export(model, example_inputs)
prepared = prepare_pt2e(exported, quantizer)
# Calibrate with representative data
for batch in calibration_data:
prepared(batch)
quantized = convert_pt2e(prepared)
torchao: A library for architecture optimization, including quantization, sparsity, and low-precision training. Offers simple APIs:
# Conceptual torchao usage
import torchao
torchao.quantize_(model, torchao.quantization.int8_weight_only())
Custom Operators (torch.library)
When PyTorch's built-in ops don't cover your needs, you can define custom operators with proper integration into autograd, torch.compile, and torch.export.
import torch
from torch.library import Library, impl
# Create a library namespace for your custom ops
my_lib = Library("myops", "DEF")
# Define the op signature
my_lib.define("my_relu(Tensor x) -> Tensor")
# Register a CPU implementation
@impl(my_lib, "my_relu", "CPU")
def my_relu_cpu(x):
return x.clamp(min=0)
# Register a Meta (shape-only) implementation for torch.compile
@impl(my_lib, "my_relu", "Meta")
def my_relu_meta(x):
return torch.empty_like(x)
# Use the custom op
x = torch.randn(5)
result = torch.ops.myops.my_relu(x)
Custom Autograd for Custom Ops
# Register autograd formula
def my_relu_backward(ctx, grad_output):
x, = ctx.saved_tensors
return grad_output * (x > 0).float()
torch.library.impl_abstract("myops::my_relu", my_relu_meta)
See custom_operators.py for a complete example.
C++ Extensions
For performance-critical code, you can write custom C++ (and CUDA) extensions:
from torch.utils.cpp_extension import load
# JIT compilation: compiles C++ code on first use
my_extension = load(
name="my_extension",
sources=["my_extension.cpp"],
verbose=True,
)
The C++ side uses the PyTorch C++ API (LibTorch):
#include <torch/extension.h>
torch::Tensor my_add(torch::Tensor a, torch::Tensor b) {
return a + b;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("my_add", &my_add, "Custom add");
}
Profiling Deep Dive
PyTorch's profiler helps identify performance bottlenecks.
Basic Profiling
import torch
from torch.profiler import profile, record_function, ProfilerActivity
with profile(
activities=[ProfilerActivity.CPU],
record_shapes=True,
profile_memory=True,
) as prof:
with record_function("model_inference"):
output = model(input_tensor)
# Print a summary table sorted by CPU time
print(prof.key_averages().table(sort_by="cpu_time_total", row_limit=10))
Chrome Trace
Export a trace viewable in Chrome's chrome://tracing or Perfetto UI:
prof.export_chrome_trace("trace.json")
record_function
Annotate specific code regions for fine-grained profiling:
with record_function("my_attention"):
attn_output = attention(q, k, v)
with record_function("my_ffn"):
ffn_output = feed_forward(attn_output)
TensorBoard Integration
with profile(
schedule=torch.profiler.schedule(wait=1, warmup=1, active=3, repeat=1),
on_trace_ready=torch.profiler.tensorboard_trace_handler("./log_dir"),
record_shapes=True,
profile_memory=True,
with_stack=True,
) as prof:
for step, batch in enumerate(dataloader):
output = model(batch)
loss = criterion(output, target)
loss.backward()
optimizer.step()
prof.step()
See profiling.py for runnable examples.
Memory Profiling
Basic Memory Tracking
# Note: these require CUDA, shown for reference
torch.cuda.memory_allocated() # current memory used by tensors
torch.cuda.max_memory_allocated() # peak memory since last reset
torch.cuda.memory_reserved() # total memory reserved by allocator
torch.cuda.reset_peak_memory_stats()
Finding Memory Leaks
Common causes:
- Storing tensors in a list that grows: solution is to
.detach()or
store .item() for scalar values
- Not clearing gradients: call
optimizer.zero_grad()each step - Keeping computation graph alive: use
.detach()orwith torch.no_grad()
# BAD: keeps entire computation graph in memory
losses = []
for batch in dataloader:
loss = model(batch).sum()
losses.append(loss) # holds onto graph!
# GOOD: detach the scalar value
losses = []
for batch in dataloader:
loss = model(batch).sum()
losses.append(loss.item()) # just a float, no graph
Debugging Techniques
Anomaly Detection
Detects the operation that produced a NaN or Inf gradient:
with torch.autograd.detect_anomaly():
output = model(input)
loss = criterion(output, target)
loss.backward() # will print a traceback if NaN/Inf is detected
Warning: this is SLOW. Only use for debugging.
Gradient Checking
Verify your autograd implementation by comparing with finite differences:
from torch.autograd import gradcheck
func = MyCustomFunction.apply
input = torch.randn(3, 4, dtype=torch.float64, requires_grad=True)
assert gradcheck(func, input, eps=1e-6, atol=1e-4)
Common Error Messages and Fixes
"one of the variables needed for gradient computation has been modified by an inplace operation"
- Cause: in-place operation (like
x += 1) on a tensor that requires grad - Fix: use
x = x + 1(out-of-place) instead
"Trying to backward through the graph a second time"
- Cause: calling
.backward()twice withoutretain_graph=True - Fix: either use
retain_graph=Trueor restructure to avoid double backward
"Expected all tensors to be on the same device"
- Cause: mixing CPU and GPU tensors in one operation
- Fix: ensure all tensors are on the same device with
.to(device)
"RuntimeError: mat1 and mat2 shapes cannot be multiplied"
- Cause: shape mismatch in linear layers
- Fix: print shapes before the operation to identify the mismatch
See debugging_tips.py for practical debugging examples.
torch.fx: Symbolic Tracing and Graph Transformation
torch.fx symbolically traces a model to produce a graph IR (intermediate representation) that you can analyze and transform.
Basic Tracing
import torch.fx
class MyModel(nn.Module):
def forward(self, x):
x = torch.relu(x)
x = x + 1
return x
model = MyModel()
traced = torch.fx.symbolic_trace(model)
print(traced.graph) # shows the operations as a graph
Writing Custom Passes
def replace_relu_with_gelu(module):
"""Replace all ReLU calls with GELU."""
traced = torch.fx.symbolic_trace(module)
for node in traced.graph.nodes:
if node.op == "call_function" and node.target == torch.relu:
node.target = torch.nn.functional.gelu
traced.graph.lint() # validate the graph
traced.recompile()
return traced
Use Cases
- Quantization: Analyze the graph to determine where to insert quant/dequant nodes
- Fusion: Merge compatible operations (e.g., Conv + BN)
- Shape inference: Propagate shapes through the graph without running data
- Visualization: Understand model structure programmatically
Meta Device: Shape Inference Without Memory
The meta device lets you analyze models without allocating real memory. Tensors on the meta device have shapes and dtypes but no data:
# Create a model on the meta device (no memory allocated)
with torch.device("meta"):
model = nn.Linear(1000, 1000)
# model.weight.shape == (1000, 1000) but uses 0 bytes
# Analyze input/output shapes
x = torch.empty(32, 1000, device="meta")
out = model(x)
print(out.shape) # torch.Size([32, 1000])
Use cases:
- Model analysis: Count parameters and compute shapes for huge models
that don't fit in memory
- Architecture prototyping: Verify shapes without waiting for memory
allocation
- Deferred initialization: Create model structure on meta, then
materialize weights on the target device
# Count parameters of a huge model without any memory
with torch.device("meta"):
huge_model = nn.Sequential(
nn.Linear(10000, 10000),
nn.ReLU(),
nn.Linear(10000, 10000),
)
total_params = sum(p.numel() for p in huge_model.parameters())
memory_gb = total_params * 4 / 1e9 # FP32 = 4 bytes
print(f"Parameters: {total_params:,}, Memory: {memory_gb:.2f} GB")
Summary
| Feature | Use Case | Key API |
|---|---|---|
| vmap | Batch any function | torch.func.vmap |
| grad | Functional gradients | torch.func.grad |
| jacrev/jacfwd | Jacobian matrices | torch.func.jacrev |
| Per-sample grads | DP-SGD, influence functions | vmap(grad(...)) |
| Sparse tensors | Graphs, sparse data | torch.sparse_coo_tensor |
| Complex numbers | FFT, signal processing | torch.complex, torch.fft |
| Custom ops | Extending PyTorch | torch.library |
| Profiling | Performance optimization | torch.profiler |
| Anomaly detection | Debugging NaN/Inf | torch.autograd.detect_anomaly |
| torch.fx | Graph transforms | torch.fx.symbolic_trace |
| Meta device | Shape analysis | torch.device("meta") |
Files in This Module
functorch_transforms.py— vmap, grad, jacrev, hessian demonstrationsper_sample_gradients.py— Per-sample gradient computation with vmap+gradcustom_operators.py— Defining custom ops with torch.libraryprofiling.py— Profiler usage, timing, and analysissparse_and_complex.py— Sparse tensors, complex numbers, and FFTdebugging_tips.py— Anomaly detection, gradient flow checking, and common fixes
📓 Open Notebook — Interactive version of this module
Source Files
functorch_transforms.py— Functorch transforms (torch.func): vmap, grad, jacrev, hessianper_sample_gradients.py— Per-sample gradients with vmap + gradcustom_operators.py— Custom operators with torch.libraryprofiling.py— Profiling deep dive — profiler usage, timing, and analysissparse_and_complex.py— Sparse tensors, complex numbers, and FFTdebugging_tips.py— Debugging techniques — anomaly detection, gradient flow, and common fixes
Module 14: Testing and Reproducibility
Writing tests and ensuring reproducibility are essential skills that separate hobby projects from production-quality deep learning code. This module covers PyTorch's testing framework, reproducibility techniques, and benchmarking.
PyTorch's Testing Framework
PyTorch has its own testing infrastructure built on top of Python's unittest. The key class is TestCase from torch.testing._internal.common_utils, which provides tensor-aware assertions and convenient utilities.
Basic Test Structure
from torch.testing._internal.common_utils import run_tests, TestCase
import torch
class TestMyFeature(TestCase):
def test_addition(self):
a = torch.tensor([1.0, 2.0, 3.0])
b = torch.tensor([4.0, 5.0, 6.0])
result = a + b
expected = torch.tensor([5.0, 7.0, 9.0])
self.assertEqual(result, expected)
def test_shape(self):
x = torch.randn(3, 4, 5)
self.assertEqual(x.shape, (3, 4, 5))
def test_dtype(self):
x = torch.zeros(5, dtype=torch.float32)
self.assertEqual(x.dtype, torch.float32)
if __name__ == "__main__":
run_tests()
assertEqual for Tensors
PyTorch's assertEqual is smarter than the standard library version. For tensors, it checks:
- Shape equality
- Dtype equality
- Value equality (with configurable tolerance for floating point)
# Exact equality (for integer tensors)
self.assertEqual(torch.tensor([1, 2, 3]), torch.tensor([1, 2, 3]))
# Approximate equality (for float tensors) — uses default tolerances
self.assertEqual(
torch.tensor([1.0, 2.0]),
torch.tensor([1.0 + 1e-7, 2.0 - 1e-7]),
)
# Custom tolerances
self.assertEqual(a, b, atol=1e-4, rtol=1e-4)
Useful Assertions
# Check that a function raises a specific exception
with self.assertRaises(RuntimeError):
torch.tensor([1, 2]) + torch.tensor([1, 2, 3])
# Check error message content
with self.assertRaisesRegex(RuntimeError, "size mismatch"):
bad_operation()
# Boolean checks
self.assertTrue(torch.all(x > 0))
self.assertFalse(torch.any(torch.isnan(x)))
Parametrized Tests
Use the @parametrize decorator to run a test with multiple inputs:
from torch.testing._internal.common_utils import parametrize
class TestOps(TestCase):
@parametrize("dtype", [torch.float32, torch.float64])
@parametrize("size", [(2, 3), (4, 5)])
def test_zeros(self, dtype, size):
x = torch.zeros(size, dtype=dtype)
self.assertEqual(x.sum().item(), 0.0)
self.assertEqual(x.dtype, dtype)
self.assertEqual(x.shape, size)
This generates 4 test cases (2 dtypes x 2 sizes), each with a descriptive name.
Device-Generic Tests
For testing across CPU and (optionally) GPU:
from torch.testing._internal.common_device_type import (
instantiate_device_type_tests,
dtypes,
)
class TestMyOp(TestCase):
@dtypes(torch.float32, torch.float64)
def test_my_op(self, device, dtype):
x = torch.randn(10, device=device, dtype=dtype)
result = my_op(x)
self.assertEqual(result.device.type, device)
self.assertEqual(result.dtype, dtype)
instantiate_device_type_tests(TestMyOp, globals())
This creates separate test classes for CPU and CUDA (if available), testing each dtype on each device.
OpInfo Framework
PyTorch uses the OpInfo framework for systematic operator testing. Each operator has an OpInfo entry that describes:
- The operator function
- Valid input dtypes
- Sample inputs (for testing)
- Reference implementations (for correctness checking)
- Gradient test configurations
While you probably won't need to write OpInfo entries unless contributing to PyTorch core, understanding the concept helps you write better tests:
# The idea: describe an op's properties declaratively, then auto-generate tests
# OpInfo("torch.add",
# dtypes=floating_types_and(torch.half),
# sample_inputs_func=sample_inputs_add,
# supports_out=True,
# )
Reproducibility
Non-deterministic behavior makes debugging nearly impossible. Here's how to control randomness in PyTorch.
Setting All Seeds
import random
import numpy as np
import torch
def set_seed(seed=42):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
# Also seeds all CUDA devices:
torch.cuda.manual_seed_all(seed)
Deterministic Mode
Even with fixed seeds, some operations have non-deterministic GPU implementations for performance. To enforce full determinism:
torch.use_deterministic_algorithms(True)
# Also set the CUBLAS workspace config for full CUDA determinism:
import os
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
Warning: deterministic mode may be slower and some operations will raise errors if no deterministic implementation exists.
torch.backends Settings
# CuDNN: control convolution algorithm selection
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
# benchmark=True auto-selects the fastest algorithm, but this selection is
# non-deterministic. Set to False for reproducibility.
DataLoader Reproducibility
def seed_worker(worker_id):
"""Ensure each DataLoader worker has a different but reproducible seed."""
worker_seed = torch.initial_seed() % 2**32
np.random.seed(worker_seed)
random.seed(worker_seed)
g = torch.Generator()
g.manual_seed(42)
loader = DataLoader(
dataset,
batch_size=32,
shuffle=True,
num_workers=4,
worker_init_fn=seed_worker,
generator=g,
)
See reproducibility.py for a complete setup.
Benchmarking
Reliable benchmarking requires attention to detail: warmup, synchronization, statistical rigor.
torch.utils.benchmark.Timer
The recommended way to benchmark PyTorch code:
from torch.utils.benchmark import Timer
t = Timer(
stmt="torch.mm(A, B)",
globals={"A": torch.randn(256, 256), "B": torch.randn(256, 256), "torch": torch},
label="Matrix Multiplication",
sub_label="256x256",
)
# blocked_autorange: automatically determines the number of runs
result = t.blocked_autorange(min_run_time=1.0)
print(result) # prints median, IQR, and other statistics
Why use Timer instead of time.time()?
- Warmup: Automatically warms up JIT compilation, CUDA initialization
- Statistics: Reports median and IQR, not just mean
- Synchronization: Handles CUDA synchronization correctly
- Isolation: Minimizes interference from other processes
Comparing Implementations
from torch.utils.benchmark import Compare
results = []
for size in [64, 128, 256, 512]:
A = torch.randn(size, size)
B = torch.randn(size, size)
for label, stmt in [("mm", "torch.mm(A, B)"), ("@", "A @ B")]:
t = Timer(
stmt=stmt,
globals={"A": A, "B": B, "torch": torch},
label="matmul",
sub_label=label,
description=f"{size}x{size}",
)
results.append(t.blocked_autorange(min_run_time=0.5))
compare = Compare(results)
compare.print()
Common Benchmarking Mistakes
- No warmup: First runs include JIT compilation and caching overhead
- No CUDA synchronization: CUDA ops are asynchronous — time.time() measures
only the kernel launch, not execution
- Too few runs: A single measurement is noisy
- Benchmarking in training mode: BatchNorm and Dropout add overhead
See benchmarking.py for complete examples.
Common Testing Patterns
Testing Numerical Correctness
def test_softmax_correctness(self):
x = torch.randn(5, 10)
result = F.softmax(x, dim=-1)
# Property: sums to 1 along the softmax dimension
self.assertTrue(torch.allclose(result.sum(dim=-1), torch.ones(5), atol=1e-6))
# Property: all values are in [0, 1]
self.assertTrue((result >= 0).all())
self.assertTrue((result <= 1).all())
Testing Gradient Correctness
from torch.autograd import gradcheck
def test_custom_op_gradient(self):
func = my_custom_op
# Use float64 for numerical gradient checking (better precision)
input = torch.randn(3, 4, dtype=torch.float64, requires_grad=True)
self.assertTrue(gradcheck(func, input, eps=1e-6, atol=1e-4))
Testing with Approximate Equality
# torch.testing.assert_close: the modern way to check approximate equality
torch.testing.assert_close(actual, expected, atol=1e-5, rtol=1e-5)
# For very loose checks (e.g., stochastic operations)
self.assertTrue(torch.allclose(result, expected, atol=0.1, rtol=0.1))
Testing Model Invariants
def test_model_deterministic_eval(self):
"""In eval mode, same input should produce same output."""
model = MyModel()
model.eval()
x = torch.randn(4, 10)
with torch.no_grad():
out1 = model(x)
out2 = model(x)
self.assertEqual(out1, out2)
def test_model_output_shape(self):
"""Output shape should be (batch_size, num_classes)."""
model = MyModel(num_classes=10)
for batch_size in [1, 4, 16]:
x = torch.randn(batch_size, 3, 32, 32)
out = model(x)
self.assertEqual(out.shape, (batch_size, 10))
Summary
| Topic | Key Tool | When to Use |
|---|---|---|
| Basic testing | TestCase, assertEqual | Every project |
| Parametrized tests | @parametrize | Testing across configs |
| Device tests | instantiate_device_type_tests | Cross-device testing |
| Reproducibility | set_seed(), deterministic mode | Debugging, CI |
| Benchmarking | torch.utils.benchmark.Timer | Performance comparison |
| Gradient checks | torch.autograd.gradcheck | Custom autograd |
Files in This Module
test_example.py— Complete test file using PyTorch's TestCasereproducibility.py— Full reproducibility setup and verificationbenchmarking.py— Benchmarking with torch.utils.benchmark
📓 Open Notebook — Interactive version of this module
Source Files
test_example.py— Example test file using PyTorch's TestCasereproducibility.py— Reproducibility in PyTorch — complete setup for reproducible experimentsbenchmarking.py— Benchmarking with torch.utils.benchmark
Module 15: Practical PyTorch Utilities
The Hidden Toolkit Most Tutorials Never Teach
PyTorch ships with a rich set of utility modules that most tutorials skip entirely. These are the tools that separate a beginner from a practitioner: weight parametrization, pruning, normalization techniques, sequence packing, nested tensors, and model fusion. This module covers them all.
Table of Contents
- torch.nn.utils.parametrize — Weight Constraints
- torch.nn.utils.prune — Model Pruning
- Spectral Norm & Weight Norm
- torch.nn.utils.rnn — Sequence Packing
- Conv-BN Fusion
- torch.nested — Nested (Jagged) Tensors
- torch.nn.utils.clip_grad — Gradient Clipping Internals
- parameters_to_vector & skip_init
1. Weight Parametrization
What it is: A framework to apply constraints or transformations to module parameters. Instead of manually enforcing constraints in forward(), you register a parametrization that automatically transforms the raw weight into a constrained version.
Why it matters: Enforcing constraints like orthogonality, symmetry, or positivity on weights is common in research and production. Without parametrization, you'd need hacky workarounds.
How It Works
import torch
import torch.nn as nn
import torch.nn.utils.parametrize as P
class Symmetric(nn.Module):
"""Parametrization that makes a matrix symmetric."""
def forward(self, X):
return X.triu() + X.triu(1).transpose(-1, -2)
linear = nn.Linear(5, 5)
P.register_parametrization(linear, "weight", Symmetric())
# Now linear.weight is ALWAYS symmetric
print(linear.weight) # Symmetric!
print(torch.allclose(linear.weight, linear.weight.T)) # True
# The raw unconstrained parameter is stored as:
print(linear.parametrizations.weight.original)
Key Concept: Original vs Parametrized
When you register a parametrization:
- The original unconstrained tensor is stored at
module.parametrizations.<name>.original - Accessing
module.<name>runs the parametrization on the original and returns the result - The optimizer updates the original (unconstrained) parameter
- The parametrization is applied on every access (or cached)
Built-in Parametrizations
from torch.nn.utils import parametrizations
# Orthogonal weight matrix (useful for RNNs, preventing vanishing/exploding gradients)
linear = nn.Linear(5, 5)
parametrizations.orthogonal(linear, "weight")
# linear.weight is now always orthogonal: W^T W = I
# Spectral normalization (stabilize GANs and training)
conv = nn.Conv2d(3, 64, 3)
parametrizations.spectral_norm(conv, "weight")
# Constrains the spectral norm (largest singular value) of the weight to 1
# Weight normalization (decouple magnitude from direction)
linear = nn.Linear(10, 5)
parametrizations.weight_norm(linear, "weight")
# Reparametrizes: w = g * (v / ||v||)
Custom Parametrization Example — Positive Weights
class Positive(nn.Module):
"""Ensures weights are always positive via softplus."""
def forward(self, X):
return torch.nn.functional.softplus(X)
linear = nn.Linear(3, 3)
P.register_parametrization(linear, "weight", Positive())
print(linear.weight) # All positive!
print((linear.weight > 0).all()) # True
Caching for Efficiency
If you use a parametrized weight multiple times in forward() (e.g., RNNs sharing the recurrent kernel), use caching to avoid recomputation:
with P.cached():
output = model(input) # Parametrizations computed once, cached
Removing Parametrizations
# Remove and keep the parametrized (constrained) value
P.remove_parametrizations(linear, "weight", leave_parametrized=True)
# Remove and go back to unconstrained original
P.remove_parametrizations(linear, "weight", leave_parametrized=False)
2. Model Pruning
What it is: Removing (zeroing out) weights from a neural network to make it smaller and faster.
Why it matters: Pruned models can be 2-10x smaller with minimal accuracy loss. Critical for edge deployment.
Pruning Strategies
| Method | What It Does |
|---|---|
random_unstructured | Zero out random individual weights |
l1_unstructured | Zero out weights with smallest L1 magnitude |
random_structured | Zero out entire channels/neurons randomly |
ln_structured | Zero out channels with smallest Ln norm |
global_unstructured | Prune across all layers by global ranking |
How Pruning Works in PyTorch
- The original weight is moved to
weight_orig - A binary mask
weight_maskis created - A forward hook computes
weight = weight_orig * weight_maskbefore each forward pass
import torch.nn.utils.prune as prune
linear = nn.Linear(10, 5)
# Prune 30% of weights (smallest magnitude)
prune.l1_unstructured(linear, name="weight", amount=0.3)
print(linear.weight) # Pruned weight (has zeros)
print(linear.weight_mask) # Binary mask
print(linear.weight_orig) # Original weight
# Count sparsity
zeros = (linear.weight == 0).sum().item()
total = linear.weight.numel()
print(f"Sparsity: {zeros}/{total} = {zeros/total:.1%}")
Global Pruning (Prune Across All Layers)
model = nn.Sequential(
nn.Linear(100, 64),
nn.ReLU(),
nn.Linear(64, 32),
nn.ReLU(),
nn.Linear(32, 10),
)
# Collect all prunable parameters
parameters_to_prune = [
(model[0], "weight"),
(model[2], "weight"),
(model[4], "weight"),
]
# Globally prune 40% of weights (by L1 magnitude across ALL layers)
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=0.4,
)
Making Pruning Permanent
# Remove the pruning reparametrization (bake the mask into the weight)
prune.remove(linear, "weight")
# Now linear.weight IS the pruned weight directly (no more mask/orig)
3. Spectral Norm & Weight Norm
Spectral Normalization
Controls the Lipschitz constant of a layer by normalizing weights by their spectral norm (largest singular value). Essential for stable GAN training.
Math: $\bar{W} = W / \sigma(W)$ where $\sigma(W)$ is the largest singular value.
from torch.nn.utils import spectral_norm
# Apply spectral norm
conv = nn.Conv2d(3, 64, 3, padding=1)
conv = spectral_norm(conv, name="weight")
# The spectral norm is estimated via power iteration (efficient)
# No full SVD needed — just one vector update per forward pass
Weight Normalization
Decouples weight magnitude from direction: $w = g \cdot \frac{v}{\|v\|}$
The optimizer can separately learn the magnitude $g$ and direction $v$, which often leads to faster convergence.
from torch.nn.utils import weight_norm
linear = nn.Linear(10, 5)
linear = weight_norm(linear, name="weight")
# Now has: linear.weight_g (magnitude) and linear.weight_v (direction)
print(linear.weight_g.shape) # (5, 1)
print(linear.weight_v.shape) # (5, 10)
4. Sequence Packing for RNNs
Problem: Sequences in a batch have different lengths. Padding wastes computation — the RNN processes pad tokens unnecessarily.
Solution: Pack sequences so the RNN only processes real tokens.
from torch.nn.utils.rnn import (
pack_padded_sequence,
pad_packed_sequence,
pad_sequence,
pack_sequence,
)
# Variable-length sequences
seqs = [torch.randn(5, 10), # length 5
torch.randn(3, 10), # length 3
torch.randn(8, 10)] # length 8
# Step 1: Pad to same length
padded = pad_sequence(seqs, batch_first=True) # (3, 8, 10)
lengths = torch.tensor([5, 3, 8])
# Step 2: Pack (sorts by length internally)
packed = pack_padded_sequence(padded, lengths, batch_first=True, enforce_sorted=False)
# Step 3: Feed to RNN
rnn = nn.LSTM(10, 20, batch_first=True)
output_packed, (h_n, c_n) = rnn(packed)
# Step 4: Unpack
output_padded, output_lengths = pad_packed_sequence(output_packed, batch_first=True)
print(f"Output: {output_padded.shape}") # (3, 8, 20)
Why Packing Matters
Without packing: RNN processes all 8 timesteps for all 3 sequences (24 steps). With packing: RNN processes 5+3+8=16 steps total. 33% less computation.
For long sequences with high variance in length, the savings are much larger.
5. Conv-BN Fusion
What it is: Merging a Conv2d + BatchNorm2d into a single Conv2d for faster inference.
Why it matters: During inference, BatchNorm is a fixed affine transform. Fusing it into the convolution eliminates one entire layer with zero accuracy loss.
from torch.nn.utils.fusion import fuse_conv_bn_eval
conv = nn.Conv2d(3, 64, 3, padding=1)
bn = nn.BatchNorm2d(64)
# Train as normal...
# Then for inference:
conv.eval()
bn.eval()
fused_conv = fuse_conv_bn_eval(conv, bn)
# fused_conv is a single Conv2d that produces identical output
# but is faster (one layer instead of two)
x = torch.randn(1, 3, 32, 32)
print(torch.allclose(fused_conv(x), bn(conv(x)), atol=1e-5)) # True
This is one of the most common inference optimizations and is done automatically by torch.compile and ONNX optimizers.
6. Nested Tensors
What it is: A tensor that can hold sequences of different lengths without padding. Also called "jagged tensors" or "ragged tensors."
Why it matters: Eliminates wasted computation on padding tokens in NLP/attention. Flash Attention can directly consume nested tensors for variable-length sequences.
import torch
from torch.nested import nested_tensor, as_nested_tensor
# Create a nested tensor from variable-length sequences
nt = nested_tensor([
torch.randn(3, 8), # sequence of length 3, dim 8
torch.randn(5, 8), # sequence of length 5, dim 8
torch.randn(2, 8), # sequence of length 2, dim 8
])
print(f"Type: {type(nt)}")
print(f"Nested size: {nt.size()}")
# The first dim is the batch, subsequent dims may vary
# Convert to padded tensor when needed
padded = torch.nested.to_padded_tensor(nt, padding=0.0)
print(f"Padded shape: {padded.shape}") # (3, 5, 8) — padded to max length
# Convert back
nt2 = as_nested_tensor(padded)
Nested Tensors with SDPA
The real power is using nested tensors with F.scaled_dot_product_attention — Flash Attention handles the variable lengths natively, avoiding wasted computation on padding:
# Instead of padding and masking:
# attn = F.scaled_dot_product_attention(Q_padded, K_padded, V_padded, attn_mask=mask)
# With nested tensors, no padding needed:
# Q_nested, K_nested, V_nested are NestedTensors
# attn = F.scaled_dot_product_attention(Q_nested, K_nested, V_nested)
# Flash Attention processes only real tokens!
7. Gradient Clipping Internals
We covered gradient clipping in Module 05, but here's the deeper story:
from torch.nn.utils import clip_grad_norm_, clip_grad_value_, get_total_norm
model = nn.Linear(10, 5)
loss = model(torch.randn(3, 10)).sum()
loss.backward()
# Get the total gradient norm BEFORE clipping
total_norm = get_total_norm(model.parameters(), norm_type=2.0)
print(f"Total gradient norm: {total_norm:.4f}")
# Clip by norm (scales all gradients proportionally if norm exceeds max)
clipped_norm = clip_grad_norm_(model.parameters(), max_norm=1.0, norm_type=2.0)
print(f"Returned (pre-clip) norm: {clipped_norm:.4f}")
# Clip by value (clamps each gradient element independently)
clip_grad_value_(model.parameters(), clip_value=0.5)
Key insight: clip_grad_norm_ returns the total norm before clipping — useful for monitoring gradient health during training.
8. Parameter Utilities
parameters_to_vector / vector_to_parameters
Flatten all model parameters into a single vector (useful for L-BFGS, evolutionary methods, or model comparison):
from torch.nn.utils import parameters_to_vector, vector_to_parameters
model = nn.Sequential(nn.Linear(10, 5), nn.Linear(5, 2))
# Flatten all parameters
vec = parameters_to_vector(model.parameters())
print(f"All params as vector: {vec.shape}") # (77,)
# Modify and write back
vec *= 0.5
vector_to_parameters(vec, model.parameters())
skip_init — Create Modules Without Initializing Weights
For very large models, the default weight initialization can be slow and wasteful (especially if you're loading a checkpoint immediately):
from torch.nn.utils import skip_init
# Normal: allocates + initializes weights (slow for large models)
linear = nn.Linear(10000, 10000)
# Skip init: allocates uninitialized memory (fast)
linear = skip_init(nn.Linear, 10000, 10000)
# Then load your checkpoint:
# linear.load_state_dict(torch.load("checkpoint.pt"))
Summary
| Utility | What It Does | When to Use |
|---|---|---|
parametrize | Constrain weights (orthogonal, symmetric, positive) | Research, stable training |
prune | Zero out weights by magnitude/structure | Model compression, edge deploy |
spectral_norm | Normalize by largest singular value | GAN training stability |
weight_norm | Decouple magnitude and direction | Faster convergence |
pack_padded_sequence | Efficient variable-length RNN processing | Any RNN with variable lengths |
fuse_conv_bn_eval | Merge Conv+BN for inference | Inference optimization |
nested_tensor | Variable-length batches without padding | Attention, NLP, Flash Attention |
clip_grad_norm_ | Prevent exploding gradients | Any training with transformers |
skip_init | Skip weight initialization | Loading large pretrained models |
Further Reading
torch.nn.utils.parametrizedocs: pytorch.org/docs/stable/generated/torch.nn.utils.parametrize.register_parametrization- Pruning tutorial: pytorch.org/tutorials/intermediate/pruning_tutorial
- NestedTensor: pytorch.org/docs/stable/nested.html
📓 Open Notebook — Interactive version
Source Files
parametrization.py— Weight parametrization — enforcing constraints like symmetry, orthogonality, and positivity on parameterspruning.py— Model pruning — making neural networks smaller by removing weightssequence_packing_and_nested.py— Sequence packing & nested tensors — efficient variable-length processingconv_bn_fusion.py— Conv-BN fusion & inference optimization utilities
Module 16: Activation Checkpointing — Trading Memory for Compute
Day 2 of the incremental learning series
The Problem: Activations Eat Your GPU Memory
When training a neural network, PyTorch stores all intermediate activations (outputs of each layer) during the forward pass. These are needed for the backward pass to compute gradients. For large models, activations consume far more memory than the model parameters themselves.
Example memory breakdown for a 7B parameter Transformer:
Model parameters: 14 GB (7B × 2 bytes in BF16)
Optimizer state: 28 GB (Adam: 2× params in FP32)
Activations: ~60 GB (scales with batch_size × seq_len × layers)
Total: ~102 GB — doesn't fit on an 80GB GPU!
Activation checkpointing solves this by not saving activations during forward. During backward, it recomputes them on the fly. This trades ~33% more compute for ~60% less memory.
Table of Contents
- How Activation Checkpointing Works
- Basic Usage: torch.utils.checkpoint
- checkpoint_sequential for Sequential Models
- Selective Activation Checkpointing (SAC)
- CheckpointPolicy Options
- Integration with torch.compile
- Practical Guidelines
1. How Activation Checkpointing Works
Normal Training (No Checkpointing)
Forward pass: Input → [Layer 1] → a₁ → [Layer 2] → a₂ → [Layer 3] → a₃ → Loss
save a₁ save a₂ save a₃
Backward pass: Uses saved a₁, a₂, a₃ to compute gradients
Memory: O(N) — stores all N layer activations
With Activation Checkpointing
Forward pass: Input → [Layer 1] → [Layer 2] → [Layer 3] → Loss
(discard) (discard) (discard)
Backward pass:
Need a₃ → Recompute: Input → Layer 1 → Layer 2 → Layer 3 → a₃ ✓
Need a₂ → Recompute: Input → Layer 1 → Layer 2 → a₂ ✓
Need a₁ → Recompute: Input → Layer 1 → a₁ ✓
Memory: O(1) per checkpointed segment
Compute: ~1.33× (one extra forward pass)
In practice, you checkpoint segments of the model (e.g., each Transformer layer), not the entire model. This gives a good memory/compute tradeoff.
2. Basic Usage
import torch
from torch.utils.checkpoint import checkpoint
class TransformerLayer(torch.nn.Module):
def __init__(self, d_model):
super().__init__()
self.attn = torch.nn.MultiheadAttention(d_model, 8, batch_first=True)
self.norm1 = torch.nn.LayerNorm(d_model)
self.ffn = torch.nn.Sequential(
torch.nn.Linear(d_model, 4 * d_model),
torch.nn.GELU(),
torch.nn.Linear(4 * d_model, d_model),
)
self.norm2 = torch.nn.LayerNorm(d_model)
def forward(self, x):
x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]
x = x + self.ffn(self.norm2(x))
return x
class TransformerModel(torch.nn.Module):
def __init__(self, d_model=512, n_layers=12):
super().__init__()
self.layers = torch.nn.ModuleList([
TransformerLayer(d_model) for _ in range(n_layers)
])
def forward(self, x, use_checkpoint=False):
for layer in self.layers:
if use_checkpoint:
# Checkpoint each layer — its activations are NOT saved
x = checkpoint(layer, x, use_reentrant=False)
else:
x = layer(x)
return x
Key Parameter: use_reentrant=False
Always use use_reentrant=False (the modern, recommended path):
use_reentrant=True(legacy): Uses a different autograd mechanism, has subtle bugs with certain opsuse_reentrant=False(recommended): Works correctly with all ops,torch.compile, and distributed training
3. checkpoint_sequential
For nn.Sequential models, there's a convenience wrapper:
from torch.utils.checkpoint import checkpoint_sequential
model = torch.nn.Sequential(
TransformerLayer(512),
TransformerLayer(512),
TransformerLayer(512),
TransformerLayer(512),
)
x = torch.randn(8, 32, 512, requires_grad=True)
# Divide into 2 segments — each segment is checkpointed
output = checkpoint_sequential(model, segments=2, input=x, use_reentrant=False)
segments controls how many checkpointed groups to split the sequential into. More segments = more memory savings but more recomputation.
4. Selective Activation Checkpointing (SAC)
The problem with basic checkpointing: It recomputes EVERYTHING, including expensive operations like matrix multiplies and attention. Ideally, we'd save expensive outputs and only recompute cheap ones (like activations, norms).
SAC lets you choose per-operation whether to save or recompute:
from torch.utils.checkpoint import (
checkpoint,
CheckpointPolicy,
create_selective_checkpoint_contexts,
)
# Policy function: decides per-op whether to save or recompute
def policy_fn(ctx, op, *args, **kwargs):
# Save expensive ops (matmul, attention)
if op in (torch.ops.aten.mm.default,
torch.ops.aten.bmm.default,
torch.ops.aten._scaled_dot_product_flash_attention.default):
return CheckpointPolicy.MUST_SAVE
# Recompute cheap ops (relu, add, norm, etc.)
return CheckpointPolicy.PREFER_RECOMPUTE
# Use with checkpoint's context_fn parameter
context_fn = create_selective_checkpoint_contexts(policy_fn)
x = checkpoint(
layer, x,
use_reentrant=False,
context_fn=context_fn,
)
Shortcut: Pass a List of Ops to Save
# Instead of a policy function, just list the ops you want to save
ops_to_save = [
torch.ops.aten.mm.default,
torch.ops.aten.bmm.default,
]
context_fn = create_selective_checkpoint_contexts(ops_to_save)
x = checkpoint(layer, x, use_reentrant=False, context_fn=context_fn)
5. CheckpointPolicy Options
| Policy | Behavior |
|---|---|
MUST_SAVE | Always save this op's output (never recompute) |
PREFER_SAVE | Save unless torch.compile decides otherwise |
MUST_RECOMPUTE | Always recompute (never save) |
PREFER_RECOMPUTE | Recompute unless torch.compile decides otherwise |
MUST_CPU_OFFLOAD | Save to CPU during forward, reload to GPU during backward |
PREFER_CPU_OFFLOAD | Offload unless torch.compile decides otherwise |
The PREFER_ variants allow torch.compile to override the decision based on its global optimization analysis. The MUST_ variants are strict.
6. Integration with torch.compile
Activation checkpointing works with torch.compile. When compiled, the compiler can make even smarter decisions about what to save vs recompute:
model = TransformerModel(d_model=512, n_layers=12)
# Compile with checkpointing
compiled_model = torch.compile(model)
x = torch.randn(8, 32, 512)
output = compiled_model(x, use_checkpoint=True)
When using PREFER_* policies with compiled code, the compiler may override your suggestions if it determines a different strategy is more efficient.
7. Practical Guidelines
When to Use Activation Checkpointing
| Scenario | Recommendation |
|---|---|
| Model fits in GPU memory | Don't checkpoint (unnecessary overhead) |
| OOM with desired batch size | Checkpoint every Transformer layer |
| Still OOM | Add selective checkpointing (save matmuls, recompute rest) |
| Still OOM | Combine with FSDP2, gradient accumulation, or CPU offload |
Rules of Thumb
- Checkpoint at the Transformer layer granularity — each
TransformerEncoderLayeror decoder layer is one checkpoint boundary - Always use
use_reentrant=False— the legacy reentrant mode has known issues - With SAC, save matmuls and attention — these are 10-100x more expensive than activations/norms
- Memory savings: ~50-70% activation memory reduction with basic checkpointing
- Compute overhead: ~30% more training time (one extra forward pass per checkpointed segment)
- Stacks with everything: Works with AMP, DDP, FSDP2, torch.compile
Memory Estimation
Without checkpointing:
Activation memory ≈ batch_size × seq_len × hidden_dim × num_layers × 2 bytes (BF16)
With checkpointing (per layer):
Activation memory ≈ batch_size × seq_len × hidden_dim × 2 × 2 bytes
(Only stores input/output of each checkpointed segment)
Savings = (num_layers - 2) / num_layers ≈ 90%+ for 24+ layer models
Further Reading
- PyTorch Activation Checkpointing Tutorial
- Selective Activation Checkpointing for torch.compile
- Source code:
torch/utils/checkpoint.py
📓 Open Notebook — Interactive version
Source Files
checkpointing_basics.py— Activation checkpointing — trading memory for compute with basic and selective checkpointing
Module 17: torch.compile Decorators & Control APIs
Day 3 of the incremental learning series
Beyond torch.compile() — Fine-Grained Compilation Control
Module 08 taught you the basics of torch.compile. This module dives into the decorator and control APIs that give you precise control over what gets compiled, how shapes are handled, and how to debug compilation issues.
Table of Contents
- Compiler Stances — Global Compilation Behavior
- torch.compiler.disable — Skip Compilation
- allow_in_graph / disallow_in_graph — Control Tracing
- substitute_in_graph — Replace Functions for Tracing
- mark_dynamic / mark_static — Shape Control
- graph_break / error_on_graph_break — Break Control
- assume_constant_result — Constant Folding
- comptime — Compile-Time Debugging
- torch._dynamo.explain — Understanding Compilation
- TORCH_LOGS — Logging and Debugging
- What's New: Recent Upstream Changes
1. Compiler Stances
Stances control the global compilation behavior — useful for debugging and gradual adoption:
import torch
torch.compiler.set_stance("default") # Normal compilation
torch.compiler.set_stance("force_eager") # Skip all compilation
torch.compiler.set_stance("eager_on_recompile") # Compile once, eager on recompile
torch.compiler.set_stance("fail_on_recompile") # Error on recompilation
torch.compiler.set_stance("eager_then_compile") # Eager first, compile on second call
# Context manager form
with torch.compiler.set_stance("force_eager"):
output = compiled_model(x) # Runs eagerly despite being compiled
| Stance | Behavior | Use Case |
|---|---|---|
default | Normal compilation | Production |
force_eager | Skip all compilation | Debugging, profiling eager |
eager_on_recompile | Compile once, eager on recompile | Avoid compile-time storms |
fail_on_recompile | Error on recompilation | CI, catch shape issues |
eager_then_compile | Eager first call, compile on second | Warmup tolerance |
2. Disable — Skip Compilation for Specific Functions
@torch.compiler.disable
def preprocessing(x):
"""Uses Python features that don't compile well."""
results = []
for i in range(x.shape[0]):
if x[i].item() > 0:
results.append(x[i] * 2)
else:
results.append(x[i])
return torch.stack(results)
@torch.compile
def model_fn(x):
x = preprocessing(x) # Runs eagerly (disabled)
return x.relu().sum() # This part is compiled
Non-recursive disable (only disables the decorated function, not functions it calls):
@torch.compiler.disable(recursive=False)
def outer(x):
return inner(x) # inner() will still be compiled
3. allow_in_graph / disallow_in_graph
@torch.compiler.allow_in_graph # Opaque node in graph
@torch._dynamo.disallow_in_graph # Forces graph break
@torch._dynamo.forbid_in_graph # Raises error during tracing
4. substitute_in_graph — Replace Functions for Tracing
def original_fn(x):
result = x.tolist() # Not traceable!
return torch.tensor(result)
def traceable_fn(x):
return x.clone()
torch.compiler.substitute_in_graph(original_fn, traceable_fn)
5. mark_dynamic / mark_static — Shape Control
x = torch.randn(batch_size, 512)
torch._dynamo.mark_dynamic(x, 0) # Dim 0 is dynamic — no recompile for different batch sizes
torch._dynamo.mark_static(x, 1) # Dim 1 is always 512
torch._dynamo.mark_static_address(weight) # Data pointer won't change
6. Graph Break Control
torch._dynamo.graph_break() # Force a graph break
torch.compile(fullgraph=True) # Error on any break
@torch._dynamo.error_on_graph_break
def must_be_one_graph(x):
return x + 1
7. assume_constant_result
@torch._dynamo.assume_constant_result
def get_config():
return load_config()["lr"] # Folded to constant at compile time
8. comptime — Compile-Time Debugging
from torch._dynamo.comptime import comptime
@torch.compile
def fn(x):
comptime.breakpoint() # pdb during COMPILATION
return x + 1
# In pdb: ctx.print_locals(), ctx.print_graph(), ctx.print_bt()
9. torch._dynamo.explain
explanation = torch._dynamo.explain(my_fn)(torch.randn(10))
print(explanation) # Shows graphs, breaks, guards
10. TORCH_LOGS — Logging and Debugging
TORCH_LOGS="graph_breaks" python train.py # Graph break reasons
TORCH_LOGS="guards,recompiles" python train.py # Guard failures
TORCH_LOGS="graph_code" python train.py # Captured FX graph
TORCH_LOGS="output_code" python train.py # Generated Triton/C++
TORCH_LOGS="+dynamo" python train.py # Full debug
TORCH_TRACE=/tmp/trace python train.py # Structured tracing
11. What's New: Recent Upstream Changes (June 4-8, 2026)
194 commits landed on PyTorch main in the last 4 days:
QuACK GEMM Kernels Vendored
PyTorch now vendors the QuACK library from Dao-AILab — high-performance CuTeDSL GEMM epilogue adapters with fused RMSNorm. Located at torch/_vendor/quack/.
CUPTI Monitor — Continuous GPU Profiling
New torch.profiler._cupti_monitor for continuous CUPTI activity monitoring across the entire program (not just a profiling window).
Optimized _foreach_mm — Grouped GEMMs
New Python override dispatching to nvmath cublasLt grouped GEMM (bf16) or CUTLASS. At torch/_native/ops/foreach_mm/.
AArch64 torch.compile
Armv9-A target support — compiled models now work on ARM servers and edge devices.
DTensor Autogen Ops
Auto-generated sharding strategies expanding DTensor op coverage. At torch/distributed/tensor/_ops/autogen.py.
NCCL Symmetric Memory Registration
External NCCL comm registration for symmetric memory at torch/distributed/_symmetric_memory/_nccl.py.
Inductor Heuristics Module
Refactored Triton template heuristics into torch/_inductor/heuristics/.
Quick Reference
| API | What It Does |
|---|---|
set_stance() | Global compilation behavior |
@disable | Skip compilation for a function |
@allow_in_graph | Treat as opaque graph node |
substitute_in_graph() | Replace with traceable version |
mark_dynamic() | Declare dynamic dimension |
mark_static() | Declare static dimension |
fullgraph=True | Error on any graph break |
graph_break() | Force a graph break |
explain() | Get compilation report |
comptime.breakpoint() | Debug during compilation |
CompileCounter | Count compilations in tests |
EagerAndRecordGraphs | Inspect captured FX graphs |
No dedicated notebook — covered in Module 08 notebook
Source Files
compile_control.py— torch.compile decorators & control APIs — fine-grained control over what gets compiled and how
Module 18: torch.package — Self-Contained Model Packaging
Day 4 of the incremental learning series
The Problem: "It Works on My Machine"
You train a model. You save model.pt. You send it to a colleague. It fails because:
- They don't have the same version of your custom modules
- An import path changed between your environments
- A dependency you forgot about isn't installed on their machine
torch.package solves this by bundling the model and all its Python dependencies into a single .pt archive.
Table of Contents
- What is torch.package?
- PackageExporter — Creating Packages
- PackageImporter — Loading Packages
- Module Actions: intern, extern, mock, deny
- Packaging Models with Weights
- Inspecting Package Contents
- Re-Packaging and Dependency Analysis
- torch.package vs torch.save vs torch.export
- Practical Workflow
- What's New Upstream (June 8-9, 2026)
1. What is torch.package?
torch.package creates a hermetic zip archive containing:
- Your model's Python source code (the actual
.pyfiles) - The model's pickled state (weights, buffers, config)
- A manifest of external dependencies
When someone loads the package, it uses its own import system — code is loaded from inside the archive, not from the local Python installation. This means:
- The exact code you packaged runs, regardless of what's installed locally
- Only explicitly listed external dependencies are loaded from the system
- No "accidental" dependencies can sneak in
┌──────────────────────────────────┐
│ my_model.pt (zip) │
├──────────────────────────────────┤
│ .data/ │
│ model.pkl (pickled) │
│ weights.pt (tensors) │
│ my_module/ │
│ model.py (source) │
│ layers.py (source) │
│ config.py (source) │
│ extern_modules (manifest) │
│ torch │
│ numpy │
└──────────────────────────────────┘
2. PackageExporter — Creating Packages
from torch.package import PackageExporter
# Create a package
with PackageExporter("my_model.pt") as exporter:
# INTERN: Include this module's source inside the package
exporter.intern("my_module.**")
# EXTERN: This module is expected to exist on the target machine
exporter.extern("torch.**")
exporter.extern("numpy.**")
# MOCK: Replace this module with a stub (for unused optional deps)
exporter.mock("matplotlib.**")
# Save the model object
exporter.save_pickle("model", "model.pkl", my_model)
The Four Module Actions
| Action | What It Does | When to Use |
|---|---|---|
intern(pattern) | Bundle the module's source code into the package | Your own code, custom modules |
extern(pattern) | Expect the module to be installed on the target machine | PyTorch, NumPy, standard libs |
mock(pattern) | Replace with a stub that returns MockedObject | Unused optional dependencies |
deny(pattern) | Error if this module is encountered | Known-bad dependencies |
Patterns use glob syntax: "my_module.**" matches my_module and all submodules.
3. PackageImporter — Loading Packages
from torch.package import PackageImporter
# Load a package
importer = PackageImporter("my_model.pt")
# Load the pickled model
model = importer.load_pickle("model", "model.pkl")
# The model runs using code FROM the package, not your local installation
output = model(torch.randn(1, 3, 224, 224))
# Import a module from the package (hermetic import)
my_config = importer.import_module("my_module.config")
Key property: The loaded code runs in an isolated namespace. If my_module/model.py inside the package says import my_module.layers, it loads layers.py from the package, not from your filesystem.
4. Module Actions in Detail
intern — Bundle Source Code
# Include specific modules
exporter.intern("my_project.models.**")
exporter.intern("my_project.utils.**")
# Include everything in your project
exporter.intern("my_project.**")
What happens: The .py source files are copied into the zip archive. When loaded, Python reads them from the archive.
extern — External Dependencies
# Standard externals
exporter.extern("torch.**")
exporter.extern("torchvision.**")
exporter.extern("numpy.**")
# Stdlib is automatically handled, but you can be explicit:
exporter.extern("os")
exporter.extern("json")
What happens: A list of external modules is saved in extern_modules. When loading, these are imported from the system Python.
mock — Stub Out Dependencies
# Mock out modules that aren't needed at inference time
exporter.mock("wandb.**") # Logging library
exporter.mock("matplotlib.**") # Plotting
exporter.mock("tensorboard.**") # TensorBoard
What happens: A lightweight _mock module replaces the real one. Any attribute access on a mocked module returns MockedObject.
deny — Prevent Inclusion
# Error if these are encountered
exporter.deny("secret_module.**")
exporter.deny("credentials.**")
5. Packaging Models with Weights
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super().__init__()
self.fc1 = nn.Linear(input_dim, hidden_dim)
self.relu = nn.ReLU()
self.fc2 = nn.Linear(hidden_dim, output_dim)
def forward(self, x):
return self.fc2(self.relu(self.fc1(x)))
# Train your model...
model = MyModel(784, 256, 10)
# Package it with all dependencies
with PackageExporter("my_model_package.pt") as exporter:
exporter.intern("__main__") # Include the current module
exporter.extern("torch.**")
exporter.extern("numpy.**")
# Save model
exporter.save_pickle("model", "model.pkl", model)
# You can also save arbitrary data
exporter.save_pickle("config", "config.pkl", {
"input_dim": 784,
"hidden_dim": 256,
"output_dim": 10,
"version": "1.0",
})
# Save raw text/binary files
exporter.save_text("info", "README.txt", "My model v1.0")
# Load on another machine
importer = PackageImporter("my_model_package.pt")
model = importer.load_pickle("model", "model.pkl")
config = importer.load_pickle("config", "config.pkl")
readme = importer.load_text("info", "README.txt")
6. Inspecting Package Contents
importer = PackageImporter("my_model_package.pt")
# View the file structure
file_structure = importer.file_structure()
print(file_structure)
# Prints a tree of all files in the archive
# List all extern modules
print(file_structure.has_file("extern_modules"))
You can also inspect with standard zip tools since .pt files are zip archives:
unzip -l my_model_package.pt
python -m zipfile -l my_model_package.pt
7. Re-Packaging
You can load a package and re-export it (e.g., to add/remove dependencies):
importer = PackageImporter("model_v1.pt")
with PackageExporter("model_v2.pt", importer=(importer,)) as exporter:
exporter.intern("my_module.**")
exporter.extern("torch.**")
# Load the old model and save it in the new package
model = importer.load_pickle("model", "model.pkl")
exporter.save_pickle("model", "model.pkl", model)
8. torch.package vs torch.save vs torch.export
| Feature | torch.save | torch.package | torch.export |
|---|---|---|---|
| Saves weights | Yes | Yes | Yes |
| Saves code | No | Yes (source) | Yes (graph) |
| Hermetic loading | No | Yes | Yes |
| Python control flow | N/A | Full | Limited |
| Works cross-version | Fragile | Robust | Robust |
| Deployment target | Python | Python | C++, mobile, ONNX |
| File format | pickle | zip (with source) | PT2 archive |
| Speed | Fast | Medium | Compile required |
When to use each:
torch.save: Quick checkpoints during trainingtorch.package: Ship Python models with all dependencies, research sharingtorch.export: Production deployment, C++ inference, mobile
9. Practical Workflow
Research → Deployment Pipeline
# 1. Researcher trains model (research/train.py)
model = train_my_model()
# 2. Researcher packages model with all custom code
with PackageExporter("model_v1.pt") as pe:
pe.intern("my_research_code.**")
pe.extern("torch.**")
pe.extern("torchvision.**")
pe.mock("wandb.**") # Don't need logging in production
pe.mock("matplotlib.**")
pe.save_pickle("model", "model.pkl", model)
# 3. Engineer loads on a different machine (no my_research_code installed!)
importer = PackageImporter("model_v1.pt")
model = importer.load_pickle("model", "model.pkl")
output = model(input_data) # Just works!
Tips
- Always extern
torch— it must match the installed version - Mock unused dependencies — logging, visualization, experiment tracking
- Test the package — load it in a clean environment to verify
- Version your packages — include version info as saved text/pickle
- Inspect before shipping — use
file_structure()to verify contents
10. Upstream Updates (June 8-9, 2026)
Recent PyTorch main commits (since last update):
- Inductor CUTLASS GELU fusion — Folding decomposed GELU back into native CUTLASS functor for better performance (
#185824) - Inductor TP pattern fusion — Fusing slice-cat tensor parallel collective patterns (
#184911) - Dynamo Python 3.15 support — Build dynamo with Python 3.15, including updated
IMPORT_NAMEgeneration (#186402) - Dynamo operator support — Added
divmod,remainder,true_divide,floor_divideoperators (#185652-#185655) - Deterministic topk —
torch.topknow respectstorch.use_deterministic_algorithms()(#186653) - XPU oneDNN LSTM — Intel GPU LSTM inference via oneDNN primitives (
#185531) - Stable ABI generator — New
torch/csrc/stable/generator.hfor stable C API
Further Reading
- Source:
torch/package/package_exporter.py,torch/package/package_importer.py - PyTorch docs: torch.package
- Tutorial: Loading a TorchScript Model in C++ (comparison)
No dedicated notebook — see Practical Workflow above
Source Files
packaging_models.py— torch.package — bundling models and Python source code into hermetic archives
Module 19: __torch_function__ & __torch_dispatch__ — Tensor Subclassing
Day 5 of the incremental learning series
Why This Matters
Every time you call torch.add(x, y) or x + y, PyTorch checks: does this tensor have a custom dispatch protocol? Two protocols exist:
__torch_function__— Python-level override. Intercepts any PyTorch function call. Like NumPy's__array_function__.__torch_dispatch__— Lower-level override. Intercepts at the ATen operator level (after decompositions). More powerful, used by DTensor, FakeTensor, and torch.compile internals.
These are the extension points that power:
- DTensor (distributed tensor) — sharding logic via
__torch_dispatch__ - FakeTensor (torch.compile) — shape-only tensors via
__torch_dispatch__ - Logging/profiling — intercept all operations without modifying model code
- Custom tensor types — sparse, quantized, masked tensors
- Unit conversion — tensors that carry physical units
Table of Contents
__torch_function__— Python-Level Override- TorchFunctionMode — Override Without Subclassing
__torch_dispatch__— ATen-Level Override- TorchDispatchMode — Mode-Based Dispatch
- Practical Examples
- When to Use Which Protocol
- Upstream Updates (June 9-10, 2026)
1. __torch_function__ — Python-Level Override
When you define __torch_function__ on a class, PyTorch calls it instead of the normal implementation for any torch function that receives your object as an argument.
import torch
class ScaledTensor:
"""A tensor wrapper that tracks a scaling factor."""
def __init__(self, data, scale=1.0):
self.data = data
self.scale = scale
def __repr__(self):
return f"ScaledTensor(data={self.data}, scale={self.scale})"
@classmethod
def __torch_function__(cls, func, types, args=(), kwargs=None):
"""Called for any torch.* function involving this type."""
if kwargs is None:
kwargs = {}
# Extract ScaledTensors from args, replace with raw data
new_args = []
scale = 1.0
for a in args:
if isinstance(a, ScaledTensor):
new_args.append(a.data)
scale = a.scale
else:
new_args.append(a)
# Call the original function on raw tensors
result = func(*new_args, **kwargs)
# Wrap the result back
if isinstance(result, torch.Tensor):
return ScaledTensor(result, scale)
return result
# Usage
x = ScaledTensor(torch.tensor([1.0, 2.0, 3.0]), scale=0.5)
y = ScaledTensor(torch.tensor([4.0, 5.0, 6.0]), scale=0.5)
z = torch.add(x, y) # Calls ScaledTensor.__torch_function__!
print(z) # ScaledTensor(data=tensor([5., 7., 9.]), scale=0.5)
w = torch.mul(x, 2) # Also intercepted
print(w) # ScaledTensor(data=tensor([2., 4., 6.]), scale=0.5)
How the Protocol Works
- PyTorch checks if any argument has
__torch_function__ - If yes, it calls
__torch_function__(func, types, args, kwargs)where: func— the original function (e.g.,torch.add)types— tuple of types that implement__torch_function__args— positional argumentskwargs— keyword arguments- Your implementation decides what to do and returns the result
2. TorchFunctionMode — Override Without Subclassing
Modes let you override all torch operations within a context manager — no tensor subclass needed:
from torch.overrides import TorchFunctionMode
class LoggingMode(TorchFunctionMode):
"""Logs every torch operation."""
def __torch_function__(self, func, types, args=(), kwargs=None):
if kwargs is None:
kwargs = {}
print(f" Called: {func.__module__}.{func.__name__}")
return func(*args, **kwargs)
# All torch ops inside the context are logged
with LoggingMode():
x = torch.randn(3, 4) # Logged
y = x + 1 # Logged
z = torch.relu(y) # Logged
w = z.mean() # Logged
Use Cases for Modes
- Logging/debugging — see every operation a model performs
- Profiling — count operations, measure shapes
- Mocking — override factory functions (torch.randn, torch.zeros)
- Validation — check all inputs are on the correct device
3. __torch_dispatch__ — ATen-Level Override
__torch_dispatch__ intercepts at a lower level — after Python function dispatch, at the ATen operator level. This is where the real computation happens.
import torch
from torch.utils._python_dispatch import return_and_correct_aliasing
class LoggingTensor(torch.Tensor):
"""A tensor subclass that logs all ATen operations."""
@staticmethod
def __new__(cls, data):
return torch.Tensor._make_subclass(cls, data)
@classmethod
def __torch_dispatch__(cls, func, types, args, kwargs=None):
"""Called for every ATen operator."""
if kwargs is None:
kwargs = {}
# Unwrap LoggingTensors to plain tensors
def unwrap(t):
return t.elem if isinstance(t, LoggingTensor) else t
print(f" dispatch: {func.__name__}")
# Call the actual ATen op
result = func(*args, **kwargs)
return result
x = LoggingTensor(torch.randn(3, 4))
y = x + 1 # Dispatches through __torch_dispatch__
z = y.relu() # Also dispatched
Key Differences from __torch_function__
| Feature | __torch_function__ | __torch_dispatch__ |
|---|---|---|
| Level | Python API | ATen operators |
| Input | torch.add, torch.nn.functional.relu | aten.add.Tensor, aten.relu.default |
| Decomposition | Before | After (sees primitive ops) |
| Used by | Custom wrappers, logging | DTensor, FakeTensor, torch.compile |
| Subclass required | No (can use any class) | Yes (must subclass torch.Tensor) |
4. TorchDispatchMode — Mode-Based Dispatch
Like TorchFunctionMode, but at the ATen operator level:
from torch.utils._python_dispatch import TorchDispatchMode
class CountOps(TorchDispatchMode):
"""Count all ATen operations in a scope."""
def __init__(self):
super().__init__()
self.ops = {}
def __torch_dispatch__(self, func, types, args, kwargs=None):
name = str(func.name())
self.ops[name] = self.ops.get(name, 0) + 1
if kwargs is None:
kwargs = {}
return func(*args, **kwargs)
# Count ops in a forward pass
counter = CountOps()
model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.ReLU(),
torch.nn.Linear(20, 5),
)
with counter:
output = model(torch.randn(4, 10))
print("Operations performed:")
for op, count in sorted(counter.ops.items()):
print(f" {op}: {count}x")
5. Practical Examples
Example 1: Device-Checking Mode
class DeviceCheckMode(TorchFunctionMode):
"""Error if any tensor is on the wrong device."""
def __init__(self, expected_device):
self.expected_device = torch.device(expected_device)
def __torch_function__(self, func, types, args=(), kwargs=None):
if kwargs is None:
kwargs = {}
for a in args:
if isinstance(a, torch.Tensor) and a.device != self.expected_device:
raise RuntimeError(
f"{func.__name__}: tensor on {a.device}, "
f"expected {self.expected_device}"
)
return func(*args, **kwargs)
# This catches CPU/GPU mismatches early
with DeviceCheckMode("cpu"):
x = torch.randn(3, 4) # OK
y = x + 1 # OK
Example 2: Shape Logging Mode
class ShapeTracer(TorchDispatchMode):
"""Track input/output shapes of all ops."""
def __init__(self):
super().__init__()
self.traces = []
def __torch_dispatch__(self, func, types, args, kwargs=None):
if kwargs is None:
kwargs = {}
result = func(*args, **kwargs)
in_shapes = [a.shape for a in args if isinstance(a, torch.Tensor)]
out_shape = result.shape if isinstance(result, torch.Tensor) else "N/A"
self.traces.append((func.name(), in_shapes, out_shape))
return result
Example 3: Tensor with Units (Physics)
class UnitTensor:
"""Tensor that tracks physical units (e.g., meters, seconds)."""
def __init__(self, data, unit=""):
self.data = data
self.unit = unit
def __repr__(self):
return f"{self.data} [{self.unit}]"
@classmethod
def __torch_function__(cls, func, types, args=(), kwargs=None):
if kwargs is None:
kwargs = {}
tensors = [a for a in args if isinstance(a, UnitTensor)]
raw_args = [a.data if isinstance(a, UnitTensor) else a for a in args]
result = func(*raw_args, **kwargs)
if func == torch.mul and len(tensors) == 2:
unit = f"{tensors[0].unit}*{tensors[1].unit}"
elif func == torch.div and len(tensors) == 2:
unit = f"{tensors[0].unit}/{tensors[1].unit}"
else:
unit = tensors[0].unit if tensors else ""
if isinstance(result, torch.Tensor):
return UnitTensor(result, unit)
return result
distance = UnitTensor(torch.tensor(100.0), "m")
time = UnitTensor(torch.tensor(9.58), "s")
speed = torch.div(distance, time)
print(f"Speed: {speed}") # 10.44 [m/s]
6. When to Use Which Protocol
| Scenario | Use |
|---|---|
| Log/trace all torch function calls | TorchFunctionMode |
| Custom tensor wrapper (non-subclass) | __torch_function__ |
| Override ATen ops for a tensor subclass | __torch_dispatch__ |
| Count/profile ops at ATen level | TorchDispatchMode |
| Build a new tensor type (like DTensor) | __torch_dispatch__ |
| Intercept factory functions (torch.randn) | TorchFunctionMode |
| Works with torch.compile | __torch_dispatch__ (preferred) |
The Dispatch Stack
User code: torch.nn.functional.relu(x)
|
v
__torch_function__ <- Python-level, sees relu
|
v
Decompositions <- relu -> clamp(x, min=0)
|
v
__torch_dispatch__ <- ATen-level, sees aten.clamp.default
|
v
C++ dispatcher <- Routes to CPU/CUDA/etc. kernel
7. Upstream Updates (June 9-10, 2026)
Recent PyTorch main commits:
- FSDP2 separate reduce-scatter group — Opt-in all-gather/reduce-scatter overlap via
set_separate_reduce_scatter_group(#186335) - Activation offloading pinned memory pool — Dedicated pinned memory pool for activation offloading ops (
#186162) - Activation offloading stride preservation — Preserves original tensor strides across offload/reload (
#186396) - BERT SDPA pattern on CUDA — Enables BERT attention pattern for SDPA on CUDA (
#184417) - DTensor group_norm fix — Fixes crash when weight=None in group_norm under DTensor (
#184819) - Pipeline parallel backward fix — Fixes None gradient handling in pipeline backward send/recv (
#182182) - Torch.cuda.stream round-trip — Dynamo now correctly handles
torch.cuda.streamcontext managers across graph breaks (#184487) - TORCH_TRACE fork-safety — Structured tracing logs now preserved across forks (
#184772) - Open registration profiler — Activity profiler support for custom backend devices via open registration (
new test files)
Quick Reference
# __torch_function__ -- Python-level override (any class)
class MyType:
@classmethod
def __torch_function__(cls, func, types, args, kwargs=None):
...
# TorchFunctionMode -- scope-based override (no subclass needed)
class MyMode(TorchFunctionMode):
def __torch_function__(self, func, types, args=(), kwargs=None):
...
# __torch_dispatch__ -- ATen-level override (tensor subclass)
class MyTensor(torch.Tensor):
@classmethod
def __torch_dispatch__(cls, func, types, args, kwargs=None):
...
# TorchDispatchMode -- scope-based ATen override
class MyDispatchMode(TorchDispatchMode):
def __torch_dispatch__(self, func, types, args, kwargs=None):
...
Further Reading
- Source:
torch/overrides.py(torch_function),torch/utils/_python_dispatch.py(torch_dispatch) - Extending PyTorch docs
__torch_function__protocol
No dedicated notebook — see examples in torch_function_examples.py
Source Files
torch_function_examples.py— __torch_function__ and __torch_dispatch__ — overriding how PyTorch operations work on custom types
Module 20: torch.backends — Performance Tuning
Prerequisites: Module 07 (Training Pipelines)
Time: ~2 hours
Level: Intermediate → Advanced
Overview
PyTorch's torch.backends module is a configuration layer that controls hardware-specific optimizations. These settings determine which algorithms run under the hood for convolutions, matrix multiplications, attention, and parallelism—often yielding 2–10x speedups with a single line of code.
Most tutorials never mention these knobs. This module changes that.
Files in This Module
| File | Description |
|---|---|
backends_tuning.py | Runnable script demonstrating all backend settings |
1. What Are torch.backends?
torch.backends exposes runtime configuration for the hardware libraries PyTorch uses:
torch.backends
├── cudnn # NVIDIA cuDNN (convolutions, RNNs)
├── cuda # NVIDIA CUDA (matmul, SDPA)
├── mkldnn # Intel oneDNN (CPU conv, linear)
├── mkl # Intel MKL (BLAS/LAPACK)
├── openmp # OpenMP (CPU threading)
├── opt_einsum # Optimized einsum contraction
└── mps # Apple Metal (M1/M2/M3)
Each backend exposes flags you can toggle at runtime. No recompilation needed.
Key principle: backends control the how, not the what. The mathematical operation stays the same; the algorithm, precision, or parallelism strategy changes.
2. torch.backends.cudnn
cuDNN is NVIDIA's deep neural network library. It provides optimized implementations for convolutions, pooling, normalization, and RNNs.
2.1 cudnn.enabled
torch.backends.cudnn.enabled # default: True
When True, PyTorch uses cuDNN for supported operations. Disabling it falls back to slower native implementations. You almost never want to disable this.
2.2 cudnn.benchmark
torch.backends.cudnn.benchmark = True # default: False
What it does: Before the first convolution at each input size, cuDNN runs multiple algorithm variants and selects the fastest. Results are cached for the session.
When it helps:
- Fixed input sizes (standard training with constant batch size and image dimensions)
- Repeated convolutions with the same shapes
- Training CNNs (ResNet, EfficientNet, etc.)
When it hurts:
- Variable input sizes (NLP with different sequence lengths, object detection with varying image sizes)
- Short-lived scripts (benchmarking overhead > savings)
- First iteration is slower (paying the auto-tuning cost)
# Typical training setup for CNNs with fixed input
torch.backends.cudnn.benchmark = True
# Disable for variable-size inputs
torch.backends.cudnn.benchmark = False
2.3 cudnn.deterministic
torch.backends.cudnn.deterministic = True # default: False
Forces cuDNN to use deterministic algorithms. Non-deterministic algorithms are often faster because they can exploit parallelism without worrying about reduction order.
Tradeoffs:
| Deterministic=False | Deterministic=True | |
|---|---|---|
| Speed | Faster | Slower (sometimes 2–3x) |
| Reproducibility | Run-to-run variance | Bit-exact results |
| Use case | Normal training | Debugging, CI, research requiring exact reproduction |
For full determinism, also call torch.use_deterministic_algorithms(True).
2.4 cudnn.allow_tf32
torch.backends.cudnn.allow_tf32 = True # default: True (PyTorch 2.x)
Controls whether cuDNN can use TF32 precision for convolutions on Ampere+ GPUs. See Section 9 for details on TF32.
3. torch.backends.cuda
Configuration for CUDA operations beyond cuDNN (matrix multiplications, attention).
3.1 matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True # default: True (PyTorch 2.x)
Allows cuBLAS to use TF32 for float32 matrix multiplications. On Ampere (A100) and later, this can double matmul throughput with minimal precision loss.
3.2 matmul.allow_fp16_reduced_precision_reduction
torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction = True
Controls whether fp16 GEMMs can use reduced precision for internal accumulation. Faster but potentially less accurate for very large matrices.
3.3 matmul.allow_bf16_reduced_precision_reduction
torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = True
Same as above but for bfloat16 operations. Relevant for training on Ampere+ GPUs where bf16 is preferred over fp16.
3.4 Flash SDP (Scaled Dot-Product Attention)
# Check if Flash Attention is enabled
torch.backends.cuda.flash_sdp_enabled() # True/False
# Enable/disable Flash Attention
torch.backends.cuda.enable_flash_sdp(True)
# Other SDPA backends
torch.backends.cuda.enable_mem_efficient_sdp(True)
torch.backends.cuda.enable_math_sdp(True)
Flash Attention is the fastest SDPA implementation for most cases. You might disable it for debugging or to force a specific backend.
3.5 preferred_blas_library
torch.backends.cuda.preferred_blas_library() # get current
torch.backends.cuda.preferred_blas_library("cublas") # set preference
# Options: "cublas", "cublaslt", "hipblaslt" (AMD)
Select which BLAS library handles matrix multiplications. cuBLASLt supports more fused operations and epilogues.
4. torch.backends.mkldnn
Intel oneDNN (formerly MKL-DNN) provides optimized CPU implementations for convolutions, linear layers, and batch normalization.
torch.backends.mkldnn.enabled # default: True on x86
# Check if available
torch.backends.mkldnn.is_available()
When enabled, PyTorch automatically routes supported operations through oneDNN for faster CPU execution. This is especially impactful for inference on Intel CPUs.
5. torch.backends.mkl
Intel Math Kernel Library (MKL) provides optimized BLAS/LAPACK routines.
# Enable verbose mode to see which MKL routines are called
torch.backends.mkl.verbose(torch.backends.mkl.VERBOSE_ON)
# VERBOSE_OFF, VERBOSE_ON
Use case: Profiling CPU performance to confirm MKL is being used for linear algebra operations.
6. torch.backends.openmp
Controls OpenMP threading for CPU parallelism.
import torch
# Get/set number of threads
torch.get_num_threads() # current thread count
torch.set_num_threads(4) # set to 4 threads
# Also controlled via environment variable (before import):
# OMP_NUM_THREADS=4
# MKL_NUM_THREADS=4
Guidelines:
- For training: set to number of physical cores (not hyperthreads)
- For inference with batching: reduce threads to allow concurrent requests
- For DataLoader workers: reduce to avoid oversubscription
import os
os.environ["OMP_NUM_THREADS"] = "4" # must be set BEFORE importing torch
7. torch.backends.opt_einsum
Optimized path planning for torch.einsum operations.
torch.backends.opt_einsum.enabled # default: True if opt_einsum installed
torch.backends.opt_einsum.strategy # default: "auto"
# Strategies: "auto", "greedy", "optimal", "branch-all", "branch-2", "dp"
What it does: For complex einsum expressions with 3+ tensors, finding the optimal contraction order is NP-hard. opt_einsum uses heuristics to find near-optimal orderings that can be orders of magnitude faster.
# Example: without optimization, this could be O(N^5)
# With optimal contraction order, it's O(N^3)
result = torch.einsum("ij,jk,kl->il", A, B, C)
Strategies:
| Strategy | Speed | Quality | Use case |
|---|---|---|---|
"greedy" | Fast | Good | Default for most cases |
"optimal" | Slow | Best | Small expressions (<10 indices) |
"dp" | Medium | Good | Balanced for larger expressions |
"auto" | Adaptive | Best tradeoff | Recommended |
8. torch.backends.mps
Apple Metal Performance Shaders backend for M1/M2/M3/M4 chips.
torch.backends.mps.is_available() # True on Apple Silicon with macOS 12.3+
torch.backends.mps.is_built() # True if PyTorch was compiled with MPS
MPS provides GPU acceleration on Apple hardware. While not as fast as CUDA for large models, it enables GPU training on Apple laptops and desktops.
if torch.backends.mps.is_available():
device = torch.device("mps")
tensor = torch.randn(1000, 1000, device=device)
9. TF32 Precision
What is TF32?
TF32 (TensorFloat-32) is a 19-bit floating point format introduced with NVIDIA Ampere (A100):
Format comparison:
FP32: 1 sign + 8 exponent + 23 mantissa = 32 bits
TF32: 1 sign + 8 exponent + 10 mantissa = 19 bits
FP16: 1 sign + 5 exponent + 10 mantissa = 16 bits
BF16: 1 sign + 8 exponent + 7 mantissa = 16 bits
TF32 has the range of FP32 (8-bit exponent) with the precision of FP16 (10-bit mantissa). It's used internally by tensor cores—inputs are read as FP32, rounded to TF32 for computation, and results are accumulated in FP32.
How TF32 Affects Operations
| Operation | Setting | Speedup (A100) | Precision Loss |
|---|---|---|---|
| Conv2d | cudnn.allow_tf32=True | ~2x | ~0.1% relative |
| matmul | cuda.matmul.allow_tf32=True | ~2–3x | ~0.1% relative |
| Linear | Via matmul setting | ~2–3x | ~0.1% relative |
Enable/Disable TF32
# Enable TF32 everywhere (recommended for training)
torch.backends.cudnn.allow_tf32 = True
torch.backends.cuda.matmul.allow_tf32 = True
# Disable TF32 for full FP32 precision (validation, debugging)
torch.backends.cudnn.allow_tf32 = False
torch.backends.cuda.matmul.allow_tf32 = False
When to Disable TF32
- Numerical validation (comparing against reference implementations)
- Scientific computing requiring full FP32 precision
- Debugging convergence issues
- Unit tests checking exact numerical equality
10. torch.set_float32_matmul_precision()
A high-level API that controls matmul precision globally:
torch.set_float32_matmul_precision("highest") # No TF32, full FP32
torch.set_float32_matmul_precision("high") # TF32 on Ampere+
torch.set_float32_matmul_precision("medium") # TF32 + reduced precision reductions
| Level | Matmul Precision | Speed | Use Case |
|---|---|---|---|
"highest" | Full FP32 | Baseline | Debugging, validation |
"high" | TF32 on Ampere+ | ~2–3x faster | Standard training |
"medium" | TF32 + BF16 reductions | Fastest | Large-batch training |
# Recommended: set once at the top of your training script
torch.set_float32_matmul_precision("high")
This is what torch.compile hints at when it logs: "TensorFloat32 tensor cores for float32 matrix multiplication available but not enabled."
11. torch.backends.flags() Context Manager
Temporarily override backend settings within a scope:
with torch.backends.flags(
cudnn_benchmark=True,
cudnn_deterministic=False,
cudnn_enabled=True,
allow_tf32=True
):
# These settings active only inside this block
output = model(input)
# Original settings restored here
Use cases:
- Temporarily disabling TF32 for a validation step
- Enabling benchmark mode for a specific layer
- Writing tests that need specific backend states
12. Performance Checklist
Training Settings (Maximum Speed)
# Set at top of training script
torch.backends.cudnn.benchmark = True
torch.backends.cudnn.deterministic = False
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")
Inference Settings (Maximum Throughput)
# Set before inference
torch.backends.cudnn.benchmark = True # if input sizes are fixed
torch.backends.cudnn.deterministic = False
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")
torch.set_num_threads(num_physical_cores)
Debugging/Reproducibility Settings
# Set for exact reproducibility
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
torch.use_deterministic_algorithms(True)
torch.manual_seed(42)
Quick Reference Table
| Setting | Training | Inference | Debug |
|---|---|---|---|
cudnn.benchmark | ✅ (fixed sizes) | ✅ (fixed sizes) | ❌ |
cudnn.deterministic | ❌ | ❌ | ✅ |
cudnn.allow_tf32 | ✅ | ✅ | ❌ |
cuda.matmul.allow_tf32 | ✅ | ✅ | ❌ |
float32_matmul_precision | "high" | "high" | "highest" |
cudnn.enabled | ✅ | ✅ | ✅ |
use_deterministic_algorithms | ❌ | ❌ | ✅ |
13. Upstream Updates (June 10–11, 2026)
Recent PyTorch commits that affect backend behavior and performance:
cuDNN SDPA d=256 Support (#185553)
cuDNN's SDPA backend now supports head dimension d=256, enabling Flash Attention-like performance for larger attention heads without falling back to the math kernel. This benefits models using large head dimensions (e.g., certain MoE architectures).
ARM scatter/gather Optimization (#156161)
Scatter and gather operations now use optimized ARM NEON intrinsics on aarch64, improving CPU performance on ARM servers (AWS Graviton, Apple M-series) by 2–4x for these ops.
Dynamo Polyfills for itertools (#186240)
torch.compile now handles itertools.chain, itertools.islice, and other itertools functions as graph-safe polyfills, reducing graph breaks in models that use standard library iteration patterns.
DTensor Single-Dim Strategies Migration (#186667)
Internal migration of DTensor sharding strategies to a single-dimension representation, improving compilation speed and reducing memory for distributed models using DeviceMesh.
AOTI torch.cond/while_loop Support (#184736)
AOTInductor (the ahead-of-time compiler) now supports torch.cond and torch.while_loop control flow operations, enabling export of models with conditional branches and loops.
Summary
torch.backends — what to remember:
1. cudnn.benchmark = True → auto-tune conv algorithms (fixed sizes)
2. allow_tf32 = True → 2–3x matmul/conv speedup on Ampere+
3. set_float32_matmul_precision("high") → same as TF32, cleaner API
4. set_num_threads(N) → match physical cores for CPU
5. opt_einsum.strategy → optimize multi-tensor contractions
6. deterministic = True → reproducibility at the cost of speed
7. backends.flags() → scope backend settings temporarily
Further Reading
Notebook: 20_backends_tuning.ipynb
Source Files
backends_tuning.py— Runnable script demonstrating all backend settings
Module 21: CUDA Graphs — Eliminating CPU Launch Overhead
Day 7 of the incremental learning series
Table of Contents
- What Are CUDA Graphs?
- Why CUDA Graphs Matter
- Basic API: torch.cuda.CUDAGraph
- The Static Inputs Requirement
- Warmup — Why It's Mandatory
- CUDA Graph Pools
- torch.compile with CUDA Graphs
- Limitations — What Breaks
- CUDA Graphs with AMP
- torch.cuda.make_graphed_callables
- Practical Patterns
- When to Use / Not Use
- Upstream Updates (June 11–12, 2026)
- Further Reading
1. What Are CUDA Graphs?
When you run PyTorch code on a GPU, every operation — matrix multiply, activation, copy — is a kernel launch. Each launch requires the CPU to:
- Prepare kernel arguments
- Submit work to the CUDA driver
- The driver queues it on the GPU stream
- Return control to the CPU
For large kernels (e.g., a big GEMM), the GPU execution time dwarfs this overhead. But for small/medium operations, the CPU launch latency (5–15 microseconds per kernel) can dominate total runtime.
CUDA Graphs solve this by recording a sequence of GPU operations into a graph during a capture phase, then replaying the entire graph with a single CPU-side launch.
Normal execution: CUDA Graph replay:
CPU: launch K1 CPU: launch graph ─────────────────┐
wait... │
launch K2 GPU: K1 → K2 → K3 → K4 → K5 │
wait... (all pre-recorded, no waits) │
launch K3 │
wait... Result: 1 CPU launch instead of 5 │
launch K4 ────────────────────────────────────┘
wait...
launch K5
GPU: ▓░░▓░░▓░░▓░░▓ GPU: ▓▓▓▓▓
(gaps = idle) (no gaps)
The graph captures:
- Which kernels to run and in what order
- Memory addresses of all inputs and outputs
- Kernel launch parameters
It does not capture tensor values — only the operations and where they read/write.
2. Why CUDA Graphs Matter
The CPU Bottleneck
Modern GPUs execute small kernels in microseconds. If each kernel takes 3 µs on the GPU but 10 µs for the CPU to launch, you spend 77% of your time on CPU overhead.
Without CUDA Graphs (100 kernels):
CPU overhead: 100 × 10 µs = 1,000 µs
GPU compute: 100 × 3 µs = 300 µs
Total: 1,300 µs
GPU utilization: 23%
With CUDA Graphs (100 kernels):
CPU overhead: 1 × 10 µs = 10 µs
GPU compute: 100 × 3 µs = 300 µs
Total: 310 µs
GPU utilization: 97%
Speedup: 4.2x
When Speedups Are Largest
| Scenario | Typical Speedup |
|---|---|
| Small model inference (ResNet-18, batch=1) | 3–10x |
| Medium model inference (BERT, batch=8) | 2–5x |
| Large model inference (GPT-2, batch=32) | 1.2–2x |
| Training (full step) | 1.1–1.5x |
The pattern: the smaller the model and the more kernels per unit of compute, the bigger the win.
3. Basic API: torch.cuda.CUDAGraph
Minimal Example
import torch
model = torch.nn.Linear(512, 512).cuda()
static_input = torch.randn(64, 512, device="cuda")
# Step 1: Warmup (covered in Section 5)
with torch.no_grad():
for _ in range(3):
_ = model(static_input)
# Step 2: Capture
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
static_output = model(static_input)
# Step 3: Replay
static_input.copy_(torch.randn(64, 512, device="cuda"))
g.replay()
print(static_output) # Contains result for the new input
What Happens During Capture
When you enter torch.cuda.graph(g):
- PyTorch switches to a capturing stream
- Every CUDA operation is recorded (not executed normally)
- Memory allocations inside the block come from a special graph pool
- On exit, the graph is finalized and ready for replay
What Happens During Replay
g.replay() submits the entire recorded sequence to the GPU in one shot. The GPU executes all kernels using the same memory addresses that were captured.
This is why inputs must be static — the graph hardcodes memory pointers.
4. The Static Inputs Requirement
This is the most important concept to understand. CUDA Graphs capture memory addresses, not tensor values. When you replay, the GPU reads from and writes to the exact same addresses.
Correct Pattern
# Pre-allocate (addresses are fixed)
static_input = torch.zeros(batch_size, features, device="cuda")
static_output = None
# Capture
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
static_output = model(static_input)
# For each new input: copy data IN, replay, read data OUT
for batch in dataloader:
static_input.copy_(batch.cuda()) # Copy INTO the pre-allocated tensor
g.replay() # GPU reads from same address
result = static_output.clone() # Copy OUT (or use in-place)
Wrong Pattern
# DON'T DO THIS — creates new tensor each iteration
for batch in dataloader:
new_input = batch.cuda() # New memory address each time!
g.replay() # Graph still reads from OLD address
# Result: stale data, wrong answers
Why This Design?
Recording memory addresses (not values) is what makes replay so fast. If the graph had to remap pointers each time, it would lose most of its advantage. The tradeoff: you manage input/output buffers manually.
5. Warmup — Why It's Mandatory
Before capturing a CUDA Graph, you must run the model at least once (usually 3 times for safety). Warmup triggers:
| Lazy Initialization | Why It Matters |
|---|---|
| cuDNN algorithm selection | First conv/GEMM benchmarks multiple algorithms |
| CUDA context creation | First CUDA call initializes the driver context |
| Memory allocator warmup | Caching allocator builds its pool |
| JIT kernel compilation | Some ops compile PTX on first use |
| cuBLAS handle creation | First matmul creates the handle |
Warmup Pattern
model = MyModel().cuda().eval()
static_input = torch.randn(B, C, H, W, device="cuda")
# Warmup: run several forward passes
with torch.no_grad():
for _ in range(3):
_ = model(static_input)
# Now safe to capture
torch.cuda.synchronize()
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
static_output = model(static_input)
What Happens Without Warmup?
If you skip warmup, the capture records the lazy initialization itself — memory allocations, algorithm searches, handle creation. This means:
- The graph becomes bloated with one-time setup
- Replaying re-executes setup on every call (pointless overhead)
- Some lazy ops allocate memory dynamically, which is forbidden inside capture and will raise a
RuntimeError
6. CUDA Graph Pools
When multiple graphs share the same model or intermediate buffers, you can share their memory pools to avoid redundant allocations.
Default: Separate Pools
g1 = torch.cuda.CUDAGraph()
g2 = torch.cuda.CUDAGraph()
with torch.cuda.graph(g1):
out1 = model(input1)
with torch.cuda.graph(g2):
out2 = model(input2)
# g1 and g2 each have their own memory pool
Shared Pool
g1 = torch.cuda.CUDAGraph()
g2 = torch.cuda.CUDAGraph()
# Capture both with the same pool
with torch.cuda.graph(g1):
out1 = model(input1)
# Share g1's pool — no extra memory allocated for g2
with torch.cuda.graph(g2, pool=g1.pool()):
out2 = model(input2)
Pool sharing works when graphs don't execute concurrently. The shared pool means they reuse the same memory, so overlapping execution would corrupt data.
When to Share Pools
- Multiple batch-size variants of the same model
- Encoder and decoder graphs that run sequentially
- A/B model comparisons that never overlap
7. torch.compile with CUDA Graphs
The easiest way to use CUDA Graphs is through torch.compile:
model = MyModel().cuda()
# reduce-overhead mode automatically uses CUDA Graphs
compiled = torch.compile(model, mode="reduce-overhead")
# Just use it normally — graphs are managed for you
output = compiled(input_tensor)
How It Works: cudagraph_trees
Under the hood, reduce-overhead mode uses Inductor's cudagraph_trees system:
torch.compile(mode="reduce-overhead")
└─ Dynamo traces the Python code
└─ AOTAutograd generates forward/backward
└─ Inductor generates optimized CUDA code
└─ cudagraph_trees wraps each compiled region in a graph
cudagraph_trees manages:
- Automatic warmup — runs the compiled code once before capturing
- Graph caching — one graph per unique input shape
- Memory management — pools are handled internally
- Fallback — if a region can't be graphed, it runs eagerly
Advantages Over Manual Graphs
Manual CUDAGraph | torch.compile(mode="reduce-overhead") |
|---|---|
| You manage static inputs | Inputs handled automatically |
| You do warmup | Warmup is automatic |
| Entire model must be graphable | Per-region graphs (partial capture) |
| No fusion | Kernel fusion + graphs combined |
| Fixed shapes only | Multiple shape variants cached |
Checking What Got Graphed
import torch._dynamo as dynamo
compiled = torch.compile(model, mode="reduce-overhead")
explanation = dynamo.explain(compiled)(sample_input)
print(explanation)
8. Limitations — What Breaks
CUDA Graphs capture a fixed sequence of GPU operations. Anything that deviates from this at replay time will fail or produce wrong results.
Hard Failures (RuntimeError)
These operations cannot be captured and will raise errors:
| Operation | Why It Fails |
|---|---|
print(tensor) inside graph | Requires CPU sync |
tensor.item() | Transfers data to CPU |
tensor.cpu() | Cross-device copy |
torch.tensor([1, 2, 3]) | CPU tensor creation |
torch.cuda.synchronize() | Blocks the stream |
| Dynamic memory allocation | Graph can't record variable-size allocs |
Silent Failures (Wrong Results)
These won't crash but will produce incorrect output:
| Pattern | Problem |
|---|---|
Data-dependent control flow (if x > 0) | Condition was recorded at capture time |
| Dynamic shapes (varying batch size) | Graph hardcodes tensor dimensions |
| In-place ops on non-static tensors | Writes to wrong addresses |
| Random ops without manual seed | Same random values on every replay |
Operations That Prevent Capture
# These will fail during capture:
# 1. CPU sync
with torch.cuda.graph(g):
out = model(x)
print(out.sum().item()) # ERROR: .item() syncs to CPU
# 2. Dynamic allocation
with torch.cuda.graph(g):
out = model(x)
mask = out > 0 # OK so far
filtered = out[mask] # ERROR: output size depends on data
# 3. CPU tensor creation
with torch.cuda.graph(g):
scale = torch.tensor(2.0) # ERROR: creates CPU tensor
out = model(x) * scale
# Fix: pre-allocate the scale on GPU
scale = torch.tensor(2.0, device="cuda")
with torch.cuda.graph(g):
out = model(x) * scale # OK: scale already on GPU
NCCL and Distributed
Most NCCL collective operations (all-reduce, all-gather, etc.) are not compatible with CUDA Graphs. This is why CUDA Graphs are primarily used for inference and single-GPU training.
Exception: PyTorch's torch.distributed has experimental support for graphing some collectives on newer NCCL versions.
9. CUDA Graphs with AMP
Automatic Mixed Precision works inside CUDA Graph capture, but you must set up the autocast context inside the capture block:
model = MyModel().cuda()
static_input = torch.randn(64, 512, device="cuda")
# Warmup with AMP
with torch.no_grad(), torch.cuda.amp.autocast():
for _ in range(3):
_ = model(static_input)
# Capture with AMP
g = torch.cuda.CUDAGraph()
with torch.cuda.amp.autocast():
with torch.cuda.graph(g):
static_output = model(static_input)
# Replay (autocast context not needed — types are baked into the graph)
static_input.copy_(new_data)
g.replay()
Key Points
- The autocast context must be inside capture (or wrapping it) so the graph records the mixed-precision kernel variants
- After capture, replay uses whatever dtypes were recorded — no autocast needed at replay time
- GradScaler is harder — avoid it with CUDA Graphs. If you need training with AMP + graphs, prefer
torch.compile(mode="reduce-overhead")
10. torch.cuda.make_graphed_callables
make_graphed_callables is a convenience wrapper that handles warmup, capture, and static-input management for nn.Module or simple callables:
model = MyModel().cuda()
sample_input = torch.randn(64, 512, device="cuda")
# Wrap the model — handles warmup and capture automatically
graphed_model = torch.cuda.make_graphed_callables(
model,
sample_args=(sample_input,),
num_warmup_iters=3,
)
# Use like a normal callable
output = graphed_model(sample_input)
Multiple Callables
You can graph multiple modules together, sharing a pool:
encoder = Encoder().cuda()
decoder = Decoder().cuda()
graphed_encoder, graphed_decoder = torch.cuda.make_graphed_callables(
(encoder, decoder),
sample_args=(
(encoder_input,),
(decoder_input,),
),
)
When to Use
- Quick experiments where you want CUDA Graph speedups without manual buffer management
- Models with simple input/output signatures
- Inference-only paths
Caveats
- Only supports fixed shapes (like manual graphs)
- The returned callable replaces the original forward pass
- Multiple return values need careful handling
11. Practical Patterns
Pattern 1: Inference Server
The most common CUDA Graph use case — a model serving predictions with fixed batch size:
class GraphedInferenceServer:
def __init__(self, model, batch_size, input_dim):
self.model = model.cuda().eval()
self.static_input = torch.zeros(
batch_size, input_dim, device="cuda"
)
self.static_output = None
self.graph = torch.cuda.CUDAGraph()
# Warmup
with torch.no_grad():
for _ in range(3):
_ = self.model(self.static_input)
# Capture
with torch.no_grad():
with torch.cuda.graph(self.graph):
self.static_output = self.model(self.static_input)
def predict(self, input_tensor):
self.static_input.copy_(input_tensor)
self.graph.replay()
return self.static_output.clone()
Pattern 2: Multiple Batch Sizes
For serving with variable batch sizes, capture one graph per batch size:
class MultiBatchServer:
def __init__(self, model, batch_sizes, input_dim):
self.model = model.cuda().eval()
self.graphs = {}
for bs in batch_sizes:
static_in = torch.zeros(bs, input_dim, device="cuda")
# Warmup
with torch.no_grad():
for _ in range(3):
_ = self.model(static_in)
g = torch.cuda.CUDAGraph()
with torch.no_grad():
with torch.cuda.graph(g):
static_out = self.model(static_in)
self.graphs[bs] = (g, static_in, static_out)
def predict(self, input_tensor):
bs = input_tensor.shape[0]
g, static_in, static_out = self.graphs[bs]
static_in.copy_(input_tensor)
g.replay()
return static_out.clone()
Pattern 3: Partial Graph Capture for Training
Full training steps are hard to graph (optimizer state updates, gradient scaling). Instead, graph just the forward pass:
model = MyModel().cuda()
optimizer = torch.optim.Adam(model.parameters())
static_input = torch.randn(32, 512, device="cuda")
static_target = torch.randn(32, 10, device="cuda")
# Warmup
for _ in range(3):
out = model(static_input)
loss = torch.nn.functional.mse_loss(out, static_target)
loss.backward()
optimizer.zero_grad()
# Capture forward + backward (not optimizer step)
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
out = model(static_input)
loss = torch.nn.functional.mse_loss(out, static_target)
loss.backward()
# Training loop
for batch_input, batch_target in dataloader:
static_input.copy_(batch_input)
static_target.copy_(batch_target)
g.replay() # Forward + backward replayed
optimizer.step() # Optimizer runs eagerly (safe)
optimizer.zero_grad()
12. When to Use / Not Use
Decision Tree
Is your model running on NVIDIA GPU?
├─ No → CUDA Graphs not applicable
└─ Yes
│
Are input shapes fixed (or a small set of fixed shapes)?
├─ No → Use torch.compile (handles dynamic shapes)
└─ Yes
│
Is CPU launch overhead significant?
(small/medium model, many small kernels, high throughput needed)
├─ No → Graphs won't help much, use torch.compile for fusion
└─ Yes
│
Is it inference only?
├─ Yes → CUDA Graphs are ideal
│ Use torch.compile(mode="reduce-overhead")
│ or manual CUDAGraph API
└─ No (training)
│
Can you graph forward+backward only?
├─ Yes → Partial graph capture (optimizer runs eagerly)
└─ No → torch.compile(mode="reduce-overhead") is safer
Quick Reference
| Scenario | Recommendation |
|---|---|
| Inference, fixed shapes, max throughput | Manual CUDAGraph or reduce-overhead |
| Inference, variable shapes | torch.compile(mode="default") — dynamic shape support |
| Training, single GPU | torch.compile(mode="reduce-overhead") — handles optimizer |
| Training, multi-GPU (DDP/FSDP) | torch.compile(mode="default") — NCCL compatibility |
| Model has data-dependent control flow | torch.compile with graph breaks |
| Quick prototyping | make_graphed_callables |
13. Upstream Updates (June 11–12, 2026)
Recent PyTorch development activity relevant to CUDA Graphs and the broader ecosystem:
Version Bump to 2.14.0a0 (#187070)
The main branch has been bumped to version 2.14.0a0+, signaling the start of the next development cycle. This is the version that includes the latest CUDA Graph improvements and new features described below.
FlexGEMM Higher-Order Op (torch/_higher_order_ops/flex_gemm.py)
A new higher-order operation for flexible GEMM execution has been added. FlexGEMM allows user-defined epilogues on GEMM results, similar to how FlexAttention allows custom attention score modifications. This interacts with CUDA Graphs through the reduce-overhead compilation path.
c10d Window Interfaces for One-Sided Communication (#186299)
New window-based interfaces in PyTorch's distributed backend (c10d) enable one-sided MPI-style communication patterns (put, get, accumulate). While most NCCL collectives remain incompatible with CUDA Graphs, these new primitives expand the distributed toolkit.
cuSOLVER Workspace Optimization (#181998)
Workspace allocation for cuSOLVER operations has been optimized, reducing memory overhead for linear algebra operations. This is relevant to CUDA Graphs because workspace allocations during graph capture can cause issues — the optimization makes these allocations more predictable and graph-friendly.
Dynamo itertools.permutations Polyfill (#186937)
torch._dynamo now supports itertools.permutations during tracing, allowing more Python code to be captured without graph breaks. Fewer graph breaks mean larger graphable regions, which directly benefits CUDA Graph capture through torch.compile(mode="reduce-overhead").
uint16/uint32/uint64 Test Coverage Extended (#183473)
Test infrastructure now covers unsigned integer types more broadly. While not directly related to CUDA Graphs, this improves the reliability of quantized models that may be served with CUDA Graph-accelerated inference.
FlexAttention INDEX_DTYPE for Pointer Arithmetic (#185264)
FlexAttention now uses a dedicated INDEX_DTYPE for pointer arithmetic in block-sparse patterns, improving numerical stability on different GPU architectures. FlexAttention kernels are common targets for CUDA Graph capture in Transformer inference.
14. Further Reading
- CUDA Graphs documentation (PyTorch)
- NVIDIA CUDA Graphs guide
- torch.compile reduce-overhead mode
- make_graphed_callables API
- Accelerating PyTorch with CUDA Graphs (NVIDIA blog)
Notebook: 21_cuda_graphs.ipynb
Source Files
cuda_graphs.py— CUDA Graphs — capture, replay, static inputs, benchmarking, torch.compile reduce-overhead
Module 22: LLM Training Recipes — Building Blocks of Modern Language Models
Day 8 of the incremental learning series
Table of Contents
- RoPE (Rotary Position Embeddings)
- KV Cache
- Grouped-Query Attention (GQA)
- Sliding Window Attention
- RMSNorm
- SwiGLU / SiLU FFN
- Weight Tying
- BFloat16 Training
- Gradient Accumulation for Large Batch
- Complete Mini-LLM Training Setup
- Upstream Updates (June 12–15, 2026)
1. RoPE (Rotary Position Embeddings)
Why RoPE?
Traditional positional encodings (sinusoidal or learned) are added to token embeddings before attention. This has downsides:
- Learned embeddings have a fixed maximum length
- Sinusoidal embeddings don't interact with attention scores directly
- Neither encodes relative position naturally
RoPE (Su et al., 2021) applies position information as a rotation to the query and key vectors inside the attention computation. The result: attention scores naturally depend on the relative distance between tokens.
The Math
Given a head dimension d, RoPE defines frequency bands:
θ_i = 10000^(-2i/d) for i = 0, 1, ..., d/2 - 1
For position m, the rotation angles are:
angles_m = [m·θ_0, m·θ_1, ..., m·θ_{d/2-1}]
The rotation is applied by treating consecutive pairs of dimensions as 2D vectors and rotating them:
[x_{2i}, x_{2i+1}] → [x_{2i}·cos(m·θ_i) - x_{2i+1}·sin(m·θ_i),
x_{2i}·sin(m·θ_i) + x_{2i+1}·cos(m·θ_i)]
Equivalently, using complex numbers: view each pair as a complex number z = x_{2i} + j·x_{2i+1}, then multiply by e^{j·m·θ_i}.
Implementation
import torch
def precompute_freqs_cis(dim: int, max_seq_len: int, theta: float = 10000.0):
"""Precompute the complex exponentials for RoPE."""
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
t = torch.arange(max_seq_len)
freqs = torch.outer(t, freqs) # (seq_len, dim/2)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # e^(i*freq)
return freqs_cis
def apply_rotary_emb(xq, xk, freqs_cis):
"""Apply rotary embeddings to Q and K tensors."""
# Reshape to complex: (batch, seq, heads, dim) -> (batch, seq, heads, dim/2)
xq_complex = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
xk_complex = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
# Reshape freqs for broadcasting: (seq, dim/2) -> (1, seq, 1, dim/2)
freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(2)
# Rotate
xq_out = torch.view_as_real(xq_complex * freqs_cis).flatten(-2)
xk_out = torch.view_as_real(xk_complex * freqs_cis).flatten(-2)
return xq_out.type_as(xq), xk_out.type_as(xk)
Why RoPE Enables Length Generalization
Because RoPE encodes relative position through rotation differences, models trained at one sequence length can often extrapolate to longer sequences (especially with techniques like NTK-aware scaling or YaRN). The rotation dot-product Re[q_m · conj(k_n)] depends only on m - n.
2. KV Cache
The Problem
During autoregressive generation, token t needs to attend to all previous tokens 0..t-1. Naively, this means recomputing K and V for all prior tokens at every step — O(n²) total work for n tokens.
The Solution
Cache the K and V projections from all previous steps. At each decode step, only compute Q, K, V for the new token, append K and V to the cache, then attend over the full cached sequence.
How It Works Step by Step
- Prefill (prompt processing): Process the full prompt in one forward pass. Store all K, V in the cache.
- Decode (token generation): For each new token:
- Compute Q, K, V for just that token
- Append new K, V to cache → cache grows by 1
- Compute attention: Q (1 token) × K^T (all cached) → scores → softmax → × V
Implementation
class KVCache:
def __init__(self, max_batch, max_seq_len, n_kv_heads, head_dim, dtype=torch.float16):
shape = (max_batch, max_seq_len, n_kv_heads, head_dim)
self.k_cache = torch.zeros(shape, dtype=dtype)
self.v_cache = torch.zeros(shape, dtype=dtype)
self.seq_len = 0
def update(self, k, v):
"""Append new K, V to cache. k, v shape: (batch, new_len, heads, dim)"""
new_len = k.shape[1]
self.k_cache[:, self.seq_len:self.seq_len + new_len] = k
self.v_cache[:, self.seq_len:self.seq_len + new_len] = v
self.seq_len += new_len
return self.k_cache[:, :self.seq_len], self.v_cache[:, :self.seq_len]
Memory Calculation
cache_size = 2 × n_layers × seq_len × n_kv_heads × head_dim × dtype_size
Example: Llama 2 7B (32 layers, 32 KV heads, dim=128, fp16, seq=4096)
= 2 × 32 × 4096 × 32 × 128 × 2 bytes = 2 GB
With GQA (8 KV heads instead of 32): only 512 MB — a 4× reduction.
Impact on Generation Speed
Without cache: generating n tokens takes O(n²) total computation. With cache: generating n tokens takes O(n) total computation (each step is O(1) for the new token's QKV, O(seq_so_far) for attention).
For a 2048-token generation, KV cache provides ~1000× speedup in total compute.
3. Grouped-Query Attention (GQA)
What Is GQA?
Standard multi-head attention uses the same number of Q, K, and V heads. GQA uses fewer K/V heads than Q heads. Multiple Q heads share the same K/V head.
MHA: 32 Q heads, 32 KV heads (standard)
GQA: 32 Q heads, 8 KV heads (Llama 2 70B, Llama 3)
MQA: 32 Q heads, 1 KV head (extreme)
Why GQA?
- Smaller KV cache — proportional reduction (4× with 8 KV heads for 32 Q heads)
- Less memory bandwidth — KV cache read is often the bottleneck during decode
- Minimal quality loss — GQA with 8 heads achieves near-MHA quality
Implementation: repeat_kv
To use GQA with standard attention, expand the KV heads to match Q heads:
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
"""Repeat KV heads to match Q heads. x: (batch, seq, n_kv_heads, head_dim)"""
if n_rep == 1:
return x
batch, seq_len, n_kv_heads, head_dim = x.shape
x = x.unsqueeze(3).expand(batch, seq_len, n_kv_heads, n_rep, head_dim)
return x.reshape(batch, seq_len, n_kv_heads * n_rep, head_dim)
4. Sliding Window Attention
Concept
Instead of attending to all previous tokens, only attend to the last W tokens (the "window"). Tokens beyond the window cannot be directly attended to.
Standard causal: token t attends to tokens 0..t
Sliding window: token t attends to tokens max(0, t-W)..t
Why?
- Memory: O(n × W) instead of O(n²)
- Compute: O(n × W) instead of O(n²)
- Information still propagates through layers: after L layers, token at position t has indirect access to tokens at position t - L×W
Used in Mistral 7B (W=4096) and Mixtral.
Implementation
def make_sliding_window_mask(seq_len: int, window_size: int) -> torch.Tensor:
"""Create a sliding window causal mask."""
mask = torch.full((seq_len, seq_len), float('-inf'))
for i in range(seq_len):
start = max(0, i - window_size + 1)
mask[i, start:i + 1] = 0.0
return mask
Or more efficiently:
def make_sliding_window_mask(seq_len: int, window_size: int) -> torch.Tensor:
row_idx = torch.arange(seq_len).unsqueeze(1)
col_idx = torch.arange(seq_len).unsqueeze(0)
# Causal: col <= row; Window: row - col < window_size
valid = (col_idx <= row_idx) & (row_idx - col_idx < window_size)
mask = torch.where(valid, 0.0, float('-inf'))
return mask
5. RMSNorm
Why Not LayerNorm?
LayerNorm computes:
y = (x - mean(x)) / sqrt(var(x) + ε) * γ + β
RMSNorm removes the mean subtraction and bias — just normalizes by the root-mean-square:
y = x / RMS(x) * γ
where RMS(x) = sqrt(mean(x²) + ε)
Why Faster?
- No mean computation (one less reduction)
- No bias parameter
- Empirically, centering doesn't help much for Transformer layers
Used in: Llama, Llama 2, Llama 3, Mistral, Gemma, and most modern LLMs.
PyTorch Implementation
class RMSNorm(torch.nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.weight = torch.nn.Parameter(torch.ones(dim))
def forward(self, x):
rms = torch.sqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps)
return x / rms * self.weight
PyTorch also provides torch.nn.RMSNorm as a built-in (since 2.4+):
norm = torch.nn.RMSNorm(dim, eps=1e-6)
6. SwiGLU / SiLU FFN
The Standard FFN
FFN(x) = ReLU(x @ W1) @ W2
Two weight matrices, ReLU activation. Simple.
The Modern FFN (SwiGLU)
FFN(x) = (SiLU(x @ W_gate) ⊙ (x @ W1)) @ W2
Three weight matrices:
W_gate(dim → hidden): produces the gating signalW1(dim → hidden): produces the valueW2(hidden → dim): projects back
⊙ is element-wise multiplication. SiLU(x) = x * sigmoid(x).
Why Three Matrices?
The gating mechanism (GLU = Gated Linear Unit) allows the network to learn which dimensions to activate. SiLU provides smooth gradients (unlike ReLU). The combination of gating + smooth activation empirically trains better for LLMs.
Implementation
class SwiGLU(torch.nn.Module):
def __init__(self, dim: int, hidden_dim: int):
super().__init__()
self.w1 = torch.nn.Linear(dim, hidden_dim, bias=False)
self.w2 = torch.nn.Linear(hidden_dim, dim, bias=False)
self.w_gate = torch.nn.Linear(dim, hidden_dim, bias=False)
def forward(self, x):
return self.w2(torch.nn.functional.silu(self.w_gate(x)) * self.w1(x))
Hidden Dimension Convention
In Llama models: hidden_dim = int(2/3 4 dim) rounded to a multiple of 256. The 2/3 factor compensates for the extra gate projection (3 matrices of size 2/3 ≈ 2 matrices of size 1).
7. Weight Tying
Concept
class LLM(nn.Module):
def __init__(self, vocab_size, dim):
super().__init__()
self.embedding = nn.Embedding(vocab_size, dim)
self.output = nn.Linear(dim, vocab_size, bias=False)
# Tie weights
self.output.weight = self.embedding.weight
Impact
For a model with vocab_size=32000 and dim=4096:
- Embedding matrix: 32000 × 4096 × 2 bytes (fp16) = 250 MB
- Without tying: 500 MB for embedding + output
- With tying: 250 MB total — 50% savings on these layers
For vocab-heavy models (large vocabulary relative to model dimension), this can save ~30% of total parameters.
8. BFloat16 Training
Why bf16 Over fp16?
| Property | fp16 | bf16 |
|---|---|---|
| Exponent bits | 5 | 8 |
| Mantissa bits | 10 | 7 |
| Max value | ~65504 | ~3.4×10³⁸ |
| Min normal | ~6×10⁻⁵ | ~1.2×10⁻³⁸ |
| Precision | Higher | Lower |
| Overflow risk | High | Very low |
| Loss scaling needed | Yes | No |
For LLM training, bf16 wins because:
- No loss scaler needed — bf16 has the same dynamic range as fp32
- Simpler code — no GradScaler, no inf checks
- Better stability — gradients rarely overflow
Usage
# bf16 autocast — no GradScaler needed
with torch.amp.autocast('cuda', dtype=torch.bfloat16):
output = model(input)
loss = criterion(output, target)
loss.backward()
optimizer.step()
Compare with fp16 which requires:
scaler = torch.amp.GradScaler()
with torch.amp.autocast('cuda', dtype=torch.float16):
output = model(input)
loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
9. Gradient Accumulation for Large Batch
The Problem
LLMs benefit from large effective batch sizes (often 1M+ tokens). But a single GPU can only fit a small micro-batch (e.g., 4 sequences of 2048 tokens = 8K tokens).
The Solution
Accumulate gradients over multiple micro-batches before stepping:
accumulation_steps = 128 # 128 × 8K = 1M tokens effective batch
optimizer.zero_grad()
for i, batch in enumerate(dataloader):
with torch.amp.autocast('cuda', dtype=torch.bfloat16):
loss = model(batch) / accumulation_steps # Normalize loss
loss.backward() # Gradients accumulate
if (i + 1) % accumulation_steps == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
optimizer.zero_grad()
Key details:
- Divide loss by
accumulation_stepsto average gradients (equivalent to a large batch mean) - Gradient clipping after accumulation, before step
- Memory: only one micro-batch is active at a time (constant memory regardless of effective batch)
10. Complete Mini-LLM Training Setup
Putting all techniques together into a minimal but complete LLM:
class MiniLLM(nn.Module):
"""Minimal LLM with all modern techniques."""
def __init__(self, vocab_size=32000, dim=512, n_layers=6,
n_heads=8, n_kv_heads=4, max_seq_len=1024):
super().__init__()
self.embedding = nn.Embedding(vocab_size, dim)
self.layers = nn.ModuleList([
TransformerBlock(dim, n_heads, n_kv_heads) for _ in range(n_layers)
])
self.norm = nn.RMSNorm(dim)
self.output = nn.Linear(vocab_size, dim, bias=False)
self.output.weight = self.embedding.weight # Weight tying
self.freqs_cis = precompute_freqs_cis(dim // n_heads, max_seq_len)
class TransformerBlock(nn.Module):
def __init__(self, dim, n_heads, n_kv_heads):
super().__init__()
self.attention = GQAAttention(dim, n_heads, n_kv_heads)
self.ffn = SwiGLU(dim, int(2/3 * 4 * dim))
self.norm1 = nn.RMSNorm(dim)
self.norm2 = nn.RMSNorm(dim)
def forward(self, x, freqs_cis, mask=None, cache=None):
x = x + self.attention(self.norm1(x), freqs_cis, mask, cache)
x = x + self.ffn(self.norm2(x))
return x
Training loop:
model = torch.compile(MiniLLM())
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.1)
for step, batch in enumerate(dataloader):
with torch.amp.autocast('cuda', dtype=torch.bfloat16):
logits = model(batch['input_ids'])
loss = F.cross_entropy(
logits.view(-1, vocab_size),
batch['labels'].view(-1)
) / accumulation_steps
loss.backward()
if (step + 1) % accumulation_steps == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
optimizer.zero_grad()
Generation with KV cache:
@torch.no_grad()
def generate(model, prompt_tokens, max_new_tokens=100, temperature=0.8, top_k=50):
caches = [KVCache(...) for _ in range(n_layers)]
# Prefill
logits = model.forward_with_cache(prompt_tokens, caches)
# Decode
for _ in range(max_new_tokens):
next_logits = logits[:, -1] / temperature
# Top-k filtering
topk_vals, topk_idx = next_logits.topk(top_k)
next_logits = torch.full_like(next_logits, float('-inf'))
next_logits.scatter_(1, topk_idx, topk_vals)
probs = F.softmax(next_logits, dim=-1)
next_token = torch.multinomial(probs, 1)
logits = model.forward_with_cache(next_token, caches)
yield next_token
See llm_training_loop.py for the complete runnable implementation.
11. Upstream Updates (June 12–15, 2026)
Recent PyTorch commits and features relevant to LLM training:
DTensor: _StridedShard to Shard via all-to-all (#170915)
Converts _StridedShard placements to Shard using all-to-all collective, enabling more efficient tensor parallelism redistribution. Previously, strided-shard tensors required falling back to full gather + reshard.
Symmetric Memory NCCL EP Support
Symmetric memory allocator now supports NCCL Extensible Parallelism (EP) endpoints, enabling custom collective algorithms that bypass the NCCL ring/tree topology for specific communication patterns (e.g., expert parallelism in MoE models).
Version 2.14.0a0
The development version has been bumped to 2.14.0a0. Key features in flight:
- FlexGEMM epilogue templates for fused post-GEMM operations
- Extended Dynamo polyfills (
itertools.permutationsnow supported) - Bug fixes for
argmin/argmaxon boolean tensors
FlexGEMM Epilogue Templates
Inductor now supports FlexGEMM epilogue templates, allowing users to fuse post-matmul operations (bias add, activation, scaling) into a single CUTLASS kernel call. Reduces memory traffic for FFN layers.
Dynamo Polyfills: itertools.permutations
torch._dynamo can now trace through itertools.permutations, converting it to a constant at trace time. This enables more Python patterns in torch.compile-d code without graph breaks.
argmin/argmax Boolean Fix
Fixed incorrect results from torch.argmin and torch.argmax on boolean tensors, which previously could return wrong indices due to an optimization that assumed numeric ordering.
MPS: log_sigmoid Metal Migration
The log_sigmoid operation on Apple Silicon (MPS backend) has been migrated from the CPU fallback to a native Metal shader, providing significant speedup for MPS users training models with sigmoid-based losses.
Key Takeaways
| Technique | What It Does | Memory Impact | Speed Impact |
|---|---|---|---|
| RoPE | Relative position via rotation | None | Slight compute |
| KV Cache | Cache K/V for generation | +cache memory | ~1000× faster decode |
| GQA | Fewer KV heads | 4× less KV cache | Faster decode |
| Sliding Window | Limit attention span | O(n·W) vs O(n²) | Faster for long seq |
| RMSNorm | Faster normalization | Same | ~10% faster norm |
| SwiGLU | Gated FFN with SiLU | +50% FFN params | Better convergence |
| Weight Tying | Share embed/output | -30% for vocab-heavy | Same |
| bf16 | Wider exponent range | Half vs fp32 | 2× throughput |
| Grad Accumulation | Large effective batch | Constant | Linear in steps |
Further Reading
- Llama 2 Paper — GQA, RoPE, SwiGLU, RMSNorm in practice
- RoFormer (Su et al., 2021) — Original RoPE paper
- Mistral 7B — Sliding window attention
- GLU Variants (Shazeer, 2020) — SwiGLU and friends
- torchtune — PyTorch's official LLM fine-tuning library
Notebook: 22_llm_recipes.ipynb
Source Files
rope_embeddings.py— RoPE — precompute freqs, apply rotary embeddings, position encoding visualizationkv_cache.py— KV Cache — pre-allocated cache, prefill/decode phases, GQA repeat_kv, benchmarkingllm_training_loop.py— Complete mini-LLM with RoPE, GQA, SwiGLU, RMSNorm, KV cache, bf16, gradient accumulation
torch.fx — Graph-Level Model Transformation
Table of Contents
- What is torch.fx?
- Symbolic Tracing
- The FX Graph IR
- Graph Inspection
- Graph Transformation — Adding, Removing, Replacing Nodes
- Pattern Matching and Replacement
- Practical Pass Examples
- ShapeProp
- graph.lint()
- torch.fx.Interpreter
- torch.fx.Transformer
- FX in the Compilation Stack
- Upstream Updates (June 15–16, 2026)
- Summary
1. What is torch.fx?
torch.fx is a Python-to-Python transformation framework for PyTorch. It lets you:
- Capture a PyTorch module's forward logic into a graph-based intermediate representation (IR)
- Inspect and modify that graph programmatically
- Generate a new, executable PyTorch module from the modified graph
┌─────────────┐ symbolic_trace ┌──────────────┐ transform ┌──────────────┐
│ nn.Module │ ──────────────────▶ │ FX Graph IR │ ─────────────▶ │ FX Graph IR │
│ (Python) │ │ (nodes) │ │ (modified) │
└─────────────┘ └──────────────┘ └──────┬───────┘
│
graph.recompile()
│
┌──────▼───────┐
│ GraphModule │
│ (executable) │
└──────────────┘
Where FX is used:
- torch.compile — Dynamo captures FX graphs, AOTAutograd transforms them, Inductor lowers them
- Quantization — FX graph mode quantization inserts observers and rewrites ops
- Distributed — Tensor Parallel, Pipeline Parallel, and FSDP use FX passes for graph partitioning
- Custom optimizations — fuse ops, eliminate dead code, add profiling
The key advantage over tracing with torch.jit.trace is that FX operates at the Python level — the graph is pure Python, the transformations are pure Python, and the output is a standard nn.Module subclass.
2. Symbolic Tracing
Basic Tracing
import torch
import torch.nn as nn
import torch.fx
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.linear1 = nn.Linear(10, 20)
self.linear2 = nn.Linear(20, 5)
def forward(self, x):
x = self.linear1(x)
x = torch.relu(x)
x = self.linear2(x)
return x
model = MyModel()
traced = torch.fx.symbolic_trace(model)
symbolic_trace executes forward() with Proxy values instead of real tensors. Proxies record every operation, building a graph that represents the computation.
The returned traced is a GraphModule — a subclass of nn.Module with:
traced.graph— theGraphobject (node-level IR)traced.code— auto-generated Python source codetraced.forward— the compiled forward method
What Gets Captured
print(traced.code)
# def forward(self, x):
# linear1 = self.linear1(x); x = None
# relu = torch.relu(linear1); linear1 = None
# linear2 = self.linear2(relu); relu = None
# return linear2
Limitations of Symbolic Tracing
Symbolic tracing has fundamental limitations because it traces with proxy values, not real data:
| Limitation | Example | Why It Fails |
|---|---|---|
| Data-dependent control flow | if x.sum() > 0: | Proxy doesn't have a real value to branch on |
| Dynamic shapes | x[:, :n] where n is runtime-determined | Proxy doesn't have concrete shapes |
| Non-torch Python ops | print(x.shape), list comprehensions over tensors | Proxies don't support arbitrary Python |
Non-forward methods | self.helper(x) not called from forward | Only forward is traced |
Workarounds:
# 1. Use torch.fx.wrap to mark functions as "leaf" (not traced into)
@torch.fx.wrap
def my_custom_op(x):
if x.sum() > 0:
return x * 2
return x
# 2. Use concrete_args to fix certain inputs
traced = torch.fx.symbolic_trace(model, concrete_args={"training": False})
For models with dynamic control flow, torch.compile (Dynamo) is preferred — it handles breaks and recompiles automatically.
3. The FX Graph IR
An FX Graph is a DAG (directed acyclic graph) of Node objects. Each node represents one operation.
Node Operations
node.op | Meaning | node.target | Example |
|---|---|---|---|
placeholder | Function input | Parameter name | x |
get_attr | Access self.attr | Attribute path (string) | self.weight |
call_function | Free function call | The function itself | torch.relu, operator.add |
call_method | Method on a value | Method name (string) | .view(), .relu() |
call_module | Call a submodule | Module path (string) | self.linear1 |
output | Return value | "output" | The return node |
Node Anatomy
Every node has:
op— one of the six operations abovename— unique identifier (e.g.,"relu","linear1")target— what to call (function, method name, module path, or attribute path)args— positional arguments (tuple of Nodes or constants)kwargs— keyword arguments (dict)users— dict of downstream nodes that consume this node's output
graph = traced.graph
for node in graph.nodes:
print(f"op={node.op:15s} name={node.name:15s} target={node.target}")
print(f" args={node.args} kwargs={node.kwargs}")
print(f" users={list(node.users.keys())}")
The Graph as a Linked List
FX nodes form a doubly-linked list in execution order. You can iterate forward (graph.nodes) or navigate with node.prev and node.next.
placeholder:x → call_module:linear1 → call_function:relu → call_module:linear2 → output
Users and Dependencies
for node in graph.nodes:
if node.op == "call_function" and node.target == torch.relu:
# Who uses the output of relu?
for user in node.users:
print(f"relu output used by: {user.name} ({user.op})")
# What are relu's inputs?
for arg in node.args:
if isinstance(arg, torch.fx.Node):
print(f"relu input from: {arg.name} ({arg.op})")
4. Graph Inspection
Print Tabular
traced.graph.print_tabular()
Output:
opcode name target args kwargs
------------- ------- ----------------------- ------------ --------
placeholder x x () {}
call_module linear1 linear1 (x,) {}
call_function relu <built-in function relu> (linear1,) {}
call_module linear2 linear2 (relu,) {}
output output output (linear2,) {}
Counting Operations
from collections import Counter
op_counts = Counter()
for node in graph.nodes:
if node.op == "call_function":
op_counts[node.target.__name__] += 1
elif node.op == "call_module":
module = traced.get_submodule(node.target)
op_counts[type(module).__name__] += 1
elif node.op == "call_method":
op_counts[node.target] += 1
print(op_counts) # Counter({'Linear': 2, 'relu': 1})
Accessing the Generated Code
# Human-readable Python
print(traced.code)
# The graph itself
print(traced.graph)
# Serializable format
print(traced.graph.python_code(root_module="self"))
Finding Specific Patterns
def find_linear_chains(graph_module):
"""Find consecutive Linear → Linear without activation."""
chains = []
for node in graph_module.graph.nodes:
if node.op != "call_module":
continue
mod = graph_module.get_submodule(node.target)
if not isinstance(mod, nn.Linear):
continue
for user in node.users:
if user.op == "call_module":
user_mod = graph_module.get_submodule(user.target)
if isinstance(user_mod, nn.Linear):
chains.append((node, user))
return chains
5. Graph Transformation — Adding, Removing, Replacing Nodes
Context Managers for Insertion
FX provides context managers that control where new nodes are inserted:
graph = traced.graph
# Insert before a specific node
with graph.inserting_before(target_node):
new_node = graph.call_function(torch.neg, args=(some_node,))
# Insert after a specific node
with graph.inserting_after(target_node):
new_node = graph.call_function(torch.abs, args=(some_node,))
Replacing a Node's Function
for node in graph.nodes:
if node.op == "call_function" and node.target == torch.relu:
node.target = torch.nn.functional.gelu
After any modification, recompile:
graph.lint() # validate correctness
traced.recompile() # regenerate code from graph
Erasing Nodes
# Must replace all uses first
node_to_remove.replace_all_uses_with(replacement_node)
graph.erase_node(node_to_remove)
The node can only be erased if it has zero users. Call replace_all_uses_with first.
Inserting a New Submodule Call
# Add a new module to the GraphModule
traced.add_module("new_bn", nn.BatchNorm1d(20))
# Insert a call_module node
with graph.inserting_after(linear1_node):
bn_node = graph.call_module("new_bn", args=(linear1_node,))
# Rewire: everything that used linear1's output now uses bn's output
linear1_node.replace_all_uses_with(bn_node)
# But bn itself still needs linear1 as input
bn_node.args = (linear1_node,)
graph.lint()
traced.recompile()
6. Pattern Matching and Replacement
replace_pattern
torch.fx.subgraph_rewriter.replace_pattern finds subgraph patterns and replaces them:
from torch.fx import subgraph_rewriter
# Define the pattern to match
def pattern(x):
x = torch.add(x, x)
x = torch.relu(x)
return x
# Define the replacement
def replacement(x):
return torch.nn.functional.gelu(torch.mul(x, 2))
# Apply
replaced = subgraph_rewriter.replace_pattern(traced, pattern, replacement)
print(f"Replaced {len(replaced)} matches")
How Pattern Matching Works
- The pattern function is symbolically traced into a small graph
- FX searches the target graph for subgraphs that are structurally isomorphic
- Matched subgraphs are spliced out and replaced with the replacement graph
- Input/output edges are rewired automatically
Limitations
- Pattern matching is structural, not semantic —
torch.add(x, y)won't matchx + y(which becomesoperator.add) - The pattern must be traceable by
symbolic_trace - Wildcards aren't directly supported — every node in the pattern must match
7. Practical Pass Examples
Pass 1: Replace ReLU with GELU
def replace_relu_with_gelu(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
for node in gm.graph.nodes:
# Handle call_function: torch.relu or F.relu
if node.op == "call_function" and node.target in (
torch.relu, torch.nn.functional.relu
):
node.target = torch.nn.functional.gelu
# Handle call_module: nn.ReLU instances
elif node.op == "call_module":
mod = gm.get_submodule(node.target)
if isinstance(mod, nn.ReLU):
# Replace the module itself
parent_name, _, attr_name = node.target.rpartition(".")
parent = gm.get_submodule(parent_name) if parent_name else gm
setattr(parent, attr_name, nn.GELU())
gm.graph.lint()
gm.recompile()
return gm
Pass 2: Add Timing Instrumentation
import time
def add_timing(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
graph = gm.graph
for node in list(graph.nodes):
if node.op in ("call_function", "call_module", "call_method"):
with graph.inserting_before(node):
start = graph.call_function(time.perf_counter, args=())
with graph.inserting_after(node):
end = graph.call_function(time.perf_counter, args=())
graph.call_function(
print,
args=(f"{node.name}: ",),
)
graph.lint()
gm.recompile()
return gm
Pass 3: Fuse Consecutive Linear Layers
When two nn.Linear layers have no activation between them, they can be algebraically fused: W2(W1·x + b1) + b2 = (W2·W1)·x + (W2·b1 + b2).
def fuse_linear_layers(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
graph = gm.graph
for node in list(graph.nodes):
if node.op != "call_module":
continue
mod1 = gm.get_submodule(node.target)
if not isinstance(mod1, nn.Linear):
continue
# Check single user, also a Linear
users = list(node.users.keys())
if len(users) != 1 or users[0].op != "call_module":
continue
next_node = users[0]
mod2 = gm.get_submodule(next_node.target)
if not isinstance(mod2, nn.Linear):
continue
# Fuse: W_fused = W2 @ W1, b_fused = W2 @ b1 + b2
with torch.no_grad():
W_fused = mod2.weight @ mod1.weight
b_fused = mod2.weight @ mod1.bias + mod2.bias if mod1.bias is not None else mod2.bias
fused = nn.Linear(mod1.in_features, mod2.out_features)
fused.weight = nn.Parameter(W_fused)
fused.bias = nn.Parameter(b_fused)
# Replace in graph
gm.add_module(f"fused_{node.name}_{next_node.name}", fused)
with graph.inserting_before(node):
fused_node = graph.call_module(
f"fused_{node.name}_{next_node.name}",
args=node.args,
)
next_node.replace_all_uses_with(fused_node)
graph.erase_node(next_node)
graph.erase_node(node)
graph.lint()
gm.recompile()
return gm
Pass 4: Dead Code Elimination
def eliminate_dead_code(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
gm.graph.eliminate_dead_code()
gm.recompile()
return gm
FX has built-in dead code elimination. A node is "dead" if it has no users and no side effects. The eliminate_dead_code() method removes all such nodes.
Pass 5: Constant Folding
If a subgraph depends only on constants (parameters, no placeholders), it can be evaluated once and replaced with the result:
def constant_fold(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
graph = gm.graph
for node in list(graph.nodes):
if node.op != "call_function":
continue
# Check if all args are constants or get_attr
if all(
not isinstance(a, torch.fx.Node) or a.op == "get_attr"
for a in node.args
):
# Evaluate the node with real values
interp = torch.fx.Interpreter(gm)
# ... fold constant into a get_attr node
pass
graph.lint()
gm.recompile()
return gm
In practice, torch._inductor.constant_folding provides a production-grade implementation.
8. ShapeProp
ShapeProp propagates tensor metadata (shape, dtype, device) through the graph by running the graph with real inputs and recording the output metadata at each node.
from torch.fx.passes.shape_prop import ShapeProp
model = MyModel()
traced = torch.fx.symbolic_trace(model)
# Run shape propagation with a sample input
sample = torch.randn(4, 10)
ShapeProp(traced).propagate(sample)
# Now every node has shape metadata
for node in traced.graph.nodes:
if "tensor_meta" in node.meta:
meta = node.meta["tensor_meta"]
print(f"{node.name:15s} shape={meta.shape} dtype={meta.dtype}")
Output:
x shape=torch.Size([4, 10]) dtype=torch.float32
linear1 shape=torch.Size([4, 20]) dtype=torch.float32
relu shape=torch.Size([4, 20]) dtype=torch.float32
linear2 shape=torch.Size([4, 5]) dtype=torch.float32
This is essential for optimization passes that need to know tensor dimensions — e.g., deciding whether to fuse operations based on their sizes.
9. graph.lint()
graph.lint() validates the graph's structural integrity:
traced.graph.lint()
It checks:
- Every node's
argsandkwargsreference valid nodes in the same graph - There is exactly one
outputnode - Placeholder nodes come before all other nodes
- No cycles exist
- All
call_moduletargets exist on the root module - Node users are consistent (if A uses B, then A is in B.users)
Always call graph.lint() after any graph transformation. It catches bugs early — an invalid graph will produce cryptic errors at execution time.
10. torch.fx.Interpreter
The Interpreter executes a GraphModule node-by-node, giving you hooks to customize behavior at each step.
class ProfilingInterpreter(torch.fx.Interpreter):
def __init__(self, module):
super().__init__(module)
self.profiling_results = {}
def run_node(self, node):
start = time.perf_counter()
result = super().run_node(node)
elapsed = time.perf_counter() - start
self.profiling_results[node.name] = elapsed
return result
interp = ProfilingInterpreter(traced)
output = interp.run(torch.randn(4, 10))
for name, t in sorted(interp.profiling_results.items(), key=lambda x: -x[1]):
print(f"{name:20s} {t*1000:.3f} ms")
Interpreter Methods You Can Override
| Method | Called When | Use Case |
|---|---|---|
run_node(node) | Every node | Profiling, logging, error handling |
call_function(target, args, kwargs) | call_function nodes | Mock functions, replace ops |
call_method(target, args, kwargs) | call_method nodes | Intercept method calls |
call_module(target, args, kwargs) | call_module nodes | Swap modules, add hooks |
placeholder(target, args, kwargs) | Input nodes | Modify inputs |
get_attr(target, args, kwargs) | Attribute access | Intercept param loads |
output(target, args, kwargs) | Return node | Post-process outputs |
Shape Inference Interpreter
class ShapeInterpreter(torch.fx.Interpreter):
def __init__(self, module):
super().__init__(module)
self.node_shapes = {}
def run_node(self, node):
result = super().run_node(node)
if isinstance(result, torch.Tensor):
self.node_shapes[node.name] = result.shape
return result
11. torch.fx.Transformer
Transformer is a higher-level API for node-by-node graph rewriting. You subclass it and override methods per op type. It creates a new graph (instead of modifying in-place).
class ReLUToGELU(torch.fx.Transformer):
def call_function(self, target, args, kwargs):
if target == torch.relu:
target = torch.nn.functional.gelu
return super().call_function(target, args, kwargs)
def call_module(self, target, args, kwargs):
mod = self.fetch_attr(target)
if isinstance(mod, nn.ReLU):
return super().call_function(
torch.nn.functional.gelu, args, kwargs
)
return super().call_module(target, args, kwargs)
transformed = ReLUToGELU(traced).transform()
Transformer vs. Manual Graph Manipulation
| Aspect | Transformer | Manual (graph.nodes iteration) |
|---|---|---|
| Creates new graph | Yes | No (in-place) |
| Node remapping | Automatic | Manual |
| Easier for per-node transforms | Yes | No |
| Better for structural changes | No | Yes (inserting/removing) |
| Risk of dangling references | Low | Higher |
Use Transformer when your pass maps each node to zero or more nodes. Use manual manipulation when you need to analyze graph structure (chains, patterns) before deciding what to change.
12. FX in the Compilation Stack
torch.compile Pipeline
torch.compile(model)
│
┌────────▼────────┐
│ TorchDynamo │ Captures Python bytecode → FX Graph
│ (Python → FX) │ Handles control flow via graph breaks
└────────┬────────┘
│ FX Graph (ATen-level ops)
┌────────▼────────┐
│ AOTAutograd │ Joint forward+backward graph
│ (FX → FX) │ Partitions into fwd/bwd graphs
└────────┬────────┘
│ FX Graph (decomposed ATen ops)
┌────────▼────────┐
│ Inductor │ Lowers FX Graph → Triton/C++ code
│ (FX → code) │ Fusion, scheduling, code generation
└─────────────────┘
Dynamo Produces FX Graphs
Unlike symbolic_trace, Dynamo operates at the bytecode level. It:
- Handles control flow by inserting graph breaks
- Supports dynamic shapes
- Captures the actual operations executed (not just proxy-traced forward)
def dynamo_backend(gm: torch.fx.GraphModule, example_inputs):
"""Custom backend receives an FX GraphModule."""
print("Received FX graph:")
gm.graph.print_tabular()
return gm # return as-is for debugging
model = torch.compile(MyModel(), backend=dynamo_backend)
model(torch.randn(4, 10))
Inductor FX Passes
Inductor applies many FX passes before code generation. They live in torch/_inductor/fx_passes/:
| Pass | What It Does |
|---|---|
decompositions.py | Break complex ops into primitives |
fuse_attention.py | Pattern-match and fuse attention |
group_batch_fusion.py | Batch small ops together |
joint_graph.py | Optimizations on the joint fwd+bwd graph |
post_grad.py | Post-autograd optimizations |
pre_grad.py | Pre-autograd optimizations |
Writing a Custom Inductor Pass
from torch._inductor import config
def my_custom_pass(gm: torch.fx.GraphModule):
for node in gm.graph.nodes:
# Your optimization here
pass
gm.graph.lint()
gm.recompile()
return gm
# Register as a post-grad pass
config.post_grad_custom_post_pass = my_custom_pass
13. Upstream Updates (June 15-16, 2026)
Recent changes in the PyTorch repository that touch FX, Dynamo, Inductor, and related infrastructure:
TokenSwitch for Distributed Token Routing (#178712)
New TokenSwitch primitive for distributed expert-parallel token routing. Uses FX graph representation for expressing token dispatch and combine patterns across devices.
Dynamo O(N^2) Decomposition Fix (#177927)
Fixed a performance regression where Dynamo's decomposition pass exhibited O(N^2) behavior on large graphs. The fix avoids redundant node iteration during decomposition table lookup.
Dynamo Native itertools Variables (#186973, #186974)
Replaced polyfill implementations of itertools.product, itertools.chain, and related functions with native Dynamo variable tracking. This eliminates graph breaks when models use itertools in traced code and avoids unnecessary Python overhead.
Inductor NVGEMM Disk Cache (#187013)
Added persistent disk caching for NVGEMM (NVIDIA GEMM library) autotuning results. Previously, autotuning was repeated on every process restart. The disk cache persists winning kernel configurations across runs, significantly reducing warm-up time for workloads heavy in matrix multiplications.
DTensor single_dim_strategy for Reduction Ops (#179201)
Extended DTensor's single_dim_strategy to handle reduction operations. This enables more efficient sharding strategies when reductions operate on a single dimension, improving distributed training throughput for models with dimension-specific reductions.
MPS Metal Kernel Migrations
Several Metal kernel migrations for Apple Silicon:
index_add— moved from MPSGraph to native Metal shader for better performancelogical_not— native Metal implementation replacing MPSGraph path- Faster reduction kernels — optimized Metal shaders for sum/mean/max operations
Dynamo nb_inv Slot Support (#185641)
Added support for the __invert__ / nb_inv numeric slot in Dynamo's variable tracker. Models using bitwise inversion (~x) on custom types no longer cause graph breaks.
14. Summary
FX Concepts at a Glance
┌─────────────────────────────────────────────────┐
│ torch.fx │
│ │
│ symbolic_trace ──▶ Graph (Nodes) ──▶ GraphModule│
│ │
│ Node ops: │
│ placeholder, get_attr, call_function, │
│ call_method, call_module, output │
│ │
│ Transform APIs: │
│ ├── Manual: inserting_before/after, erase │
│ ├── replace_pattern (subgraph rewriter) │
│ ├── Interpreter (execute with hooks) │
│ └── Transformer (node-level rewrite) │
│ │
│ Validation: graph.lint(), ShapeProp │
│ │
│ In torch.compile: │
│ Dynamo → AOTAutograd → Inductor │
│ (all use FX Graphs internally) │
└─────────────────────────────────────────────────┘
When to Use Each API
| Goal | API |
|---|---|
| Inspect model structure | symbolic_trace + iterate graph.nodes |
| Simple op replacement | Manual node iteration, change node.target |
| Structural transforms (fuse, split) | Manual inserting_before/after + erase_node |
| Pattern-based replacement | replace_pattern |
| Per-node behavior (profiling, logging) | Interpreter subclass |
| Clean per-node transforms | Transformer subclass |
| Production optimization passes | Custom Inductor passes |
Key Rules
- Always call
graph.lint()after modifying a graph - Always call
gm.recompile()after modifying the graph of aGraphModule - Erase nodes bottom-up — a node can only be erased when it has zero users
replace_all_uses_withbefore erasing a node with userssymbolic_trace≠torch.compile— symbolic trace is simpler but less powerful; torch.compile (Dynamo) handles control flow and dynamic shapes
Further Reading
- torch.fx Official Docs — API reference
- torch.fx Technical Overview — design philosophy
- FX Graph Mode Quantization — quantization with FX
- Building a Custom Backend — receive FX graphs from torch.compile
- Inductor Deep Dive — how Inductor uses FX
Notebook: 23_fx_transforms.ipynb
Source Files
[README.md](README.md)— This guide — torch.fx theory, IR, passes, patterns[fx_basics.py](fx_basics.py)— Symbolic tracing, graph inspection, ShapeProp[graph_passes.py](graph_passes.py)— Graph transformations, pattern matching, Interpreter, Transformer
torch.masked — First-Class Missing Data in PyTorch
Table of Contents
- The Problem: Missing Data & Masking
- What is MaskedTensor?
- Creating MaskedTensors
- Masked Reductions
- Masked Softmax
- Masked Log Softmax and Normalize
- MaskedTensor Semantics
- Practical Example: Padded Sequence Mean
- Practical Example: Masked Attention
- MaskedTensor vs Manual Masking
- Current Limitations
- Upstream Updates (June 16, 2026)
- Summary
1. The Problem: Missing Data & Masking
Missing or invalid data appears in virtually every domain of deep learning:
| Domain | Scenario | What's "Missing" |
|---|---|---|
| NLP | Padded sequences in a batch | Positions beyond each sequence's true length |
| Vision | Irregular shapes, masked regions | Pixels outside the region of interest |
| Tabular | Incomplete records | Columns with no observed value |
| Attention | Causal masks, padding masks | Future tokens, padding positions |
The standard workarounds all have drawbacks:
Approach 1: Sentinel Values
# Replace missing values with 0 — but 0 is a valid number!
padded = torch.zeros(batch_size, max_len)
for i, seq in enumerate(sequences):
padded[i, :len(seq)] = seq
mean = padded.mean(dim=1) # WRONG — includes padding zeros
The mean is diluted by the padding zeros. For a sequence of length 3 padded to length 10, you compute sum / 10 instead of sum / 3.
Approach 2: masked_fill with -inf
# Common for attention: fill masked positions with -inf before softmax
scores = query @ key.T
scores = scores.masked_fill(mask == 0, float('-inf'))
probs = torch.softmax(scores, dim=-1)
# Works — but produces NaN if an entire row is masked
This works for softmax specifically, but is brittle. Different operations need different sentinel values (-inf for softmax, 0 for sum, +inf for min), and you must remember which to use.
Approach 3: Manual Boolean Masks
# Correct but verbose
mask = torch.arange(max_len).unsqueeze(0) < lengths.unsqueeze(1)
masked_sum = (data * mask.float()).sum(dim=1)
masked_mean = masked_sum / mask.float().sum(dim=1)
This is correct, but you carry two tensors (data and mask) through every operation, manually applying the mask at each step. It's easy to forget, and bugs are subtle.
The core issue: PyTorch operations don't natively understand that some elements are "not there." You must manually propagate this information, and every operation needs its own masking logic.
2. What is MaskedTensor?
MaskedTensor is a tensor subclass that bundles data and a boolean mask into a single object. The mask is a first-class citizen — operations automatically respect it.
from torch.masked import MaskedTensor
data = torch.tensor([1.0, 2.0, 3.0, 0.0, 0.0])
mask = torch.tensor([True, True, True, False, False])
mt = MaskedTensor(data, mask)
# MaskedTensor(
# [ 1.0000, 2.0000, 3.0000, --, --]
# )
Key properties:
mt.get_data()— returns the underlying data tensormt.get_mask()— returns the boolean maskTruemeans valid,Falsemeans masked/missing- Masked elements display as
--in the repr - Operations propagate the mask automatically
MaskedTensor is currently a prototype feature (as of PyTorch 2.14). Import it with:
from torch.masked import MaskedTensor
The torch.masked module also provides standalone masked operations that work on regular tensors with explicit mask arguments — useful even without MaskedTensor.
3. Creating MaskedTensors
From Data + Mask
The most common pattern: pair a data tensor with a boolean mask tensor of the same shape.
import torch
from torch.masked import MaskedTensor
# 1D
data = torch.tensor([10.0, 20.0, 30.0, 0.0])
mask = torch.tensor([True, True, True, False])
mt = MaskedTensor(data, mask)
# 2D — batch of sequences with padding
data = torch.tensor([
[1.0, 2.0, 3.0, 0.0, 0.0],
[4.0, 5.0, 0.0, 0.0, 0.0],
[6.0, 7.0, 8.0, 9.0, 0.0],
])
lengths = torch.tensor([3, 2, 4])
mask = torch.arange(5).unsqueeze(0) < lengths.unsqueeze(1)
# mask:
# tensor([[ True, True, True, False, False],
# [ True, True, False, False, False],
# [ True, True, True, True, False]])
mt = MaskedTensor(data, mask)
From Padded Sequences
When working with nn.utils.rnn.pad_sequence, you already have the lengths — just build the mask:
sequences = [torch.randn(3), torch.randn(5), torch.randn(2)]
padded = torch.nn.utils.rnn.pad_sequence(sequences, batch_first=True)
lengths = torch.tensor([3, 5, 2])
mask = torch.arange(padded.size(1)).unsqueeze(0) < lengths.unsqueeze(1)
mt = MaskedTensor(padded, mask)
Mask Requirements
- Shape: mask must be the same shape as data (broadcastable masks are not supported)
- Dtype: must be
torch.bool - Convention:
True= valid,False= masked
4. Masked Reductions
The torch.masked module provides reduction functions that correctly ignore masked elements. These work with plain tensors + mask arguments — you don't need MaskedTensor to use them.
torch.masked.sum
data = torch.tensor([
[1.0, 2.0, 3.0, 0.0, 0.0],
[4.0, 5.0, 0.0, 0.0, 0.0],
])
mask = torch.tensor([
[True, True, True, False, False],
[True, True, False, False, False],
])
# Regular sum includes padding zeros (happens to be correct for sum, but misleading)
regular_sum = data.sum(dim=1) # tensor([6., 9.])
# Masked sum — explicitly only sums valid elements
masked_sum = torch.masked._ops.sum(data, dim=1, mask=mask)
torch.masked.mean
This is where masking matters most — mean divides by the count of valid elements, not the total.
# Regular mean includes padding → WRONG
regular_mean = data.mean(dim=1) # tensor([1.2, 1.8]) (divides by 5)
# Masked mean → CORRECT
# Row 0: (1+2+3)/3 = 2.0
# Row 1: (4+5)/2 = 4.5
masked_mean = torch.masked._ops.mean(data, dim=1, mask=mask)
Other Masked Reductions
| Function | What it does |
|---|---|
torch.masked._ops.sum | Sum of valid elements |
torch.masked._ops.mean | Mean of valid elements (divides by valid count) |
torch.masked._ops.amax | Maximum of valid elements |
torch.masked._ops.amin | Minimum of valid elements |
torch.masked._ops.prod | Product of valid elements |
torch.masked._ops.norm | Norm over valid elements |
torch.masked._ops.var | Variance over valid elements |
torch.masked._ops.std | Standard deviation over valid elements |
All follow the same signature: func(data, dim, *, mask).
5. Masked Softmax
Softmax over masked data is one of the most common needs in attention mechanisms. The manual approach uses masked_fill with -inf:
Manual Approach
scores = torch.randn(2, 5)
mask = torch.tensor([
[True, True, True, False, False],
[True, True, True, True, False],
])
# Step 1: Fill masked positions with -inf
filled = scores.masked_fill(~mask, float('-inf'))
# Step 2: Softmax — exp(-inf) = 0, so masked positions get probability 0
probs = torch.softmax(filled, dim=1)
This works, but has a subtle problem: if an entire row is masked, softmax([-inf, -inf, ...]) produces NaN.
torch.masked.softmax
probs = torch.masked.softmax(scores, dim=1, mask=mask)
This handles the edge cases correctly and produces zero for masked positions without the NaN risk.
# Under the hood (simplified):
# 1. Replace masked positions with -inf
# 2. Compute softmax
# 3. Replace masked positions with 0 in the output
# 4. Handle all-masked rows gracefully
Comparison
scores = torch.tensor([[0.5, 1.2, 0.3, 0.0, 0.0]])
mask = torch.tensor([[True, True, True, False, False]])
# Manual
manual = torch.softmax(scores.masked_fill(~mask, float('-inf')), dim=1)
# tensor([[0.2753, 0.5545, 0.2253, 0.0000, 0.0000]])
# — masked positions are 0 because exp(-inf)=0, but sum of valid = 1.0525 ≠ 1
# torch.masked.softmax
masked = torch.masked.softmax(scores, dim=1, mask=mask)
# tensor([[0.2615, 0.5269, 0.2141, 0.0000, 0.0000]])
# — valid positions sum to ~1.0 (properly normalized over valid only)
The key difference: torch.masked.softmax normalizes over valid elements only, so the probabilities of valid positions sum to 1.0.
6. Masked Log Softmax and Normalize
torch.masked.log_softmax
Log-softmax is used in NLL loss and related computations. The masked version correctly computes log(softmax(x)) only over valid positions:
log_probs = torch.masked.log_softmax(scores, dim=1, mask=mask)
Masked positions in the output are set to 0.0 (or -inf depending on implementation), ensuring they don't contribute to downstream loss computations.
torch.masked.normalize
Normalize a tensor over valid elements only:
# L2 normalize each row, ignoring masked positions
normalized = torch.masked.normalize(data, ord=2.0, dim=1, mask=mask)
This computes the norm using only valid elements, then divides each valid element by that norm. Masked positions remain unchanged in the output.
7. MaskedTensor Semantics
When you use MaskedTensor directly, operations follow specific rules for mask propagation.
Unary Operations — Preserve Mask
Applying a unary function to a MaskedTensor keeps the same mask:
mt = MaskedTensor(
torch.tensor([1.0, -2.0, 3.0, 0.0]),
torch.tensor([True, True, True, False])
)
result = mt.abs()
# Data: [1.0, 2.0, 3.0, ???]
# Mask: [True, True, True, False]
# Masked element is still masked — the abs() was applied only to valid elements
Other unary ops that preserve masks: neg(), exp(), log(), sin(), cos(), sqrt(), relu(), etc.
Binary Operations — Intersection (AND) of Masks
When combining two MaskedTensors, an element is valid only if it's valid in both inputs:
a = MaskedTensor(
torch.tensor([1.0, 2.0, 3.0]),
torch.tensor([True, True, False])
)
b = MaskedTensor(
torch.tensor([10.0, 20.0, 30.0]),
torch.tensor([True, False, True])
)
result = a + b
# Data: [11.0, ???, ???]
# Mask: [True, False, False]
# Position 0: both valid → valid (1+10=11)
# Position 1: a valid, b masked → masked
# Position 2: a masked → masked
This is the conservative (safe) choice: if either input is missing, the output is missing.
Reductions — Collapse Mask
Reductions over a dimension collapse the mask along that dimension. The output position is valid if any input along the reduction axis was valid:
mt = MaskedTensor(
torch.tensor([[1.0, 2.0, 0.0],
[4.0, 0.0, 0.0]]),
torch.tensor([[True, True, False],
[True, False, False]])
)
# Sum over dim=1
# Row 0: sum of [1.0, 2.0] = 3.0 (2 valid elements)
# Row 1: sum of [4.0] = 4.0 (1 valid element)
result = mt.sum(dim=1)
8. Practical Example: Padded Sequence Mean
Computing the true mean of variable-length sequences is a common task. Let's compare three approaches.
Setup
# Three sequences of different lengths, padded to max_len=5
data = torch.tensor([
[3.0, 1.0, 4.0, 0.0, 0.0], # length 3, true mean = 2.667
[2.0, 7.0, 0.0, 0.0, 0.0], # length 2, true mean = 4.5
[5.0, 3.0, 2.0, 8.0, 0.0], # length 4, true mean = 4.5
])
lengths = torch.tensor([3, 2, 4])
Approach 1: Naive Mean (WRONG)
naive_mean = data.mean(dim=1)
# tensor([1.6000, 1.8000, 3.6000])
# All wrong! Divides by 5 instead of the actual lengths.
# Sequence 0: (3+1+4+0+0)/5 = 1.6, should be (3+1+4)/3 = 2.667
Approach 2: Manual Masking (Correct but Verbose)
mask = torch.arange(5).unsqueeze(0) < lengths.unsqueeze(1)
masked_sum = (data * mask.float()).sum(dim=1)
masked_count = mask.float().sum(dim=1)
manual_mean = masked_sum / masked_count
# tensor([2.6667, 4.5000, 4.5000]) ← Correct!
This works but requires 4 lines and careful bookkeeping. In a larger pipeline, you must thread the mask through every operation.
Approach 3: torch.masked API (Clean)
mask = torch.arange(5).unsqueeze(0) < lengths.unsqueeze(1)
masked_mean = torch.masked._ops.mean(data, dim=1, mask=mask)
# tensor([2.6667, 4.5000, 4.5000]) ← Correct!
One function call, no manual bookkeeping. The division by the valid count is handled internally.
9. Practical Example: Masked Attention
Attention mechanisms frequently need masking: padding masks (ignore padding tokens), causal masks (prevent attending to future positions), or combined masks.
Padding Mask in Attention
batch_size, seq_len, d_model = 2, 6, 8
query = torch.randn(batch_size, seq_len, d_model)
key = torch.randn(batch_size, seq_len, d_model)
value = torch.randn(batch_size, seq_len, d_model)
lengths = torch.tensor([4, 6]) # sequence 0 has 4 real tokens, sequence 1 has 6
# Build padding mask: [batch, 1, 1, seq_len] for broadcasting
pad_mask = torch.arange(seq_len).unsqueeze(0) < lengths.unsqueeze(1)
pad_mask = pad_mask.unsqueeze(1).unsqueeze(2) # [batch, 1, 1, seq_len]
# Attention scores
scores = (query @ key.transpose(-2, -1)) / (d_model ** 0.5)
# scores shape: [batch, seq_len, seq_len]
# Apply padding mask — masked positions get -inf
scores_2d_mask = pad_mask.squeeze(1) # [batch, 1, seq_len]
scores = scores.masked_fill(~scores_2d_mask, float('-inf'))
attn_weights = torch.softmax(scores, dim=-1)
# NaN appears in rows where all positions are masked
attn_weights = attn_weights.nan_to_num(0.0) # cleanup
Using torch.masked.softmax
mask_2d = torch.arange(seq_len).unsqueeze(0) < lengths.unsqueeze(1)
mask_3d = mask_2d.unsqueeze(1).expand(-1, seq_len, -1)
scores = (query @ key.transpose(-2, -1)) / (d_model ** 0.5)
attn_weights = torch.masked.softmax(scores, dim=-1, mask=mask_3d)
No masked_fill, no nan_to_num. The masked softmax handles everything, including the all-masked-row edge case.
10. MaskedTensor vs Manual Masking
A side-by-side comparison for common operations:
| Operation | Manual Masking | torch.masked API |
|---|---|---|
| Sum | (data * mask.float()).sum(dim) | torch.masked._ops.sum(data, dim, mask=mask) |
| Mean | (data * mask.float()).sum(dim) / mask.sum(dim) | torch.masked._ops.mean(data, dim, mask=mask) |
| Max | data.masked_fill(~mask, -inf).max(dim) | torch.masked._ops.amax(data, dim, mask=mask) |
| Min | data.masked_fill(~mask, inf).min(dim) | torch.masked._ops.amin(data, dim, mask=mask) |
| Softmax | softmax(data.masked_fill(~mask, -inf), dim) | torch.masked.softmax(data, dim, mask=mask) |
| Normalize | Compute norm manually, divide | torch.masked.normalize(data, ord, dim, mask=mask) |
When to Use Each
*Use torch.masked. functions when:**
- You need masked reductions (sum, mean, amax, etc.)
- You need masked softmax / log_softmax
- You want cleaner, less error-prone code
Use MaskedTensor when:
- You want automatic mask propagation through a pipeline
- You're prototyping and want to verify your masking logic
Stick with manual masking when:
- Performance is critical and you need full control
- You need operations not yet supported by MaskedTensor
- You're working with complex multi-mask scenarios
11. Current Limitations
MaskedTensor and torch.masked are in prototype status. Be aware of:
Not All Ops Are Supported
MaskedTensor works with a subset of PyTorch operations. Unsupported ops will raise errors:
# These work:
mt.sum(), mt.mean(), mt.abs(), mt + mt, mt * 2
# These may not work (as of 2.14):
# torch.nn.functional.linear(mt, weight) — not all nn.functional ops supported
# mt.view(...) — some shape ops may not be supported
Performance Overhead
MaskedTensor is a Python tensor subclass, which means:
- Extra Python dispatch overhead on every operation
- Not yet optimized by
torch.compilein all cases - For hot loops, manual masking may be faster
No Gradient Through Mask
The mask itself is not differentiable — it's a fixed boolean tensor. You cannot learn which elements to mask.
Sparse Mask Support
MaskedTensor supports sparse masks (COO and CSR) for memory efficiency when most elements are masked:
sparse_mask = mask.to_sparse()
mt = MaskedTensor(data, sparse_mask)
This can save memory when the mask is mostly False (most elements are masked).
API Stability
The torch.masked API may change between releases. Pin your PyTorch version for reproducibility, and check the release notes when upgrading.
12. Upstream Updates (June 16, 2026)
Recent PyTorch developments relevant to masking and related systems:
Dynamo O(N²) Decomposition Fix (#177927)
A performance bug was fixed where Dynamo's decomposition of certain ops had quadratic complexity in the number of elements. This affects any workload using torch.compile with masked operations, as the decomposed ops could include masking logic.
TokenSwitch for Distributed Token Routing (#178712)
A new TokenSwitch primitive for Mixture-of-Experts models enables efficient token routing across devices. This is relevant to masking because token routing inherently involves masking — tokens are assigned to specific experts, and the routing mask determines which tokens go where.
Native Itertools Variables in Dynamo (#186973, #186974)
Dynamo now handles Python itertools constructs (like itertools.chain, itertools.product) as native variables, reducing graph breaks. This benefits masked operations that iterate over mask patterns or dynamically construct masks in compiled code.
NVGEMM Disk Cache (#187013)
NVIDIA's GEMM kernel auto-tuning results are now cached to disk, avoiding re-tuning on subsequent runs. While not directly mask-related, this improves the startup time of any compiled workload, including those using masked operations.
MPS Metal Kernel Migrations
Ongoing work to migrate MPS (Apple Silicon) kernels from Objective-C++ to native Metal shaders. This improves performance of operations on Apple hardware, including masked operations on MPS devices.
DTensor Reduction Strategies (#179201)
New reduction strategies for DTensor (Distributed Tensor) improve how reductions are performed across devices. Since masked reductions are a key use case, this work may eventually enable efficient distributed masked operations.
13. Summary
Key Takeaways
| Concept | Description |
|---|---|
| The Problem | Missing data is everywhere: padding, irregular shapes, missing values. Manual masking is verbose and error-prone. |
| MaskedTensor | A tensor subclass pairing data + boolean mask. Operations respect the mask automatically. |
| torch.masked.softmax | Softmax that correctly ignores masked positions and normalizes over valid elements only. |
| Masked Reductions | torch.masked._ops.sum/mean/amax/amin — correct reductions that ignore masked elements. |
| Mask Convention | True = valid, False = masked/missing. |
| Unary Ops | Preserve the mask. |
| Binary Ops | AND the masks (both must be valid). |
| Prototype Status | Not all ops supported. API may change. Use for new prototyping, not critical production paths. |
Quick Reference
from torch.masked import MaskedTensor
# Create
mt = MaskedTensor(data, mask)
# Inspect
mt.get_data() # underlying data tensor
mt.get_mask() # boolean mask tensor
# Masked reductions (work on plain tensors too)
torch.masked._ops.sum(data, dim=1, mask=mask)
torch.masked._ops.mean(data, dim=1, mask=mask)
torch.masked._ops.amax(data, dim=1, mask=mask)
torch.masked._ops.amin(data, dim=1, mask=mask)
torch.masked._ops.prod(data, dim=1, mask=mask)
torch.masked._ops.var(data, dim=1, mask=mask)
# Masked softmax / log_softmax / normalize
torch.masked.softmax(data, dim=1, mask=mask)
torch.masked.log_softmax(data, dim=1, mask=mask)
torch.masked.normalize(data, ord=2, dim=1, mask=mask)
Further Reading
- torch.masked Official Docs — API reference
- MaskedTensor Overview — creation, semantics, sparsity
- MaskedTensor RFC — original design proposal
- Attention Mechanisms — where masked softmax is most commonly used
- Tensor Subclassing — how MaskedTensor is implemented under the hood
Notebook: 24_masked_tensor.ipynb
Source Files
[README.md](README.md)— This guide — torch.masked API, MaskedTensor, semantics[masked_tensor_basics.py](masked_tensor_basics.py)— Masked reductions, softmax, padded sequences, mask propagation
Custom Triton Kernels — GPU Programming in Python
Table of Contents
- What is Triton?
- Why Custom Triton Kernels?
- Triton Programming Model
- Hello World: Vector Addition
- Fused Add + ReLU
- Fused Softmax
- Matrix Multiplication
- Integrating Triton Kernels with PyTorch
- Autotuning
- Grid Functions
- Common Patterns
- Triton vs CUDA
- How TorchInductor Uses Triton
- Upstream Updates (June 2026)
- Summary & Next Steps
1. What is Triton?
Triton is OpenAI's open-source language and compiler for writing GPU kernels in Python. It sits between the ease of PyTorch and the raw power of CUDA C++:
Ease of use: PyTorch > Triton > CUDA C++
Performance: CUDA C++ ≈ Triton > PyTorch (eager)
Key facts about Triton:
- Python syntax — you write GPU kernels that look like Python (with NumPy-like operations), but they compile to PTX/SASS and run directly on NVIDIA GPUs.
- Near-peak performance — Triton's compiler handles tiling, shared memory, register allocation, and memory coalescing automatically. Well-written Triton kernels achieve 80-95% of hand-tuned CUDA performance.
- PyTorch uses Triton internally — TorchInductor (the
torch.compilebackend) generates Triton code for fused operations. When youtorch.compilea model, the generated kernels are Triton. - Custom kernel integration — you can write your own Triton kernels and register them as PyTorch custom ops, complete with autograd support, shape inference, and
torch.compilecompatibility.
Installation
Triton ships with PyTorch on Linux (CUDA builds). You can also install it standalone:
pip install triton
Note: Triton requires an NVIDIA GPU (Compute Capability 7.0+, i.e., Volta or later). All examples in this module detect GPU availability and provide explanations when running on CPU.
2. Why Custom Triton Kernels?
The Memory Bandwidth Problem
Modern GPUs have enormous compute throughput (e.g., A100: 312 TFLOPS for FP16) but relatively limited memory bandwidth (e.g., A100: 2 TB/s). For many operations, the bottleneck is not compute — it is moving data between GPU global memory and the compute units.
Consider a simple y = relu(x + bias):
Eager PyTorch (2 separate kernels):
Kernel 1: Read x, Read bias → Compute add → Write temp (2 reads + 1 write)
Kernel 2: Read temp → Compute relu → Write y (1 read + 1 write)
Total memory traffic: 3 reads + 2 writes = 5 memory ops
Fused Triton kernel (1 kernel):
Kernel: Read x, Read bias → Compute add+relu → Write y (2 reads + 1 write)
Total memory traffic: 2 reads + 1 write = 3 memory ops
The fused version does 40% less memory traffic. For larger fusion chains (common in Transformers), the savings are even greater.
Use Cases
| Use Case | Example |
|---|---|
| Fuse operations | Combine elementwise ops, reductions, and activations into one kernel |
| Custom ops | Implement operations that don't exist in PyTorch (novel attention variants, custom normalizations) |
| Eliminate overhead | Remove Python/dispatch overhead by running everything in a single GPU launch |
| Prototyping | Iterate on GPU kernel ideas 10x faster than CUDA C++ |
| Match Inductor | Write kernels that equal or beat what torch.compile generates |
3. Triton Programming Model
Block-Based Execution
Triton programs run as a grid of blocks (called "programs"). Each block processes a chunk of data independently and in parallel:
Data: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]
Block 0: [0, 1, 2, 3] pid=0, processes indices 0-3
Block 1: [4, 5, 6, 7] pid=1, processes indices 4-7
Block 2: [8, 9, 10, 11] pid=2, processes indices 8-11
Block 3: [12, 13, 14, 15] pid=3, processes indices 12-15
Key Concepts
| Concept | Description |
|---|---|
@triton.jit | Decorator that compiles a Python function into a GPU kernel |
tl.program_id(axis) | Returns the index of the current block along the given axis |
BLOCK_SIZE | Number of elements each block processes (a tl.constexpr) |
tl.arange(0, N) | Creates a range [0, 1, ..., N-1] within a block (like torch.arange) |
tl.load(ptr, mask) | Loads data from GPU memory. The mask handles out-of-bounds indices |
tl.store(ptr, val, mask) | Stores data to GPU memory with an optional mask |
| Grid | Total number of blocks to launch — grid = (num_blocks,) |
The tl.constexpr Annotation
Parameters marked as tl.constexpr are compile-time constants. Triton compiles a separate kernel for each unique value. This allows the compiler to make aggressive optimizations (unrolling, constant folding):
@triton.jit
def my_kernel(x_ptr, BLOCK_SIZE: tl.constexpr):
# BLOCK_SIZE is known at compile time
# Triton can fully unroll loops over BLOCK_SIZE
offsets = tl.arange(0, BLOCK_SIZE)
Memory Model
Unlike CUDA, Triton automatically manages shared memory (SRAM). When you tl.load data, the compiler decides whether to stage it through shared memory for reuse. You focus on the algorithm; the compiler handles the memory hierarchy.
4. Hello World: Vector Addition
The simplest Triton kernel — adding two vectors element by element:
import triton
import triton.language as tl
@triton.jit
def add_kernel(
x_ptr, # Pointer to first input vector
y_ptr, # Pointer to second input vector
out_ptr, # Pointer to output vector
n, # Total number of elements
BLOCK_SIZE: tl.constexpr, # Elements per block (compile-time)
):
# Step 1: Which block am I?
pid = tl.program_id(0)
# Step 2: Compute which indices this block handles
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
# Step 3: Mask to avoid out-of-bounds access
mask = offsets < n
# Step 4: Load inputs (masked — out-of-bounds loads return 0)
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
# Step 5: Compute
result = x + y
# Step 6: Store output (masked — out-of-bounds stores are skipped)
tl.store(out_ptr + offsets, result, mask=mask)
Line-by-Line Explanation
pid = tl.program_id(0)— Gets the block index along axis 0. For a 1D grid with 4 blocks, this returns 0, 1, 2, or 3.
- *
offsets = pidBLOCK_SIZE + tl.arange(0, BLOCK_SIZE)* — Computes the global indices this block processes. Block 0 gets[0, 1, ..., BS-1], block 1 gets[BS, BS+1, ..., 2BS-1], etc.
mask = offsets < n— Creates a boolean mask. The last block may extend past the array — the mask prevents reading/writing garbage.
tl.load(x_ptr + offsets, mask=mask)— Loads elements from GPU memory. Pointer arithmetic in Triton is element-wise (like C). The mask ensures out-of-bounds addresses are not accessed.
result = x + y— Standard addition. This happens in registers — no memory traffic.
tl.store(out_ptr + offsets, result, mask=mask)— Writes results back to GPU memory.
Launching the Kernel
import torch
def triton_add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
assert x.is_cuda and y.is_cuda
output = torch.empty_like(x)
n = x.numel()
BLOCK_SIZE = 1024
grid = (triton.cdiv(n, BLOCK_SIZE),) # Ceiling division
add_kernel[grid](x, y, output, n, BLOCK_SIZE=BLOCK_SIZE)
return output
# Usage
x = torch.randn(100_000, device='cuda')
y = torch.randn(100_000, device='cuda')
z = triton_add(x, y)
assert torch.allclose(z, x + y)
The kernelgrid syntax launches the kernel over the grid. triton.cdiv(n, BLOCK_SIZE) computes ceil(n / BLOCK_SIZE) — the number of blocks needed.
5. Fused Add + ReLU
Fusion is Triton's killer feature. Instead of two kernel launches, we do everything in one pass:
@triton.jit
def fused_add_relu_kernel(
x_ptr, y_ptr, out_ptr, n,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
# Fused: add + relu in one pass (no intermediate memory write)
result = tl.maximum(x + y, 0.0)
tl.store(out_ptr + offsets, result, mask=mask)
Why This is Faster
PyTorch eager:
temp = x + y # Kernel 1: read x,y → write temp (allocate temp tensor!)
out = relu(temp) # Kernel 2: read temp → write out
# 2 kernel launches, 1 temporary allocation, 5 memory transactions
Triton fused:
out = max(x + y, 0) # 1 kernel: read x,y → write out
# 1 kernel launch, 0 temporary allocations, 3 memory transactions
For a 10M element tensor in FP32, the temporary tensor alone is 40 MB of wasted memory bandwidth. On an A100 (2 TB/s bandwidth), that's ~20 microseconds of pure overhead eliminated.
6. Fused Softmax
A real-world kernel: computing softmax over rows of a matrix in a single pass.
The Algorithm
For each row x:
max_val = max(x)— for numerical stabilityx = x - max_val— shiftx = exp(x)— exponentiatesum_val = sum(x)— normalizeout = x / sum_val
In eager PyTorch, this involves multiple intermediate tensors. In Triton, we do it all in registers:
@triton.jit
def softmax_kernel(
input_ptr, output_ptr,
n_cols,
input_row_stride, output_row_stride,
BLOCK_SIZE: tl.constexpr,
):
# Each block processes one row
row_idx = tl.program_id(0)
# Pointers to the start of this row
row_start_ptr = input_ptr + row_idx * input_row_stride
col_offsets = tl.arange(0, BLOCK_SIZE)
mask = col_offsets < n_cols
# Load the entire row into SRAM
row = tl.load(row_start_ptr + col_offsets, mask=mask, other=float('-inf'))
# Compute softmax in registers
row_max = tl.max(row, axis=0)
numerator = tl.exp(row - row_max)
denominator = tl.sum(numerator, axis=0)
softmax_out = numerator / denominator
# Write back
out_start_ptr = output_ptr + row_idx * output_row_stride
tl.store(out_start_ptr + col_offsets, softmax_out, mask=mask)
Key Details
- One block per row — the grid size equals the number of rows.
other=float('-inf')— masked-out positions get negative infinity, soexp(-inf) = 0and they don't contribute to the sum.- Everything in registers/SRAM — the row is loaded once, all computation happens locally, and the result is written once. Eager PyTorch would create temporaries for each step.
Launching
def triton_softmax(x: torch.Tensor) -> torch.Tensor:
n_rows, n_cols = x.shape
BLOCK_SIZE = triton.next_power_of_2(n_cols)
output = torch.empty_like(x)
grid = (n_rows,)
softmax_kernel[grid](
x, output,
n_cols,
x.stride(0), output.stride(0),
BLOCK_SIZE=BLOCK_SIZE,
)
return output
Limitation: This simple kernel requires BLOCK_SIZE >= n_cols (the whole row must fit in one block). For very wide rows, you'd need a two-pass approach. In practice, Triton's maximum block size (up to 64K elements depending on dtype) handles most use cases.
7. Matrix Multiplication
Matrix multiplication demonstrates Triton's tiling model. We compute C = A @ B where A is (M, K) and B is (K, N):
@triton.jit
def matmul_kernel(
a_ptr, b_ptr, c_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
# 2D grid: each block computes a BLOCK_M x BLOCK_N tile of C
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
# Offsets for this tile
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Accumulator (initialized to zero)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Loop over K dimension in tiles of BLOCK_K
for k in range(0, K, BLOCK_K):
offs_k = k + tl.arange(0, BLOCK_K)
# Load tiles of A and B
a = tl.load(
a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak,
mask=(offs_m[:, None] < M) & (offs_k[None, :] < K),
other=0.0,
)
b = tl.load(
b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn,
mask=(offs_k[:, None] < K) & (offs_n[None, :] < N),
other=0.0,
)
# Accumulate: BLOCK_M x BLOCK_K @ BLOCK_K x BLOCK_N
acc += tl.dot(a, b)
# Store the output tile
tl.store(
c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn,
acc,
mask=(offs_m[:, None] < M) & (offs_n[None, :] < N),
)
Tiling Explained
Matrix C (M x N):
┌─────────────────────────┐
│ Block(0,0) │ Block(0,1) │ ← Each block computes a BLOCK_M x BLOCK_N tile
│ │ │
├────────────┼────────────┤
│ Block(1,0) │ Block(1,1) │
│ │ │
└─────────────────────────┘
For each tile of C, we iterate over K in chunks of BLOCK_K:
acc += A_tile @ B_tile (repeated K/BLOCK_K times)
- 2D grid —
grid = (M // BLOCK_M, N // BLOCK_N). Each block is identified by(pid_m, pid_n). tl.dot— hardware-accelerated matrix multiply on Tensor Cores (FP16/BF16/TF32).- Accumulator in FP32 — even if inputs are FP16, we accumulate in FP32 for precision.
- Shared memory is implicit — Triton automatically stages loaded tiles through SRAM. You never manually manage
__shared__memory.
8. Integrating Triton Kernels with PyTorch
Raw Triton kernels are useful, but to work with PyTorch's autograd, torch.compile, and other features, you need to register them as custom ops.
Step 1: Define the Kernel
@triton.jit
def _fused_gelu_kernel(x_ptr, out_ptr, n, BLOCK_SIZE: tl.constexpr):
pid = tl.program_id(0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n
x = tl.load(x_ptr + offsets, mask=mask)
# GELU approximation: 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))
out = 0.5 * x * (1.0 + tl.math.tanh(0.7978845608 * (x + 0.044715 * x * x * x)))
tl.store(out_ptr + offsets, out, mask=mask)
Step 2: Register as a Custom Op
@torch.library.custom_op("mylib::fused_gelu", mutates_args=())
def fused_gelu(x: torch.Tensor) -> torch.Tensor:
output = torch.empty_like(x)
n = x.numel()
grid = (triton.cdiv(n, 1024),)
_fused_gelu_kernel[grid](x, output, n, BLOCK_SIZE=1024)
return output
Step 3: Add Meta (Fake Tensor) Implementation
For torch.compile to trace through your op, it needs to know the output shape without running the kernel:
@fused_gelu.register_fake
def fused_gelu_fake(x: torch.Tensor) -> torch.Tensor:
return torch.empty_like(x)
Step 4: Add Autograd Support
def fused_gelu_setup_context(ctx, inputs, output):
(x,) = inputs
ctx.save_for_backward(x)
def fused_gelu_backward(ctx, grad_output):
(x,) = ctx.saved_tensors
# GELU derivative (could also be a Triton kernel)
grad_input = grad_output * (
0.5 * (1.0 + torch.tanh(0.7978845608 * (x + 0.044715 * x**3)))
+ 0.5 * x * (1.0 - torch.tanh(0.7978845608 * (x + 0.044715 * x**3))**2)
* 0.7978845608 * (1.0 + 3.0 * 0.044715 * x**2)
)
return grad_input
fused_gelu.register_autograd(fused_gelu_backward, setup_context=fused_gelu_setup_context)
Step 5: Use with torch.compile
@torch.compile
def model_forward(x):
return torch.ops.mylib.fused_gelu(x) # Seamlessly compiled
x = torch.randn(1024, requires_grad=True, device='cuda')
y = model_forward(x)
y.sum().backward() # Autograd works!
9. Autotuning
Different GPUs and problem sizes perform best with different block sizes. Triton provides built-in autotuning:
@triton.autotune(
configs=[
triton.Config({'BLOCK_SIZE': 128}),
triton.Config({'BLOCK_SIZE': 256}),
triton.Config({'BLOCK_SIZE': 512}),
triton.Config({'BLOCK_SIZE': 1024}),
triton.Config({'BLOCK_SIZE': 2048}),
],
key=['n'], # Re-tune when 'n' changes
)
@triton.jit
def add_kernel_autotuned(
x_ptr, y_ptr, out_ptr, n,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
tl.store(out_ptr + offsets, x + y, mask=mask)
How It Works
- First call — Triton benchmarks all configs and picks the fastest one for the given
n. - Subsequent calls with the same
n— uses the cached best config. - Different
n— re-benchmarks (since different sizes may have different optimal configs).
Matmul Autotuning
For more complex kernels, you can tune multiple parameters simultaneously:
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64, 'BLOCK_K': 32}, num_warps=4),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 32}, num_warps=4),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 32}, num_warps=8),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 32}, num_warps=8),
],
key=['M', 'N', 'K'],
)
@triton.jit
def matmul_kernel_autotuned(a_ptr, b_ptr, c_ptr, M, N, K, ...):
...
The num_warps parameter controls how many CUDA warps (groups of 32 threads) execute each block. More warps can hide memory latency but increase register pressure.
10. Grid Functions
The grid tells Triton how many blocks to launch. For simple kernels, you compute it directly:
grid = (triton.cdiv(n, BLOCK_SIZE),)
kernel[grid](...)
With autotuning, BLOCK_SIZE is chosen at runtime, so you need a lambda grid:
grid = lambda meta: (triton.cdiv(n, meta['BLOCK_SIZE']),)
kernel[grid](x_ptr, y_ptr, out_ptr, n)
The meta dict contains all tl.constexpr parameters. The lambda is called after autotuning selects a config.
2D Grids
For matmul-style kernels with two-dimensional tiling:
grid = lambda meta: (
triton.cdiv(M, meta['BLOCK_M']),
triton.cdiv(N, meta['BLOCK_N']),
)
Grid Considerations
| Factor | Guidance |
|---|---|
| Too few blocks | GPU SMs sit idle. Aim for at least num_SMs * 4 blocks |
| Too many blocks | Minor overhead from scheduling. Generally harmless |
| Block size too small | Instruction overhead dominates. Use 256+ for elementwise |
| Block size too large | Register spill, reduced occupancy |
11. Common Patterns
Reduction (Sum)
@triton.jit
def sum_kernel(x_ptr, out_ptr, n, BLOCK_SIZE: tl.constexpr):
# Single block reduction (for n <= BLOCK_SIZE)
offsets = tl.arange(0, BLOCK_SIZE)
mask = offsets < n
x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
total = tl.sum(x, axis=0)
tl.store(out_ptr, total)
For large arrays, you need a two-pass approach: each block reduces a chunk, then a second kernel reduces the partial sums.
Elementwise with Multiple Inputs
@triton.jit
def fused_bias_dropout_relu(
x_ptr, bias_ptr, out_ptr, n,
p_drop, # dropout probability
seed, # random seed
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n
x = tl.load(x_ptr + offsets, mask=mask)
bias = tl.load(bias_ptr + offsets % x.shape[0], mask=mask)
# Fused: bias + dropout + relu
x = x + bias
random = tl.rand(seed, offsets)
x = tl.where(random > p_drop, x / (1 - p_drop), 0.0)
x = tl.maximum(x, 0.0)
tl.store(out_ptr + offsets, x, mask=mask)
Online Softmax (Numerically Stable, Two-Pass in Registers)
The fused softmax kernel in Section 6 uses the standard approach. An online softmax computes max and sum in a single pass using the log-sum-exp trick, which is more register-efficient for very long rows.
Tiled Operations
For 2D operations (convolutions, attention), use 2D indexing:
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Use offs_m[:, None] and offs_n[None, :] for 2D indexing
12. Triton vs CUDA
| Aspect | Triton | CUDA C++ |
|---|---|---|
| Language | Python | C++ |
| Iteration speed | Fast (Python workflow, auto-compile) | Slow (compile, link, debug cycle) |
| Shared memory | Automatic (compiler-managed) | Manual (__shared__, bank conflict avoidance) |
| Thread-level control | Block-level only | Full warp/thread control |
| Performance | ~80-95% of hand-tuned CUDA | 100% (by definition) |
| Warp primitives | Limited (tl.atomic_*, basic reductions) | Full (__shfl_*, warp vote, cooperative groups) |
| Tensor Cores | Via tl.dot (automatic) | Via wmma or mma.sync (manual) |
| Portability | NVIDIA GPUs (AMD ROCm support WIP) | NVIDIA GPUs |
| Debugging | Print + assert | CUDA-GDB, Nsight Compute |
When Triton is Sufficient
- Elementwise operations (any complexity)
- Reductions (sum, max, mean, argmax)
- Matrix multiply (including with epilogues like bias + activation)
- Softmax, layer norm, RMS norm
- Attention variants (forward pass)
- Most operations you'd find in a Transformer
When You Need CUDA
- Warp-level primitives (warp shuffle, ballot)
- Complex thread synchronization patterns
- Custom memory access patterns (e.g., cross-warp communication)
- Extreme tuning for specific GPU architectures
- Operations on non-NVIDIA hardware (outside ROCm support)
13. How TorchInductor Uses Triton
When you call torch.compile(model), TorchInductor:
- Traces the model with Dynamo to get an FX graph
- Lowers operations to a scheduling IR
- Fuses compatible operations into groups
- Generates Triton kernels for each fused group
- Compiles the Triton kernels to PTX
- Caches the compiled kernels for reuse
Viewing Generated Triton Code
import torch
# Set environment variable BEFORE running
# TORCH_LOGS="output_code" python my_script.py
# Or programmatically:
import torch._logging
torch._logging.set_logs(output_code=True)
@torch.compile
def f(x, y):
return torch.relu(x + y)
x = torch.randn(1024, device='cuda')
y = torch.randn(1024, device='cuda')
f(x, y) # Check logs for generated Triton code
The generated code looks like:
# (Simplified example of Inductor-generated Triton)
@triton.jit
def triton_(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr):
xoffset = tl.program_id(0) * XBLOCK
xindex = xoffset + tl.arange(0, XBLOCK)
xmask = xindex < xnumel
x0 = xindex
tmp0 = tl.load(in_ptr0 + x0, xmask)
tmp1 = tl.load(in_ptr1 + x0, xmask)
tmp2 = tmp0 + tmp1
tmp3 = tl.maximum(tmp2, 0)
tl.store(out_ptr0 + x0, tmp3, xmask)
Notice it automatically fused add + relu — the same optimization we wrote by hand in Section 5!
Your Kernels + Inductor
When you register a Triton kernel as a custom_op with a register_fake implementation, Inductor can:
- Schedule your kernel alongside its generated kernels
- Fuse operations before/after your kernel (if applicable)
- Apply autotuning and caching
If you don't register as a custom op, your kernel appears as a graph break to Dynamo.
14. Upstream Updates (June 2026)
Recent PyTorch commits relevant to this module's topics (June 16-17, 2026):
| PR | Area | Summary |
|---|---|---|
| #187402 | Optimizers | Muon optimizer: spectral_unclamped scaling — new scaling strategy for the Muon optimizer that avoids clamping spectral norms, improving convergence for certain architectures |
| #186300 | Distributed | c10d abort hooks and pre/post collective hooks — new extensibility points for distributed collectives: register callbacks before and after collectives, and abort hooks for cleanup |
| #187387 | Distributed | Public torch.distributed.set_timeout — exposes a public API for setting distributed operation timeouts, replacing internal-only mechanisms |
| #183838 | Inductor | Unbacked FlexAttention predicates — Inductor now supports FlexAttention score_mod/mask_mod with unbacked SymInt predicates, enabling more dynamic attention patterns |
| #187406 | Testing | torchfuzz ~190 ops coverage expansion — the torchfuzz fuzzing framework now covers approximately 190 PyTorch operators, up from the initial set |
| #186976 | Dynamo | object() support — Dynamo can now trace through code that creates and compares object() sentinels, eliminating a common source of graph breaks |
These updates reflect the ongoing evolution of PyTorch's compilation stack, distributed infrastructure, and testing tooling — all areas that interact with custom Triton kernel development.
15. Summary & Next Steps
What We Learned
| Topic | Key Takeaway |
|---|---|
| Triton | Write GPU kernels in Python with near-CUDA performance |
| Programming model | Grid of blocks, program_id, BLOCK_SIZE, load/store with masks |
| Fusion | Combine operations to eliminate memory bandwidth waste |
| Softmax | Practical kernel: load row, compute in registers, write once |
| Matmul | Tiled approach with tl.dot for Tensor Core utilization |
| PyTorch integration | custom_op → register_fake → register_autograd pipeline |
| Autotuning | @triton.autotune automatically finds the best config |
| TorchInductor | Generates Triton code from torch.compile — your kernels can interact with it |
When to Write Custom Triton Kernels
- torch.compile already fuses your ops — check first!
TORCH_LOGS="output_code"shows what Inductor generates. Often it's already optimal. - Custom logic — novel attention, custom normalization, or domain-specific ops that PyTorch doesn't support natively.
- Squeeze the last 10% — when profiling shows a specific kernel is the bottleneck and you can beat Inductor's generated code.
Further Resources
- Triton Documentation — official tutorials and API reference
- Triton GitHub — source code and examples
- PyTorch Custom Operators —
torch.libraryAPI reference - TorchInductor Deep Dive — how
torch.compileworks under the hood - Training Pipelines — where custom kernels fit in the training loop
Notebook: 25_triton_kernels.ipynb
Source Files
[README.md](README.md)— This guide — Triton programming model, kernels, PyTorch integration, autotuning[triton_basics.py](triton_basics.py)— Vector add, fused add+ReLU, fused softmax kernels with benchmarks[triton_with_pytorch.py](triton_with_pytorch.py)— torch.library registration, autograd, torch.compile, autotuning
GPU Memory Profiling & Optimization — Every Byte Accounted For
Table of Contents
- Where Does GPU Memory Go?
- torch.cuda.memory_allocated / memory_reserved
- torch.cuda.memory_summary()
- Peak Memory Tracking
- torch.cuda.memory_stats()
- Memory Snapshots
- Finding Memory Leaks
- torch.cuda.empty_cache()
- Memory Optimization Techniques
- torch.profiler for Memory
- Memory-Efficient Attention
- Practical: Estimating Memory Before Training
- Upstream Updates (June 17–18, 2026)
1. Where Does GPU Memory Go?
Before optimizing, you need to understand what consumes GPU memory during training. Every byte falls into one of these categories:
1.1 CUDA Context Overhead
The CUDA runtime itself consumes memory just by being initialized — typically ~300–800 MB depending on GPU architecture and driver version:
import torch
torch.cuda.init() # triggers context creation
print(torch.cuda.memory_reserved()) # ~300-800 MB before any tensors
This is unavoidable. A freshly initialized CUDA context on an A100 typically uses ~300 MB, while an H100 may use ~500 MB.
1.2 Model Parameters
Each parameter stores weights in the model's dtype:
Parameter Memory = num_params × bytes_per_element
fp32: 4 bytes/param → 1B params = 4 GB
bf16: 2 bytes/param → 1B params = 2 GB
fp16: 2 bytes/param → 1B params = 2 GB
1.3 Optimizer State
Optimizers store per-parameter state tensors. This is often the largest memory consumer:
| Optimizer | State per Parameter | Total for N params (fp32) |
|---|---|---|
| SGD (no momentum) | 0 | 0 |
| SGD + momentum | 1× (momentum buffer) | 4N bytes |
| Adam / AdamW | 2× (first moment m, second moment v) | 8N bytes |
| Adam + master weights (bf16 training) | 2× moments + 1× master copy | 12N bytes |
| Adafactor | ~1× (row/col factors) | ~4N bytes |
Adam dominates: for a model with N fp32 params, Adam needs 8N bytes of optimizer state — twice the parameter memory.
1.4 Gradients
During backward, each parameter accumulates gradients in the parameter's dtype (or fp32 for mixed precision):
Gradient Memory = num_params × bytes_per_element
= same as parameter memory (1×)
1.5 Activations (Forward Pass Intermediates)
The biggest variable. Activations saved for backward scale with:
Activation Memory ∝ batch_size × seq_len × hidden_dim × num_layers
For a Transformer with L layers, hidden size H, sequence length S, batch size B:
Per-layer activations ≈ 2 × B × S × H × bytes_per_element (input + output)
Attention scores ≈ B × num_heads × S × S × bytes_per_element
Total activations ≈ L × (per-layer + attention) activations
Activations grow linearly with batch size and number of layers, and quadratically with sequence length (for standard attention).
1.6 Fragmentation
PyTorch uses a caching allocator. Memory that has been freed but not returned to CUDA appears as "reserved but not allocated." Fragmentation occurs when freed blocks can't be coalesced into contiguous chunks for new allocations.
1.7 Complete Memory Budget Example: 7B Model
Model: 7B parameters, bf16, Adam optimizer, batch=4, seq=2048, 32 layers
Parameters: 7B × 2 bytes = 14 GB (bf16)
Gradients: 7B × 2 bytes = 14 GB (bf16)
Adam state: 7B × 4 bytes × 2 = 56 GB (fp32 moments: m + v)
Master weights: 7B × 4 bytes = 28 GB (fp32 copy for mixed precision)
Activations: ~12-20 GB (depends on architecture details)
CUDA context: ~0.5 GB
─────────────────────────────────────────
Total: ~125-133 GB → needs multiple GPUs or FSDP
This is why a 7B model requires at least 2× 80 GB GPUs (A100/H100) for full-parameter fine-tuning with Adam. Techniques like LoRA, 8-bit optimizers, and gradient checkpointing reduce this dramatically.
2. torch.cuda.memory_allocated / memory_reserved
These two functions tell you the current state of GPU memory:
import torch
# memory_allocated: bytes actively used by tensors
allocated = torch.cuda.memory_allocated() # in bytes
# memory_reserved: bytes held by PyTorch's caching allocator
reserved = torch.cuda.memory_reserved() # in bytes
# The gap = cached but not actively used (fragmentation + cache)
gap = reserved - allocated
What Each Means
memory_allocated()— memory occupied by tensors that currently exist. When youdela tensor, this number drops.memory_reserved()— total memory PyTorch has claimed from CUDA. This only drops when you callempty_cache()or PyTorch returns blocks to CUDA under memory pressure.
x = torch.randn(1000, 1000, device='cuda')
print(f"Allocated: {torch.cuda.memory_allocated() / 1e6:.1f} MB")
print(f"Reserved: {torch.cuda.memory_reserved() / 1e6:.1f} MB")
del x
# Allocated drops, but reserved stays the same
print(f"After del - Allocated: {torch.cuda.memory_allocated() / 1e6:.1f} MB")
print(f"After del - Reserved: {torch.cuda.memory_reserved() / 1e6:.1f} MB")
torch.cuda.empty_cache()
# Now reserved drops too
print(f"After empty_cache - Reserved: {torch.cuda.memory_reserved() / 1e6:.1f} MB")
Helper for Human-Readable Output
def print_memory(tag=""):
alloc = torch.cuda.memory_allocated() / 1024**2
res = torch.cuda.memory_reserved() / 1024**2
print(f"[{tag}] Allocated: {alloc:.1f} MB | Reserved: {res:.1f} MB | Gap: {res-alloc:.1f} MB")
3. torch.cuda.memory_summary()
For a detailed breakdown, memory_summary() prints a formatted table:
print(torch.cuda.memory_summary())
This prints a table with columns like:
| | Cur Usage | Peak Usage | Tot Alloc | Tot Freed |
| Allocated Bytes | 4.00 MB | 16.00 MB | 128.00 MB | 124.00 MB |
| Active Bytes | 4.00 MB | 16.00 MB | 128.00 MB | 124.00 MB |
| Reserved Bytes | 20.00 MB | 20.00 MB | 20.00 MB | 0.00 MB |
| Inactive Split | 0.00 MB | ... | ... | ... |
Key Rows Explained
| Row | Meaning |
|---|---|
| Allocated Bytes | Memory occupied by live tensors |
| Active Bytes | Same as allocated (non-released blocks) |
| Reserved Bytes | Total memory held by the caching allocator |
| Inactive Split Bytes | Freed memory within split blocks (fragmentation indicator) |
| Allocation count | How many cudaMalloc calls were made |
| Active allocs | Number of currently live tensor allocations |
Reading the Table for Fragmentation
If Inactive Split Bytes is large relative to Reserved Bytes, you have significant fragmentation. This means memory was allocated, some tensors were freed, but the freed chunks are sandwiched between live allocations and can't be coalesced.
# Common usage: log memory state at key points
model = build_model().cuda()
print(torch.cuda.memory_summary(abbreviated=True))
output = model(input_batch)
print(torch.cuda.memory_summary(abbreviated=True))
loss = criterion(output, target)
loss.backward()
print(torch.cuda.memory_summary(abbreviated=True))
4. Peak Memory Tracking
max_memory_allocated / max_memory_reserved
Track the high-water mark of GPU memory usage:
# Reset counters before your experiment
torch.cuda.reset_peak_memory_stats()
# Run your training step
output = model(batch)
loss = criterion(output, targets)
loss.backward()
optimizer.step()
# Check peak usage during that step
peak_alloc = torch.cuda.max_memory_allocated() / 1024**3
peak_res = torch.cuda.max_memory_reserved() / 1024**3
print(f"Peak allocated: {peak_alloc:.2f} GB")
print(f"Peak reserved: {peak_res:.2f} GB")
reset_peak_memory_stats()
Resets the peak counters without clearing any cached memory. Use this to compare memory usage between different configurations:
# Experiment 1: batch_size=32
torch.cuda.reset_peak_memory_stats()
train_step(model, batch_32)
peak_bs32 = torch.cuda.max_memory_allocated()
# Experiment 2: batch_size=64
torch.cuda.reset_peak_memory_stats()
train_step(model, batch_64)
peak_bs64 = torch.cuda.max_memory_allocated()
print(f"BS=32 peak: {peak_bs32/1e9:.2f} GB")
print(f"BS=64 peak: {peak_bs64/1e9:.2f} GB")
print(f"Ratio: {peak_bs64/peak_bs32:.2f}x")
5. torch.cuda.memory_stats()
For programmatic access to all memory statistics (useful for logging to TensorBoard/W&B):
stats = torch.cuda.memory_stats()
# Key fields
print(f"Current allocated: {stats['allocated_bytes.all.current'] / 1e9:.2f} GB")
print(f"Peak allocated: {stats['allocated_bytes.all.peak'] / 1e9:.2f} GB")
print(f"Current reserved: {stats['reserved_bytes.all.current'] / 1e9:.2f} GB")
print(f"Peak reserved: {stats['reserved_bytes.all.peak'] / 1e9:.2f} GB")
print(f"Active allocations: {stats['active.all.current']}")
print(f"Peak active allocs: {stats['active.all.peak']}")
print(f"Total alloc calls: {stats['allocation.all.current']}")
Important stat keys
| Key Pattern | Description |
|---|---|
allocated_bytes.all.current | Current allocated memory (bytes) |
allocated_bytes.all.peak | Peak allocated memory |
reserved_bytes.all.current | Current reserved memory |
reserved_bytes.all.peak | Peak reserved memory |
active.all.current | Number of currently active allocations |
active.all.peak | Peak number of simultaneous allocations |
inactive_split_bytes.all.current | Current fragmentation (inactive splits) |
num_alloc_retries | How many times allocator retried after cudaMalloc failure |
num_ooms | Number of OOM errors caught by the allocator |
Logging to Training Loop
def log_memory_stats(step, writer):
stats = torch.cuda.memory_stats()
writer.add_scalar('memory/allocated_gb',
stats['allocated_bytes.all.current'] / 1e9, step)
writer.add_scalar('memory/peak_allocated_gb',
stats['allocated_bytes.all.peak'] / 1e9, step)
writer.add_scalar('memory/reserved_gb',
stats['reserved_bytes.all.current'] / 1e9, step)
writer.add_scalar('memory/fragmentation_mb',
stats['inactive_split_bytes.all.current'] / 1e6, step)
writer.add_scalar('memory/num_alloc_retries',
stats['num_alloc_retries'], step)
6. Memory Snapshots
Memory snapshots capture a complete record of every allocation and deallocation, including Python stack traces. This is the most powerful debugging tool for memory issues.
Recording Snapshots
# Start recording allocation history
torch.cuda.memory._record_memory_history(max_entries=100_000)
# Run your code
model = MyModel().cuda()
output = model(batch)
loss = criterion(output, target)
loss.backward()
optimizer.step()
# Save the snapshot
torch.cuda.memory._dump_snapshot("snapshot.pickle")
# Stop recording
torch.cuda.memory._record_memory_history(enabled=None)
Visualizing with PyTorch Memory Viz
PyTorch provides an interactive HTML visualizer. There are two ways to use it:
Option 1: Use the online tool
Upload snapshot.pickle to pytorch.org/memory_viz.
Option 2: Generate HTML locally
from torch.cuda._memory_viz import segment_plot, trace_plot
# Read the snapshot
import pickle
with open("snapshot.pickle", "rb") as f:
snapshot = pickle.load(f)
# Generate plots
with open("segment_plot.html", "w") as f:
f.write(segment_plot(snapshot))
with open("trace_plot.html", "w") as f:
f.write(trace_plot(snapshot))
What Snapshots Show
- Segment plot: shows how memory blocks are allocated over time, colored by which Python call site allocated them. Reveals fragmentation patterns.
- Trace plot: timeline of allocations/deallocations. Shows exactly where peak memory occurs and what's live at that point.
- Stack traces: for each allocation, you get the full Python stack trace, so you know exactly which line of code created each tensor.
Best Practices for Snapshots
- Record during a single training step — recording for many steps generates huge files
- Use
max_entriesto limit snapshot size - Call
torch.cuda.empty_cache()before recording to reduce noise from cached blocks - Compare snapshots between configurations to identify which optimization helped
7. Finding Memory Leaks
A memory leak in PyTorch means tensors that should be garbage-collected are kept alive by unintentional references. The symptom is memory_allocated() growing monotonically across training steps.
Common Cause 1: Accumulating Tensors in Lists
# BAD — stores computation graph for every step
all_losses = []
for batch in dataloader:
loss = model(batch).sum()
all_losses.append(loss) # holds entire computation graph!
# GOOD — detach before storing
all_losses = []
for batch in dataloader:
loss = model(batch).sum()
all_losses.append(loss.item()) # scalar, no graph reference
Common Cause 2: Not Detaching from Computation Graph
# BAD — hidden_state retains the graph from the forward pass
hidden_state = model.get_hidden(batch)
# hidden_state.grad_fn exists, keeping all intermediates alive
# GOOD — detach to break the graph
hidden_state = model.get_hidden(batch).detach()
Common Cause 3: Closures Holding References
# BAD — closure captures `output` tensor
def make_callback(output):
def callback():
print(output.shape) # keeps output alive indefinitely
return callback
# GOOD — capture only what you need
def make_callback(shape):
def callback():
print(shape)
return callback
callback = make_callback(output.shape)
del output
Common Cause 4: Global Variables and Caches
# BAD — module-level cache grows without bound
_cache = {}
def forward(x, key):
result = model(x)
_cache[key] = result # never cleared!
return result
# GOOD — use bounded cache or WeakRef
from weakref import WeakValueDictionary
_cache = WeakValueDictionary()
Detection Pattern
def detect_memory_leak(model, dataloader, num_steps=10):
"""Run a few steps and check if memory grows."""
torch.cuda.reset_peak_memory_stats()
torch.cuda.empty_cache()
memory_readings = []
for i, batch in enumerate(dataloader):
if i >= num_steps:
break
output = model(batch)
loss = output.sum()
loss.backward()
model.zero_grad(set_to_none=True)
torch.cuda.synchronize()
memory_readings.append(torch.cuda.memory_allocated())
# Check for monotonic growth
for i in range(1, len(memory_readings)):
if memory_readings[i] > memory_readings[0] * 1.1:
print(f"WARNING: Memory grew from {memory_readings[0]/1e6:.1f} MB "
f"to {memory_readings[i]/1e6:.1f} MB in {i} steps")
return True
print("No memory leak detected")
return False
8. torch.cuda.empty_cache()
What It Does
empty_cache() releases all unused cached memory held by the PyTorch caching allocator back to CUDA. It does NOT free tensors that are still alive.
# Before: reserved = 2 GB, allocated = 500 MB
torch.cuda.empty_cache()
# After: reserved ≈ 500 MB, allocated = 500 MB (unchanged)
When to Use
- Between training phases — after loading a model but before the first forward pass
- After deleting large tensors — when you know you won't need that memory pattern again
- Before memory-critical operations — to maximize available contiguous memory
- Between experiments — when switching from one model/config to another
# Good: between phases
model_1 = train_phase_1(...)
del model_1
torch.cuda.empty_cache() # return memory to CUDA before loading phase 2
model_2 = load_phase_2_model()
When NOT to Use
- During training — calling
empty_cache()every step forces reallocation overhead - As a "fix" for OOM — it won't help if your tensors genuinely don't fit
- Inside tight loops — the caching allocator exists to avoid expensive
cudaMalloc/cudaFreecalls
# BAD — defeats the purpose of the caching allocator
for batch in dataloader:
output = model(batch)
loss = criterion(output, target)
loss.backward()
optimizer.step()
torch.cuda.empty_cache() # forces realloc every step — slower!
9. Memory Optimization Techniques
9.1 Gradient Checkpointing (Trade Compute for Memory)
Instead of storing all intermediate activations during forward, recompute them during backward:
from torch.utils.checkpoint import checkpoint
class CheckpointedTransformerBlock(nn.Module):
def __init__(self, block):
super().__init__()
self.block = block
def forward(self, x):
return checkpoint(self.block, x, use_reentrant=False)
Memory savings: reduces activation memory from O(L) to O(√L) for L layers. Cost: ~33% more compute (one extra forward pass per checkpointed segment).
For SAC (Selective Activation Checkpointing), see Module 16.
9.2 Mixed Precision Training (Halve Activation Memory)
Using bf16 or fp16 instead of fp32 halves the memory for activations and parameters:
# Using torch.autocast
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
output = model(batch)
loss = criterion(output, target)
# Or cast model directly
model = model.to(dtype=torch.bfloat16)
Memory savings: activations and parameters cut in half. Optimizer state may still be fp32.
9.3 Gradient Accumulation (Smaller Per-Step Batch)
Process a large effective batch using smaller micro-batches:
accumulation_steps = 4
optimizer.zero_grad()
for i, batch in enumerate(dataloader):
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
output = model(batch)
loss = criterion(output, target) / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
Memory savings: peak activation memory scales with micro_batch_size, not effective_batch_size. Using 4× accumulation = 4× smaller activation footprint.
9.4 In-Place Operations
Some operations can be done in-place to avoid allocating new tensors:
# Out-of-place: allocates new tensor
x = x + 1 # new tensor
x = F.relu(x) # new tensor
# In-place: modifies tensor directly
x.add_(1) # no new allocation
x = F.relu(x, inplace=True) # reuses memory
Warning: in-place operations on tensors that require gradients can break autograd:
# This will raise an error if x requires grad and is needed for backward
x.add_(1) # RuntimeError: a leaf Variable that requires grad has been used in an in-place operation
Use in-place operations only on tensors that don't need gradients or that aren't saved for backward.
9.5 del + gc.collect() for Large Intermediates
Explicitly delete large tensors and trigger garbage collection:
import gc
# After using a large intermediate
large_features = extract_features(data) # several GB
predictions = classify(large_features)
del large_features
gc.collect()
torch.cuda.empty_cache() # optional: return blocks to CUDA
9.6 torch.cuda.empty_cache() Between Phases
See Section 8 above for detailed guidance.
9.7 CPU Offloading
FSDP2 can offload parameters and gradients to CPU between forward/backward:
from torch.distributed._composable.fsdp import fully_shard, CPUOffloadPolicy
policy = CPUOffloadPolicy(pin_memory=True)
for layer in model.layers:
fully_shard(layer, offload_policy=policy)
fully_shard(model, offload_policy=policy)
Memory savings: massive — only one layer's parameters on GPU at a time. Cost: CPU↔GPU transfer overhead. pin_memory=True helps with async transfers.
9.8 Reducing Optimizer Memory
| Approach | Memory vs. Adam | Trade-off |
|---|---|---|
| SGD + momentum | 50% less | May need different hyperparameters |
| 8-bit Adam (bitsandbytes) | 75% less | Slight accuracy impact |
| Adafactor | ~50% less | Uses row/column factorization |
| LoRA (low-rank adaptation) | 90%+ less | Only trains adapter weights |
| GaLore | ~65% less | Projects gradients to low-rank space |
# 8-bit Adam with bitsandbytes
import bitsandbytes as bnb
optimizer = bnb.optim.Adam8bit(model.parameters(), lr=1e-4)
Summary Table
| Technique | Saves | Cost | Typical Reduction |
|---|---|---|---|
| Gradient checkpointing | Activations | ~33% more compute | 60-70% activation memory |
| Mixed precision (bf16) | Parameters + activations | None (sometimes better) | ~50% |
| Gradient accumulation | Activations | None | Proportional to accum steps |
| In-place operations | Intermediate tensors | Autograd limitations | 5-15% |
del + gc.collect() | Named intermediates | Manual management | Variable |
| CPU offloading | Parameters + optimizer | Transfer overhead | Up to 90% GPU memory |
| 8-bit optimizers | Optimizer state | Slight accuracy impact | 75% optimizer memory |
10. torch.profiler for Memory
The PyTorch profiler can track memory allocations alongside compute:
from torch.profiler import profile, ProfilerActivity, schedule, tensorboard_trace_handler
with profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
profile_memory=True, # enable memory profiling
record_shapes=True,
with_stack=True, # include Python stack traces
schedule=schedule(wait=1, warmup=1, active=3, repeat=1),
on_trace_ready=tensorboard_trace_handler('./log/memory_profile'),
) as prof:
for step, batch in enumerate(dataloader):
output = model(batch)
loss = criterion(output, target)
loss.backward()
optimizer.step()
optimizer.zero_grad()
prof.step()
Memory Timeline
With profile_memory=True, the profiler records every allocation and deallocation. In TensorBoard, this appears as a memory timeline showing:
- Memory curve: total allocated memory over time
- Allocation spikes: where peak memory occurs (usually during backward)
- Memory events: individual tensor allocations, tagged with their size and operator
Identifying Allocation Spikes
# Print the top memory-consuming operations
print(prof.key_averages().table(
sort_by="self_cuda_memory_usage",
row_limit=20
))
This table shows which operations allocate the most GPU memory — the targets for optimization.
Export for Chrome Trace Viewer
prof.export_chrome_trace("trace_with_memory.json")
# Open in chrome://tracing — memory events appear alongside compute
11. Memory-Efficient Attention
Standard dot-product attention has O(N²) memory complexity in sequence length:
Standard attention: stores full N×N attention matrix
Memory = batch × heads × seq_len × seq_len × bytes_per_element
Example: batch=8, heads=32, seq=4096, bf16
= 8 × 32 × 4096 × 4096 × 2 = 8.6 GB ← just for attention scores!
Flash Attention: O(N) Memory
Flash Attention (used by default in F.scaled_dot_product_attention) never materializes the full N×N matrix. Instead, it computes attention in tiles:
import torch.nn.functional as F
# This automatically uses Flash Attention when available
output = F.scaled_dot_product_attention(query, key, value)
# Memory: O(N) instead of O(N²)
# For seq=4096: ~100× less memory for the attention computation
Memory Impact on Long Sequences
| Sequence Length | Standard Attention | Flash Attention | Savings |
|---|---|---|---|
| 512 | 16 MB | 0.25 MB | 64× |
| 2048 | 256 MB | 1 MB | 256× |
| 4096 | 1 GB | 2 MB | 512× |
| 16384 | 16 GB | 8 MB | 2048× |
| 65536 | 256 GB (impossible) | 32 MB | — |
Flash Attention is what makes long-context LLMs (100k+ tokens) feasible.
FlexAttention
For custom attention patterns with memory efficiency:
from torch.nn.attention.flex_attention import flex_attention, create_block_mask
def causal_mask(b, h, q_idx, kv_idx):
return q_idx >= kv_idx
block_mask = create_block_mask(causal_mask, B=1, H=1, Q_LEN=4096, KV_LEN=4096)
output = flex_attention(query, key, value, block_mask=block_mask)
FlexAttention maintains O(N) memory while supporting arbitrary attention patterns. See Module 09 for full coverage.
12. Practical: Estimating Memory Before Training
Using meta Device
The meta device creates tensors with shapes and dtypes but no actual storage — perfect for estimating memory without a GPU:
with torch.device('meta'):
model = MyLargeModel(config)
# Count parameters
total_params = sum(p.numel() for p in model.parameters())
param_bytes = sum(p.numel() * p.element_size() for p in model.parameters())
print(f"Parameters: {total_params:,} ({param_bytes / 1e9:.2f} GB)")
Full Memory Estimator
def estimate_training_memory(
model_params: int,
dtype_bytes: int = 2, # 2 for bf16, 4 for fp32
optimizer: str = "adam", # "sgd", "adam", "adam_8bit"
batch_size: int = 1,
seq_len: int = 2048,
hidden_dim: int = 4096,
num_layers: int = 32,
num_heads: int = 32,
use_flash_attn: bool = True,
gradient_checkpointing: bool = False,
cuda_context_gb: float = 0.5,
) -> dict:
"""Estimate GPU memory needed for training."""
# Parameters
param_gb = model_params * dtype_bytes / 1e9
# Gradients (same dtype as params)
grad_gb = param_gb
# Optimizer state
if optimizer == "sgd":
optim_gb = model_params * 4 / 1e9 # momentum buffer in fp32
elif optimizer == "adam":
optim_gb = model_params * 4 * 2 / 1e9 # m + v in fp32
if dtype_bytes < 4:
optim_gb += model_params * 4 / 1e9 # master weights
elif optimizer == "adam_8bit":
optim_gb = model_params * 1 * 2 / 1e9 # 8-bit m + v
# Activations (rough estimate for Transformer)
bytes_per_activation = dtype_bytes
per_layer_act = 2 * batch_size * seq_len * hidden_dim * bytes_per_activation
if use_flash_attn:
attn_mem = batch_size * num_heads * seq_len * 64 * bytes_per_activation
else:
attn_mem = batch_size * num_heads * seq_len * seq_len * bytes_per_activation
effective_layers = num_layers
if gradient_checkpointing:
effective_layers = int(num_layers ** 0.5) # sqrt(L) with checkpointing
act_gb = effective_layers * (per_layer_act + attn_mem) / 1e9
total = param_gb + grad_gb + optim_gb + act_gb + cuda_context_gb
return {
"parameters_gb": param_gb,
"gradients_gb": grad_gb,
"optimizer_gb": optim_gb,
"activations_gb": act_gb,
"cuda_context_gb": cuda_context_gb,
"total_gb": total,
}
"Will This Fit on My GPU?" Calculator
def will_it_fit(total_memory_gb: float, gpu: str = "A100-80GB") -> dict:
gpu_memory = {
"RTX-3090": 24, "RTX-4090": 24, "A100-40GB": 40,
"A100-80GB": 80, "H100-80GB": 80, "H200-141GB": 141,
}
available = gpu_memory.get(gpu, 80)
usable = available * 0.90 # ~10% overhead for CUDA/driver
fits = total_memory_gb <= usable
headroom = usable - total_memory_gb
return {
"gpu": gpu,
"gpu_memory_gb": available,
"usable_gb": usable,
"estimated_gb": total_memory_gb,
"fits": fits,
"headroom_gb": headroom if fits else 0,
"gpus_needed": max(1, int(total_memory_gb / usable) + 1) if not fits else 1,
}
Example: 13B Model
estimate = estimate_training_memory(
model_params=13_000_000_000,
dtype_bytes=2, # bf16
optimizer="adam",
batch_size=4,
seq_len=2048,
hidden_dim=5120,
num_layers=40,
num_heads=40,
)
# Parameters: 26.0 GB
# Gradients: 26.0 GB
# Optimizer: 104.0 GB (Adam fp32 moments + master weights)
# Activations: ~30 GB
# Total: ~187 GB → needs 3× A100-80GB with FSDP
13. Upstream Updates (June 17–18, 2026)
Recent changes to the PyTorch codebase relevant to memory profiling, FX infrastructure, and profiler internals:
FX Canonicalize Pass
A new canonicalization pass at torch/fx/passes/canonicalize.py normalizes FX graph node ordering for deterministic graph comparisons. This helps when comparing graph structures before and after memory optimization passes, ensuring that semantically equivalent graphs produce identical canonical forms regardless of insertion order.
ShapesSpec Binding Helpers
New utilities in torch/fx/experimental/_spec_binding.py provide structured binding for shape specifications in FX graphs. These helpers enable cleaner expression of dynamic shape constraints, which is relevant for memory estimation — knowing tensor shapes at trace time allows accurate activation memory predictions.
Profiler pattern_matcher Removed
The profiler's internal pattern_matcher module was refactored out. Pattern-based profiling analysis (identifying common patterns like conv-bn-relu fusion opportunities) has been reorganized into more targeted analysis passes, reducing profiler overhead for memory-focused profiling sessions.
Profiler RecordFunction Drain Fix (#187483)
A fix for RecordFunction callback drain ordering in the profiler. Previously, when profiling memory-intensive workloads, callback cleanup could trigger additional allocations during drain, inflating peak memory measurements. The fix ensures callbacks are drained in the correct order, giving more accurate memory profiles.
_scaled_mm_v2 Swizzled Scales Test (#186948)
New tests for _scaled_mm_v2 with swizzled quantization scales. This is relevant to FP8 training memory optimization — swizzled scales enable more efficient memory layout for quantized matrix multiplications, reducing both memory footprint and fragmentation from scale tensors.
Dynamo Canonicalize output_graph Node Order (#181775)
Dynamo now canonicalizes the node order in output_graph, ensuring deterministic compilation output. For memory profiling, this means memory allocation patterns are reproducible across runs, making it easier to isolate the impact of individual optimizations.
Dynamic Spec Error Messages (#187143)
Improved error messages for dynamic shape specification mismatches. When memory estimation relies on symbolic shapes (e.g., via torch.export with dynamic shapes), clearer error messages help diagnose why estimated memory budgets differ from actual usage due to unexpected shape specialization.
Best Practices Checklist
Before training a large model, work through this checklist:
- Estimate first: Use meta device to compute parameter + optimizer memory. Don't guess.
- Pick your precision: bf16 is free performance and memory. Use it unless you have a reason not to.
- Enable Flash Attention: ensure
F.scaled_dot_product_attentionis routing to the flash kernel. - Consider gradient checkpointing: if activations dominate, checkpoint every N layers.
- Right-size your batch: use gradient accumulation to decouple batch size from memory.
- Profile before optimizing: use
memory_summary()or snapshots to find the actual bottleneck. - Monitor during training: log
memory_allocated()per step to catch leaks early. - Choose your optimizer wisely: Adam costs 2× parameters in state. SGD or 8-bit Adam costs much less.
- Test at scale gradually: increase batch size / seq_len incrementally, checking peak memory each time.
- Use FSDP for multi-GPU: shards parameters, gradients, and optimizer state across GPUs.
Further Resources
- PyTorch Memory Management — official CUDA memory docs
- Memory Snapshot Visualizer — interactive memory snapshot viewer
- Training a 1T Model — scaling techniques for massive models
- Flash Attention Paper — the algorithm behind O(N) attention memory
- Module 07 — Training Pipelines — mixed precision and gradient accumulation
- Module 16 — Activation Checkpointing — SAC details
Notebook: 26_memory_profiling.ipynb
Source Files
[README.md](README.md)— This guide — GPU memory anatomy, profiling tools, optimization techniques[memory_tools.py](memory_tools.py)— Meta-device estimation, memory monitoring, "will it fit?" calculator[memory_optimization.py](memory_optimization.py)— Gradient checkpointing, mixed precision, accumulation, in-place ops
Multi-GPU Inference Patterns — Serving Large Models at Scale
Table of Contents
- Why Multi-GPU Inference?
- Strategy 1: Tensor Parallel Inference
- Strategy 2: Pipeline Parallel Inference
- Strategy 3: Simple Model Sharding with device_map
- KV Cache Across GPUs
- Continuous Batching
- Quantized Inference
- torch.compile for Inference
- AOTInductor for Production
- Benchmarking Inference
- Decision Tree: Choosing a Strategy
- Upstream Updates (June 18–19, 2026)
1. Why Multi-GPU Inference?
Modern large language models simply do not fit in the memory of a single GPU. A 70B parameter model in FP16 requires 140 GB of memory just for the weights — exceeding even the 80 GB available on an A100 or H100:
Model Size (FP16 weights only):
7B parameters → 14 GB (fits 1× A100-80GB)
13B parameters → 26 GB (fits 1× A100-80GB)
34B parameters → 68 GB (fits 1× A100-80GB, tight)
70B parameters → 140 GB (needs 2× A100-80GB minimum)
175B parameters → 350 GB (needs 5× A100-80GB)
405B parameters → 810 GB (needs 11× A100-80GB)
Even models that fit on one GPU can benefit from multi-GPU inference for two reasons:
1.1 Latency Reduction
Tensor parallelism splits each matrix multiplication across GPUs, reducing the per-operation compute time. For latency-sensitive serving (chatbots, real-time APIs), splitting a 7B model across 2 GPUs can halve per-token generation time.
1.2 Throughput Scaling
Pipeline parallelism and continuous batching allow you to process more requests simultaneously. While GPU 0 generates tokens for request A, GPU 1 processes request B's prefill. Throughput scales roughly linearly with GPU count.
1.3 Latency vs Throughput Tradeoff
┌─────────────────────────────────┐
│ Inference Goals │
├────────────────┬────────────────┤
│ Low Latency │ High Throughput │
├────────────────┼────────────────┤
│ Tensor Parallel│ Pipeline Parallel│
│ CUDA Graphs │ Continuous Batch │
│ reduce-overhead│ Larger batches │
│ Single request │ Many concurrent │
│ focus │ requests │
└────────────────┴────────────────┘
For most production LLM serving, you want both: TP within a node for latency, PP across nodes for throughput.
2. Strategy 1: Tensor Parallel Inference
Tensor Parallelism (TP) splits individual weight matrices across GPUs. Every GPU participates in every layer, each holding a shard of every weight matrix. Communication (all-reduce) happens once per layer.
2.1 How TP Works for Transformers
In a Transformer layer, the key operations are linear projections. TP splits these column-wise or row-wise:
Column-wise (split output dim): Row-wise (split input dim):
┌──────────┐ ┌──────────┐
│ W full │ │ W full │
│ [d, 4d] │ │ [4d, d] │
└──────────┘ └──────────┘
↓ split columns ↓ split rows
┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐
│W_0 │ │W_1 │ ← each on │W_0 │ │W_1 │
│[d,2d]│ │[d,2d]│ one GPU │[2d,d]│ │[2d,d]│
└─────┘ └─────┘ └─────┘ └─────┘
ColwiseParallel: each GPU computes a subset of the output features. No communication during the matmul. Used for QKV projections and FFN up-projections.
RowwiseParallel: each GPU holds a subset of input features. Requires an all-reduce after the matmul to sum partial results. Used for output projections and FFN down-projections.
2.2 TP with torch.distributed.tensor.parallel
import torch
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import (
ColwiseParallel,
RowwiseParallel,
parallelize_module,
)
mesh = init_device_mesh("cuda", (world_size,), mesh_dim_names=("tp",))
# Define parallelization plan for a Transformer block
plan = {
# QKV: split output columns across GPUs
"attn.qkv_proj": ColwiseParallel(),
# Output projection: split input rows, all-reduce output
"attn.out_proj": RowwiseParallel(),
# FFN up: split output columns
"ffn.up_proj": ColwiseParallel(),
"ffn.gate_proj": ColwiseParallel(),
# FFN down: split input rows, all-reduce output
"ffn.down_proj": RowwiseParallel(),
}
for layer in model.layers:
parallelize_module(layer, mesh["tp"], plan)
2.3 Communication Cost
Each Transformer layer with TP requires 2 all-reduce operations (one after attention output projection, one after FFN down projection). For L layers:
Total all-reduce calls = 2 × L
Per all-reduce data = batch × seq_len × hidden_dim × dtype_bytes
On NVLink (900 GB/s bidirectional on H100), this overhead is small for large hidden dims. On PCIe (64 GB/s), TP across more than 2 GPUs becomes communication-bound.
2.4 TP Best Practices
- Use TP within a single node (NVLink connectivity)
- TP degree should divide
num_headsandnum_kv_headsevenly - TP=2 for 7-13B models, TP=4 for 34-70B, TP=8 for 70B+ on a single node
- Combine with PP for models spanning multiple nodes
3. Strategy 2: Pipeline Parallel Inference
Pipeline Parallelism (PP) assigns different layers to different GPUs. GPU 0 holds layers 0-15, GPU 1 holds layers 16-31. Data flows sequentially through the pipeline.
3.1 How PP Works
Request → [GPU 0: Layers 0-15] → [GPU 1: Layers 16-31] → Output
embed, norm layers 16-31, head
PP=4 example (70B, 80 layers):
GPU 0: embed + layers 0-19 (20 layers)
GPU 1: layers 20-39 (20 layers)
GPU 2: layers 40-59 (20 layers)
GPU 3: layers 60-79 + head (20 layers)
3.2 Pipeline with Micro-Batching
Without micro-batching, PP has terrible utilization — only one GPU is active at a time. Micro-batching fills the pipeline:
Time →
GPU 0: [batch0] [batch1] [batch2] [batch3] idle idle idle idle
GPU 1: idle [batch0] [batch1] [batch2] [batch3] idle idle idle
GPU 2: idle idle [batch0] [batch1] [batch2] [batch3] idle idle
GPU 3: idle idle idle [batch0] [batch1] [batch2] [batch3] idle
Pipeline bubble = (PP_degree - 1) / (PP_degree + num_microbatches - 1). With 4 GPUs and 8 micro-batches, the bubble is 3/11 ≈ 27%. More micro-batches shrink the bubble.
3.3 PP Communication
PP communicates only at pipeline stage boundaries — the hidden states between consecutive layer groups. This is point-to-point (send/recv), not all-reduce:
Communication per stage boundary:
data = batch × seq_len × hidden_dim × dtype_bytes
For 70B (hidden=8192, bf16), batch=1, seq=4096:
= 1 × 4096 × 8192 × 2 = 64 MB per boundary
Much less communication than TP, making PP suitable for cross-node distribution.
3.4 PP vs TP Tradeoffs
| Aspect | Tensor Parallel | Pipeline Parallel |
|---|---|---|
| Communication | All-reduce per layer | Point-to-point between stages |
| Latency | Lower (all GPUs active per token) | Higher (pipeline bubble) |
| Throughput | Limited by TP comm | Scales with micro-batches |
| Best interconnect | NVLink (intra-node) | Works on PCIe/InfiniBand |
| Memory balance | Even (same layers on all GPUs) | Can be uneven (first/last stage) |
4. Strategy 3: Simple Model Sharding with device_map
The simplest multi-GPU approach: manually assign model components to different devices. No distributed communication library needed.
4.1 Manual device_map
model = MyLLM(config)
# Assign layers to GPUs
model.embed.to('cuda:0')
model.layers[:16].to('cuda:0')
model.layers[16:].to('cuda:1')
model.head.to('cuda:1')
4.2 Forward Pass with Cross-Device Transfer
The forward pass must move activations between devices at the boundary:
def forward(self, input_ids):
# Phase 1: on cuda:0
x = self.embed(input_ids.to('cuda:0'))
for layer in self.layers[:16]:
x = layer(x)
# Transfer activations to cuda:1
x = x.to('cuda:1')
# Phase 2: on cuda:1
for layer in self.layers[16:]:
x = layer(x)
x = self.norm(x)
logits = self.head(x)
return logits
4.3 Limitations
- No parallelism within a layer: only one GPU is active at any time during a single request
- Sequential execution: GPU 0 sits idle while GPU 1 runs its layers
- Good for: fitting large models for batch inference where latency is not critical
- Bad for: real-time serving (doubles latency compared to TP)
4.4 Balancing Memory Across Devices
Not all layers are the same size. The embedding and output head can be large (vocab_size × hidden_dim). Balance by assigning more transformer layers to GPUs without the embedding/head:
def compute_device_map(model, num_gpus):
"""Assign layers to GPUs, balancing parameter memory."""
param_sizes = {}
for name, param in model.named_parameters():
device_key = name.split('.')[0]
param_sizes[device_key] = param_sizes.get(device_key, 0) + param.numel() * param.element_size()
total = sum(param_sizes.values())
per_gpu = total / num_gpus
device_map = {}
current_gpu, current_load = 0, 0
for name, size in param_sizes.items():
device_map[name] = f'cuda:{current_gpu}'
current_load += size
if current_load >= per_gpu and current_gpu < num_gpus - 1:
current_gpu += 1
current_load = 0
return device_map
5. KV Cache Across GPUs
During autoregressive generation, the KV cache stores past key and value tensors to avoid recomputation. Multi-GPU inference shards this cache differently depending on the parallelism strategy.
5.1 KV Cache in Tensor Parallel
With TP, each GPU holds a shard of the attention heads. The KV cache is naturally sharded — each GPU caches only its head shard:
TP=2, 32 heads total:
GPU 0: caches heads 0-15 → cache_size / 2
GPU 1: caches heads 16-31 → cache_size / 2
With GQA (8 KV heads):
GPU 0: caches KV heads 0-3 → cache_size / 2
GPU 1: caches KV heads 4-7 → cache_size / 2
5.2 KV Cache in Pipeline Parallel
With PP, each stage caches only its own layers' KV pairs. The cache is split by layers, not by heads:
PP=2, 32 layers total:
GPU 0 (layers 0-15): caches 16 layers × full heads
GPU 1 (layers 16-31): caches 16 layers × full heads
5.3 KV Cache Memory Estimation
Per-token KV cache memory:
= 2 × num_layers × num_kv_heads × head_dim × dtype_bytes
For Llama-70B (80 layers, 8 KV heads, head_dim=128, bf16):
= 2 × 80 × 8 × 128 × 2 = 327,680 bytes ≈ 320 KB per token
For 4096 tokens:
= 320 KB × 4096 = 1.28 GB per sequence
For batch_size=32 concurrent requests:
= 1.28 GB × 32 = 41 GB for KV cache alone
5.4 Pre-Allocation Strategies
Pre-allocating KV cache avoids memory fragmentation during serving:
def preallocate_kv_cache(num_layers, num_kv_heads, head_dim, max_seq_len,
max_batch, dtype=torch.bfloat16, device='cuda'):
"""Pre-allocate KV cache buffers to avoid fragmentation."""
cache = []
for _ in range(num_layers):
k = torch.zeros(max_batch, num_kv_heads, max_seq_len, head_dim,
dtype=dtype, device=device)
v = torch.zeros(max_batch, num_kv_heads, max_seq_len, head_dim,
dtype=dtype, device=device)
cache.append((k, v))
return cache
Pre-allocation trades unused memory for deterministic allocation patterns, which is critical for CUDA Graphs and production serving.
6. Continuous Batching
Traditional batching waits for all sequences in a batch to finish before starting new ones. Continuous batching (also called in-flight batching) allows new requests to join the batch as soon as any request finishes.
6.1 The Problem with Static Batching
Static batching (batch=4):
Request A: ████████████████████████████████ (128 tokens)
Request B: ████████ (32 tokens) ← idle after 32 tokens
Request C: ████████████████ (64 tokens) ← idle after 64 tokens
Request D: ████████████ (48 tokens) ← idle after 48 tokens
GPU utilization: only ~68% — short requests waste GPU cycles
6.2 Continuous Batching
Continuous batching:
Request A: ████████████████████████████████
Request B: ████████ E: ████████████████████████
Request C: ████████████████ F: ████████████████
Request D: ████████████ G: ████████████████████
GPU utilization: ~95% — new requests fill slots immediately
6.3 Implementation Concepts
The key data structures for continuous batching:
class ContinuousBatchScheduler:
"""Simplified continuous batching scheduler."""
def __init__(self, max_batch_size, max_seq_len):
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.active_requests = {} # request_id -> RequestState
self.waiting_queue = [] # requests awaiting scheduling
def step(self):
"""Run one generation step for all active requests."""
finished = []
for req_id, state in self.active_requests.items():
if state.is_finished():
finished.append(req_id)
for req_id in finished:
del self.active_requests[req_id]
while (len(self.active_requests) < self.max_batch_size
and self.waiting_queue):
new_req = self.waiting_queue.pop(0)
self.active_requests[new_req.id] = new_req
This is the core idea behind serving engines like vLLM, TensorRT-LLM, and SGLang. Production implementations add PagedAttention for efficient KV cache management.
7. Quantized Inference
Quantization reduces model precision to fit larger models on fewer GPUs and improve throughput.
7.1 Precision Comparison
Precision Bits/Param 70B Model Size Quality Impact
─────────────────────────────────────────────────────────
FP32 32 bits 280 GB Baseline
FP16/BF16 16 bits 140 GB Negligible
INT8 8 bits 70 GB Minimal (<1% degradation)
INT4 4 bits 35 GB Small (1-3% degradation)
NF4 4 bits 35 GB Very small with double quant
7.2 Dynamic Quantization (INT8)
import torch.ao.quantization as quant
model = MyLLM(config)
model.eval()
# Dynamic quantization: weights quantized statically, activations quantized dynamically
quantized_model = torch.ao.quantization.quantize_dynamic(
model,
{torch.nn.Linear}, # quantize Linear layers
dtype=torch.qint8,
)
Dynamic quantization is the simplest approach — no calibration data needed. Weights are quantized to INT8 at load time, activations are quantized on-the-fly during inference.
7.3 Weight-Only Quantization (INT4/INT8)
Weight-only quantization keeps activations in FP16/BF16 but stores weights in lower precision:
from torchao.quantization import quantize_, int4_weight_only, int8_weight_only
model = MyLLM(config).to(dtype=torch.bfloat16)
# INT4 weight-only quantization
quantize_(model, int4_weight_only(group_size=128))
# Or INT8 weight-only
quantize_(model, int8_weight_only())
7.4 Combining Quantization with TP
Quantize first, then apply tensor parallelism:
# 1. Load and quantize
model = load_model(config)
quantize_(model, int4_weight_only(group_size=128))
# 2. Apply TP
mesh = init_device_mesh("cuda", (tp_size,))
for layer in model.layers:
parallelize_module(layer, mesh, tp_plan)
# 3. Compile for peak performance
model = torch.compile(model, mode="max-autotune")
This combination is how production systems serve 70B models on 2 GPUs: INT4 reduces the 140 GB to 35 GB (fits on 2 × 24 GB GPUs), and TP splits computation for lower latency.
8. torch.compile for Inference
torch.compile applies kernel fusion and optimization for significant inference speedups.
8.1 Compile Modes for Inference
# Maximum throughput — longer compile time, best steady-state performance
model = torch.compile(model, mode="max-autotune")
# Lowest latency — uses CUDA Graphs to eliminate kernel launch overhead
model = torch.compile(model, mode="reduce-overhead")
# Default balance
model = torch.compile(model)
8.2 mode="reduce-overhead" and CUDA Graphs
reduce-overhead mode wraps the compiled model in CUDA Graphs, eliminating CPU kernel launch overhead:
model = model.eval().cuda()
model = torch.compile(model, mode="reduce-overhead")
# Warmup: first few calls trigger compilation + graph capture
with torch.no_grad():
for _ in range(3):
_ = model(warmup_input)
# Steady state: near-zero CPU overhead per forward pass
with torch.no_grad():
output = model(real_input) # runs via CUDA Graph replay
8.3 Static Shapes for Best Performance
CUDA Graphs require static input shapes. For inference, pad inputs to fixed lengths:
def pad_to_static(input_ids, pad_id=0, max_len=2048):
"""Pad inputs to static shape for CUDA Graph compatibility."""
batch, seq = input_ids.shape
if seq < max_len:
padding = torch.full((batch, max_len - seq), pad_id,
dtype=input_ids.dtype, device=input_ids.device)
input_ids = torch.cat([input_ids, padding], dim=1)
return input_ids
8.4 Combining torch.compile with TP
# Apply TP first, then compile each rank's model
for layer in model.layers:
parallelize_module(layer, mesh, tp_plan)
model = torch.compile(model, mode="max-autotune")
Each TP rank compiles independently. The compiled graph includes the communication operations (all-reduce), so they are fused into the overall execution plan.
9. AOTInductor for Production
AOTInductor (Ahead-of-Time Inductor) pre-compiles PyTorch models into shared libraries (.so files) that can be loaded and run from C++ without any Python dependency.
9.1 Export and Compile
import torch
from torch._export import aot_compile
model = MyLLM(config).eval().cuda()
example_input = torch.randint(0, 32000, (1, 2048), device='cuda')
# Export to a .so file
so_path = aot_compile(
model,
args=(example_input,),
options={"max_autotune": True},
)
print(f"Compiled to: {so_path}")
9.2 Load in C++ (No Python)
#include <torch/csrc/inductor/aoti_runner/model_container_runner_cuda.h>
int main() {
auto runner = std::make_unique<torch::inductor::AOTIModelContainerRunnerCuda>(
"model.so"
);
auto input = torch::randint(0, 32000, {1, 2048},
torch::dtype(torch::kLong).device(torch::kCUDA));
auto outputs = runner->run({input});
auto logits = outputs[0];
return 0;
}
9.3 PT2 Archive Format
The newer packaging approach bundles model + weights + metadata:
import torch
from torch._inductor.package import package_aoti
model = MyLLM(config).eval()
example = (torch.randint(0, 32000, (1, 2048)),)
ep = torch.export.export(model, example)
# Package to .pt2 archive
package_aoti("model.pt2", ep)
# Load and run (Python)
runner = torch._inductor.package.load_package("model.pt2")
output = runner(input_ids)
9.4 Benefits for Production
| Aspect | Python (torch.compile) | AOTInductor (.so) |
|---|---|---|
| Startup time | Compile on first call | Pre-compiled, instant |
| Python GIL | Yes, limits concurrency | No Python needed |
| Deployment | Needs Python + PyTorch | Just the .so + libtorch |
| Debugging | Full Python stack traces | C++ debugging |
| Use case | Development, prototyping | Production serving |
10. Benchmarking Inference
10.1 Key Metrics
| Metric | Definition | Target |
|---|---|---|
| TTFT (Time to First Token) | Time from request to first generated token | < 500ms |
| ITL (Inter-Token Latency) | Time between consecutive generated tokens | < 30ms |
| Throughput | Total tokens generated per second across all requests | Maximize |
| p50 / p95 / p99 latency | Percentile latency across requests | p99 < 2× p50 |
10.2 TTFT vs ITL
TTFT includes the prefill phase (processing the entire prompt). ITL measures the decode phase (generating one token at a time):
Request lifecycle:
[Prompt arrives] → [Prefill: process all prompt tokens] → [First token]
TTFT ──────────────────────────┘
[First token] → [Second token] → [Third token] → ... → [EOS]
├─── ITL ────┤├─── ITL ────┤
Prefill is compute-bound (one large matmul). Decode is memory-bandwidth-bound (one token at a time, must read all weights).
10.3 Measuring with torch.utils.benchmark
import torch.utils.benchmark as benchmark
model = model.eval().cuda()
input_ids = torch.randint(0, 32000, (1, 512), device='cuda')
# Prefill latency
timer_prefill = benchmark.Timer(
stmt='model(input_ids)',
globals={'model': model, 'input_ids': input_ids},
num_threads=1,
)
result_prefill = timer_prefill.blocked_autorange(min_run_time=5.0)
print(f"Prefill (512 tokens): {result_prefill.median * 1000:.1f} ms")
# Decode latency (single token)
single_token = torch.randint(0, 32000, (1, 1), device='cuda')
timer_decode = benchmark.Timer(
stmt='model(single_token)',
globals={'model': model, 'single_token': single_token},
num_threads=1,
)
result_decode = timer_decode.blocked_autorange(min_run_time=5.0)
print(f"Decode (1 token): {result_decode.median * 1000:.2f} ms")
print(f"Max tokens/sec: {1.0 / result_decode.median:.0f}")
10.4 Throughput Benchmarking
import time
def benchmark_throughput(model, prompts, max_new_tokens=128):
"""Measure end-to-end throughput in tokens/second."""
total_tokens = 0
start = time.perf_counter()
for prompt in prompts:
output = generate(model, prompt, max_new_tokens=max_new_tokens)
total_tokens += output.shape[-1]
elapsed = time.perf_counter() - start
throughput = total_tokens / elapsed
print(f"Throughput: {throughput:.0f} tokens/sec")
print(f"Total time: {elapsed:.2f}s for {total_tokens} tokens")
return throughput
11. Decision Tree: Choosing a Strategy
Model fits on 1 GPU?
├─ YES → Use torch.compile
│ ├─ Latency-sensitive? → mode="reduce-overhead" (CUDA Graphs)
│ └─ Throughput? → mode="max-autotune" + larger batches
│
└─ NO → How many GPUs needed?
│
├─ 2-4 GPUs (within 1 node, NVLink)
│ └─ Tensor Parallel
│ Best single-request latency
│ Each GPU holds a shard of every layer
│
├─ 4-8 GPUs (1-2 nodes)
│ └─ TP + PP hybrid
│ TP within node, PP across nodes
│ Balance latency and throughput
│
└─ Need maximum throughput?
└─ Pipeline Parallel + continuous batching
Fill pipeline with micro-batches
New requests join as old ones finish
Additional considerations:
├─ Production deployment? → AOTInductor (.so, no Python)
├─ Memory constrained? → INT4 quantization first, then TP
└─ Simple setup needed? → device_map sharding (no dist init)
Quick Reference
| Scenario | Recommended Strategy | Why |
|---|---|---|
| 7B model, 1 GPU | torch.compile(mode="reduce-overhead") | Fits easily, maximize latency |
| 7B model, high QPS | Continuous batching on 1 GPU | Maximize throughput |
| 70B model, 2 GPUs | TP=2 + INT4 quantization | Fits with INT4, TP for latency |
| 70B model, 4 GPUs | TP=4 (FP16) or TP=2 (INT8) | Full precision or quant + TP |
| 70B model, 8 GPUs | TP=4 + PP=2 | TP within node, PP across |
| 405B model, 16 GPUs | TP=8 + PP=2 + INT8 | Multi-node, hybrid parallelism |
12. Upstream Updates (June 18–19, 2026)
Recent changes to the PyTorch codebase relevant to multi-GPU inference, distributed systems, and compiler infrastructure:
CUPTI Profiler Refactored
The CUPTI profiler has been refactored into a dedicated torch/profiler/_cupti/ package, separating CUPTI-specific logic from the general profiler infrastructure. This improves maintainability and makes it easier to profile multi-GPU inference workloads where per-GPU profiling data needs independent collection and aggregation.
DTensor logspace Support (#186398)
torch.logspace now supports DTensor, enabling logarithmically-spaced tensor creation across distributed meshes. Useful for creating learning rate schedules or quantization scales directly on sharded tensors without manual gather/scatter.
ShapesSpec/ParamsSpec in Non-Strict Export (#187602)
New ShapesSpec and ParamsSpec support in non-strict export mode allows more flexible shape specifications when exporting models for AOTInductor deployment. This is particularly relevant for inference models with dynamic batch sizes or sequence lengths that need to be compiled ahead of time.
Distributed Backend Accessors Exposed (#187494)
Backend accessors for distributed communication are now publicly exposed, making it easier to query and configure the communication backend (NCCL, Gloo) programmatically. Useful for inference servers that need to dynamically select backends based on available hardware.
set_timeout on FakeProcessGroup (#187693)
FakeProcessGroup now supports set_timeout, enabling better testing of multi-GPU inference code without actual GPUs. Test timeouts for distributed operations can be configured independently, catching hangs in TP/PP initialization logic during unit tests.
L2-Aware Two-Pass Variance Heuristic (#183661)
A new L2-cache-aware heuristic for two-pass variance computation improves kernel selection in Inductor. For inference workloads that compute layer normalization or RMS normalization across TP shards, this heuristic selects kernels that better utilize GPU L2 cache, reducing memory bandwidth pressure.
Best Practices Checklist
Before deploying a multi-GPU inference system:
- Estimate model memory: use meta device to compute weight + KV cache memory before allocating GPUs.
- Choose quantization first: INT4/INT8 can reduce GPU count by 2-4×. Always quantize before considering more GPUs.
- Match parallelism to hardware: TP within NVLink nodes, PP across nodes. Never TP across PCIe.
- Pre-allocate KV cache: avoids fragmentation and enables CUDA Graphs.
- Use torch.compile:
reduce-overheadfor latency,max-autotunefor throughput. - Benchmark all three metrics: TTFT, ITL, and throughput. Optimizing one can hurt another.
- Consider AOTInductor for production: eliminates Python overhead and GIL contention.
- Implement continuous batching: static batching wastes 30-50% of GPU cycles.
- Monitor GPU utilization: all GPUs should be >80% utilized in steady state.
- Test with realistic workloads: generation length distributions affect throughput significantly.
Further Resources
- PyTorch Tensor Parallel docs — official TP API reference
- PyTorch Pipeline Parallel docs — official PP API reference
- vLLM — production LLM serving with PagedAttention
- Module 10 — Distributed Training — DDP, FSDP2, DeviceMesh, TP, PP
- Module 11 — Export & Deployment — torch.export, AOTInductor, NativeRT
- Module 22 — LLM Recipes — RoPE, KV Cache, GQA, SwiGLU
Notebook: 27_multi_gpu_inference.ipynb
Source Files
[README.md](README.md)— This guide — multi-GPU inference strategies, quantization, benchmarking[inference_patterns.py](inference_patterns.py)— Model size estimation, device_map sharding, KV cache sizing, decision tree[model_sharding.py](model_sharding.py)— Manual sharding, TP/PP patterns, continuous batching, AOTInductor workflow
torch.utils.benchmark Deep Dive — Measuring Performance Correctly
Table of Contents
- Why Proper Benchmarking Matters
- torch.utils.benchmark.Timer
- blocked_autorange()
- Measurement Object
- Compare — Side-by-Side Tables
- Benchmarking torch.compile
- Benchmarking with Different Shapes
- num_threads — Controlling CPU Parallelism
- Fuzzer — Random Test Configurations
- Callgrind — Instruction Counts
- Common Pitfalls
- Practical Recipes
- Upstream Updates (June 2026)
1. Why Proper Benchmarking Matters
Most PyTorch benchmarks you'll find online are wrong. Here's why:
time.time() Is Wrong for GPU Code
import time
import torch
x = torch.randn(1000, 1000, device='cuda')
start = time.time()
y = x @ x # launches kernel but DOESN'T wait for it
elapsed = time.time() - start # measures launch time (~10μs), NOT compute time
CUDA operations are asynchronous — torch.matmul returns immediately after launching the GPU kernel. The CPU continues while the GPU computes. time.time() only captures how long the CPU took to enqueue the work.
timeit Doesn't Sync CUDA Either
import timeit
# Still wrong — timeit uses time.perf_counter() internally, no CUDA sync
timeit.timeit(lambda: x @ x, number=100)
What Distorts Benchmark Results
| Factor | Effect | Fix |
|---|---|---|
| No CUDA sync | Measures launch time, not compute time | torch.cuda.synchronize() |
| Cold start | First call initializes CUDA context (~1-3s) | Warmup iterations |
| JIT compilation | torch.compile first call is slow | Separate warmup phase |
| cuDNN autotuning | First convolution triggers autotuner | torch.backends.cudnn.benchmark = True before warmup |
| Garbage collection | GC pauses inject random latency spikes | Disable GC during measurement |
| CPU frequency scaling | Dynamic clocks cause variance | Pin CPU frequency or use instruction counts |
| Memory caching | CUDA caching allocator reuses memory | Consistent allocation patterns |
torch.utils.benchmark handles all of this automatically.
2. torch.utils.benchmark.Timer
The core API for all PyTorch benchmarking:
from torch.utils.benchmark import Timer
t = Timer(
stmt="x @ y",
setup="x = torch.randn(1000, 1000); y = torch.randn(1000, 1000)",
)
print(t.timeit(100)) # fixed number of runs
print(t.blocked_autorange()) # auto-determine run count
Constructor Parameters
| Parameter | Type | Description |
|---|---|---|
stmt | str | The code to benchmark (can be multi-line) |
setup | str | Code run once before measurement (imports, tensor creation) |
globals | dict | Variables accessible in stmt and setup |
num_threads | int | CPU threads to use (controls torch.set_num_threads) |
label | str | Row label for Compare tables |
sub_label | str | Sub-row label for Compare tables |
description | str | Column label for Compare tables |
env | str | Environment name (for cross-environment comparison) |
timer | callable | Custom timer function (default: timeit.default_timer) |
Using globals vs setup
# Option 1: setup string (self-contained, but limited)
t = Timer(
stmt="x.mm(y)",
setup="import torch; x = torch.randn(256, 256); y = torch.randn(256, 256)"
)
# Option 2: globals dict (more flexible — use existing Python objects)
x = torch.randn(256, 256)
y = torch.randn(256, 256)
t = Timer(
stmt="x.mm(y)",
globals={"x": x, "y": y}
)
Use globals when your setup is complex or involves objects that can't be easily expressed as a string.
Multi-Statement Benchmarks
t = Timer(
stmt="""
y = model(x)
loss = criterion(y, target)
loss.backward()
""",
globals={"model": model, "criterion": criterion, "x": x, "target": target}
)
timeit() — Fixed Iteration Count
result = t.timeit(number=100) # run stmt exactly 100 times
Returns a Measurement object. Good when you know how many iterations you want, but bad for comparing fast vs slow operations (the fast one may need more iterations for stable results).
3. blocked_autorange()
The recommended way to benchmark. It automatically determines the right number of iterations:
result = t.blocked_autorange(min_run_time=1.0)
How It Works
- Adaptive warmup: Runs increasing numbers of iterations (1, 2, 4, 8, ...) until a single block takes ≥
min_run_timeseconds - Measurement: Runs multiple blocks at the determined iteration count
- Aggregation: Reports median of block times (not mean)
Why Median Over Mean
Run times: [1.2ms, 1.1ms, 1.3ms, 1.1ms, 15.2ms, 1.2ms]
Mean: 3.5ms ← distorted by one GC pause
Median: 1.2ms ← robust to outliers
The mean is distorted by occasional outliers (GC pauses, context switches, thermal throttling). The median gives a more reliable estimate of typical performance.
Parameters
| Parameter | Default | Description |
|---|---|---|
min_run_time | 2.0 | Minimum total wall time in seconds |
callback | None | Called after each block with intermediate results |
4. Measurement Object
Both timeit() and blocked_autorange() return a Measurement object:
result = t.blocked_autorange()
Key Attributes
| Attribute | Type | Description |
|---|---|---|
result.mean | float | Mean time per execution (seconds) |
result.median | float | Median time per execution (seconds) |
result.times | List[float] | All measured times (per execution) |
result.number_per_run | int | Number of stmt executions per block |
result.raw_times | List[float] | Raw block times (total, not per execution) |
result.iqr | float | Interquartile range |
result.significant_figures | int | Stable digits across measurements |
String Representation
print(result)
# Output:
# <torch.utils.benchmark.utils.common.Measurement object at 0x...>
# x @ y
# 1.23 ms
# 1 measurement, 100 runs, 1 thread
The string representation automatically scales the units (ns, μs, ms, s) and includes the IQR to indicate measurement stability.
Comparing Measurements
# Access raw timing data
for t in result.times:
print(f" {t * 1e3:.3f} ms")
# Mean vs median
print(f"Mean: {result.mean * 1e3:.3f} ms")
print(f"Median: {result.median * 1e3:.3f} ms")
print(f"IQR: {result.iqr * 1e3:.3f} ms")
5. Compare — Side-by-Side Tables
The Compare class renders a formatted table comparing multiple benchmarks:
from torch.utils.benchmark import Timer, Compare
results = []
for n in [64, 256, 1024]:
for impl in ["mm", "matmul", "einsum"]:
if impl == "mm":
stmt = "torch.mm(x, y)"
elif impl == "matmul":
stmt = "x @ y"
else:
stmt = "torch.einsum('ij,jk->ik', x, y)"
t = Timer(
stmt=stmt,
setup=f"import torch; x = torch.randn({n},{n}); y = torch.randn({n},{n})",
label="matmul",
sub_label=f"[{n}x{n}]",
description=impl,
)
results.append(t.blocked_autorange(min_run_time=0.5))
compare = Compare(results)
compare.print()
Output Format
[----------- matmul -----------]
| mm | matmul | einsum
1 threads: ----+---------+--------+-------
[64x64] | 5.2 | 5.3 | 12.1
[256x256] | 42.0 | 42.1 | 55.3
[1024x1024] | 850.1 | 851.0 | 870.2
Times are in microseconds (us).
Label Hierarchy
The three label fields control table layout:
label: Groups rows into sections (the[--- label ---]header)sub_label: Individual rows within a sectiondescription: Column headers
Colorized Output
compare.colorize() adds terminal colors highlighting the fastest implementation per row:
compare = Compare(results)
compare.colorize() # enables green/red coloring
compare.print() # fastest in green, slowest in red
Trimming Significant Figures
compare.trim_significant_figures() # reduces digits to match precision
compare.print()
6. Benchmarking torch.compile
torch.compile has a critical subtlety: the first call triggers compilation (which can take seconds). You must separate warmup from measurement:
Correct Methodology
import torch
from torch.utils.benchmark import Timer
model = torch.nn.Linear(1024, 1024)
x = torch.randn(32, 1024)
# Eager baseline
eager_timer = Timer(
stmt="model(x)",
globals={"model": model, "x": x},
label="Linear(1024, 1024)",
sub_label="batch=32",
description="eager",
)
# Compiled model — warmup OUTSIDE the Timer
compiled_model = torch.compile(model)
for _ in range(3): # warmup: triggers compilation
compiled_model(x)
compiled_timer = Timer(
stmt="compiled_model(x)",
globals={"compiled_model": compiled_model, "x": x},
label="Linear(1024, 1024)",
sub_label="batch=32",
description="compiled",
)
results = [
eager_timer.blocked_autorange(),
compiled_timer.blocked_autorange(),
]
from torch.utils.benchmark import Compare
Compare(results).print()
Comparing Compile Modes
results = []
for mode in [None, "reduce-overhead", "max-autotune"]:
compiled = torch.compile(model, mode=mode)
for _ in range(3):
compiled(x) # warmup
desc = mode or "default"
t = Timer(
stmt="fn(x)",
globals={"fn": compiled, "x": x},
label="compile modes",
sub_label="Linear(1024,1024)",
description=desc,
)
results.append(t.blocked_autorange())
Common Mistake: Including Compilation Time
# WRONG: compilation happens inside Timer, polluting measurements
t = Timer(
stmt="torch.compile(model)(x)",
globals={"model": model, "x": x},
)
7. Benchmarking with Different Shapes
A common task: sweep over input sizes to understand scaling behavior.
import torch
from torch.utils.benchmark import Timer, Compare
results = []
sizes = [128, 256, 512, 1024, 2048, 4096]
for n in sizes:
x = torch.randn(n, n)
y = torch.randn(n, n)
for desc, stmt in [("mm", "x @ y"), ("bmm", "x.unsqueeze(0) @ y.unsqueeze(0)")]:
t = Timer(
stmt=stmt,
globals={"x": x, "y": y},
label="Matrix multiply",
sub_label=f"[{n}x{n}]",
description=desc,
)
results.append(t.blocked_autorange(min_run_time=0.5))
Compare(results).print()
Analyzing Scaling
# Check if runtime scales as expected (O(n^3) for matmul)
for i in range(1, len(sizes)):
ratio = results[i*2].median / results[(i-1)*2].median
size_ratio = (sizes[i] / sizes[i-1]) ** 3
print(f"{sizes[i]:4d} vs {sizes[i-1]:4d}: "
f"time ratio={ratio:.1f}x, theoretical={size_ratio:.1f}x")
8. num_threads — Controlling CPU Parallelism
CPU benchmarks can vary wildly depending on thread count. The num_threads parameter pins the thread count:
from torch.utils.benchmark import Timer
x = torch.randn(1000, 1000)
y = torch.randn(1000, 1000)
results = []
for nthreads in [1, 2, 4, 8]:
t = Timer(
stmt="x @ y",
globals={"x": x, "y": y},
num_threads=nthreads,
label="matmul",
sub_label="[1000x1000]",
description=f"{nthreads} threads",
)
results.append(t.blocked_autorange())
from torch.utils.benchmark import Compare
Compare(results).print()
Why This Matters
- Reproducibility: Without pinning, thread count may vary between machines
- Fair comparison: Two implementations should use the same thread count
- Scaling analysis: See how an operation scales with core count
- Production relevance: Server may limit threads per process
Getting Default Thread Count
default_threads = torch.get_num_threads() # returns current setting
print(f"Default: {default_threads} threads")
9. Fuzzer — Random Test Configurations
For thorough benchmarking, you want to test across a range of random configurations. torch.utils.benchmark provides a Fuzzer that generates randomized parameters:
from torch.utils.benchmark import Fuzzer, FuzzedParameter, FuzzedTensor
fuzzer = Fuzzer(
parameters=[
FuzzedParameter("n", minval=4, maxval=16, distribution="loguniform"),
FuzzedParameter("m", minval=4, maxval=16, distribution="loguniform"),
],
tensors=[
FuzzedTensor("x", size=("n", "m"), probability_contiguous=0.6),
FuzzedTensor("y", size=("m", "n"), probability_contiguous=0.6),
],
seed=42,
)
results = []
for tensors, tensor_params, params in fuzzer.take(10):
n, m = int(params["n"]), int(params["m"])
t = Timer(
stmt="x @ y",
globals=tensors,
label="matmul",
sub_label=f"[{n}x{m}]",
description="torch.mm",
)
results.append(t.blocked_autorange(min_run_time=0.2))
FuzzedParameter Options
| Parameter | Description |
|---|---|
minval, maxval | Range for generated values |
distribution | "uniform" or "loguniform" |
FuzzedTensor Options
| Parameter | Description |
|---|---|
size | Tuple of parameter names or ints |
probability_contiguous | Probability the tensor is contiguous (0.0-1.0) |
min_elements | Minimum total elements |
max_elements | Maximum total elements |
dtype | Tensor dtype |
Why Use Fuzzer?
- Avoid cherry-picking: Testing only powers of 2 can hit cache-aligned fast paths
- Find edge cases: Non-contiguous tensors, odd sizes, small inputs
- Statistical rigor: Random configurations give a more realistic performance picture
10. Callgrind — Instruction Counts
Wall-clock time has inherent noise (OS scheduling, thermal throttling, other processes). For micro-benchmarks where you need deterministic results, use instruction counting via Valgrind's Callgrind tool:
from torch.utils.benchmark import Timer
t = Timer(
stmt="x @ y",
setup="import torch; x = torch.randn(128, 128); y = torch.randn(128, 128)",
)
# Requires Valgrind to be installed
stats = t.collect_callgrind(number=100)
print(stats)
What It Measures
Instead of wall-clock time, Callgrind counts the number of CPU instructions executed. This is:
- Deterministic: Same input → same count, every time
- Noise-free: No interference from other processes, scheduling, or frequency scaling
- Reproducible: Results are identical across runs
CallgrindStats
stats = t.collect_callgrind(number=100)
# Total instruction count
print(f"Total instructions: {stats.counts()}")
# Filter by function pattern
fn_counts = stats.as_standardized().stats(inclusive=True)
FunctionCounts
# Get per-function instruction counts
fn_counts = stats.as_standardized().stats(inclusive=False)
for fn in fn_counts[:10]:
print(fn)
When to Use Callgrind
| Use Case | Wall Clock | Callgrind |
|---|---|---|
| Comparing two implementations | ✓ | ✓ |
| Micro-benchmarks (< 1μs) | Noisy | ✓ |
| CI regression testing | Noisy | ✓ |
| Production-like workloads | ✓ | Slow |
| GPU benchmarks | ✓ | ✗ |
Limitation: Callgrind only measures CPU instructions. It cannot measure GPU kernel performance.
11. Common Pitfalls
Pitfall 1: Not Warming Up
# BAD — first call initializes CUDA context (1-3 seconds)
t = Timer(stmt="x @ y", globals={"x": x_cuda, "y": y_cuda})
result = t.timeit(1) # includes CUDA init!
# GOOD — blocked_autorange handles warmup automatically
result = t.blocked_autorange()
For torch.compile, warmup is even more critical:
# BAD — compilation time included
compiled = torch.compile(model)
t = Timer(stmt="fn(x)", globals={"fn": compiled, "x": x})
result = t.blocked_autorange() # first block includes compilation!
# GOOD — warmup compiled model before benchmarking
compiled = torch.compile(model)
for _ in range(3):
compiled(x) # trigger compilation
t = Timer(stmt="fn(x)", globals={"fn": compiled, "x": x})
result = t.blocked_autorange()
Pitfall 2: Not Syncing GPU
# BAD — time.time() doesn't wait for GPU
import time
start = time.time()
y = model(x_cuda) # async!
print(f"{time.time() - start:.3f}s") # measures CPU launch time only
# GOOD — torch.utils.benchmark syncs automatically
t = Timer(stmt="model(x)", globals={"model": model, "x": x_cuda})
print(t.blocked_autorange()) # inserts torch.cuda.synchronize()
Pitfall 3: Garbage Collection Interference
# BAD — GC pauses inject random latency
result = t.timeit(100) # GC may fire mid-measurement
# GOOD — Timer disables GC during measurement by default
result = t.blocked_autorange() # GC disabled automatically
Pitfall 4: Too Few Iterations
# BAD — single measurement is noisy
result = t.timeit(1)
# GOOD — enough iterations for statistical stability
result = t.blocked_autorange(min_run_time=2.0) # runs for at least 2 seconds
Pitfall 5: Benchmarking In-Place vs Out-of-Place
# Unfair comparison — in-place doesn't allocate
x = torch.randn(1000, 1000)
# Out-of-place: allocates new tensor each time
t1 = Timer(stmt="x + 1", globals={"x": x}, description="out-of-place")
# In-place: no allocation
t2 = Timer(stmt="x.add_(1)", globals={"x": x}, description="in-place")
# In-place will be faster partly because it avoids allocation overhead
Pitfall 6: Not Controlling CPU Affinity / Frequency Scaling
For reproducible CPU benchmarks:
# Pin thread count
t = Timer(stmt="x @ y", globals={"x": x, "y": y}, num_threads=1)
# On Linux, also consider: taskset, cpufreq-set for pinning CPU core and frequency
12. Practical Recipes
Recipe 1: Compare Two Model Implementations
import torch
import torch.nn as nn
from torch.utils.benchmark import Timer, Compare
class ModelV1(nn.Module):
def __init__(self, dim):
super().__init__()
self.linear1 = nn.Linear(dim, dim * 4)
self.linear2 = nn.Linear(dim * 4, dim)
self.relu = nn.ReLU()
def forward(self, x):
return self.linear2(self.relu(self.linear1(x)))
class ModelV2(nn.Module):
def __init__(self, dim):
super().__init__()
self.linear1 = nn.Linear(dim, dim * 4)
self.linear2 = nn.Linear(dim * 4, dim)
self.silu = nn.SiLU()
def forward(self, x):
return self.linear2(self.silu(self.linear1(x)))
dim = 512
batch = 64
x = torch.randn(batch, dim)
v1, v2 = ModelV1(dim), ModelV2(dim)
results = []
for name, model in [("ReLU-FFN", v1), ("SiLU-FFN", v2)]:
t = Timer(
stmt="model(x)",
globals={"model": model, "x": x},
label="FFN forward",
sub_label=f"dim={dim}",
description=name,
)
results.append(t.blocked_autorange())
Compare(results).print()
Recipe 2: Profile Scaling Behavior
import torch
from torch.utils.benchmark import Timer, Compare
results = []
for batch in [1, 4, 16, 64, 256]:
for seq_len in [128, 512, 2048]:
x = torch.randn(batch, seq_len, 768)
w = torch.randn(768, 768)
t = Timer(
stmt="x @ w",
globals={"x": x, "w": w},
label="Projection",
sub_label=f"batch={batch}, seq={seq_len}",
description="matmul",
)
results.append(t.blocked_autorange(min_run_time=0.5))
Compare(results).print()
Recipe 3: Benchmark Custom Triton Kernel vs PyTorch Op
import torch
from torch.utils.benchmark import Timer, Compare
# Assuming you have a custom Triton kernel
# from my_kernels import triton_softmax
x = torch.randn(1024, 1024)
results = []
for desc, stmt in [
("torch", "torch.softmax(x, dim=-1)"),
("manual", "(x - x.max(dim=-1, keepdim=True).values).exp().div_("
"(x - x.max(dim=-1, keepdim=True).values).exp().sum(dim=-1, keepdim=True))"),
]:
t = Timer(
stmt=stmt,
globals={"x": x},
label="softmax",
sub_label="[1024x1024]",
description=desc,
)
results.append(t.blocked_autorange())
Compare(results).print()
Recipe 4: Measure torch.compile Speedup Properly
import torch
import torch.nn as nn
from torch.utils.benchmark import Timer, Compare
class TransformerBlock(nn.Module):
def __init__(self, d_model=512, nhead=8):
super().__init__()
self.attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_model * 4),
nn.GELU(),
nn.Linear(d_model * 4, d_model),
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
def forward(self, x):
x = self.norm1(x + self.attn(x, x, x)[0])
x = self.norm2(x + self.ffn(x))
return x
model = TransformerBlock()
x = torch.randn(8, 128, 512)
# Eager
eager_timer = Timer(
stmt="model(x)", globals={"model": model, "x": x},
label="TransformerBlock", sub_label="[8,128,512]", description="eager",
)
# Compiled — warmup first!
compiled = torch.compile(model)
for _ in range(3):
compiled(x)
compiled_timer = Timer(
stmt="fn(x)", globals={"fn": compiled, "x": x},
label="TransformerBlock", sub_label="[8,128,512]", description="compiled",
)
results = [eager_timer.blocked_autorange(), compiled_timer.blocked_autorange()]
compare = Compare(results)
compare.colorize()
compare.print()
speedup = results[0].median / results[1].median
print(f"\ntorch.compile speedup: {speedup:.2f}x")
13. Upstream Updates (June 2026)
Recent PyTorch changes relevant to benchmarking and performance:
| PR | Feature | Impact |
|---|---|---|
| #187218 | FlexGEMM BMM support | New batched matmul paths to benchmark |
| #187605 | Dynamo RangeVariable symbolic specialization | Changed compile behavior for range-based loops |
| #187494 | Distributed backend accessors | Cleaner backend switching for distributed benchmarks |
| #187602 | ShapesSpec in non-strict export | Better shape control for exported model benchmarks |
| #186398 | DTensor logspace | New distributed tensor op to benchmark |
| N/A | CUPTI profiler refactoring into _cupti/ package | Cleaner profiling infrastructure, separate from benchmarking |
FlexGEMM BMM Support
FlexGEMM now supports batched matrix multiplication, providing an alternative to cuBLAS for certain workloads. Benchmark with:
from torch.utils.benchmark import Timer, Compare
results = []
for batch in [1, 8, 32]:
x = torch.randn(batch, 256, 256)
y = torch.randn(batch, 256, 256)
t = Timer(
stmt="torch.bmm(x, y)",
globals={"x": x, "y": y},
label="BMM",
sub_label=f"batch={batch}",
description="bmm",
)
results.append(t.blocked_autorange())
Compare(results).print()
Dynamo RangeVariable Symbolic Specialization
torch.compile now handles range() variables differently, specializing on symbolic values. This can affect benchmarks that use range-based iteration in compiled code:
@torch.compile
def loop_fn(x, n):
for i in range(n):
x = x + 1
return x
Summary
| Tool | Purpose | When to Use |
|---|---|---|
Timer.timeit(N) | Fixed N iterations | Quick checks, known iteration count |
Timer.blocked_autorange() | Auto iterations, robust stats | Most benchmarks (recommended default) |
Compare | Side-by-side formatted table | Comparing implementations, shapes, configs |
Fuzzer | Random test configurations | Thorough coverage, avoiding bias |
Callgrind | Deterministic instruction counts | Micro-benchmarks, CI regression tests |
Key Rules
- Always use
torch.utils.benchmark— nevertime.time()or rawtimeit - Warmup torch.compile before measuring — compilation time is not runtime
- Use
blocked_autorange()— it handles warmup, GC, iteration count - Pin
num_threadsfor CPU benchmarks — reproducibility requires it - Use
Comparetables — organized comparison beats ad-hoc prints - Report median, not mean — outlier resistance matters
Further Resources
- torch.utils.benchmark documentation — official API reference
- PyTorch Benchmarking Tutorial — official tutorial
- Module 07 — Training Pipelines — training loops to benchmark
- Module 08 — torch.compile — understanding compilation for proper benchmark methodology
- Module 14 — Testing & Benchmarking — related testing utilities
- Module 26 — Memory Profiling — complementary profiling tools
Notebook: 28_benchmarking.ipynb
Source Files
[README.md](README.md)— This guide — benchmarking methodology, Timer, Compare, Fuzzer, Callgrind[benchmark_basics.py](benchmark_basics.py)— Timer API, blocked_autorange, Measurement objects, Compare tables, num_threads[benchmark_advanced.py](benchmark_advanced.py)— torch.compile benchmarking, shape sweeps, Fuzzer, dtype comparison, model comparison
Mixed Precision Deep Dive — FP32, FP16, BF16, and FP8
Table of Contents
- Numerical Formats Overview
- Why Mixed Precision?
- AMP: Automatic Mixed Precision
- GradScaler
- BF16 vs FP16
- FP8 Training
- Loss Scaling Deep Dive
- Mixed Precision with torch.compile
- Mixed Precision with FSDP2
- Numerical Stability Checklist
- Precision-Performance Tradeoffs
- Upstream Updates (June 2026)
1. Numerical Formats Overview
Every floating-point number is represented as: (-1)^sign × 2^exponent × (1 + mantissa). The tradeoff is fundamental: more exponent bits = larger representable range, more mantissa bits = more decimal precision.
IEEE 754 and PyTorch Formats
| Format | Bits | Exponent | Mantissa | Range | Precision | PyTorch dtype |
|---|---|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | ±3.4e38 | High | torch.float32 |
| TF32 | 19 | 8 | 10 | ±3.4e38 | Medium | (internal) |
| BF16 | 16 | 8 | 7 | ±3.4e38 | Low | torch.bfloat16 |
| FP16 | 16 | 5 | 10 | ±65504 | Medium-low | torch.float16 |
| FP8 E4M3 | 8 | 4 | 3 | ±448 | Very low | torch.float8_e4m3fn |
| FP8 E5M2 | 8 | 5 | 2 | ±57344 | Lowest | torch.float8_e5m2 |
Key Observations
FP32 — The default. 8 exponent bits give a huge dynamic range (±3.4×10^38), and 23 mantissa bits give ~7 decimal digits of precision. This is more than enough for any training scenario but uses 4 bytes per parameter.
TF32 — NVIDIA's "Tensor Float 32". Same range as FP32 (8 exponent bits) but only 10 mantissa bits. Not a user-facing dtype — it's an internal hardware format used by tensor cores when torch.backends.cuda.matmul.allow_tf32 = True. Gives near-FP32 accuracy with FP16-like throughput for matmuls.
BF16 (Brain Float 16) — Google's format. Same 8-bit exponent as FP32 (same range!) but only 7 mantissa bits (~2 decimal digits of precision). The key insight: for neural network training, range matters more than precision. Developed at Google Brain for TPU training.
FP16 (Half) — IEEE half-precision. Only 5 exponent bits means range is limited to ±65504. Values larger than this overflow to infinity. However, 10 mantissa bits give more precision than BF16. The limited range is the reason GradScaler exists.
FP8 E4M3 — 4 exponent bits, 3 mantissa bits. Range up to ±448. Designed for the forward pass where more precision helps (activations stay in a moderate range after normalization).
FP8 E5M2 — 5 exponent bits, 2 mantissa bits. Range up to ±57344. Designed for the backward pass where gradients can vary wildly in magnitude (larger range handles this better).
Memory Implications
Parameters: 100M model
FP32: 400 MB (4 bytes × 100M)
FP16/BF16: 200 MB (2 bytes × 100M) — 2× savings
FP8: 100 MB (1 byte × 100M) — 4× savings
2. Why Mixed Precision?
Mixed precision means using lower-precision formats for computation while keeping FP32 master copies for numerical stability. The benefits:
Memory Savings
| Component | FP32 | Mixed (BF16 compute) | Savings |
|---|---|---|---|
| Model parameters (compute copy) | 4B/param | 2B/param | 2× |
| Activations | 4B/element | 2B/element | 2× |
| Gradients | 4B/param | 2B/param | 2× |
| Optimizer states (Adam) | 8B/param | 8B/param | 1× (kept in FP32) |
| Master weights | — | 4B/param | Overhead |
For a 1B parameter model with Adam:
- FP32 only: 4 + 4 + 8 = 16 GB (params + grads + optimizer)
- Mixed precision: 2 + 2 + 8 + 4 = 16 GB (BF16 params + BF16 grads + FP32 optimizer + FP32 master)
The memory win comes from activations (which scale with batch size and sequence length) and from not needing FP32 gradients during backward:
Activation memory for a Transformer layer (seq_len=2048, hidden=4096, batch=8):
FP32: ~1.6 GB per layer
BF16: ~0.8 GB per layer
Throughput Gains
NVIDIA Tensor Cores operate on lower-precision types:
- FP16/BF16: 2-3× throughput vs FP32 on A100/H100 (matmul and convolution)
- FP8: 2× throughput vs BF16 on H100 (matmul only)
- TF32: ~2× throughput vs FP32 (enabled by default on Ampere+)
The key: these speedups apply to tensor core operations (matmul, conv). Elementwise ops, reductions, and memory-bound ops see less benefit.
Accuracy Impact
With proper techniques (loss scaling for FP16, FP32 accumulation), mixed precision training converges to the same accuracy as FP32 for virtually all workloads. The neural network optimization landscape is robust to reduced precision because:
- Gradient noise from mini-batching already exceeds quantization noise
- Normalization layers keep activations in representable ranges
- FP32 master weights accumulate small updates that would be lost in FP16
3. AMP: Automatic Mixed Precision
The modern PyTorch AMP API uses torch.amp.autocast to automatically cast operations to the appropriate precision.
Basic Usage
import torch
import torch.nn as nn
model = MyModel().cuda()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
for data, target in dataloader:
data, target = data.cuda(), target.cuda()
optimizer.zero_grad()
# autocast region: eligible ops run in float16
with torch.amp.autocast('cuda', dtype=torch.float16):
output = model(data)
loss = criterion(output, target)
# backward and step happen outside autocast
loss.backward()
optimizer.step()
What autocast Does
autocast maintains two lists of operations:
Cast to FP16/BF16 (compute-intensive, benefit from tensor cores):
torch.mm,torch.matmul,torch.bmmtorch.nn.functional.lineartorch.nn.functional.conv1d/2d/3dtorch.baddbmm
Keep in FP32 (numerically sensitive):
torch.nn.functional.softmaxtorch.nn.functional.cross_entropy, all loss functionstorch.nn.functional.layer_norm,batch_norm,group_normtorch.sum,torch.mean(reductions)torch.exp,torch.log,torch.pow
Rules for mixed inputs:
- If any input is FP32 and the op is NOT in the cast-down list, it stays FP32
- If inputs are mixed (FP16 + FP32), they get promoted to the wider type
- autocast only affects CUDA ops (CPU autocast exists but has limited support)
autocast Nesting and Disabling
# Nested autocast — inner region can override dtype
with torch.amp.autocast('cuda', dtype=torch.float16):
# FP16 region
y = model.encoder(x)
with torch.amp.autocast('cuda', enabled=False):
# Force FP32 for this subcomputation
y_float = y.float()
sensitive_result = custom_numerics(y_float)
z = model.decoder(sensitive_result.half())
CPU autocast
# Limited CPU autocast (mainly for BF16 on Intel CPUs with AMX)
with torch.amp.autocast('cpu', dtype=torch.bfloat16):
output = model(data)
4. GradScaler
The Problem: Gradient Underflow in FP16
FP16 has a minimum positive subnormal of ~5.96×10^-8. Gradients in deep networks routinely have magnitudes of 10^-5 to 10^-8 — right at the edge of FP16 representability. Small gradients underflow to zero, and the model stops learning.
BF16 does NOT have this problem because it shares FP32's exponent range. GradScaler is only needed for FP16 training.
How GradScaler Works
Forward: compute loss normally
Scale: loss_scaled = loss × scale_factor (e.g., 65536)
Backward: gradients are also scaled by scale_factor (chain rule)
→ small gradients become representable in FP16
Unscale: divide gradients by scale_factor before optimizer step
Check: if any gradient is inf/nan, skip optimizer step
Update: adjust scale_factor dynamically
Complete Training Loop with GradScaler
import torch
from torch.amp import autocast, GradScaler
from torch.nn.utils import clip_grad_norm_
model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
scaler = GradScaler('cuda')
for epoch in range(num_epochs):
for data, target in dataloader:
data, target = data.cuda(), target.cuda()
optimizer.zero_grad()
with autocast('cuda', dtype=torch.float16):
output = model(data)
loss = criterion(output, target)
# Scale loss and call backward
scaler.scale(loss).backward()
# Unscale gradients for clipping
scaler.unscale_(optimizer)
clip_grad_norm_(model.parameters(), max_norm=1.0)
# Step (skips if inf/nan detected)
scaler.step(optimizer)
scaler.update()
GradScaler Internals
scaler = GradScaler(
device='cuda',
init_scale=2**16, # Initial scale factor (65536)
growth_factor=2.0, # Multiply scale by this after growth_interval
backoff_factor=0.5, # Multiply scale by this on inf/nan
growth_interval=2000, # Steps between scale increases
)
The dynamic scaling algorithm:
- Start with
init_scale(default 65536) - If
growth_intervalconsecutive steps have no inf/nan → multiply scale bygrowth_factor - If any step produces inf/nan → multiply scale by
backoff_factor, skip that step - This finds the largest scale that doesn't overflow
5. BF16 vs FP16
Why BF16 Is Preferred for LLMs
| Property | FP16 | BF16 |
|---|---|---|
| Max value | 65504 | 3.4×10^38 |
| Min positive normal | 6.1×10^-5 | 1.2×10^-38 |
| Precision (decimal digits) | ~3.3 | ~2.1 |
| GradScaler needed | Yes | No |
| Overflow risk | High | None (same as FP32) |
| Hardware support | All GPUs with tensor cores | Ampere+ (A100, H100, RTX 3090+) |
The Overflow Problem
import torch
# FP16 overflow — values > 65504 become inf
x = torch.tensor(70000.0, dtype=torch.float16)
print(x) # tensor(inf, dtype=torch.float16)
# BF16 handles it — same exponent range as FP32
x = torch.tensor(70000.0, dtype=torch.bfloat16)
print(x) # tensor(70000., dtype=torch.bfloat16)
In LLM training, logits before softmax can easily exceed 65504 (especially early in training with random initialization). BF16 handles this naturally; FP16 would produce inf and NaN gradients.
When to Use Each
Use BF16 when:
- Training LLMs or large models
- You have Ampere+ hardware (A100, H100, RTX 30/40 series)
- You want simplicity (no GradScaler)
- You don't need high precision for inference
Use FP16 when:
- Deploying on older GPUs (V100, T4) that lack BF16 tensor cores
- Running inference where overflow isn't a concern (inputs are bounded)
- Accuracy is critical and you can handle GradScaler complexity
BF16 Training Loop (No Scaler)
model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
for data, target in dataloader:
data, target = data.cuda(), target.cuda()
optimizer.zero_grad()
# BF16 autocast — no scaler needed!
with torch.amp.autocast('cuda', dtype=torch.bfloat16):
output = model(data)
loss = criterion(output, target)
loss.backward()
optimizer.step()
6. FP8 Training
FP8 is the cutting edge of low-precision training, offering 2× throughput over BF16 on H100 GPUs.
Two Complementary Formats
E4M3 (4 exponent, 3 mantissa):
- Range: ±448
- Used for the forward pass (activations are bounded after normalization)
- More mantissa bits → better precision for weight × activation products
E5M2 (5 exponent, 2 mantissa):
- Range: ±57344
- Used for the backward pass (gradients span many orders of magnitude)
- More exponent bits → handles gradient dynamic range
FP8 Tensors in PyTorch
import torch
# Create FP8 tensors (requires explicit casting)
x_fp32 = torch.randn(4, 4)
x_e4m3 = x_fp32.to(torch.float8_e4m3fn)
x_e5m2 = x_fp32.to(torch.float8_e5m2)
print(x_e4m3.dtype) # torch.float8_e4m3fn
print(x_e5m2.dtype) # torch.float8_e5m2
Scaled FP8 Matmul
Because FP8 has very limited range, scaling is required to keep values representable. PyTorch provides torch._scaled_mm for this:
# FP8 scaled matrix multiplication
# a: [M, K] in float8_e4m3fn
# b: [K, N] in float8_e4m3fn (transposed)
# scale_a, scale_b: scalar tensors
a_fp32 = torch.randn(64, 128, device='cuda')
b_fp32 = torch.randn(256, 128, device='cuda') # will be transposed
# Compute scales: scale = max_representable / absmax(tensor)
scale_a = torch.tensor(448.0 / a_fp32.abs().max(), device='cuda')
scale_b = torch.tensor(448.0 / b_fp32.abs().max(), device='cuda')
# Quantize to FP8
a_fp8 = (a_fp32 * scale_a).to(torch.float8_e4m3fn)
b_fp8 = (b_fp32 * scale_b).to(torch.float8_e4m3fn)
# Scaled matmul: result = (a_fp8 @ b_fp8.T) / (scale_a * scale_b)
result = torch._scaled_mm(
a_fp8, b_fp8.t(),
scale_a=scale_a.reciprocal(),
scale_b=scale_b.reciprocal(),
out_dtype=torch.bfloat16
)
Per-Tensor vs Block Scaling
Per-tensor scaling (shown above):
- One scale factor for the entire tensor
- Simple but lossy if values have high dynamic range within a tensor
- Used in early FP8 implementations
Block scaling (MX format):
- One scale factor per block (e.g., 32 or 128 elements)
- Finer-grained: handles intra-tensor dynamic range better
- Hardware support on H100+ with newer CUDA versions
# Conceptual block scaling
block_size = 128
scales = []
for i in range(0, tensor.numel(), block_size):
block = tensor[i:i+block_size]
scale = 448.0 / block.abs().max()
scales.append(scale)
When to Use FP8
- Hardware: H100 or newer (Ada Lovelace for inference only)
- Model size: Benefits increase with larger matmuls (≥1024 dimensions)
- Use case: LLM pretraining at scale, where 2× throughput over BF16 justifies complexity
- Maturity: Still evolving — API may change between PyTorch versions
7. Loss Scaling Deep Dive
Dynamic vs Static Scaling
Dynamic scaling (default GradScaler behavior):
- Automatically finds the right scale
- Adapts to different training phases (early training may need different scale than fine-tuning)
- Costs: occasional wasted steps when scale is too high
Static scaling (manual):
scaler = GradScaler('cuda', init_scale=1024, growth_interval=float('inf'))
# Scale stays at 1024 forever — no growth, but still backs off on inf/nan
When to use static: when you know the gradient magnitude distribution won't change (e.g., fine-tuning a pre-trained model with frozen layers).
Scale Factor Dynamics
Training starts: scale = 65536
Step 1-2000: no overflow → scale grows to 131072
Step 2001: overflow detected → scale drops to 65536, step skipped
Step 2002-4001: no overflow → scale grows to 131072
...eventually finds stable maximum scale
Handling inf/nan
When GradScaler detects overflow:
- Skip the optimizer step — corrupted gradients would harm the model
- Reduce scale — multiply by
backoff_factor(default 0.5) - Zero the gradients — they're invalid
- Continue training — next step uses the reduced scale
# Monitoring scale factor during training
for step, (data, target) in enumerate(dataloader):
# ... training step ...
if step % 100 == 0:
print(f"Step {step}: scale = {scaler.get_scale():.0f}")
Common Patterns
# Pattern: gradient clipping with GradScaler
scaler.scale(loss).backward()
scaler.unscale_(optimizer) # MUST unscale before clipping
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
# Pattern: multiple losses
with autocast('cuda', dtype=torch.float16):
loss1 = criterion1(output1, target1)
loss2 = criterion2(output2, target2)
loss = loss1 + 0.5 * loss2
scaler.scale(loss).backward() # Single backward for combined loss
# Pattern: gradient accumulation
for i, (data, target) in enumerate(dataloader):
with autocast('cuda', dtype=torch.float16):
output = model(data)
loss = criterion(output, target) / accumulation_steps
scaler.scale(loss).backward()
if (i + 1) % accumulation_steps == 0:
scaler.unscale_(optimizer)
clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
8. Mixed Precision with torch.compile
torch.amp.autocast composes cleanly with torch.compile. The compiler traces through autocast regions and can further optimize precision handling.
Basic Composition
model = MyModel().cuda()
compiled_model = torch.compile(model)
with torch.amp.autocast('cuda', dtype=torch.bfloat16):
output = compiled_model(data) # compile traces through autocast
What torch.compile Does with Precision
- Fuses cast operations — instead of casting per-op, fuses multiple casts into one kernel
- Eliminates redundant casts — if an op chain stays in one dtype, no cast needed
- Optimizes accumulation — ensures FP32 accumulators where needed (e.g., large reductions)
- Pattern-matches precision — recognizes patterns like "cast → matmul → cast back" and uses tensor core instructions directly
torch.compile + GradScaler
model = MyModel().cuda()
compiled_model = torch.compile(model, mode='max-autotune')
scaler = GradScaler('cuda')
for data, target in dataloader:
optimizer.zero_grad()
with autocast('cuda', dtype=torch.float16):
output = compiled_model(data)
loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
set_float32_matmul_precision
This interacts with torch.compile by controlling TF32 usage:
# 'highest' — pure FP32 matmul (slowest, most precise)
# 'high' — TF32 for internal compute (default on Ampere+)
# 'medium' — reduced precision (BF16 accumulation for large matmuls)
torch.set_float32_matmul_precision('high')
compiled = torch.compile(model)
# Now matmuls inside compiled regions use TF32 for internal compute
9. Mixed Precision with FSDP2
FSDP2 (fully_shard) has first-class mixed precision support through MixedPrecisionPolicy.
MixedPrecisionPolicy
from torch.distributed._composable.fsdp import fully_shard, MixedPrecisionPolicy
# Define precision policy
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16, # Parameters stored/computed in BF16
reduce_dtype=torch.float32, # All-reduce in FP32 for stability
)
# Apply to model
for layer in model.layers:
fully_shard(layer, mp_policy=mp_policy)
fully_shard(model, mp_policy=mp_policy)
What Each Setting Controls
| Setting | Effect | Recommendation |
|---|---|---|
param_dtype | Cast parameters to this dtype for forward/backward | torch.bfloat16 |
reduce_dtype | Dtype for gradient all-reduce communication | torch.float32 for stability |
Why FP32 Reduce?
When averaging gradients across GPUs, small gradients can lose significant bits in FP16/BF16 addition. FP32 reduce ensures that:
- Small gradient contributions from each worker aren't lost
- The final averaged gradient is as accurate as possible
- Master weight updates (which accumulate many small deltas) stay precise
Full FSDP2 Mixed Precision Example
import torch
import torch.distributed as dist
from torch.distributed._composable.fsdp import fully_shard, MixedPrecisionPolicy
def train_fsdp_mixed_precision():
dist.init_process_group("nccl")
rank = dist.get_rank()
device = torch.device(f"cuda:{rank}")
model = LargeModel().to(device)
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
)
# Shard with mixed precision
for block in model.transformer_blocks:
fully_shard(block, mp_policy=mp_policy)
fully_shard(model, mp_policy=mp_policy)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
for data, target in dataloader:
optimizer.zero_grad()
# No autocast needed — FSDP handles casting via mp_policy
output = model(data.to(device))
loss = criterion(output, target.to(device))
loss.backward()
optimizer.step()
10. Numerical Stability Checklist
Common issues and solutions when training with mixed precision:
Gradient Underflow (FP16 only)
Symptom: Loss stops decreasing, gradients are all zero. Diagnosis: Check (grad == 0).float().mean() — if >50%, you have underflow. Solutions:
- Use GradScaler (standard fix)
- Switch to BF16 (eliminates the problem)
- Increase learning rate (larger gradients)
Loss Explosion / NaN
Symptom: Loss suddenly becomes inf or NaN. Diagnosis: Check for FP16 overflow in logits or intermediate values. Solutions:
- Switch to BF16 (larger range)
- Add gradient clipping:
clip_grad_norm_(model.parameters(), 1.0) - Reduce learning rate
- Check for numerical instability in custom layers
NaN in Softmax
Symptom: NaN output from attention or classification head. Root cause: Large logit values overflow in exp() during softmax. Solutions:
- Keep softmax in FP32 (autocast does this automatically)
- Use numerically stable softmax:
softmax(x - x.max()) - If using custom attention, ensure FP32 for the softmax step
Accumulation Errors
Symptom: Reduced accuracy compared to FP32 baseline, especially for large models. Root cause: Summing many FP16/BF16 values loses precision (catastrophic cancellation). Solutions:
- Use FP32 accumulators for reductions (PyTorch does this for matmul by default)
- Keep layer norm, batch norm in FP32 (autocast handles this)
- For custom kernels: accumulate in FP32, cast result to BF16
Optimizer State Precision
Symptom: Training diverges after many steps. Root cause: Adam's running averages (m, v) lose precision in FP16. Solution: Always keep optimizer states in FP32. This is the default — never manually cast optimizer states to FP16.
Debugging Checklist
def check_precision_health(model, loss, step):
"""Call periodically during training to catch issues early."""
# Check for NaN/inf in loss
if torch.isnan(loss) or torch.isinf(loss):
print(f"Step {step}: Loss is {loss.item()}")
return False
# Check gradient statistics
total_norm = 0.0
num_zero = 0
num_params = 0
for p in model.parameters():
if p.grad is not None:
total_norm += p.grad.data.float().norm().item() ** 2
num_zero += (p.grad == 0).sum().item()
num_params += p.grad.numel()
total_norm = total_norm ** 0.5
zero_frac = num_zero / max(num_params, 1)
if zero_frac > 0.5:
print(f"Step {step}: {zero_frac:.1%} of gradients are zero (underflow?)")
if total_norm > 100:
print(f"Step {step}: Gradient norm = {total_norm:.1f} (explosion?)")
return True
11. Precision-Performance Tradeoffs
Expected Speedups (A100/H100)
| Precision | Matmul Throughput | Memory | Use Case |
|---|---|---|---|
| FP32 | 1× (baseline) | 4B/param | Debugging, validation |
| TF32 | ~2× | 4B/param | Default (transparent) |
| FP16 + scaler | 2-3× | 2B/param | Older GPUs (V100, T4) |
| BF16 | 2-3× | 2B/param | LLM training (standard) |
| FP8 (H100) | 4-6× vs FP32 | 1B/param | Large-scale LLM training |
When Each Format Is Appropriate
FP32 only:
- Debugging numerical issues
- Tiny models where speed doesn't matter
- Reference implementations for validation
BF16 (most common for training):
- LLM pretraining
- Fine-tuning large models
- Any Ampere+ GPU workload
- Default recommendation for new projects
FP16 + GradScaler:
- V100 / T4 deployments (no BF16 support)
- Inference on all GPUs (no overflow risk with bounded inputs)
- ONNX export (better ecosystem support)
FP8:
- H100 clusters running LLM pretraining
- When 2× over BF16 throughput justifies the engineering effort
- Models with matmul-heavy architectures (Transformers)
Real-World Performance Numbers
Approximate speedups for a Transformer forward pass (batch=32, seq=2048, d=4096):
A100 GPU:
FP32: 1.0× (baseline)
TF32 (default): 1.8×
BF16 autocast: 2.5×
BF16 + compile: 3.2×
H100 GPU:
BF16: 1.0× (new baseline)
FP8: 1.6-2.0×
FP8 + compile: 2.2-2.5×
Memory Savings in Practice
For a 7B parameter model (LLaMA-like):
Parameters Activations* Optimizer Total
FP32: 28 GB ~40 GB 56 GB ~124 GB
BF16 (mixed): 14 GB ~20 GB 56 GB** ~90 GB
FP8 (experimental): 7 GB ~10 GB 56 GB** ~73 GB
* Activations for batch=4, seq=4096 (approximate)
** Optimizer always in FP32 for stability
12. Upstream Updates (June 2026)
Recent PyTorch commits relevant to mixed precision and performance:
SymmMem all_gather_offset (#187642)
Adds all_gather_offset to SymmMem for parameter-contiguous all-gather operations. Enables overlapping communication with computation in FSDP-style training by gathering only the offset portion of symmetrically allocated memory. Relevant for mixed-precision distributed training where parameter shards may be in different dtypes.
all_to_all_nd for Ulysses-style attention (#178230)
Introduces N-dimensional all-to-all collective supporting Ulysses-style sequence parallel attention. This enables efficient attention computation across devices where KV pairs are distributed, working with BF16 attention tensors for memory efficiency.
MPS FlexAttention KV batch broadcasting (#187722)
Extends FlexAttention on Apple Silicon (MPS backend) with KV batch broadcasting support. Allows K/V tensors with batch_size=1 to broadcast across query batches — critical for inference with KV cache in mixed-precision (float16 on MPS).
Dynamo virtual iterator simplification (#187103)
Simplifies virtual iterator handling in TorchDynamo, reducing graph breaks in training loops that use custom iterators. Fewer graph breaks = more operations within a single compiled region = better opportunity for precision-related fusion optimizations.
Native DSL RMSNorm fix for misaligned pointers (#186235)
Fixes a Native DSL RMSNorm implementation that could produce incorrect results with misaligned memory pointers. RMSNorm operates in FP32 for stability during mixed-precision training — a memory alignment bug here could silently corrupt the normalization, leading to training divergence.
Summary
| Concept | Key Takeaway |
|---|---|
| BF16 | Default for training on Ampere+ GPUs — same range as FP32, no scaler needed |
| FP16 + GradScaler | Required for older GPUs — GradScaler prevents gradient underflow |
| FP8 | Cutting edge — 2× over BF16 on H100, requires careful scaling |
| autocast | Automatically handles per-op precision — just wrap forward pass |
| GradScaler | Only needed for FP16 — dynamically scales loss to prevent underflow |
| FSDP2 MixedPrecisionPolicy | Compute in BF16, reduce in FP32 for distributed stability |
| torch.compile | Fuses casts, eliminates redundant precision changes |
Further Resources
- PyTorch AMP documentation — official autocast and GradScaler reference
- NVIDIA Mixed Precision Training — hardware perspective
- Module 07 — Training Pipelines — AMP in complete training loops
- Module 08 — torch.compile — compilation with mixed precision
- Module 20 — Backends Tuning — TF32 and matmul precision settings
- Module 10 — Distributed Training — FSDP2 and mixed precision at scale
- Module 28 — Benchmarking — measuring precision-performance tradeoffs
Notebook: 29_mixed_precision.ipynb
Source Files
[README.md](README.md)— This guide — numerical formats, AMP, GradScaler, BF16, FP8, FSDP2 mixed precision[precision_formats.py](precision_formats.py)— Dtype exploration, range/precision, memory comparison, conversion errors[mixed_precision_training.py](mixed_precision_training.py)— AMP training loops, GradScaler, BF16 vs FP16 comparison, torch.compile integration
Module 30: Debugging PyTorch Models
Prerequisites: Module 07 — Training Pipelines, Module 08 — torch.compile Time: ~2 hours Files: debugging_toolkit.py, compile_debugging.py
Table of Contents
- The Debugging Mindset
- Anomaly Detection
- NaN/Inf Detection
- Gradient Flow Checking
- Shape Debugging
- Device Mismatch
- TORCH_SHOW_CPP_STACKTRACES
- Debugging torch.compile
- Common Error Messages and Fixes
- Memory Debugging
- Performance Debugging
- Reproducibility for Bug Reports
- Upstream Updates (June 20–22, 2026)
1. The Debugging Mindset
Debugging PyTorch models requires a systematic approach. Random changes and trial-and-error waste hours. Follow this protocol:
Reproduce → Isolate → Identify → Fix → Verify
The Protocol
Step 1: Reproduce — Create a minimal reproducer that triggers the bug every time. Strip away everything unnecessary: smaller batch size, fewer layers, synthetic data. A 20-line script that reproduces the bug is worth more than a 500-line training loop that "sometimes fails."
Step 2: Isolate — Narrow down where the problem occurs. Is it in the forward pass? Backward pass? Data loading? A specific layer? Use binary search: comment out half the model, check if the bug persists.
Step 3: Identify — Once isolated, understand why it happens. Read error messages carefully — PyTorch gives detailed tracebacks. Check tensor shapes, dtypes, devices, and values at the failure point.
Step 4: Fix — Apply the minimal fix. Don't rewrite working code around the bug.
Step 5: Verify — Run the original failing case AND related cases. Confirm the fix doesn't break other things.
The Minimal Repro
Always start here. A good minimal repro:
import torch
import torch.nn as nn
# Smallest model that triggers the bug
model = nn.Linear(10, 5)
# Simplest input that triggers the bug
x = torch.randn(2, 10)
# Exact sequence that fails
loss = model(x).sum()
loss.backward()
Strip data loading, logging, checkpointing, distributed — anything not needed to trigger the bug. If the bug disappears when you simplify, you've already learned something about its cause.
2. Anomaly Detection
PyTorch's autograd anomaly detection catches problems during the backward pass that are otherwise silent or produce cryptic errors later.
Enabling Anomaly Detection
# Context manager (preferred)
with torch.autograd.detect_anomaly():
output = model(input)
loss = criterion(output, target)
loss.backward()
# Global setting
torch.autograd.set_detect_anomaly(True)
# ... training code ...
torch.autograd.set_detect_anomaly(False)
What It Catches
| Problem | Without detect_anomaly | With detect_anomaly |
|---|---|---|
| NaN in backward | Silent propagation | RuntimeError with traceback |
| In-place op on grad tensor | Cryptic error later | Immediate error at the op |
| Double backward without retain_graph | Confusing error | Clear traceback to the first backward |
Example: Catching NaN in Backward
import torch
import torch.nn as nn
class BuggyModel(nn.Module):
def forward(self, x):
# log(0) produces -inf, gradient becomes NaN
return torch.log(x)
model = BuggyModel()
x = torch.zeros(5, requires_grad=True) # log(0) = -inf
with torch.autograd.detect_anomaly():
out = model(x)
out.sum().backward() # Raises RuntimeError with full traceback
Performance Cost
Anomaly detection adds significant overhead (2-5x slower) because it:
- Records the full forward-pass traceback for every operation
- Validates every gradient in the backward pass
Rule: Enable only during debugging. Never in production training.
# Good: enable only when investigating a bug
debug_mode = os.environ.get("DEBUG", "0") == "1"
torch.autograd.set_detect_anomaly(debug_mode)
3. NaN/Inf Detection
NaN (Not a Number) and Inf values are the most common silent killers in training. They propagate through computations and corrupt all downstream values.
Manual Checks
def check_tensor(t, name="tensor"):
"""Check a tensor for NaN/Inf values."""
if torch.isnan(t).any():
print(f"WARNING: NaN detected in {name}")
print(f" Shape: {t.shape}, NaN count: {torch.isnan(t).sum().item()}")
return False
if torch.isinf(t).any():
print(f"WARNING: Inf detected in {name}")
print(f" Shape: {t.shape}, Inf count: {torch.isinf(t).sum().item()}")
return False
return True
Hook-Based Automatic Detection
Register hooks to catch NaN/Inf as they appear, without modifying model code:
def nan_hook(module, input, output):
"""Forward hook that detects NaN/Inf in module outputs."""
if isinstance(output, torch.Tensor):
if torch.isnan(output).any() or torch.isinf(output).any():
raise RuntimeError(
f"NaN/Inf detected in output of {module.__class__.__name__}\n"
f" Output shape: {output.shape}\n"
f" NaN count: {torch.isnan(output).sum().item()}\n"
f" Inf count: {torch.isinf(output).sum().item()}"
)
# Register on all modules
for name, module in model.named_modules():
module.register_forward_hook(nan_hook)
Common Causes
| Cause | Example | Fix |
|---|---|---|
| Learning rate too high | Gradients explode → weights overflow | Reduce LR, use gradient clipping |
| log(0) | torch.log(probabilities) where some are 0 | torch.log(x + 1e-8) or torch.clamp(x, min=1e-8) |
| Division by zero | x / norm where norm is 0 | x / (norm + 1e-8) |
| Softmax overflow | Very large logits → exp overflow | Use log_softmax instead of log(softmax(x)) |
| sqrt(0) gradient | torch.sqrt(x) at x=0 has infinite gradient | torch.sqrt(x + 1e-8) |
| Unstable loss | Cross-entropy with raw probabilities | Use F.cross_entropy (numerically stable) |
Gradient NaN Detection
def grad_nan_hook(module, grad_input, grad_output):
"""Backward hook that detects NaN in gradients."""
for i, grad in enumerate(grad_output):
if grad is not None and torch.isnan(grad).any():
raise RuntimeError(
f"NaN gradient in {module.__class__.__name__} "
f"grad_output[{i}], shape={grad.shape}"
)
for module in model.modules():
module.register_full_backward_hook(grad_nan_hook)
4. Gradient Flow Checking
Training may fail silently when gradients vanish or explode. The model "trains" but loss never decreases, or it diverges suddenly.
Check That Gradients Exist
# After loss.backward()
for name, param in model.named_parameters():
if param.requires_grad:
if param.grad is None:
print(f"NO GRADIENT: {name}")
elif param.grad.norm() == 0:
print(f"ZERO GRADIENT: {name}")
Gradient Flow Checker
def check_gradient_flow(named_parameters):
"""Print gradient statistics for all parameters."""
print(f"{'Layer':<40} {'Grad Norm':<12} {'Grad Mean':<12} {'Grad Max':<12}")
print("-" * 76)
for name, param in named_parameters:
if param.requires_grad and param.grad is not None:
grad = param.grad
print(f"{name:<40} {grad.norm():<12.6f} "
f"{grad.mean():<12.2e} {grad.abs().max():<12.2e}")
Detecting Vanishing Gradients
Symptoms: loss plateaus, early layers have near-zero gradients, later layers have normal gradients.
def detect_vanishing_gradients(model, threshold=1e-7):
"""Detect layers with vanishing gradients."""
vanishing = []
for name, param in model.named_parameters():
if param.grad is not None and param.grad.norm() < threshold:
vanishing.append((name, param.grad.norm().item()))
if vanishing:
print("WARNING: Vanishing gradients detected:")
for name, norm in vanishing:
print(f" {name}: grad_norm = {norm:.2e}")
return vanishing
Detecting Exploding Gradients
Symptoms: loss becomes NaN/Inf suddenly, gradient norms grow exponentially each step.
def detect_exploding_gradients(model, threshold=100.0):
"""Detect layers with exploding gradients."""
exploding = []
for name, param in model.named_parameters():
if param.grad is not None and param.grad.norm() > threshold:
exploding.append((name, param.grad.norm().item()))
if exploding:
print("WARNING: Exploding gradients detected:")
for name, norm in exploding:
print(f" {name}: grad_norm = {norm:.2e}")
return exploding
Fix: Gradient Clipping
# Clip by norm (most common)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# Clip by value
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)
5. Shape Debugging
Shape mismatches are the most common PyTorch error. They produce clear error messages, but finding which operation caused the mismatch in a large model can be tricky.
Strategy 1: Print Shapes at Each Step
class DebugModel(nn.Module):
def __init__(self):
super().__init__()
self.layers = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 10),
)
def forward(self, x):
print(f"Input: {x.shape}")
for i, layer in enumerate(self.layers):
x = layer(x)
print(f"After layer {i} ({layer.__class__.__name__}): {x.shape}")
return x
Strategy 2: Shape-Logging Hook
def shape_hook(name):
"""Create a hook that logs input/output shapes for a module."""
def hook(module, input, output):
in_shapes = [x.shape if isinstance(x, torch.Tensor) else type(x) for x in input]
out_shape = output.shape if isinstance(output, torch.Tensor) else type(output)
print(f"{name}: input={in_shapes} → output={out_shape}")
return hook
# Register on all modules
for name, module in model.named_modules():
if not list(module.children()): # leaf modules only
module.register_forward_hook(shape_hook(name))
Strategy 3: Model Surgery (Isolate the Layer)
When a model is too large to debug as a whole, run each layer individually:
x = torch.randn(batch_size, channels, height, width)
for name, module in model.named_children():
try:
x = module(x)
print(f"✓ {name}: output shape = {x.shape}")
except Exception as e:
print(f"✗ {name}: FAILED — {e}")
print(f" Input shape was: {x.shape}")
break
Strategy 4: torchinfo
from torchinfo import summary
summary(model, input_size=(1, 3, 224, 224))
This prints a table showing each layer's input/output shape, parameter count, and multiply-accumulate operations.
Common Shape Errors
| Error | Cause | Fix |
|---|---|---|
mat1 and mat2 shapes cannot be multiplied | Linear layer input size wrong | Check in_features matches input dim |
Expected 4D input (got 2D) | Conv2d needs (B, C, H, W) | x.unsqueeze(0).unsqueeze(0) or fix data |
size mismatch, m1: [32 x 512], m2: [256 x 10] | Flatten size doesn't match Linear input | Calculate correct flatten size |
6. Device Mismatch
The error Expected all tensors to be on the same device means you're mixing CPU and CUDA tensors in an operation.
Systematic Fix
# Step 1: Check where tensors live
def print_devices(model, inputs):
"""Print device of all model parameters and inputs."""
print("Model parameters:")
for name, param in model.named_parameters():
print(f" {name}: {param.device}")
print("\nInputs:")
if isinstance(inputs, torch.Tensor):
print(f" input: {inputs.device}")
elif isinstance(inputs, (list, tuple)):
for i, inp in enumerate(inputs):
if isinstance(inp, torch.Tensor):
print(f" input[{i}]: {inp.device}")
Device Checker Hook
def device_check_hook(expected_device):
"""Hook that verifies all inputs/outputs are on the expected device."""
def hook(module, input, output):
for i, inp in enumerate(input):
if isinstance(inp, torch.Tensor) and inp.device != expected_device:
raise RuntimeError(
f"{module.__class__.__name__}: input[{i}] is on "
f"{inp.device}, expected {expected_device}"
)
return hook
Common Causes
- Forgot to move input to GPU:
x = x.to(device)before passing to model - Created tensor inside forward(): Use
torch.zeros(..., device=x.device)nottorch.zeros(...) - Loss target on wrong device:
target = target.to(device) - Buffer not registered: Use
self.register_buffer('name', tensor)notself.tensor = tensor
The Fix Pattern
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)
for batch in dataloader:
inputs, targets = batch
inputs = inputs.to(device)
targets = targets.to(device)
output = model(inputs)
loss = criterion(output, targets)
7. TORCH_SHOW_CPP_STACKTRACES
When PyTorch crashes at the C++ level, Python tracebacks only show you the Python call that triggered the error. The actual cause is buried in C++ code.
Enabling C++ Stacktraces
TORCH_SHOW_CPP_STACKTRACES=1 python script.py
What It Shows
Without the environment variable:
RuntimeError: CUDA error: device-side assert triggered
With TORCH_SHOW_CPP_STACKTRACES=1:
RuntimeError: CUDA error: device-side assert triggered
CUDA kernel errors might be asynchronously reported at some other API call...
C++ Stacktrace:
at aten::index_select(self, dim, index)
at torch::autograd::...
...
When to Use
- Segfaults or crashes (not Python exceptions)
- CUDA runtime errors
- Internal PyTorch assertion failures
- Errors from custom C++ extensions
Other Useful Environment Variables
# Synchronous CUDA errors (pinpoints the exact kernel)
CUDA_LAUNCH_BLOCKING=1 python script.py
# Both together for maximum debug info
TORCH_SHOW_CPP_STACKTRACES=1 CUDA_LAUNCH_BLOCKING=1 python script.py
# Disable CUDA caching allocator (for memory debugging)
PYTORCH_NO_CUDA_MEMORY_CACHING=1 python script.py
8. Debugging torch.compile
torch.compile introduces a new class of errors: graph breaks, recompilations, and backend failures. These are different from standard PyTorch bugs.
Graph Breaks: Detection
A graph break means Dynamo couldn't compile a section of your code, falling back to eager mode. This hurts performance.
# Method 1: TORCH_LOGS environment variable
# TORCH_LOGS="graph_breaks" python script.py
# Method 2: explain() API
import torch._dynamo as dynamo
def my_function(x):
x = x * 2
print(x) # This causes a graph break!
return x + 1
explanation = dynamo.explain(my_function)(torch.randn(10))
print(explanation)
# Shows: graph_break_count, break_reasons, out_guards
Graph Breaks: Common Causes and Fixes
| Cause | Example | Fix |
|---|---|---|
print() in compiled code | print(x.shape) | Remove or guard with if not torch.compiler.is_compiling() |
| Data-dependent control flow | if x.sum() > 0: | Use torch.where or torch.cond |
| Unsupported Python builtin | sorted(list) | Rewrite with torch ops |
| Non-tensor data structures | Building a list in a loop | Use tensor operations |
| Calling uncompiled functions | External library calls | Wrap or inline |
Recompilation Detection
Recompilation happens when Dynamo's guards are triggered (input shapes change, etc.):
TORCH_LOGS="recompiles" python script.py
# Programmatic detection
import torch._dynamo as dynamo
# Count compilations
compile_count = 0
def counting_compiler(gm, example_inputs):
global compile_count
compile_count += 1
return gm
compiled_fn = torch.compile(my_fn, backend=counting_compiler)
The Minifier
When torch.compile produces wrong results or crashes, the minifier creates a minimal reproduction:
import torch._dynamo.config
# Generate minimal repro after dynamo error
torch._dynamo.config.repro_after = "dynamo"
# Or after AOTAutograd
torch._dynamo.config.repro_after = "aot"
Verbose Mode
# Full compilation logs
TORCH_LOGS="dynamo" python script.py
# Inductor-generated code
TORCH_LOGS="output_code" python script.py
# Everything (very verbose)
TORCH_LOGS="+dynamo,+inductor" python script.py
Debugging Wrong Results
# Compare compiled vs eager outputs
model_eager = MyModel()
model_compiled = torch.compile(MyModel())
# Load same weights
model_compiled.load_state_dict(model_eager.state_dict())
x = torch.randn(2, 10)
out_eager = model_eager(x)
out_compiled = model_compiled(x)
print(f"Max difference: {(out_eager - out_compiled).abs().max()}")
assert torch.allclose(out_eager, out_compiled, atol=1e-5)
9. Common Error Messages and Fixes
Error Table
| # | Error | Cause | Fix |
|---|---|---|---|
| 1 | RuntimeError: CUDA out of memory | GPU memory exhausted | Reduce batch size, use gradient checkpointing, use mixed precision, call torch.cuda.empty_cache() |
| 2 | RuntimeError: Expected all tensors to be on the same device | Mixing CPU and CUDA tensors | Move all tensors to same device with .to(device) |
| 3 | RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation | In-place op on a tensor needed for backward | Replace x.add_(1) with x = x + 1, avoid in-place ops on leaf tensors |
| 4 | RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn | Calling .backward() on a non-grad tensor | Ensure inputs have requires_grad=True, check model parameters |
| 5 | torch._dynamo.exc.Unsupported | Graph break in torch.compile | See Section 8 — remove unsupported ops or use torch.compiler.is_compiling() guard |
| 6 | CUDA error: device-side assert triggered | Index out of bounds in CUDA kernel | Run with CUDA_LAUNCH_BLOCKING=1, check label indices < num_classes |
| 7 | RuntimeError: Trying to backward through the graph a second time | Calling .backward() twice without retain_graph=True | Add retain_graph=True or restructure to avoid double backward |
| 8 | RuntimeError: expected scalar type Float but found Half | Dtype mismatch between FP32 and FP16 | Use autocast or explicit .float() / .half() conversion |
| 9 | RuntimeError: mat1 and mat2 shapes cannot be multiplied | Linear layer shape mismatch | Check in_features matches the flattened input dimension |
| 10 | ValueError: optimizer got an empty parameter list | No parameters passed to optimizer | Check model.parameters() is not empty, ensure modules are registered as attributes |
| 11 | RuntimeError: Input type and weight type should be the same | Mixed dtypes (e.g., double input, float weights) | Use x = x.float() or model.double() |
| 12 | RuntimeError: expected stride to be a single integer or a list of integers | Wrong argument type to operation | Check function signature — likely passing tensor where int expected |
Detailed Examples
CUDA Out of Memory
# Diagnose
print(f"Allocated: {torch.cuda.memory_allocated() / 1e9:.2f} GB")
print(f"Reserved: {torch.cuda.memory_reserved() / 1e9:.2f} GB")
print(f"Max allocated: {torch.cuda.max_memory_allocated() / 1e9:.2f} GB")
# Fix 1: Reduce batch size
# Fix 2: Gradient checkpointing
from torch.utils.checkpoint import checkpoint
# Fix 3: Mixed precision
with torch.autocast('cuda'):
output = model(input)
# Fix 4: Clear cache (doesn't free PyTorch tensors, just cached blocks)
torch.cuda.empty_cache()
In-Place Operation Error
# BAD: in-place modification of a tensor needed for backward
x = torch.randn(5, requires_grad=True)
y = x ** 2
x.mul_(2) # In-place modification!
y.sum().backward() # ERROR
# GOOD: create a new tensor
x = torch.randn(5, requires_grad=True)
y = x ** 2
x_new = x * 2 # New tensor, x unchanged
y.sum().backward() # Works
Device-Side Assert (Index Out of Bounds)
# Common cause: label index >= num_classes
num_classes = 10
labels = torch.tensor([0, 5, 10]) # 10 is out of bounds!
output = torch.randn(3, num_classes)
loss = F.cross_entropy(output, labels) # CUDA assert!
# Fix: clamp or validate labels
assert labels.max() < num_classes, f"Label {labels.max()} >= {num_classes}"
10. Memory Debugging
Detecting Memory Leaks
A memory leak in PyTorch usually means tensors are being held alive unintentionally.
import gc
def check_memory_growth(model, dataloader, num_steps=10):
"""Check if memory grows over training steps."""
memory_log = []
for i, (x, y) in enumerate(dataloader):
if i >= num_steps:
break
output = model(x)
loss = F.cross_entropy(output, y)
loss.backward()
# Record memory AFTER backward
if torch.cuda.is_available():
mem = torch.cuda.memory_allocated()
else:
import psutil
mem = psutil.Process().memory_info().rss
memory_log.append(mem)
# Critical: zero gradients
model.zero_grad(set_to_none=True)
# Check for growth
if memory_log[-1] > memory_log[0] * 1.1:
print(f"WARNING: Memory grew from {memory_log[0]/1e6:.1f}MB "
f"to {memory_log[-1]/1e6:.1f}MB over {num_steps} steps")
return memory_log
Common Memory Leak Causes
- Storing loss history without
.item():
```python # BAD: holds entire computation graph! losses.append(loss)
# GOOD: detach the scalar value losses.append(loss.item()) ```
- Not zeroing gradients:
``python # Gradients accumulate by default optimizer.zero_grad() # or model.zero_grad(set_to_none=True) ``
- Holding references in hooks:
```python # BAD: closure holds reference to output outputs = [] def hook(m, i, o): outputs.append(o) # Keeps tensor alive!
# GOOD: detach or only store what you need def hook(m, i, o): outputs.append(o.detach().cpu()) ```
Memory Snapshot (CUDA)
# Record memory history
torch.cuda.memory._record_memory_history()
# ... run your code ...
# Save snapshot
torch.cuda.memory._dump_snapshot("memory_snapshot.pickle")
torch.cuda.memory._record_memory_history(enabled=None)
# Analyze with: https://pytorch.org/memory_viz
11. Performance Debugging
torch.profiler
from torch.profiler import profile, ProfilerActivity, schedule
with profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
schedule=schedule(wait=1, warmup=1, active=3, repeat=1),
on_trace_ready=torch.profiler.tensorboard_trace_handler('./log'),
record_shapes=True,
profile_memory=True,
with_stack=True,
) as prof:
for step, (x, y) in enumerate(dataloader):
output = model(x)
loss = criterion(output, y)
loss.backward()
optimizer.step()
optimizer.zero_grad()
prof.step()
CPU-Bound vs GPU-Bound
# Quick test: does CUDA synchronization slow things down?
import time
torch.cuda.synchronize()
start = time.time()
for _ in range(100):
output = model(x)
torch.cuda.synchronize()
elapsed = time.time() - start
# If adding synchronize() doesn't change timing much → CPU-bound
# If it significantly increases time → GPU is already the bottleneck
Data Loading Bottleneck
# If GPU utilization is low, data loading may be the bottleneck
import time
# Time data loading
load_times = []
for i, batch in enumerate(dataloader):
if i >= 10:
break
start = time.time()
x, y = batch
x = x.to(device)
load_times.append(time.time() - start)
# Time model
model_times = []
for i, batch in enumerate(dataloader):
if i >= 10:
break
x, y = batch[0].to(device), batch[1].to(device)
start = time.time()
output = model(x)
loss = criterion(output, y)
loss.backward()
torch.cuda.synchronize()
model_times.append(time.time() - start)
print(f"Avg data load: {sum(load_times)/len(load_times)*1000:.1f}ms")
print(f"Avg model step: {sum(model_times)/len(model_times)*1000:.1f}ms")
Fix data loading bottlenecks: increase num_workers, enable pin_memory=True, use persistent_workers=True, pre-process data.
12. Reproducibility for Bug Reports
When filing a bug report (or debugging your own code), reproducibility is essential.
Setting All Seeds
import torch
import numpy as np
import random
def set_all_seeds(seed=42):
"""Set all random seeds for reproducibility."""
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
np.random.seed(seed)
random.seed(seed)
# Deterministic algorithms (may be slower)
torch.use_deterministic_algorithms(True)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
Capturing Environment Info
def print_environment():
"""Print all relevant environment information for bug reports."""
import sys
import platform
print(f"Python: {sys.version}")
print(f"PyTorch: {torch.__version__}")
print(f"CUDA available: {torch.cuda.is_available()}")
if torch.cuda.is_available():
print(f"CUDA version: {torch.version.cuda}")
print(f"GPU: {torch.cuda.get_device_name(0)}")
print(f"OS: {platform.platform()}")
print(f"\nFull config:\n{torch.__config__.show()}")
Minimal Repro Template
"""
Minimal reproducer for [describe the bug].
Environment:
- PyTorch: [version]
- Python: [version]
- OS: [os]
- GPU: [gpu or CPU]
Steps to reproduce:
1. Run this script
Expected: [what should happen]
Actual: [what actually happens]
"""
import torch
torch.manual_seed(42)
# Minimal code that triggers the bug
model = ...
x = ...
output = model(x) # Bug occurs here
13. Upstream Updates (June 20–22, 2026)
Recent PyTorch commits relevant to debugging and general usage:
| PR | Title | Impact |
|---|---|---|
| #187768 | MPS FlexAttention lse return | FlexAttention on MPS now correctly returns log-sum-exp alongside attention output |
| #187758 | Sequential.__getitem__ type overloads | Better type checking when indexing nn.Sequential — clearer errors for invalid indexing |
| #187776 | SymmMem copy optimization | Optimized symmetric memory copy for distributed training — reduced latency |
| #187702 | vmap batching rule for repeat_interleave | torch.vmap now supports repeat_interleave — no more manual unbatching workaround |
| #184653 | Dynamo globals fix for unregistered modules | Fixed graph break when accessing global modules not registered as submodules — helps torch.compile debugging |
| #187778 | all_to_all_nd narrow-row throughput fix | Improved throughput for narrow-row all-to-all patterns common in MoE training |
Impact on Debugging
- #184653 is directly relevant: if you had graph breaks from accessing global module objects, this is now fixed. Upgrade PyTorch to resolve.
- #187758 improves error messages when mis-indexing Sequential — less confusing debugging.
- #187768 fixes a subtle bug where MPS FlexAttention returned wrong
lsevalues — this would show up as incorrect loss values on Apple Silicon.
Quick Reference Card
┌─────────────────────────────────────────────────────────────────────────┐
│ PyTorch Debugging Quick Reference │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ NaN/Inf Detection: │
│ torch.autograd.set_detect_anomaly(True) │
│ torch.isnan(t).any() / torch.isinf(t).any() │
│ │
│ Gradient Debugging: │
│ param.grad is None → not connected to loss │
│ param.grad.norm() ≈ 0 → vanishing gradients │
│ param.grad.norm() → ∞ → exploding gradients │
│ │
│ torch.compile Debugging: │
│ TORCH_LOGS="graph_breaks" python script.py │
│ TORCH_LOGS="recompiles" python script.py │
│ torch._dynamo.explain(fn)(inputs) │
│ │
│ C++ Errors: │
│ TORCH_SHOW_CPP_STACKTRACES=1 python script.py │
│ CUDA_LAUNCH_BLOCKING=1 python script.py │
│ │
│ Memory: │
│ torch.cuda.memory_allocated() │
│ torch.cuda.memory_summary() │
│ torch.cuda.memory._record_memory_history() │
│ │
│ Reproducibility: │
│ torch.manual_seed(42) │
│ torch.use_deterministic_algorithms(True) │
│ torch.__config__.show() │
│ │
└─────────────────────────────────────────────────────────────────────────┘
Key Takeaways
| Principle | Implementation |
|---|---|
| Always create a minimal repro | Strip to smallest code that reproduces the bug |
| Use anomaly detection | torch.autograd.set_detect_anomaly(True) — but only when debugging |
| Check NaN/Inf early | Register forward hooks on all modules |
| Monitor gradient norms | Log per-layer gradient norms each step |
| Print shapes systematically | Hooks > manual prints > torchinfo |
| Fix device mismatches at data boundary | .to(device) right after data loading |
| Use environment variables for C++ issues | TORCH_SHOW_CPP_STACKTRACES=1, CUDA_LAUNCH_BLOCKING=1 |
| Use explain() for compile issues | torch._dynamo.explain(fn)(inputs) |
| Store scalars not tensors | loss.item() not loss |
| Set all seeds for repro | torch.manual_seed, np.random.seed, random.seed |
Further Resources
- PyTorch Debugging FAQ — official troubleshooting
- torch.compile troubleshooting — Dynamo debugging guide
- Module 07 — Training Pipelines — gradient clipping and AMP
- Module 08 — torch.compile — compilation fundamentals
- Module 26 — Memory Profiling — detailed memory analysis
- Module 29 — Mixed Precision — dtype-related debugging
Notebook: 30_debugging.ipynb
Module 31 — torchao: Architecture Optimization
Prerequisites: Module 07 — Training, Module 08 — torch.compile, Module 29 — Mixed Precision
> > Time: ~3 hours | Files: quantization_basics.py, torchao_workflows.py
Table of Contents
- What is torchao?
- torchao vs torch.ao
- Installation
- The quantize_() API
- Weight-Only Quantization (INT8/INT4)
- Dynamic Quantization (INT8)
- Float8 Training and Inference
- Semi-Structured Sparsity (2:4)
- PT2E Quantization Flow
- Integration with torch.compile
- Practical: Quantize and Benchmark
- Choosing a Quantization Strategy
- Upstream Updates (June 2026)
1. What is torchao?
torchao (PyTorch Architecture Optimization) is a PyTorch-native library for making models faster and smaller through quantization, sparsity, and dtype optimization.
Repository: github.com/pytorch/ao
Key features
- Quantization: INT8, INT4, FP8 weight-only and dynamic quantization
- Sparsity: Semi-structured (2:4) sparsity with hardware acceleration
- Composability: Works seamlessly with
torch.compilefor fused, optimized kernels - Tensor subclass-based: Uses PyTorch's tensor subclass system — no graph rewrites needed
Why torchao matters
Quantization and sparsity can provide:
| Technique | Memory Reduction | Speedup | Use Case |
|---|---|---|---|
| INT8 weight-only | ~2× | 1.2–2× | Memory-bound inference (LLM serving) |
| INT4 weight-only | ~4× | 2–4× | Extremely memory-bound inference |
| INT8 dynamic | ~2× | 1.5–3× | Compute-bound batch inference |
| FP8 | ~2× | 1.5–2× | Training on H100+ |
| 2:4 sparsity | ~2× | ~2× | Prunable models on Ampere+ |
These are approximate — actual results depend on model architecture, hardware, batch size, and workload characteristics.
2. torchao vs torch.ao
PyTorch has two quantization stories. Understanding the distinction is critical:
torch.ao.quantization (Old — Being Deprecated)
# OLD approach — FX-based graph rewriting
import torch.ao.quantization as taq
model_prepared = taq.prepare(model, qconfig_mapping, example_inputs)
# ... calibrate ...
model_quantized = taq.convert(model_prepared)
- Lives in-tree at
torch.ao.quantization - Uses FX graph tracing (fragile, breaks on dynamic control flow)
- Requires
qconfig_mappingboilerplate - Not composable with
torch.compile - Being deprecated in favor of torchao
torchao (New — Recommended)
# NEW approach — tensor subclass-based
from torchao import quantize_
from torchao.quantization import int8_weight_only
quantize_(model, int8_weight_only())
# That's it. Model is quantized.
- External library at
pytorch/ao - Uses tensor subclasses — weights become quantized tensor objects
- No graph rewriting — the model structure stays the same
- Composable with
torch.compile(Inductor generates fused kernels) - Actively developed, production-ready
Migration path
torch.ao.quantization.quantize_dynamic → torchao.quantize_(model, int8_dynamic_activation_int8_weight())
torch.ao.quantization.prepare/convert → torchao.quantize_(model, int8_weight_only())
torch.ao.quantization (FX) → PT2E quantization flow (torch.export + quantize_pt2e)
3. Installation
pip install torchao
Verify:
import torchao
print(torchao.__version__)
print(f"CUDA available: {torchao.utils.TORCH_VERSION_AT_LEAST_2_3}")
torchao requires PyTorch 2.3+ and works best with PyTorch 2.6+ for the latest features.
Note: torchao quantization kernels are optimized for CUDA. CPU execution works for development/testing but won't show the full performance benefits.
4. The quantize_() API
The quantize_() function is the single entry point for all torchao quantization:
from torchao import quantize_
from torchao.quantization import (
int8_weight_only,
int4_weight_only,
int8_dynamic_activation_int8_weight,
float8_dynamic_activation_float8_weight,
)
# Weight-only quantization
quantize_(model, int8_weight_only()) # INT8 weights
quantize_(model, int4_weight_only()) # INT4 weights
# Dynamic quantization (weights + activations)
quantize_(model, int8_dynamic_activation_int8_weight())
# Float8 quantization (H100+)
quantize_(model, float8_dynamic_activation_float8_weight())
How it works
quantize_()walks the model'snn.Moduletree- For each matching module (default:
nn.Linear), it replaces the weight tensor with a quantized tensor subclass - The module itself is unchanged — it's still
nn.Linear - When the layer runs, the tensor subclass handles dequantization during matmul
- With
torch.compile, Inductor fuses the dequantize + matmul into a single kernel
# Before quantize_():
model.layer.weight # torch.float16, shape [1024, 512]
quantize_(model, int8_weight_only())
# After quantize_():
model.layer.weight # AffineQuantizedTensor (int8 storage, float16 dequant)
type(model.layer) # Still nn.Linear!
Filtering which layers get quantized
# Only quantize layers with > 1024 dimensions
def filter_fn(module, fqn):
if isinstance(module, torch.nn.Linear):
return module.in_features >= 1024
return False
quantize_(model, int8_weight_only(), filter_fn=filter_fn)
5. Weight-Only Quantization (INT8/INT4)
Weight-only quantization stores model weights in lower precision (INT8 or INT4) while keeping activations in FP16/BF16. This is the most common approach for memory-bound inference (e.g., LLM serving at batch size 1).
How quantization works
For a weight tensor W in FP16:
Quantize: W_int8 = round(W / scale) + zero_point
Dequantize: W_approx = (W_int8 - zero_point) * scale
Where:
scale = (max(W) - min(W)) / (2^bits - 1)zero_pointmaps the real zero to an integer value
INT8 weight-only
from torchao import quantize_
from torchao.quantization import int8_weight_only
quantize_(model, int8_weight_only())
# With group-wise quantization (more accurate, slightly more overhead)
quantize_(model, int8_weight_only(group_size=128))
Per-channel (default): One scale per output channel. Good accuracy, simple.
Per-group (group_size=128 or 32): One scale per group of weights. Better accuracy for large layers, slight overhead for scale storage.
INT4 weight-only
from torchao.quantization import int4_weight_only
# INT4 always uses group-wise quantization
quantize_(model, int4_weight_only(group_size=128))
quantize_(model, int4_weight_only(group_size=32)) # More accurate, more scales
Memory comparison
For a 7B parameter model (all Linear layers):
| Precision | Memory | Relative |
|---|---|---|
| FP32 | 28 GB | 1.0× |
| FP16/BF16 | 14 GB | 0.5× |
| INT8 | 7 GB | 0.25× |
| INT4 | 3.5 GB | 0.125× |
6. Dynamic Quantization (INT8)
Dynamic quantization quantizes both weights and activations. Weights are quantized ahead of time; activations are quantized on-the-fly during each forward pass.
from torchao import quantize_
from torchao.quantization import int8_dynamic_activation_int8_weight
quantize_(model, int8_dynamic_activation_int8_weight())
When to use dynamic quantization
- Compute-bound workloads: Batch inference (batch_size > 1)
- INT8 matrix multiply (GEMM) is ~2× faster than FP16 on most hardware
- The activation quantization overhead is amortized over larger batches
How it differs from weight-only
| Weight-Only | Dynamic | |
|---|---|---|
| Weights | INT8/INT4 | INT8 |
| Activations | FP16/BF16 | INT8 (computed at runtime) |
| Matmul precision | FP16 | INT8 |
| Best for | Memory-bound (batch=1) | Compute-bound (batch>1) |
| Accuracy impact | Lower | Slightly higher |
| Extra overhead | None | Per-batch activation quantization |
Calibration-free
Unlike static quantization, dynamic quantization doesn't need calibration data — activation ranges are computed per-batch. This makes deployment simpler.
7. Float8 Training and Inference
Float8 (FP8) uses 8-bit floating-point formats for both training and inference. Unlike INT8, FP8 preserves the floating-point dynamic range.
FP8 formats
- E4M3 (4 exponent, 3 mantissa): Higher precision, used for weights and activations
- E5M2 (5 exponent, 2 mantissa): Higher range, used for gradients during training
FP8 inference with torchao
from torchao import quantize_
from torchao.quantization import float8_dynamic_activation_float8_weight
# Requires H100 or newer (Hopper architecture)
quantize_(model, float8_dynamic_activation_float8_weight())
FP8 training with Float8Linear
For training, torchao provides Float8Linear which replaces nn.Linear layers:
from torchao.float8 import Float8LinearConfig, convert_to_float8_training
config = Float8LinearConfig()
convert_to_float8_training(model, config=config)
# Now train normally — forward/backward use FP8 matmuls
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
for batch in dataloader:
loss = model(batch)
loss.backward()
optimizer.step()
Scaling strategies
FP8 has limited range, so scaling is critical:
- Per-tensor scaling (default): One scale factor per tensor. Simple but less accurate.
- Per-row scaling: One scale factor per row/column. Better accuracy, used in production.
- Delayed scaling: Use statistics from previous iteration to scale current one. Reduces overhead.
Hardware requirements
| Feature | Minimum GPU |
|---|---|
| FP8 inference | H100, L40S, MI300 |
| FP8 training | H100, MI300 |
| FP8 with per-row scaling | H100 |
8. Semi-Structured Sparsity (2:4)
NVIDIA Ampere and newer GPUs have hardware support for semi-structured sparsity: exactly 2 out of every 4 consecutive values must be zero (the "2:4" pattern).
Dense: [1.2, 0.5, 3.1, 0.8, 2.0, 1.1, 0.3, 4.2]
2:4 : [1.2, 0.0, 3.1, 0.0, 2.0, 0.0, 0.3, 4.2]
↑ ↑ ↑
zeroed out zeroed zeroed
The hardware stores only the non-zero values + a 2-bit index per group of 4, achieving ~2× compression and ~2× matmul speedup.
Applying 2:4 sparsity with torchao
from torchao.sparsity import sparsify_
from torchao.sparsity import semi_structured_sparsify
# Apply 2:4 sparsity to model weights
sparsify_(model, semi_structured_sparsify())
The sparsification process
- For each group of 4 consecutive values, keep the 2 with largest magnitude
- Zero out the other 2
- Repack into the hardware sparse format
- CUDA sparse matmul kernel handles the rest
Combining sparsity and quantization
You can combine 2:4 sparsity with quantization for compounding benefits:
# First quantize, then sparsify
quantize_(model, int8_weight_only())
sparsify_(model, semi_structured_sparsify())
# Theoretical: 2× (quantization) × 2× (sparsity) = 4× improvement
Considerations
- Accuracy: Forcing 50% of weights to zero degrades accuracy. Models may need fine-tuning.
- Hardware: Requires NVIDIA A100 or newer.
- Not all layers benefit: Small layers or layers already memory-bound may not see speedup.
- Training: For best results, use sparsity-aware training (gradually introduce the sparsity pattern during training).
9. PT2E Quantization Flow
PT2E (PyTorch 2 Export) quantization is the export-based quantization flow. It uses torch.export to capture the model graph, then applies quantization transformations. This is the path for deploying quantized models to specific hardware backends.
The PT2E pipeline
Model
│
▼
torch.export() ← Capture to ExportedProgram
│
▼
prepare_pt2e() ← Insert observers for calibration
│
▼
Calibrate (run data) ← Collect activation statistics
│
▼
convert_pt2e() ← Replace observers with quantize/dequantize ops
│
▼
Quantized Model ← Ready for backend-specific compilation
Code walkthrough
import torch
from torch.ao.quantization.quantize_pt2e import prepare_pt2e, convert_pt2e
from torch.ao.quantization.quantizer.xnnpack_quantizer import (
XNNPACKQuantizer,
get_symmetric_quantization_config,
)
# 1. Export the model
exported = torch.export.export(model, example_inputs)
# 2. Create a backend-specific quantizer
quantizer = XNNPACKQuantizer().set_global(
get_symmetric_quantization_config()
)
# 3. Prepare for calibration
prepared = prepare_pt2e(exported, quantizer)
# 4. Calibrate with representative data
with torch.no_grad():
for batch in calibration_loader:
prepared(batch)
# 5. Convert to quantized model
quantized = convert_pt2e(prepared)
Available backend quantizers
| Quantizer | Target Hardware | Typical Use |
|---|---|---|
XNNPACKQuantizer | ARM CPU (mobile) | Android/iOS inference |
X86InductorQuantizer | x86 CPU | Server-side CPU inference |
QNNPackQuantizer | ARM CPU | Legacy mobile path |
When to use PT2E vs quantize_()
quantize_() (torchao) | PT2E | |
|---|---|---|
| Ease of use | One-liner | Multi-step pipeline |
| Needs calibration | No (weight-only/dynamic) | Yes (static quant) |
| Backend-specific | No (generic) | Yes (XNNPack, x86, etc.) |
| Works with compile | Yes (primary use case) | Yes (through export) |
| Best for | GPU inference, LLM serving | Mobile/edge deployment |
10. Integration with torch.compile
The key advantage of torchao over the old torch.ao is native composability with torch.compile. When you compile a torchao-quantized model, Inductor generates fused kernels that combine dequantization with the actual computation.
Basic workflow
import torch
from torchao import quantize_
from torchao.quantization import int8_weight_only
# Step 1: Quantize
quantize_(model, int8_weight_only())
# Step 2: Compile
model = torch.compile(model, mode="max-autotune")
# Step 3: Run (first call triggers compilation)
output = model(input_tensor)
What happens under the hood
Without torch.compile:
input → dequantize(int8_weight → fp16) → matmul(input, fp16_weight) → output
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
Two separate operations, intermediate fp16 tensor allocated
With torch.compile:
input → fused_int8_matmul(input, int8_weight, scale) → output
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
Single fused kernel, no intermediate allocation
The fusion eliminates:
- Memory allocation for the dequantized weight
- Memory bandwidth for reading/writing the intermediate tensor
- Kernel launch overhead for the separate dequantize op
Performance impact
A compiled quantized model is typically faster than either quantization or compilation alone:
Baseline (FP16, eager): 1.0×
FP16 + torch.compile: 1.5×
INT8 quantized (eager): 1.3×
INT8 quantized + compile: 2.5× ← Multiplicative gains
Best practices
- Quantize before compile:
quantize_()first, thentorch.compile() - Use
mode="max-autotune"for best performance (longer compilation time) - Warm up: First forward pass triggers compilation. Time the second pass onward.
- Dynamic shapes: Use
torch.compile(dynamic=True)if input shapes vary
11. Practical: Quantize and Benchmark
A complete workflow for quantizing and evaluating a model:
import torch
import torch.nn as nn
import time
class SimpleMLP(nn.Module):
def __init__(self, dim=4096, hidden=11008):
super().__init__()
self.gate = nn.Linear(dim, hidden, bias=False)
self.up = nn.Linear(dim, hidden, bias=False)
self.down = nn.Linear(hidden, dim, bias=False)
def forward(self, x):
return self.down(torch.nn.functional.silu(self.gate(x)) * self.up(x))
def measure_model_size(model):
"""Measure total parameter memory in MB."""
total = sum(
p.nelement() * p.element_size() for p in model.parameters()
)
return total / 1024 / 1024
def benchmark_inference(model, input_tensor, warmup=10, runs=100):
"""Benchmark inference latency."""
for _ in range(warmup):
model(input_tensor)
if torch.cuda.is_available():
torch.cuda.synchronize()
start = time.perf_counter()
for _ in range(runs):
model(input_tensor)
if torch.cuda.is_available():
torch.cuda.synchronize()
elapsed = (time.perf_counter() - start) / runs
return elapsed * 1000 # ms
# Create model
model = SimpleMLP().half().cuda() # FP16 on GPU
x = torch.randn(1, 4096, dtype=torch.float16, device="cuda")
print(f"Baseline size: {measure_model_size(model):.1f} MB")
print(f"Baseline latency: {benchmark_inference(model, x):.2f} ms")
# Quantize
from torchao import quantize_
from torchao.quantization import int8_weight_only
quantize_(model, int8_weight_only())
print(f"INT8 size: {measure_model_size(model):.1f} MB")
print(f"INT8 latency: {benchmark_inference(model, x):.2f} ms")
# Compile for maximum performance
model = torch.compile(model, mode="max-autotune")
print(f"INT8+compile latency: {benchmark_inference(model, x):.2f} ms")
12. Choosing a Quantization Strategy
Use this decision tree to pick the right method for your workload:
What's your goal?
│
┌────────────┴────────────┐
▼ ▼
Inference Training
│ │
┌─────────┴──────────┐ ┌────┴────┐
▼ ▼ ▼ ▼
Memory-bound? Compute-bound? BF16 FP8
(batch=1, LLMs) (batch>1) (default) (H100+)
│ │
┌───┴───┐ INT8
▼ ▼ dynamic
INT4 INT8
(max (good
savings) balance)
Also consider:
├─ Maximum speed on H100 → FP8
├─ Sparse model → 2:4 sparsity + quantization
└─ Mobile/edge deployment → PT2E + XNNPack quantizer
Quick reference
| Scenario | Method | Memory | Speedup | Accuracy |
|---|---|---|---|---|
| LLM serving (batch=1) | int4_weight_only(group_size=128) | 4× less | 2–4× | Good |
| LLM serving (batch=1, quality) | int8_weight_only() | 2× less | 1.5–2× | Very good |
| Batch inference (batch>8) | int8_dynamic_activation_int8_weight() | 2× less | 2–3× | Good |
| H100 inference | float8_dynamic_activation_float8_weight() | 2× less | 1.5–2× | Excellent |
| H100 training (large models) | Float8Linear | 2× less | 1.3–1.5× | Excellent |
| Mobile deployment | PT2E + XNNPack | 4× less | 2–4× | Good |
| Prunable model | 2:4 sparsity | 2× less | ~2× | Varies |
Accuracy considerations
Quantization always trades some accuracy for efficiency. Guidelines:
- INT8 weight-only: Usually <0.1% accuracy loss. Safe for most models.
- INT4 weight-only (group_size=128): ~0.5–1% loss. Test on your task.
- INT4 weight-only (group_size=32): ~0.2–0.5% loss. Better accuracy, more scales.
- Dynamic INT8: ~0.1–0.5% loss. Depends on activation distribution.
- FP8: Negligible loss. Closest to original precision.
- 2:4 sparsity: 1–3% loss without fine-tuning. Fine-tuning recovers most of it.
Quantization-Aware Training (QAT)
When post-training quantization isn't accurate enough, use QAT to simulate quantization during training:
from torchao.quantization import int8_weight_only
# During training, quantization is simulated (fake quantize)
# The model learns to be robust to quantization noise
# After training, apply real quantization with quantize_()
13. Upstream Updates (June 2026)
Recent PyTorch development highlights relevant to architecture optimization and production deployment:
Gloo fault tolerance support (#187381)
The Gloo collective communication backend now supports fault-tolerant operation, enabling better recovery from node failures in distributed training. This complements torchao's distributed quantization workflows where training nodes may need to recover gracefully.
CUDAGraph execution state exposed to Python (#187740)
CUDA Graph execution state is now accessible from Python, enabling better integration between CUDA Graphs and quantized model serving. This is relevant for torchao users who combine torch.compile(mode="reduce-overhead") with quantized models for maximum inference throughput.
NativeRT selectScalarOverload fix (#187059)
TorchElastic signal-failure enrichment (#187098)
TorchElastic now provides richer signal information on training failures. When running large-scale FP8 training jobs with torchao's Float8Linear, better failure diagnostics help identify whether crashes are caused by numerical issues (FP8 overflow) vs. infrastructure problems.
MPS bucket large allocations for decode (#187441)
Memory allocation improvements for Apple MPS (Metal Performance Shaders) backend during decode operations. While torchao's primary optimization targets are CUDA, this improves the experience for development and testing on Apple Silicon.
Dynamo symbolic range propagation (#187350)
Improved symbolic shape analysis in Dynamo helps torch.compile generate better code for quantized models with dynamic shapes. torchao's tensor subclasses benefit from more precise shape tracking during compilation.
XPU device info in Inductor (#187308)
Inductor now has access to XPU (Intel GPU) device information, enabling better code generation for Intel hardware. This lays groundwork for torchao quantization support on Intel discrete GPUs.
Summary
Core concepts
| Concept | Description |
|---|---|
| Quantization | Reducing numerical precision (FP16→INT8/INT4) to save memory and increase speed |
| Weight-only | Only weights are quantized; activations stay in higher precision |
| Dynamic | Both weights and activations are quantized; activation scales computed at runtime |
| Tensor subclass | torchao's approach: quantized weights are special tensor objects that handle dequant transparently |
| Semi-structured sparsity | 2:4 pattern — hardware-accelerated on NVIDIA Ampere+ |
| PT2E | Export-based quantization for backend-specific deployment |
The three-step workflow
# 1. Choose your method
from torchao.quantization import int8_weight_only
# 2. Quantize in-place
from torchao import quantize_
quantize_(model, int8_weight_only())
# 3. Compile for maximum performance
model = torch.compile(model, mode="max-autotune")
Common pitfalls
- Forgetting to compile: torchao without
torch.compileworks but misses the fused-kernel speedup - Wrong method for workload: INT4 weight-only for batch inference (should use dynamic INT8)
- Expecting GPU speedups on CPU: torchao kernels are optimized for CUDA
- Not warming up: First inference call triggers compilation — benchmark subsequent calls
- Quantizing tiny models: Overhead of quantization may exceed savings for small models
Further Resources
- torchao GitHub — source code and documentation
- torchao tutorials — official torchao documentation
- PyTorch Quantization Docs — quantization overview
- Module 07 — Training Pipelines — training fundamentals
- Module 08 — torch.compile — compilation deep dive
- Module 29 — Mixed Precision — FP16, BF16, FP8 precision
Notebook: 31_torchao.ipynb
Module 32: Efficient Data Pipelines
Prerequisites: Module 06 — Data Loading, Module 07 — Training Pipelines
Time: ~2 hours
Files:streaming_datasets.py,performance_tuning.py
Table of Contents
- Beyond Basic DataLoader
- IterableDataset
- Memory-Mapped Files
- Efficient Tokenization for LLMs
- Multi-Worker DataLoader
- Prefetching
- pin_memory and Non-Blocking Transfer
- Custom Samplers
- Distributed Data Loading
- DataLoader Performance Profiling
- Collate Optimization
- Data Pipeline Patterns
- Upstream Updates (June 23–25, 2026)
1. Beyond Basic DataLoader
Module 06 covered the fundamentals: Dataset, DataLoader, custom collate functions, and basic samplers. That's enough for most research prototypes — datasets that fit in RAM, single-GPU training, moderate throughput requirements.
Production-scale training is different. Consider:
- TB-scale datasets that can't fit in RAM (or even on a single disk)
- Streaming data from databases, object stores, or log pipelines
- Multi-GPU training where each rank must see a disjoint slice of data
- GPU utilization — if data loading can't keep up, your expensive GPUs sit idle
This module covers the patterns and tools for building data pipelines that scale.
The Data Loading Bottleneck
In a typical training loop:
┌─────────────┐ ┌──────────────┐ ┌──────────┐
│ Load Batch │────▶│ Forward/Back │────▶│ Optimize │
│ (CPU/Disk) │ │ (GPU) │ │ (GPU) │
└─────────────┘ └──────────────┘ └──────────┘
▲ │
└────────────────────────────────────────┘
If loading takes longer than compute, the GPU blocks waiting for data. The goal is to ensure the next batch is always ready before the GPU finishes the current step.
Key Metrics
| Metric | Target | Problem if missed |
|---|---|---|
| GPU utilization | >90% | Data starvation |
| Data loading time | < compute time | GPU idle cycles |
| Memory usage | Stable over time | OOM from leaks |
| Worker utilization | Balanced | Stragglers slow everything |
2. IterableDataset
A standard Dataset (map-style) requires __getitem__ and __len__. This assumes random access and known size — assumptions that break for streaming data.
IterableDataset replaces these with a single __iter__ method:
from torch.utils.data import IterableDataset, DataLoader
class LogStreamDataset(IterableDataset):
def __init__(self, log_files):
self.log_files = log_files
def __iter__(self):
for path in self.log_files:
with open(path) as f:
for line in f:
yield self.parse(line)
def parse(self, line):
# Convert raw log line to tensor
return torch.tensor([float(x) for x in line.strip().split(',')])
When to Use IterableDataset
| Use case | Map-style | Iterable |
|---|---|---|
| Data fits in RAM | ✓ | |
| Random access needed | ✓ | |
| Streaming/infinite data | ✓ | |
| Database queries | ✓ | |
| Very large file collections | ✓ | |
Need len() for progress bars | ✓ |
Worker Splitting
With num_workers > 0, each worker gets a full copy of the IterableDataset object. Without explicit splitting, every worker yields the same data — duplicated batches:
class ShardedStreamDataset(IterableDataset):
def __init__(self, file_list):
self.file_list = file_list
def __iter__(self):
worker_info = torch.utils.data.get_worker_info()
if worker_info is None:
# Single-process loading
files = self.file_list
else:
# Split files across workers
per_worker = len(self.file_list) // worker_info.num_workers
worker_id = worker_info.id
start = worker_id * per_worker
end = start + per_worker if worker_id < worker_info.num_workers - 1 else len(self.file_list)
files = self.file_list[start:end]
for path in files:
with open(path) as f:
for line in f:
yield self.process(line)
Combining with DistributedSampler
In distributed training, you need to split across both ranks and workers:
def __iter__(self):
worker_info = torch.utils.data.get_worker_info()
rank = dist.get_rank() if dist.is_initialized() else 0
world_size = dist.get_world_size() if dist.is_initialized() else 1
# First: split by rank
rank_files = self.file_list[rank::world_size]
# Then: split by worker within this rank
if worker_info is not None:
rank_files = rank_files[worker_info.id::worker_info.num_workers]
for path in rank_files:
yield from self.read_file(path)
3. Memory-Mapped Files
For datasets too large for RAM but stored on disk, memory mapping lets the OS manage paging:
import numpy as np
import torch
# Create a memory-mapped file (data stays on disk)
data = np.memmap('data.bin', dtype=np.int32, mode='r', shape=(num_tokens,))
# Access works like a normal array — OS pages data in/out
batch = torch.from_numpy(data[start:end].copy())
How It Works
Memory mapping maps a file into virtual address space without loading it into RAM. The OS loads pages on demand and evicts them under memory pressure:
Virtual Address Space Physical RAM Disk
┌──────────────┐ ┌──────────┐ ┌──────────┐
│ Page 0 │──mapped──────│ Page 0 │◀──────│ Page 0 │
│ Page 1 │──page fault──│ │ │ Page 1 │
│ Page 2 │──mapped──────│ Page 2 │◀──────│ Page 2 │
│ ... │ │ │ │ ... │
│ Page N │ └──────────┘ │ Page N │
└──────────────┘ └──────────┘
Benefits
- Near-zero startup time — no loading delay, just mmap the file
- Automatic memory management — OS handles page eviction
- Efficient for sequential access — OS prefetches sequentially
- Shared across processes — multiple workers share the same pages
torch.from_file
PyTorch has a built-in memory-mapped tensor:
# Create a storage backed by a file
storage = torch.FloatStorage.from_file('weights.bin', shared=False, size=num_elements)
tensor = torch.tensor(storage).reshape(shape)
Best Practices
- Pre-process into binary format — tokenize text, encode images, save as contiguous binary
- Use fixed-size records — enables O(1) random access by index
- Align to page boundaries — typically 4KB on Linux
- Copy before modifying —
data[i:j].copy()avoids modifying the mmap
4. Efficient Tokenization for LLMs
Tokenizing on-the-fly during training wastes compute. The pattern for LLM training:
Step 1: Pre-tokenize Offline
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3-8b")
tokens = []
for doc in documents:
tokens.extend(tokenizer.encode(doc))
tokens.append(tokenizer.eos_token_id)
# Save as binary
arr = np.array(tokens, dtype=np.uint16)
arr.tofile('train_tokens.bin')
Step 2: Pack into Fixed-Length Chunks
Rather than padding each document to max_seq_len (wasteful), concatenate all tokens and slice into chunks:
class ChunkedTokenDataset(torch.utils.data.Dataset):
def __init__(self, token_file, chunk_size=2048):
self.data = np.memmap(token_file, dtype=np.uint16, mode='r')
self.chunk_size = chunk_size
self.n_chunks = len(self.data) // chunk_size
def __len__(self):
return self.n_chunks
def __getitem__(self, idx):
start = idx * self.chunk_size
chunk = self.data[start:start + self.chunk_size].astype(np.int64)
x = torch.from_numpy(chunk[:-1])
y = torch.from_numpy(chunk[1:])
return x, y
This packing approach wastes zero tokens — every training sample is exactly chunk_size tokens with no padding.
Document Boundaries
The above approach ignores document boundaries (attention spans across documents). For models sensitive to this, insert boundary markers or use attention masks:
# Track document boundaries within each chunk
boundaries = []
pos = 0
for doc_len in doc_lengths:
pos += doc_len + 1 # +1 for EOS
if pos >= chunk_size:
break
boundaries.append(pos)
5. Multi-Worker DataLoader
num_workers > 0 spawns separate processes that load data in parallel:
Main Process Worker Processes
┌───────────────┐ ┌────────────────────┐
│ Training Loop │ │ Worker 0: load, │
│ │◀────│ transform, send │
│ GPU compute │ ├────────────────────┤
│ │◀────│ Worker 1: load, │
│ Optimizer │ │ transform, send │
│ │ ├────────────────────┤
│ │◀────│ Worker 2: load, │
│ │ │ transform, send │
└───────────────┘ └────────────────────┘
IPC via shared memory / pipes
How Workers Operate
- Each worker is a separate process (forked or spawned)
- Workers prefetch
prefetch_factorbatches each - Data transfers via shared memory (tensors) or pipes (other Python objects)
- The main process consumes batches round-robin from workers
Common Pitfalls
RNG Seeding: By default, each worker inherits the same random seed. This means augmentations are correlated across workers:
def seed_worker(worker_id):
worker_seed = torch.initial_seed() % 2**32
np.random.seed(worker_seed)
random.seed(worker_seed)
loader = DataLoader(
dataset,
num_workers=4,
worker_init_fn=seed_worker,
generator=torch.Generator().manual_seed(42),
)
File Handle Leaks: Opening files in __init__ and forking creates shared file descriptors. Open files in __iter__ or __getitem__ instead, or use worker_init_fn to open per-worker handles.
Memory Growth: Workers that accumulate state (caches, buffers) can grow memory over time. Use persistent_workers=True to avoid restarting workers each epoch — but monitor memory:
loader = DataLoader(
dataset,
num_workers=4,
persistent_workers=True, # Workers survive across epochs
)
Choosing num_workers
Rule of thumb: start with num_workers = num_cpu_cores and benchmark. Too few workers starve the GPU; too many cause contention on I/O and CPU cache thrashing.
6. Prefetching
Each worker prefetches prefetch_factor batches ahead of consumption:
loader = DataLoader(
dataset,
num_workers=4,
prefetch_factor=2, # Default: each worker prefetches 2 batches
)
How It Works
Time ──────────────────────────────────────────▶
Worker 0: [Load B0] [Load B4] [Load B8] ...
Worker 1: [Load B1] [Load B5] [Load B9] ...
Worker 2: [Load B2] [Load B6] [Load B10] ...
Worker 3: [Load B3] [Load B7] [Load B11] ...
Queue: B0 B1 B2 B3 | B4 B5 B6 B7 | ...
▲
Main process consumes from here
With prefetch_factor=2, each of the 4 workers has 2 batches in flight, so the queue holds up to 8 ready batches.
When to Increase prefetch_factor
- Slow I/O (network storage, spinning disks) — increase to 4-8
- Variable loading time — higher prefetch smooths out spikes
- Memory pressure — lower prefetch reduces memory usage
GPU Prefetching
loader = DataLoader(dataset, pin_memory=True, num_workers=4)
for batch in loader:
# Non-blocking transfer overlaps with GPU compute
x = batch.to(device, non_blocking=True)
output = model(x)
7. pin_memory and Non-Blocking Transfer
What Is Pinned Memory?
Normal (pageable) memory can be swapped to disk by the OS. GPU transfers from pageable memory require an extra copy through a staging buffer:
Pageable Memory ──copy──▶ Pinned Buffer ──DMA──▶ GPU Memory
(CPU) (PCIe)
Pinned (page-locked) memory skips the staging copy:
Pinned Memory ──────────DMA──────────▶ GPU Memory
(direct PCIe transfer)
Using pin_memory in DataLoader
loader = DataLoader(
dataset,
batch_size=64,
num_workers=4,
pin_memory=True, # Allocate batches in pinned memory
)
for data, target in loader:
# Non-blocking: starts transfer, returns immediately
data = data.to(device, non_blocking=True)
target = target.to(device, non_blocking=True)
# GPU compute can overlap with the transfer
output = model(data)
loss = criterion(output, target)
When pin_memory Helps
It helps always for GPU training. The overhead of pinning is negligible compared to the transfer speedup. The combination of pin_memory=True + non_blocking=True enables overlap of data transfer and computation.
Caveats
- Pinned memory is a limited resource — don't pin large tensors unnecessarily
non_blocking=Truerequires a CUDA stream synchronization before using the tensor on CPU again- Custom collate functions that return non-tensor objects won't benefit from pin_memory
8. Custom Samplers
Samplers control the order indices are fed to the DataLoader.
WeightedRandomSampler (Class Imbalance)
from torch.utils.data import WeightedRandomSampler
# Class counts: [9000, 500, 500] — heavily imbalanced
class_weights = [1.0/9000, 1.0/500, 1.0/500]
sample_weights = [class_weights[label] for label in all_labels]
sampler = WeightedRandomSampler(
weights=sample_weights,
num_samples=len(all_labels),
replacement=True,
)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
Curriculum Learning Sampler
Train on easy examples first, progressively introduce harder ones:
class CurriculumSampler(torch.utils.data.Sampler):
def __init__(self, difficulties, epoch=0, total_epochs=10):
self.difficulties = difficulties
self.epoch = epoch
self.total_epochs = total_epochs
def __iter__(self):
# Fraction of data available increases with epoch
fraction = min(1.0, 0.3 + 0.7 * self.epoch / self.total_epochs)
threshold = sorted(self.difficulties)[int(len(self.difficulties) * fraction) - 1]
indices = [i for i, d in enumerate(self.difficulties) if d <= threshold]
random.shuffle(indices)
return iter(indices)
def __len__(self):
fraction = min(1.0, 0.3 + 0.7 * self.epoch / self.total_epochs)
return int(len(self.difficulties) * fraction)
def set_epoch(self, epoch):
self.epoch = epoch
Hard Example Mining Sampler
Over-sample examples with high loss:
class HardExampleSampler(torch.utils.data.Sampler):
def __init__(self, dataset_size, initial_weights=None):
self.weights = initial_weights or torch.ones(dataset_size)
def update_weights(self, indices, losses):
for idx, loss in zip(indices, losses):
self.weights[idx] = loss.item()
def __iter__(self):
probs = self.weights / self.weights.sum()
indices = torch.multinomial(probs, len(self.weights), replacement=True)
return iter(indices.tolist())
def __len__(self):
return len(self.weights)
9. Distributed Data Loading
DistributedSampler
Ensures each GPU processes a disjoint subset of the data:
from torch.utils.data import DistributedSampler
sampler = DistributedSampler(
dataset,
num_replicas=world_size,
rank=rank,
shuffle=True,
drop_last=True,
)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
for epoch in range(num_epochs):
sampler.set_epoch(epoch) # CRITICAL: different shuffle each epoch
for batch in loader:
train_step(batch)
Why set_epoch Matters
Without set_epoch(), every epoch uses the same shuffle permutation. Each rank always sees the same subset of data — effectively training on 1/world_size of the dataset:
# Without set_epoch: same permutation every epoch
# Rank 0 always sees indices [0, 3, 6, 9, ...]
# Rank 1 always sees indices [1, 4, 7, 10, ...]
# With set_epoch: different permutation each epoch
# Epoch 0 Rank 0: [5, 2, 8, 1, ...]
# Epoch 1 Rank 0: [3, 9, 0, 7, ...]
Sharded Data Files
For very large datasets, pre-shard data files so each rank reads different files:
class ShardedDataset(torch.utils.data.Dataset):
def __init__(self, shard_dir, rank, world_size):
all_shards = sorted(Path(shard_dir).glob('shard_*.bin'))
self.shards = all_shards[rank::world_size]
self.data = np.concatenate([
np.memmap(s, dtype=np.int32, mode='r') for s in self.shards
])
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return torch.tensor(self.data[idx])
This avoids the DistributedSampler overhead entirely — no cross-rank coordination needed.
10. DataLoader Performance Profiling
Detecting the Bottleneck
If GPU utilization is below 90%, data loading is likely the bottleneck. Measure it:
import time
data_times = []
compute_times = []
for batch in loader:
t0 = time.perf_counter()
x, y = batch[0].to(device), batch[1].to(device)
t1 = time.perf_counter()
output = model(x)
loss = criterion(output, y)
loss.backward()
optimizer.step()
optimizer.zero_grad()
torch.cuda.synchronize()
t2 = time.perf_counter()
data_times.append(t1 - t0)
compute_times.append(t2 - t1)
avg_data = sum(data_times) / len(data_times)
avg_compute = sum(compute_times) / len(compute_times)
print(f"Data: {avg_data*1000:.1f}ms Compute: {avg_compute*1000:.1f}ms")
print(f"Data fraction: {avg_data/(avg_data+avg_compute)*100:.1f}%")
Reading the Results
| Data % | Diagnosis | Fix |
|---|---|---|
| < 10% | Compute-bound (good) | Focus on model optimization |
| 10-30% | Mild bottleneck | More workers, prefetching |
| 30-50% | Significant bottleneck | Pre-process data, faster storage |
| > 50% | Severe bottleneck | Restructure pipeline entirely |
Solutions by Severity
- Quick wins: Increase
num_workers, setpin_memory=True, increaseprefetch_factor - Medium effort: Pre-process expensive transforms offline, cache decoded images
- Significant effort: Move data to SSD/NVMe, use memory-mapped files
- Architecture change: Pre-shard data, use streaming datasets, move to WebDataset format
11. Collate Optimization
Variable-Length Sequences
The default collate pads to max length in the batch — wasteful if lengths vary widely:
# Naive: pad everything to max_len (512)
# Batch of sequences: [3, 7, 12, 490] tokens
# Padded: [512, 512, 512, 512] — 93% padding!
Bucketed Batching
Group sequences by length to minimize padding:
class BucketBatchSampler(torch.utils.data.Sampler):
def __init__(self, lengths, batch_size, bucket_boundaries=None):
self.lengths = lengths
self.batch_size = batch_size
if bucket_boundaries is None:
bucket_boundaries = [32, 64, 128, 256, 512]
# Assign each sample to a bucket
buckets = {b: [] for b in bucket_boundaries}
for idx, length in enumerate(lengths):
for boundary in bucket_boundaries:
if length <= boundary:
buckets[boundary].append(idx)
break
# Create batches within each bucket
self.batches = []
for boundary, indices in buckets.items():
random.shuffle(indices)
for i in range(0, len(indices), batch_size):
self.batches.append(indices[i:i + batch_size])
random.shuffle(self.batches)
def __iter__(self):
return iter(self.batches)
def __len__(self):
return len(self.batches)
Dynamic Batching by Token Count
Instead of a fixed number of sequences per batch, fix the total number of tokens:
class TokenBatchSampler(torch.utils.data.Sampler):
def __init__(self, lengths, max_tokens=4096):
sorted_indices = sorted(range(len(lengths)), key=lambda i: lengths[i])
self.batches = []
current_batch = []
current_max_len = 0
for idx in sorted_indices:
new_max = max(current_max_len, lengths[idx])
if new_max * (len(current_batch) + 1) > max_tokens and current_batch:
self.batches.append(current_batch)
current_batch = [idx]
current_max_len = lengths[idx]
else:
current_batch.append(idx)
current_max_len = new_max
if current_batch:
self.batches.append(current_batch)
def __iter__(self):
random.shuffle(self.batches)
return iter(self.batches)
def __len__(self):
return len(self.batches)
Custom Collate for Packed Sequences
def packed_collate(batch):
sequences, labels = zip(*batch)
lengths = [len(s) for s in sequences]
padded = torch.nn.utils.rnn.pad_sequence(sequences, batch_first=True)
mask = torch.arange(padded.size(1)).unsqueeze(0) < torch.tensor(lengths).unsqueeze(1)
return padded, torch.stack(labels), mask
12. Data Pipeline Patterns
Pattern 1: Offline Pre-processing
Process data once, save results, load during training:
Raw Data ──[preprocess.py]──▶ Processed Files ──[DataLoader]──▶ Training
(images, text) (tensors, .bin)
When to use: Expensive transforms (tokenization, image decoding, feature extraction) that produce the same result every time.
Pattern 2: On-the-Fly Augmentation
Apply random transforms during loading:
transform = transforms.Compose([
transforms.RandomResizedCrop(224),
transforms.RandomHorizontalFlip(),
transforms.ColorJitter(0.4, 0.4, 0.4),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
When to use: Stochastic augmentations that should differ each epoch.
Pattern 3: Cached Transforms
Cache the expensive part, apply cheap augmentations on-the-fly:
class CachedDataset(torch.utils.data.Dataset):
def __init__(self, base_dataset, cache_dir):
self.base = base_dataset
self.cache_dir = Path(cache_dir)
self.cache_dir.mkdir(exist_ok=True)
def __getitem__(self, idx):
cache_path = self.cache_dir / f"{idx}.pt"
if cache_path.exists():
return torch.load(cache_path, weights_only=True)
item = self.base[idx]
torch.save(item, cache_path)
return item
def __len__(self):
return len(self.base)
Pattern 4: Multi-Stage Pipeline
┌────────┐ ┌────────┐ ┌──────────┐ ┌─────────┐
│ Load │───▶│ Decode │───▶│ Augment │───▶│ Collate │
│ (I/O) │ │ (CPU) │ │ (CPU/GPU)│ │ (CPU) │
└────────┘ └────────┘ └──────────┘ └─────────┘
Workers Workers Workers Main proc
Each stage can be parallelized independently. Workers handle load + decode + augment; the main process collates and sends to GPU.
Summary of Patterns
| Pattern | I/O cost | CPU cost | Flexibility | Memory |
|---|---|---|---|---|
| Offline pre-process | Low | Low | Low | High (disk) |
| On-the-fly | High | High | High | Low |
| Cached transforms | Low (after warmup) | Low | Medium | High (disk) |
| Multi-stage | Low | Distributed | High | Medium |
13. Upstream Updates (June 23-25, 2026)
Recent PyTorch commits relevant to data pipelines and general training infrastructure:
CUDAGraph Multiple Pools in Single Graph (#187929)
Support for using multiple memory pools within a single CUDA graph capture. Previously, all allocations within a graph had to come from a single pool. This enables more flexible memory management when composing graphs from subgraphs that use different pools.
nonstrict_trace Added to torch.compiler (#187737)
A new nonstrict_trace API in torch.compiler that provides non-strict tracing semantics. This relaxes some of the constraints of strict tracing, making it easier to trace models with data-dependent control flow while still producing exportable graphs.
SymmMem Barrier for NCCL Backend (#188051)
Symmetric memory barrier support for the NCCL communication backend. This provides fine-grained synchronization primitives for distributed training, reducing the overhead of full-barrier synchronization when only local ordering guarantees are needed.
CUTLASS SiLU Epilogue Fusion (#186197)
Fuses SiLU activation into CUTLASS GEMM epilogues, eliminating a separate kernel launch for the activation function. Particularly beneficial for LLM architectures that use SwiGLU (which contains SiLU) in their feed-forward layers.
FlexGEMM Captured Tensor Epilogue Args (#187254)
FlexGEMM now supports captured tensor arguments in epilogue computations. This allows more complex epilogue patterns (like residual additions or bias terms) to be fused directly into the GEMM kernel without additional kernel launches.
Profiler NodeTimerObserver for Per-Node Timing (#186802)
A new NodeTimerObserver for the PyTorch profiler that provides per-node timing information in the execution graph. This enables fine-grained performance analysis of individual operations, making it easier to identify bottlenecks in data processing and model execution.
ROCm Origami Enabled (#186644)
ROCm Origami optimization enabled for AMD GPUs. This provides optimized kernel implementations for common operations on ROCm, improving training throughput for AMD GPU users.
Putting It All Together
A production-grade data pipeline combines several techniques:
# 1. Pre-tokenize data offline → binary files
# 2. Memory-map the binary files
# 3. Use IterableDataset with worker splitting
# 4. pin_memory + non_blocking transfer
# 5. Profile and tune num_workers + prefetch_factor
loader = DataLoader(
dataset,
batch_size=64,
num_workers=8,
pin_memory=True,
prefetch_factor=4,
persistent_workers=True,
worker_init_fn=seed_worker,
generator=torch.Generator().manual_seed(42),
)
Checklist
- [ ] Data loading time < compute time?
- [ ] GPU utilization > 90%?
- [ ] Memory usage stable over training?
- [ ] Workers properly seeded?
- [ ] Distributed: set_epoch() called each epoch?
- [ ] Variable-length data: using bucketed batching?
- [ ] Large data: using memory-mapped files or streaming?
- [ ] Expensive transforms: pre-processed offline?
Further Resources
- PyTorch DataLoader docs — official DataLoader reference
- Module 06 — Data Loading — DataLoader fundamentals
- Module 07 — Training Pipelines — training loop patterns
- Module 10 — Distributed Training — multi-GPU training
- Module 22 — LLM Recipes — LLM-specific training patterns
Notebook: 32_data_pipelines.ipynb
Module 33: Model Interpretability with Hooks
Prerequisites: Module 04 — Neural Networks, Module 07 — Training Pipelines
Time: ~2 hours
Files:hook_techniques.py,gradcam_saliency.py
Table of Contents
- What Are Hooks?
- Forward Hooks
- Forward Pre-Hooks
- Backward Hooks
- Tensor Hooks
- Activation Extraction
- Grad-CAM (Gradient-weighted Class Activation Mapping)
- Saliency Maps
- Attention Map Extraction
- Guided Backpropagation
- Practical Tips
- Upstream Updates (June 27–29, 2026)
1. What Are Hooks?
Hooks are callbacks registered on modules or tensors that execute during forward or backward passes. They let you inspect or modify activations and gradients without changing model code.
Three types of module hooks:
| Hook Type | Registration | Signature | When It Runs |
|---|---|---|---|
| Forward hook | module.register_forward_hook(fn) | fn(module, input, output) | After forward() returns |
| Forward pre-hook | module.register_forward_pre_hook(fn) | fn(module, input) | Before forward() executes |
| Backward hook | module.register_full_backward_hook(fn) | fn(module, grad_input, grad_output) | During backward() |
Plus one tensor-level hook:
| Hook Type | Registration | Signature | When It Runs |
|---|---|---|---|
| Tensor hook | tensor.register_hook(fn) | fn(grad) | When gradient is computed for that tensor |
All registration methods return a RemovableHandle. Call handle.remove() to unregister.
Why Hooks Matter
Without hooks, inspecting intermediate activations requires modifying the model's forward() method — breaking encapsulation, cluttering code, and requiring different code paths for inference vs. debugging. Hooks decouple observation from computation:
Model Code (unchanged) Observer Code (hooks)
┌─────────────────────┐ ┌──────────────────────┐
│ class MyModel: │ │ activations = {} │
│ def forward(x): │ │ │
│ x = self.conv(x)│──hook──│ store conv output │
│ x = self.relu(x)│──hook──│ store relu output │
│ x = self.fc(x) │──hook──│ store fc output │
│ return x │ │ │
└─────────────────────┘ └──────────────────────┘
2. Forward Hooks
A forward hook runs after a module's forward() completes:
def hook_fn(module, input, output):
# module: the nn.Module instance
# input: tuple of input tensors
# output: the module's return value
print(f"{module.__class__.__name__}: output shape = {output.shape}")
handle = model.layer1.register_forward_hook(hook_fn)
output = model(x) # hook_fn called when layer1 executes
handle.remove() # always clean up
Use Cases
Extract intermediate activations:
activations = {}
def save_activation(name):
def hook(module, input, output):
activations[name] = output.detach()
return hook
model.conv1.register_forward_hook(save_activation('conv1'))
model.conv2.register_forward_hook(save_activation('conv2'))
model(x)
# activations['conv1'] and activations['conv2'] now populated
Log shapes for debugging:
def shape_hook(module, input, output):
in_shape = input[0].shape if isinstance(input, tuple) else input.shape
out_shape = output.shape if hasattr(output, 'shape') else type(output)
print(f"{module.__class__.__name__}: {in_shape} -> {out_shape}")
Modify outputs (use carefully — can break autograd during training):
def clamp_hook(module, input, output):
return torch.clamp(output, -10, 10)
The with_kwargs Parameter
PyTorch 2.x added support for keyword arguments in hooks:
def hook_with_kwargs(module, input, kwargs, output):
# kwargs is a dict of keyword arguments passed to forward()
pass
handle = model.register_forward_hook(hook_with_kwargs, with_kwargs=True)
3. Forward Pre-Hooks
Pre-hooks run before the module's forward():
def pre_hook_fn(module, input):
# input is a tuple of positional arguments to forward()
# Return None to leave input unchanged, or return modified input
print(f"Input to {module.__class__.__name__}: shape={input[0].shape}")
handle = model.layer1.register_forward_pre_hook(pre_hook_fn)
Use Cases
Input validation:
def validate_input(module, input):
x = input[0]
if torch.isnan(x).any():
raise ValueError(f"NaN detected in input to {module.__class__.__name__}")
if torch.isinf(x).any():
raise ValueError(f"Inf detected in input to {module.__class__.__name__}")
Input normalization:
def normalize_input(module, input):
x = input[0]
return (x - x.mean()) / (x.std() + 1e-8),
Shape modification:
def reshape_for_conv(module, input):
x = input[0]
if x.dim() == 3:
return x.unsqueeze(1), # Add channel dimension
4. Backward Hooks
Backward hooks run during the backward pass and provide access to gradients:
def backward_hook(module, grad_input, grad_output):
# grad_input: tuple of gradients w.r.t. module inputs
# grad_output: tuple of gradients w.r.t. module outputs
print(f"{module.__class__.__name__}: grad_output norm = {grad_output[0].norm():.4f}")
handle = model.layer1.register_full_backward_hook(backward_hook)
loss = criterion(model(x), target)
loss.backward() # backward_hook called during backward pass
handle.remove()
register_full_backward_hook vs register_backward_hook
Always use register_full_backward_hook. The older register_backward_hook has known issues with modules that have multiple inputs and is deprecated.
Gradient Monitoring
def gradient_monitor(name):
def hook(module, grad_input, grad_output):
grad = grad_output[0]
stats = {
'mean': grad.mean().item(),
'std': grad.std().item(),
'norm': grad.norm().item(),
'max': grad.abs().max().item(),
'has_nan': torch.isnan(grad).any().item(),
}
print(f"[{name}] {stats}")
return hook
Gradient Modification
Backward hooks can modify gradients by returning new values:
def clip_grad_hook(module, grad_input, grad_output):
clipped = tuple(
torch.clamp(g, -1.0, 1.0) if g is not None else None
for g in grad_input
)
return clipped
5. Tensor Hooks
Tensor hooks operate on individual tensors rather than modules. They're called when the gradient for that specific tensor is computed:
x = torch.randn(3, requires_grad=True)
def tensor_hook(grad):
print(f"Gradient for x: {grad}")
return grad * 2 # Optionally modify the gradient
handle = x.register_hook(tensor_hook)
y = (x ** 2).sum()
y.backward() # prints gradient, then doubles it
handle.remove()
Use Cases
Per-parameter gradient logging:
for name, param in model.named_parameters():
param.register_hook(
lambda grad, n=name: print(f"{n}: grad norm = {grad.norm():.4f}")
)
Per-tensor gradient clipping:
for param in model.parameters():
param.register_hook(lambda grad: torch.clamp(grad, -1.0, 1.0))
Freezing specific gradients:
# Zero out gradients for specific parameters
param.register_hook(lambda grad: torch.zeros_like(grad))
6. Activation Extraction
The most common hook pattern: extract activations from target layers without modifying the model.
FeatureExtractor Class
class FeatureExtractor:
def __init__(self, model, target_layers):
self.model = model
self.features = {}
self._handles = []
for name, module in model.named_modules():
if name in target_layers:
handle = module.register_forward_hook(self._make_hook(name))
self._handles.append(handle)
def _make_hook(self, name):
def hook(module, input, output):
self.features[name] = output.detach()
return hook
def __call__(self, x):
self.features.clear()
output = self.model(x)
return output, dict(self.features)
def close(self):
for handle in self._handles:
handle.remove()
self._handles.clear()
Usage
extractor = FeatureExtractor(model, ['layer1', 'layer2.conv1', 'layer3'])
output, features = extractor(input_tensor)
print(features['layer1'].shape) # Intermediate activations
extractor.close()
Activation Statistics
Beyond raw activations, hooks can compute statistics on-the-fly:
class ActivationStats:
def __init__(self, model):
self.stats = {}
self._handles = []
for name, module in model.named_modules():
if isinstance(module, (nn.ReLU, nn.GELU, nn.SiLU)):
handle = module.register_forward_hook(self._stats_hook(name))
self._handles.append(handle)
def _stats_hook(self, name):
def hook(module, input, output):
self.stats[name] = {
'mean': output.mean().item(),
'std': output.std().item(),
'dead_fraction': (output == 0).float().mean().item(),
'max': output.max().item(),
}
return hook
def close(self):
for h in self._handles:
h.remove()
The dead_fraction metric (fraction of ReLU outputs that are exactly zero) is particularly useful — a high dead neuron fraction suggests the learning rate is too high or the initialization is poor.
7. Grad-CAM (Gradient-weighted Class Activation Mapping)
Grad-CAM visualizes which spatial regions of an input a CNN focuses on for a particular class prediction.
Algorithm
- Run forward pass, hook the last convolutional layer to capture its output activations
A - Compute the gradient of the target class score w.r.t.
A - Global-average-pool these gradients across spatial dimensions to get weights
α - Compute weighted combination:
L = ReLU(Σ αk · Ak) - Upsample
Lto input resolution
Input Image ──▶ CNN ──▶ [Last Conv Layer] ──▶ FC ──▶ Class Score
│ │
Activations A Gradient ∂y/∂A
│ │
▼ ▼
Weighted Sum ◀── GAP(gradient) = weights α
│
ReLU + Upsample
│
Heatmap
Why It Works
The global-average-pooled gradients represent the importance of each feature map channel for the target class. Weighting the activations by these importance values and taking the ReLU (we only care about features that have a positive influence) produces a coarse localization map.
Implementation
class GradCAM:
def __init__(self, model, target_layer):
self.model = model
self.activations = None
self.gradients = None
target_layer.register_forward_hook(self._save_activation)
target_layer.register_full_backward_hook(self._save_gradient)
def _save_activation(self, module, input, output):
self.activations = output.detach()
def _save_gradient(self, module, grad_input, grad_output):
self.gradients = grad_output[0].detach()
def generate(self, input_tensor, target_class=None):
self.model.eval()
output = self.model(input_tensor)
if target_class is None:
target_class = output.argmax(dim=1)
self.model.zero_grad()
one_hot = torch.zeros_like(output)
one_hot[0, target_class] = 1.0
output.backward(gradient=one_hot)
# Global average pool gradients → channel weights
weights = self.gradients.mean(dim=(2, 3), keepdim=True)
# Weighted combination of activation maps
cam = (weights * self.activations).sum(dim=1, keepdim=True)
cam = torch.relu(cam)
# Normalize to [0, 1]
cam = cam - cam.min()
cam = cam / (cam.max() + 1e-8)
# Upsample to input size
cam = torch.nn.functional.interpolate(
cam, size=input_tensor.shape[2:], mode='bilinear', align_corners=False
)
return cam.squeeze()
8. Saliency Maps
The simplest gradient-based attribution method. Shows which input pixels most affect the model's prediction.
Algorithm
- Set
input.requires_grad_(True) - Forward pass → get class score for target class
- Backward pass → compute
∂score/∂input - Take the absolute value of the gradient
- For RGB images: take the max across channels
def saliency_map(model, input_tensor, target_class):
model.eval()
input_tensor = input_tensor.clone().requires_grad_(True)
output = model(input_tensor)
score = output[0, target_class]
score.backward()
# Absolute gradient, max across channels
saliency = input_tensor.grad.abs()
if saliency.dim() == 4:
saliency = saliency.squeeze(0).max(dim=0).values
# Normalize
saliency = (saliency - saliency.min()) / (saliency.max() - saliency.min() + 1e-8)
return saliency
Interpreting Saliency Maps
- Bright regions = pixels that strongly influence the predicted class
- A good model should highlight the object, not the background
- Saliency maps are noisy — they show pixel-level sensitivity, not necessarily semantic understanding
Limitations
- Very noisy compared to Grad-CAM
- Sensitive to input perturbations
- Doesn't capture spatial coherence
- Shows sensitivity, not necessarily relevance
9. Attention Map Extraction
For Transformer models, attention weights show which tokens attend to which. Hook into MultiheadAttention to capture them:
class AttentionExtractor:
def __init__(self, model):
self.attention_maps = {}
self._handles = []
for name, module in model.named_modules():
if isinstance(module, nn.MultiheadAttention):
handle = module.register_forward_hook(self._attn_hook(name))
self._handles.append(handle)
def _attn_hook(self, name):
def hook(module, input, output):
# MHA returns (attn_output, attn_weights)
if isinstance(output, tuple) and len(output) == 2:
self.attention_maps[name] = output[1].detach()
return hook
def close(self):
for h in self._handles:
h.remove()
When calling MultiheadAttention, pass need_weights=True (default) and average_attn_weights=False to get per-head attention weights with shape (batch, num_heads, seq_len, seq_len).
Visualizing Attention
For a sequence ["The", "cat", "sat", "on", "mat"], attention weights form a matrix showing how each token attends to every other token:
The cat sat on mat
The [ 0.1 0.3 0.2 0.1 0.3 ]
cat [ 0.2 0.1 0.4 0.1 0.2 ]
sat [ 0.1 0.4 0.1 0.3 0.1 ]
on [ 0.1 0.1 0.3 0.1 0.4 ]
mat [ 0.3 0.2 0.1 0.1 0.3 ]
Each row sums to 1.0 (softmax). High values indicate strong attention from the row token to the column token.
10. Guided Backpropagation
Standard backpropagation through ReLU gates gradients based on the forward pass mask (positive activations). Guided backpropagation additionally masks out negative gradients, producing sharper attribution maps.
Algorithm
At each ReLU during backward:
- Standard backprop: pass gradient where forward activation > 0
- Guided backprop: pass gradient where forward activation > 0 AND gradient > 0
class GuidedBackprop:
def __init__(self, model):
self.model = model
self._handles = []
for module in model.modules():
if isinstance(module, nn.ReLU):
handle = module.register_full_backward_hook(self._relu_backward_hook)
self._handles.append(handle)
def _relu_backward_hook(self, module, grad_input, grad_output):
# Only pass positive gradients
return (torch.clamp(grad_output[0], min=0.0),)
def generate(self, input_tensor, target_class):
self.model.eval()
input_tensor = input_tensor.clone().requires_grad_(True)
output = self.model(input_tensor)
self.model.zero_grad()
one_hot = torch.zeros_like(output)
one_hot[0, target_class] = 1.0
output.backward(gradient=one_hot)
guided_grads = input_tensor.grad.clone()
return guided_grads
def close(self):
for h in self._handles:
h.remove()
Comparison of Methods
| Method | Granularity | Sharpness | Speed | Complexity |
|---|---|---|---|---|
| Saliency maps | Pixel | Low | Fast | Trivial |
| Grad-CAM | Region | Medium | Fast | Low |
| Guided backprop | Pixel | High | Fast | Low |
| Guided Grad-CAM | Pixel | High | Fast | Medium |
Guided Grad-CAM combines Grad-CAM with guided backprop by element-wise multiplication:
guided_gradcam = guided_grads * F.interpolate(gradcam, size=input_size)
11. Practical Tips
Always Remove Hooks
Hooks that are not removed cause memory leaks — the hook closure holds a reference to whatever it captures:
# BAD: hook leaks if exception occurs
handle = model.layer.register_forward_hook(my_hook)
output = model(x)
handle.remove()
# GOOD: use try/finally
handle = model.layer.register_forward_hook(my_hook)
try:
output = model(x)
finally:
handle.remove()
Context Manager Pattern
Wrap hook registration in a context manager for automatic cleanup:
from contextlib import contextmanager
@contextmanager
def hook_context(module, hook_fn, hook_type='forward'):
if hook_type == 'forward':
handle = module.register_forward_hook(hook_fn)
elif hook_type == 'backward':
handle = module.register_full_backward_hook(hook_fn)
elif hook_type == 'pre':
handle = module.register_forward_pre_hook(hook_fn)
try:
yield handle
finally:
handle.remove()
with hook_context(model.layer1, my_hook):
output = model(x)
# hook automatically removed
Don't Modify Outputs During Training
Returning a new value from a forward hook replaces the module's output. This can break autograd graph construction during training:
# DANGEROUS during training — breaks gradient computation
def bad_hook(module, input, output):
return output.detach() # Detaches from autograd graph!
For training, only read from hooks. Save activations with .detach() to avoid holding the entire graph in memory, but don't replace the module output.
Hook Execution Order
Multiple hooks on the same module execute in registration order:
model.layer.register_forward_hook(hook_a) # runs first
model.layer.register_forward_hook(hook_b) # runs second
If hook_a returns a modified output, hook_b receives that modified output.
Performance Considerations
- Hooks add overhead per module per forward/backward call
.detach().cpu()in hooks moves data off GPU — useful for memory but adds transfer cost- For large-scale profiling, consider disabling hooks after collecting enough data
- Hooks are not compatible with
torch.compilein all cases — test before deploying
12. Upstream Updates (June 27–29, 2026)
Recent PyTorch commits relevant to model interpretability, hooks, and inference:
Inductor CompiledArtifact Binary Extraction (#187850)
New API for extracting compiled artifact binaries from Inductor. This enables better introspection of compiled models — inspecting the actual generated code and binaries that torch.compile produces. Useful for understanding what Inductor does under the hood and debugging compilation issues.
MPS CTC Loss Backward (#188187)
Backward pass implementation for CTC loss on Apple MPS devices. Previously, CTC loss gradients had to be computed on CPU even when the forward pass ran on MPS. This enables full MPS training for speech recognition and OCR models.
MPS BatchNorm channels_last Fix (#188371)
Fixes BatchNorm computation on MPS for tensors in channels-last memory format. The previous implementation produced incorrect results when the input tensor used torch.channels_last memory layout, which is the preferred format for CNN inference.
torch._check LiteralString Enforcement (#188274)
torch._check now enforces LiteralString type for its message argument. This prevents accidental injection of dynamic strings into check messages and ensures that constraint messages are static, making them safer for graph export and compilation.
Dynamo CPython str Semantics Fix (#187775)
Fixes string operation semantics in Dynamo tracing to match CPython behavior. Previously, certain string operations during tracing could produce incorrect results or graph breaks. This improves model traceability for code that manipulates strings in control flow.
AO control_deps Ordering Fixes
Fixes ordering issues in control_deps for the Architecture Optimization (AO) library. Ensures correct operation ordering when quantization and sparsity transforms interact with control flow, particularly important for models that use conditional computation patterns.
Putting It All Together
A typical interpretability workflow:
model = load_pretrained_model()
model.eval()
# 1. Feature extraction
extractor = FeatureExtractor(model, ['features.28', 'features.14'])
output, features = extractor(image)
# 2. Grad-CAM for spatial attribution
cam = GradCAM(model, model.features[28])
heatmap = cam.generate(image, target_class=predicted_class)
# 3. Saliency for pixel-level sensitivity
saliency = saliency_map(model, image, predicted_class)
# 4. Compare and analyze
print(f"Grad-CAM highlights: {(heatmap > 0.5).sum()} pixels")
print(f"High-saliency pixels: {(saliency > 0.5).sum()}")
extractor.close()
When to Use Each Method
| Goal | Method |
|---|---|
| "What features did the model extract?" | Activation extraction |
| "Where does the model look?" | Grad-CAM |
| "Which pixels matter most?" | Saliency maps |
| "What does the model see at each level?" | Guided backprop |
| "Which tokens attend to which?" | Attention extraction |
| "Is training stable?" | Gradient monitoring hooks |
| "Are neurons dying?" | Activation statistics |
Further Resources
- PyTorch Hooks Documentation — official API reference
- Module 04 — Neural Networks —
nn.Modulefundamentals and hooks overview - Module 07 — Training Pipelines — training loop patterns
- Module 30 — Debugging — anomaly detection and gradient debugging
- Selvaraju et al., "Grad-CAM: Visual Explanations from Deep Networks" (2017) — original Grad-CAM paper
- Springenberg et al., "Striving for Simplicity" (2015) — guided backpropagation
Notebook: 33_interpretability.ipynb
Module 34: End-to-End — Fine-Tuning an LLM
Capstone Project: Tying everything together
Prerequisites: Module 04 (Neural Networks), Module 07 (Training), Module 08 (torch.compile), Module 09 (Attention), Module 16 (Activation Checkpointing), Module 22 (LLM Recipes), Module 29 (Mixed Precision)
> > Time: ~4 hours > > Files: lora_adapter.py, finetuning_pipeline.py, evaluation_and_export.py
Table of Contents
- Why Fine-Tune?
- LoRA (Low-Rank Adaptation)
- QLoRA
- Applying LoRA to a Transformer
- Data Preparation
- Training Loop with All Best Practices
- Scaling with FSDP2
- Evaluation
- Merging and Exporting
- Complete Workflow Summary
- Hyperparameter Guide
- Upstream Updates (June 29-30, 2026)
1. Why Fine-Tune?
Pretrained large language models (LLMs) are general-purpose: they learn broad linguistic patterns from trillions of tokens of web text. However, they rarely perform optimally out of the box for specific downstream tasks — medical Q&A, code generation for a proprietary API, legal document summarization, etc.
Fine-tuning adapts a pretrained model to your task with far less data and compute than training from scratch.
Full Fine-Tuning vs Parameter-Efficient Fine-Tuning
| Aspect | Full Fine-Tuning | Parameter-Efficient (PEFT) |
|---|---|---|
| Parameters updated | All (billions) | Small subset (millions) |
| Memory | Very high (full optimizer state) | Low (only adapter state) |
| Training speed | Slow | Fast |
| Risk of forgetting | Higher | Lower |
| Multiple tasks | One model per task | One base + multiple adapters |
| GPU requirement | Multi-GPU for 7B+ | Single GPU for 7B (QLoRA) |
Full fine-tuning updates every parameter in the model. For a 7B parameter model with AdamW, that means storing 7B weights + 7B gradients + 14B optimizer states (momentum + variance) = ~42B float parameters in memory.
PEFT methods freeze the pretrained weights and only train a small number of additional parameters. The dominant PEFT method today is LoRA.
2. LoRA (Low-Rank Adaptation)
The Core Idea
Instead of updating a weight matrix W ∈ R^(d×k) directly, LoRA learns a low-rank update:
W' = W + B @ A
where:
B ∈ R^(d×r)— down-projectionA ∈ R^(r×k)— up-projectionr << min(d, k)— the rank (typically 8 or 16)
The original weight W is frozen. Only A and B are trained.
Parameter Savings
For a weight matrix of shape (d, k):
- Full fine-tuning:
d × ktrainable parameters - LoRA:
r × (d + k)trainable parameters
For a typical attention projection with d = k = 4096 and r = 16:
- Full: 16,777,216 parameters
- LoRA: 131,072 parameters → 128× reduction
Initialization
Ais initialized fromN(0, 1/r)so the initial magnitude is controlledBis initialized to zeros so thatB @ A = 0at the start — the model begins as the pretrained model
Scaling Factor
A scaling factor alpha / r is applied to the LoRA output:
output = W @ x + (B @ A @ x) * (alpha / r)
Typical alpha equals r (so scaling = 1) or 2 * r.
Implementation
class LoRALinear(nn.Module):
def __init__(self, base_linear, rank=8, alpha=16):
super().__init__()
self.base = base_linear
self.base.weight.requires_grad_(False)
if self.base.bias is not None:
self.base.bias.requires_grad_(False)
d_out, d_in = base_linear.weight.shape
self.lora_A = nn.Parameter(torch.randn(rank, d_in) / rank)
self.lora_B = nn.Parameter(torch.zeros(d_out, rank))
self.scaling = alpha / rank
def forward(self, x):
base_out = self.base(x)
lora_out = (x @ self.lora_A.T @ self.lora_B.T) * self.scaling
return base_out + lora_out
See lora_adapter.py for the complete implementation.
3. QLoRA
QLoRA combines LoRA with weight quantization to reduce memory even further:
- Quantize the base model weights to INT4 or NF4 (4-bit NormalFloat)
- Add LoRA adapters in full precision (BF16)
- Train only the adapters
Memory Comparison (7B Model)
| Method | Base Weights | Adapters | Optimizer | Total |
|---|---|---|---|---|
| Full FT (FP32) | 28 GB | — | 56 GB | ~84 GB |
| Full FT (BF16) | 14 GB | — | 28 GB | ~42 GB |
| LoRA (BF16) | 14 GB | ~50 MB | ~100 MB | ~14.2 GB |
| QLoRA (NF4) | 3.5 GB | ~50 MB | ~100 MB | ~3.7 GB |
QLoRA makes fine-tuning a 7B model possible on a single 8GB consumer GPU.
Pattern
# Pseudocode for QLoRA
model = load_pretrained("llama-7b")
model = quantize_to_4bit(model) # Base weights → NF4
model = apply_lora(model, rank=16) # Adapters in BF16
train(model) # Only adapter gradients computed
The key insight: quantization errors in the base weights are compensated by the LoRA adapters during training. The adapters learn to correct for quantization noise.
4. Applying LoRA to a Transformer
Which Layers to Adapt
Not all layers benefit equally from LoRA. Common targets:
| Layer | Benefit | Typically Adapted |
|---|---|---|
| Q projection | High | Yes |
| K projection | Medium | Yes |
| V projection | High | Yes |
| O projection | Medium | Sometimes |
| FFN up-projection | Medium | Yes |
| FFN down-projection | Medium | Yes |
| Embeddings | Low | No |
| LayerNorm/RMSNorm | Low | No |
Replacing Linear Layers
def apply_lora_to_model(model, rank=8, alpha=16, target_modules=None):
"""Replace target nn.Linear layers with LoRALinear."""
if target_modules is None:
target_modules = {"q_proj", "k_proj", "v_proj", "ffn"}
for name, module in model.named_modules():
for child_name, child in module.named_children():
if isinstance(child, nn.Linear) and child_name in target_modules:
lora_layer = LoRALinear(child, rank=rank, alpha=alpha)
setattr(module, child_name, lora_layer)
Merging Adapters Back
After training, fold the LoRA weights back into the base weights for inference with zero overhead:
def merge_lora(model):
for module in model.modules():
if isinstance(module, LoRALinear):
# W_merged = W + B @ A * scaling
module.base.weight.data += (
module.lora_B @ module.lora_A * module.scaling
)
After merging, the model is identical to a regular model — no extra inference cost.
5. Data Preparation
Instruction Tuning Format
The standard format for instruction fine-tuning:
{
"instruction": "Summarize the following text in one sentence.",
"input": "PyTorch is an open-source machine learning framework...",
"output": "PyTorch is an open-source ML framework for deep learning research and production."
}
Tokenization
def format_example(example):
prompt = f"### Instruction:\n{example['instruction']}\n"
if example.get("input"):
prompt += f"### Input:\n{example['input']}\n"
prompt += f"### Response:\n{example['output']}"
return prompt
def tokenize_and_pad(text, tokenizer, max_length=512):
tokens = tokenizer.encode(text)
tokens = tokens[:max_length]
padding = max_length - len(tokens)
input_ids = tokens + [tokenizer.pad_id] * padding
labels = tokens + [-100] * padding # -100 = ignore in loss
return input_ids, labels
Setting labels=-100 for padding tokens tells CrossEntropyLoss to ignore those positions (via its ignore_index parameter).
Train/Validation Split
from torch.utils.data import random_split
dataset = InstructionDataset(data)
train_size = int(0.9 * len(dataset))
val_size = len(dataset) - train_size
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
6. Training Loop with All Best Practices
This is where the capstone brings together techniques from across the guide:
Mixed Precision (Module 29)
from torch.amp import autocast, GradScaler
scaler = GradScaler("cuda")
with autocast("cuda", dtype=torch.bfloat16):
loss = model(input_ids, labels=labels)
BF16 is preferred for LLMs because it has the same exponent range as FP32 (no overflow risk), and modern GPUs (A100, H100) have native BF16 tensor cores.
Gradient Accumulation (Module 07)
Simulate larger batch sizes without more memory:
accumulation_steps = 4
for i, batch in enumerate(dataloader):
loss = model(**batch) / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
Activation Checkpointing (Module 16)
Trade compute for memory — recompute activations during backward instead of storing them:
from torch.utils.checkpoint import checkpoint
class CheckpointedTransformerBlock(nn.Module):
def forward(self, x):
return checkpoint(self._forward_impl, x, use_reentrant=False)
torch.compile (Module 08)
Compile the model for kernel fusion and optimization:
model = torch.compile(model)
For LoRA fine-tuning, torch.compile fuses the base linear + LoRA computation, giving 10-30% speedup.
Gradient Clipping
Prevent exploding gradients, especially important for LLMs:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
Learning Rate Schedule
Cosine warmup is standard for LLM fine-tuning:
scheduler = torch.optim.lr_scheduler.OneCycleLR(
optimizer, max_lr=2e-4, total_steps=total_steps,
pct_start=0.03, anneal_strategy="cos"
)
Checkpointing
Save only the LoRA parameters (much smaller):
def save_lora_checkpoint(model, path):
lora_state = {}
for name, param in model.named_parameters():
if param.requires_grad:
lora_state[name] = param.data
torch.save(lora_state, path)
See finetuning_pipeline.py for the complete training loop.
7. Scaling with FSDP2
For models too large for a single GPU, use FSDP2 (fully_shard) from Module 10:
from torch.distributed._composable.fsdp import fully_shard, MixedPrecisionPolicy
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
)
# Shard each transformer block
for block in model.blocks:
fully_shard(block, mp_policy=mp_policy)
fully_shard(model, mp_policy=mp_policy)
Distributed Checkpointing
from torch.distributed.checkpoint import save, load
from torch.distributed.checkpoint.state_dict import (
get_model_state_dict, get_optimizer_state_dict
)
# Save
model_state = get_model_state_dict(model)
optim_state = get_optimizer_state_dict(model, optimizer)
save({"model": model_state, "optim": optim_state}, checkpoint_dir)
# Load
load({"model": model_state, "optim": optim_state}, checkpoint_dir)
FSDP2 + LoRA
When combining FSDP2 with LoRA, only the LoRA parameters participate in gradient all-reduce. The frozen base weights are still sharded across GPUs for memory efficiency, but they don't accumulate gradients:
# Apply LoRA first, then FSDP
model = create_model()
apply_lora_to_model(model, rank=16)
# FSDP shards everything, but only LoRA params have requires_grad=True
for block in model.blocks:
fully_shard(block, mp_policy=mp_policy)
fully_shard(model, mp_policy=mp_policy)
8. Evaluation
Perplexity
Perplexity measures how well the model predicts the next token. Lower is better:
PPL = exp(average cross-entropy loss)
@torch.no_grad()
def compute_perplexity(model, dataloader):
model.eval()
total_loss = 0.0
total_tokens = 0
for batch in dataloader:
logits = model(batch["input_ids"])
shift_logits = logits[:, :-1, :].contiguous()
shift_labels = batch["labels"][:, 1:].contiguous()
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1),
ignore_index=-100,
reduction="sum",
)
total_loss += loss.item()
total_tokens += (shift_labels != -100).sum().item()
return math.exp(total_loss / total_tokens)
Generation with KV Cache (Module 22)
@torch.no_grad()
def generate(model, prompt_ids, max_new_tokens=100, temperature=0.8,
top_k=50, top_p=0.9):
model.eval()
generated = list(prompt_ids)
kv_cache = None
for _ in range(max_new_tokens):
input_ids = torch.tensor([generated[-1:]]) if kv_cache else torch.tensor([generated])
logits, kv_cache = model(input_ids, kv_cache=kv_cache)
next_logits = logits[0, -1, :] / temperature
# Top-k filtering
if top_k > 0:
topk_vals, _ = torch.topk(next_logits, top_k)
next_logits[next_logits < topk_vals[-1]] = float("-inf")
# Top-p (nucleus) filtering
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(next_logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
remove = cumulative_probs > top_p
remove[..., 1:] = remove[..., :-1].clone()
remove[..., 0] = False
next_logits[sorted_indices[remove]] = float("-inf")
probs = F.softmax(next_logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1).item()
generated.append(next_token)
return generated
Sampling Strategies
| Strategy | Description | Use Case |
|---|---|---|
Greedy (temperature=0) | Always pick highest probability | Factual Q&A |
| Temperature | Scale logits before softmax | Control randomness |
| Top-k | Keep only k highest-probability tokens | Moderate diversity |
| Top-p (nucleus) | Keep smallest set with cumulative prob >= p | Dynamic vocabulary |
9. Merging and Exporting
Step 1: Merge LoRA Weights
def merge_and_unload(model):
"""Merge LoRA weights into base model and remove adapters."""
for name, module in model.named_modules():
if isinstance(module, LoRALinear):
module.base.weight.data += (
module.lora_B @ module.lora_A * module.scaling
)
# Replace LoRALinear with the merged base Linear
parent = get_parent_module(model, name)
setattr(parent, name.split(".")[-1], module.base)
return model
Step 2: Export with torch.export
merged_model = merge_and_unload(model)
merged_model.eval()
example_input = torch.randint(0, vocab_size, (1, 128))
exported = torch.export.export(merged_model, (example_input,))
torch.export.save(exported, "finetuned_model.pt2")
Step 3: Size Comparison
# Base model (BF16): ~14 GB for 7B
# LoRA checkpoint: ~50 MB (rank=16, all attention + FFN)
# Merged model (BF16): ~14 GB (same as base, but specialized)
The LoRA checkpoint is 280× smaller than the full model — you can store hundreds of task-specific adapters alongside one base model.
See evaluation_and_export.py for the complete workflow.
10. Complete Workflow Summary
┌─────────────────┐
│ Pretrained Model │
│ (frozen W) │
└────────┬────────┘
│
▼
┌─────────────────┐
│ Add LoRA │ B ∈ R^(d×r), A ∈ R^(r×k)
│ Adapters │ Only A, B are trainable
└────────┬────────┘
│
▼
┌─────────────────┐
│ Prepare Data │ Instruction format
│ (tokenize, pad) │ Labels with -100 masking
└────────┬────────┘
│
▼
┌─────────────────────────────────────────┐
│ Training Loop │
│ ┌─────────┐ ┌──────────┐ ┌───────────┐ │
│ │ BF16 │ │ Grad │ │ Activation│ │
│ │ autocast│ │ accum │ │ ckpt │ │
│ └─────────┘ └──────────┘ └───────────┘ │
│ ┌─────────┐ ┌──────────┐ ┌───────────┐ │
│ │ compile │ │ grad │ │ cosine LR │ │
│ │ │ │ clip │ │ warmup │ │
│ └─────────┘ └──────────┘ └───────────┘ │
└────────┬────────────────────────────────┘
│
▼
┌─────────────────┐
│ Evaluate │ Perplexity
│ Generate │ Temperature, top-k, top-p
└────────┬────────┘
│
▼
┌─────────────────┐
│ Merge LoRA │ W' = W + B @ A * scaling
│ (zero overhead) │
└────────┬────────┘
│
▼
┌─────────────────┐
│ Export │ torch.export → .pt2
│ Deploy │ AOTInductor / NativeRT
└─────────────────┘
11. Hyperparameter Guide
Recommended settings by model size:
| Parameter | 1B | 7B | 13B | 70B |
|---|---|---|---|---|
| LoRA rank (r) | 8 | 16 | 16 | 32 |
| LoRA alpha | 16 | 32 | 32 | 64 |
| LoRA target layers | QKV + FFN | QKV + FFN | QKV + FFN | QKV + FFN |
| Learning rate | 3e-4 | 2e-4 | 1e-4 | 5e-5 |
| Batch size (effective) | 32 | 64 | 128 | 128 |
| Grad accumulation steps | 4 | 8 | 16 | 16 |
| Max sequence length | 512 | 1024 | 2048 | 2048 |
| Warmup ratio | 0.03 | 0.03 | 0.03 | 0.03 |
| Epochs | 3 | 3 | 2 | 1 |
| Weight decay | 0.01 | 0.01 | 0.01 | 0.01 |
| Gradient clip | 1.0 | 1.0 | 1.0 | 1.0 |
| Precision | BF16 | BF16 | BF16 | BF16 |
| Method | LoRA | LoRA/QLoRA | QLoRA | QLoRA |
| GPUs needed | 1 | 1 (QLoRA) / 2 (LoRA) | 2-4 | 4-8 |
| Memory per GPU | ~6 GB | ~8 GB (QLoRA) | ~20 GB | ~40 GB |
Tips
- Start with rank 8, increase if the model underfits
- Lower learning rates for larger models — they're more sensitive
- Cosine schedule with 3% warmup works well across scales
- Gradient clipping at 1.0 is nearly universal for LLMs
- BF16 over FP16 — no loss scaling needed, same exponent range as FP32
- Evaluate every 100-500 steps on validation set for early stopping
12. Upstream Updates (June 29-30, 2026)
Recent PyTorch changes relevant to LLM fine-tuning:
FlexAttention Blocksparse Fix (#188484)
Fixed a bug where blocksparse attention masks could produce incorrect results with certain block sizes. If you use FlexAttention with custom masks for fine-tuning, update to the latest nightly.
cuDNN Heuristic Fast Path (#187212)
New fast path for cuDNN convolution heuristics that reduces kernel selection overhead. While primarily a CNN optimization, this also benefits models that combine attention with convolutional layers (e.g., vision-language models).
CUPTI Monitor Tests (#186812)
Added comprehensive tests for CUPTI-based monitoring, improving reliability of profiling during training. Use torch.profiler with the CUPTI backend for accurate kernel-level profiling during fine-tuning.
Dynamo Literal Types (#188486)
Improved handling of literal types in TorchDynamo. This fixes graph breaks that could occur when using constant values in model definitions — relevant when torch.compile encounters LoRA scaling factors or rank constants.
CUBLASLt Tunable GEMM Headers
Updated headers for CUBLASLt tunable GEMMs, enabling better autotuning of matrix multiplication kernels. The LoRA forward pass x @ A^T @ B^T benefits from tuned GEMM kernels, especially for the non-standard shapes that LoRA introduces (tall-skinny matrices with rank << hidden_dim).
Files in This Module
| File | Description | Lines |
|---|---|---|
README.md | This guide — complete theory and workflow | 450+ |
lora_adapter.py | LoRA implementation, apply/merge, QLoRA concept | 250+ |
finetuning_pipeline.py | Mini-LLM + LoRA + full training loop | 300+ |
evaluation_and_export.py | Perplexity, generation, merge, export | 200+ |
Key Takeaways
This capstone module demonstrates how the entire PyTorch ecosystem comes together for a real-world task:
- nn.Module (Module 04) — the foundation for model definition and LoRA layers
- Training loops (Module 07) — gradient accumulation, checkpointing, scheduling
- torch.compile (Module 08) — automatic kernel fusion for faster training
- Attention (Module 09) — FlexAttention, SDPA for efficient self-attention
- Activation checkpointing (Module 16) — trade compute for memory in long sequences
- LLM building blocks (Module 22) — RoPE, KV cache, RMSNorm, SwiGLU
- Mixed precision (Module 29) — BF16 for 2× memory reduction and faster matmuls
Fine-tuning is not just about LoRA — it's about combining all these techniques into a cohesive, efficient pipeline.
Further Resources
- Hu et al., "LoRA: Low-Rank Adaptation of Large Language Models" (2021) — original LoRA paper
- Dettmers et al., "QLoRA: Efficient Finetuning of Quantized LLMs" (2023) — QLoRA paper
- PyTorch FSDP2 Tutorial — official distributed training guide
- Module 10 — Distributed Training — DDP, FSDP2, tensor parallelism
- Module 11 — Export & Deployment — torch.export, AOTInductor
Notebook: 34_llm_finetuning.ipynb
Source Files
lora_adapter.py— LoRA implementation, apply/merge, QLoRA conceptfinetuning_pipeline.py— Mini-LLM + LoRA + full training loopevaluation_and_export.py— Perplexity, generation, merge, export
Module 35: PyTorch Internals — The Dispatcher
Deep Dive: How every PyTorch operation gets routed to the right kernel
Prerequisites: Module 02 (Tensors), Module 04 (Neural Networks), Module 08 (torch.compile), Module 19 (Tensor Dispatch)
> > Time: ~3 hours > > Files: dispatch_keys.py, custom_dispatch.py
Table of Contents
- What is the Dispatcher?
- The Journey of
torch.add(x, y) - Dispatch Keys
- How Dispatch Keys Are Determined
- The Priority Chain
- Fallthrough Keys
- CompositeImplicitAutograd
- CompositeExplicitAutograd
- torch.library — Registering Custom Ops
- @custom_op — The Modern API
- Viewing Dispatch Tables
- Structured Kernels
- How torch.compile Interacts
- How Autograd Uses Dispatch
- Upstream Updates (June 30 - July 1, 2026)
1. What is the Dispatcher?
The dispatcher is the central routing mechanism of PyTorch. Every single operator call — torch.add, torch.mm, tensor.relu() — flows through it. The dispatcher examines the input tensors, determines which "features" are active (autograd? autocast? vmap?), and routes to the correct kernel implementation.
This is what makes PyTorch extensible. New backends (XPU, MPS, custom hardware), autograd, torch.compile, vmap, autocast — all work by registering dispatch keys in the dispatcher. No single subsystem needs to know about the others; they all plug into the same routing table.
┌──────────────────────────────────────────────────────┐
│ User Code │
│ z = torch.add(x, y) │
└──────────────────────┬───────────────────────────────┘
│
▼
┌──────────────────────────────────────────────────────┐
│ Dispatcher │
│ │
│ 1. Collect dispatch keys from inputs │
│ 2. Walk priority chain (high → low) │
│ 3. Find first key with registered kernel │
│ 4. Execute kernel (may redispatch to lower keys) │
└──────────────────────┬───────────────────────────────┘
│
┌────────────┼────────────┐
▼ ▼ ▼
┌─────────┐ ┌─────────┐ ┌─────────┐
│Autograd │ │Autocast │ │ CPU │
│ kernel │ │ kernel │ │ kernel │
└─────────┘ └─────────┘ └─────────┘
Why Does This Matter?
Understanding the dispatcher helps you:
- Debug why an operation behaves differently than expected
- Write custom ops that integrate cleanly with autograd, compile, etc.
- Understand performance — dispatch overhead, kernel selection
- Extend PyTorch with new backends or functional transforms
2. The Journey of torch.add(x, y)
Let's trace exactly what happens when you call torch.add(x, y) where x is a CUDA tensor with requires_grad=True:
Python: torch.add(x, y)
│
├─ 1. Python binding → C++ at::add(x, y)
│
├─ 2. Dispatcher examines tensors:
│ x.dispatch_keyset() = {CUDA, AutogradCUDA}
│ y.dispatch_keyset() = {CUDA, AutogradCUDA}
│ combined = x.keys | y.keys = {CUDA, AutogradCUDA}
│
├─ 3. Walk priority chain (highest first):
│ AutogradCUDA → has kernel? YES → execute
│
├─ 4. AutogradCUDA kernel:
│ - Save x, y for backward
│ - Create AddBackward0 node
│ - Redispatch to CUDA key (exclude AutogradCUDA)
│
├─ 5. CUDA kernel:
│ - Launch element-wise add kernel on GPU
│ - Return result tensor
│
└─ 6. Result propagates back:
- Attach grad_fn to output
- Return to Python
The key insight: the Autograd kernel doesn't compute the addition itself. It records the operation for backward, then redispatches to the actual compute backend. This separation of concerns is what makes the system composable.
3. Dispatch Keys
Each tensor carries a dispatch key set — a bitset where each bit represents a "feature" or "backend" that should handle operations on that tensor.
The Full Priority Table
| Priority | Key | Purpose |
|---|---|---|
| Highest | PythonTLSSnapshot | Thread-local state snapshot |
| High | PythonDispatcher | Python-level dispatch (torch.compile) |
| FuncTorchDynamicLayerFront | Front guard for functorch | |
| Functionalize | Convert mutations to functional ops | |
| Autocast | Mixed precision dtype casting | |
| AutogradCPU | Record op for backward (CPU tensors) | |
| AutogradCUDA | Record op for backward (CUDA tensors) | |
| AutogradMPS | Record op for backward (MPS tensors) | |
| AutogradXPU | Record op for backward (XPU tensors) | |
| ADInplaceOrView | Track in-place ops and views for autograd | |
| FuncTorchBatched | vmap batching rules | |
| FuncTorchVmapMode | vmap mode (outer) | |
| BackendSelect | Route to correct backend for factory ops | |
| Low | CPU | Actual computation on CPU |
| Low | CUDA | Actual computation on CUDA |
| Low | MPS | Actual computation on MPS |
| Low | XPU | Actual computation on XPU |
| Low | Meta | Shape/dtype computation (no data) |
| Lowest | CompositeImplicitAutograd | Default decompositions (autograd-aware) |
| Lowest | CompositeExplicitAutograd | Decompositions with explicit autograd |
Viewing Keys on a Tensor
import torch
x = torch.randn(3, 3)
print(torch._C._dispatch_keys(x))
# DispatchKeySet(CPU, AutogradCPU)
y = torch.randn(3, 3, device='cuda', requires_grad=True)
print(torch._C._dispatch_keys(y))
# DispatchKeySet(CUDA, AutogradCUDA)
m = torch.randn(3, 3, device='meta')
print(torch._C._dispatch_keys(m))
# DispatchKeySet(Meta, AutogradMeta)
4. How Dispatch Keys Are Determined
Dispatch keys come from multiple sources:
From Tensor Properties
| Property | Key Added |
|---|---|
device='cpu' | CPU |
device='cuda' | CUDA |
device='mps' | MPS |
device='meta' | Meta |
requires_grad=True | AutogradCPU/CUDA/... (matches device) |
| Is a view or in-place result | ADInplaceOrView |
From Thread-Local State
| Context | Key Added |
|---|---|
Inside torch.autocast(...) | Autocast |
Inside torch.vmap(...) | FuncTorchBatched |
Inside torch._dynamo | PythonDispatcher |
Custom TorchDispatchMode active | Python |
Key Set Computation
When an op is called with multiple tensor arguments, the dispatcher computes the union of all input key sets:
# x has keys {CPU, AutogradCPU}
# y has keys {CPU, AutogradCPU}
# Combined: {CPU, AutogradCPU}
z = torch.add(x, y) # Dispatcher uses combined key set
For factory functions (no tensor inputs), BackendSelect routes based on the device argument.
5. The Priority Chain
The dispatcher walks keys from highest priority to lowest. The first key with a registered kernel for that op wins.
┌────────────────────────────────────────┐
│ Key Set: {AutogradCUDA, Autocast, CUDA}│
└────────────────────┬───────────────────┘
│
Priority walk: │
▼
Autocast ──── has kernel? ─── YES ──→ Execute
│ │
│ (if no) │ redispatch
▼ ▼
AutogradCUDA ─ has kernel? ─── YES ──→ Execute
│ │
│ (if no) │ redispatch
▼ ▼
CUDA ────── has kernel? ─── YES ──→ Execute (final)
Redispatch
After a higher-priority kernel does its work, it redispatches to the remaining keys by excluding itself:
// Inside the Autocast kernel for add:
at::AutoDispatchBelowAutocast guard; // Excludes Autocast from key set
return at::add(self, other); // Redispatches with remaining keys
This is how features compose — Autocast casts dtypes, then Autograd records the op, then the backend computes it.
6. Fallthrough Keys
Not every dispatch key has a kernel registered for every op. When a key has no kernel, the dispatcher falls through to the next key in the priority chain.
# BackendSelect only has kernels for factory ops (torch.randn, torch.empty, etc.)
# For torch.add, BackendSelect falls through to the backend key (CPU/CUDA)
Types of Fallthrough
- No registration — key is skipped entirely
- Explicit fallthrough — kernel registered that simply redispatches
- Default/catch-all — CompositeImplicitAutograd provides fallback decompositions
Example: torch.randn
torch.randn(3, 3, device='cuda')
→ BackendSelect kernel routes to CUDA
→ CUDA kernel allocates memory + fills with random values
BackendSelect is needed here because torch.randn has no input tensors — the dispatcher can't infer the backend from inputs.
7. CompositeImplicitAutograd
Ops registered at CompositeImplicitAutograd are decomposed into other ops. The autograd graph is built from the decomposed primitives — no custom backward formula needed.
# torch.addmm decomposes into mm + add:
def addmm(input, mat1, mat2, beta=1, alpha=1):
return beta * input + alpha * (mat1 @ mat2)
Since mm and add each have their own autograd formulas, the chain rule composes them automatically.
When to Use CompositeImplicit
- The decomposition is numerically stable
- Performance of the decomposed version is acceptable
- You don't need a specialized backward pass
Implications
- Works on ALL backends without backend-specific code
- Autograd "just works" through the decomposition
- torch.compile can see through the decomposition and fuse
8. CompositeExplicitAutograd
Ops at CompositeExplicitAutograd have a custom backward formula registered separately. Used when:
- The naive decomposition is numerically unstable (e.g.,
log_softmax) - A custom backward is more memory-efficient (recompute vs store)
- The mathematical gradient simplifies significantly
# log_softmax: naive decomposition has numerical issues
# Custom backward avoids computing exp() twice and is more stable
# Naive (unstable):
def log_softmax_naive(x):
return torch.log(torch.softmax(x, dim=-1))
# Actual implementation uses log-sum-exp trick for stability
Registration Pattern
# Forward registered at CompositeExplicitAutograd (works on all backends)
# Backward registered at AutogradCPU, AutogradCUDA, etc.
# This lets the forward decompose freely while backward is specialized
9. torch.library — Registering Custom Ops
The torch.library module provides the Python API to register ops in the dispatcher.
Step 1: Define the Op Schema
from torch.library import Library, impl
lib = Library("mylib", "DEF")
lib.define("my_op(Tensor x, float scale) -> Tensor")
The schema uses PyTorch's operator schema language — it specifies argument types, return types, and optional mutability annotations.
Step 2: Register Backend Implementations
@impl(lib, "my_op", "CPU")
def my_op_cpu(x, scale):
return x * scale + x.sin()
@impl(lib, "my_op", "CUDA")
def my_op_cuda(x, scale):
# Could call a custom CUDA kernel here
return x * scale + x.sin()
Step 3: Register Meta Implementation
Meta kernels compute output shape/dtype without actual data — required for torch.compile and torch.export:
@impl(lib, "my_op", "Meta")
def my_op_meta(x, scale):
return torch.empty_like(x)
Step 4: Register Autograd
class MyOpAutograd(torch.autograd.Function):
@staticmethod
def forward(ctx, x, scale):
ctx.save_for_backward(x)
ctx.scale = scale
return torch.ops.mylib.my_op(x, scale)
@staticmethod
def backward(ctx, grad_output):
x, = ctx.saved_tensors
grad_x = grad_output * (ctx.scale + x.cos())
return grad_x, None
def my_op_autograd(x, scale):
return MyOpAutograd.apply(x, scale)
lib_autograd = Library("mylib", "IMPL")
lib_autograd.impl("my_op", my_op_autograd, "AutogradCPU")
Step 5: Use the Op
x = torch.randn(4, 4, requires_grad=True)
y = torch.ops.mylib.my_op(x, scale=2.0)
y.sum().backward() # Autograd works!
10. @custom_op — The Modern API
PyTorch 2.4+ provides a simpler decorator-based API for custom ops:
@torch.library.custom_op("mylib::fast_gelu", mutates_args=())
def fast_gelu(x: torch.Tensor) -> torch.Tensor:
return x * torch.sigmoid(1.702 * x)
Register Fake (Meta) Implementation
@fast_gelu.register_fake
def fast_gelu_fake(x):
return torch.empty_like(x)
Register Autograd
def fast_gelu_setup_context(ctx, inputs, output):
x, = inputs
ctx.save_for_backward(x)
def fast_gelu_backward(ctx, grad_output):
x, = ctx.saved_tensors
sigmoid_val = torch.sigmoid(1.702 * x)
grad = sigmoid_val + 1.702 * x * sigmoid_val * (1 - sigmoid_val)
return (grad_output * grad,)
fast_gelu.register_autograd(fast_gelu_backward, setup_context=fast_gelu_setup_context)
Advantages of @custom_op
| Feature | Library API | @custom_op |
|---|---|---|
| Boilerplate | High | Low |
| Schema inference | Manual | Automatic from type hints |
| torch.compile | Manual Meta reg | register_fake |
| Autograd | Manual Function class | register_autograd |
| Composability | Manual | Built-in |
11. Viewing Dispatch Tables
PyTorch exposes tools to inspect what kernels are registered for any operator:
Dump Full Table
print(torch._C._dispatch_dump("aten::add.Tensor"))
Output shows every dispatch key and its registered kernel:
Registered Kernels:
CompositeImplicitAutograd[alias]: ...
CPU[kernel]: at::native::add(...)
CUDA[kernel]: at::native::add_cuda(...)
Meta[kernel]: at::native::add_meta(...)
AutogradCPU[autograd]: ...
AutogradCUDA[autograd]: ...
...
Check Specific Key
# Check if a key has a kernel for an op
print(torch._C._dispatch_has_kernel_for_dispatch_key(
"aten::add.Tensor", "CPU"
)) # True
List All Registered Ops
# All ops in a namespace
ops = [op for op in dir(torch.ops.aten) if not op.startswith('_')]
print(f"aten namespace has {len(ops)} ops")
12. Structured Kernels
Structured kernels are the modern pattern for implementing ops in C++. They split the kernel into two parts:
- Meta function — computes output shape/dtype, allocates output tensor
- Impl function — fills the output tensor with computed values
// Meta function (shared across all backends)
TORCH_META_FUNC(add)(const Tensor& self, const Tensor& other, const Scalar& alpha) {
// Compute output shape via broadcasting
auto output_shape = infer_size(self.sizes(), other.sizes());
set_output_raw_strided(0, output_shape, {}, self.options());
}
// CPU implementation
TORCH_IMPL_FUNC(add_out_cpu)(const Tensor& self, const Tensor& other,
const Scalar& alpha, const Tensor& result) {
// Fill result with self + alpha * other
add_kernel(kCPU, *this); // Dispatches to vectorized CPU code
}
// CUDA implementation
TORCH_IMPL_FUNC(add_out_cuda)(const Tensor& self, const Tensor& other,
const Scalar& alpha, const Tensor& result) {
add_kernel(kCUDA, *this); // Dispatches to CUDA kernel
}
Benefits
- Shape/dtype logic written once
- Output allocation handled uniformly
- Each backend only writes the compute
- Meta backend gets the meta function for free
13. How torch.compile Interacts
torch.compile operates above the dispatcher in most cases:
┌─────────────────────────────────────────────────────┐
│ TorchDynamo (Python bytecode analysis) │
│ Intercepts Python code BEFORE it hits dispatcher │
│ Captures a graph of operations │
└──────────────────────────┬──────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────┐
│ AOTAutograd │
│ Traces through Autograd dispatch keys │
│ Produces forward + backward graphs │
└──────────────────────────┬──────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────┐
│ Inductor │
│ Generates fused kernels │
│ Bypasses dispatcher entirely at runtime │
└─────────────────────────────────────────────────────┘
Key Points
- Dynamo intercepts at the Python level — it sees
torch.addcalls before they reach C++ - AOTAutograd traces through the autograd dispatch keys to produce explicit forward/backward graphs
- Inductor generates code that calls compute kernels directly — no dispatch overhead at runtime
- Custom ops with
register_fakework seamlessly — Dynamo uses the fake implementation for tracing
Dispatch Keys Relevant to Compile
PythonDispatcher— active when Dynamo is tracingFakeTensor— uses Meta kernels to track shapes during tracingProxyTorchDispatchMode— captures ops as a graph during AOTAutograd
14. How Autograd Uses Dispatch
The Autograd dispatch keys (AutogradCPU, AutogradCUDA, etc.) implement automatic differentiation as a dispatch layer:
Forward Pass
AutogradCUDA kernel for add(x, y):
1. x and y have requires_grad=True
2. Save inputs needed for backward (none for add)
3. Create AddBackward0 node
4. Exclude AutogradCUDA from key set
5. Redispatch → CUDA kernel computes result
6. Attach grad_fn to result tensor
7. Return result
Backward Pass
When .backward() is called, the engine walks the autograd graph and calls each node's backward function. This does NOT go through the dispatcher again for the top-level backward — but the individual gradient computations (e.g., grad * weight.T inside a linear backward) do go through the dispatcher.
ADInplaceOrView
This key handles a subtle case: in-place operations and views.
x = torch.randn(3, 3, requires_grad=True)
y = x.view(9) # ADInplaceOrView tracks this
y.add_(1.0) # In-place on a view — must update version counter
x.backward(...) # Still works because of proper tracking
The ADInplaceOrView key ensures that:
- View operations record the view relationship
- In-place ops increment version counters
- Autograd can detect illegal in-place modifications
15. Upstream Updates (June 30 - July 1, 2026)
Recent PyTorch changes relevant to the dispatcher:
CPU Flash SDPA Non-Contiguous Fix (#187506)
Fixed a bug where the CPU implementation of Flash Scaled Dot Product Attention produced incorrect results for non-contiguous input tensors. The dispatcher correctly routed to the CPU SDPA kernel, but the kernel itself assumed contiguous memory layout. Now properly handles strided inputs.
DTensor linspace (#187933)
Added dispatcher registration for linspace in the DTensor subsystem. DTensor implements its own dispatch key to intercept operations and distribute them across a device mesh. This PR ensures torch.linspace works correctly in distributed tensor contexts.
c10d setSequenceNumberForGroup Deprecation (#188611)
Deprecated setSequenceNumberForGroup in favor of a new sequence tracking mechanism. This affects the distributed dispatch keys (c10d) that handle collective operations. The dispatcher routes collective ops (all_reduce, broadcast, etc.) through dedicated dispatch keys.
MPS F.linear Bias Fix (#188619)
Fixed incorrect results from F.linear on MPS backend when bias is provided. The MPS dispatch key routes to Apple Metal kernels — this fix corrects the bias addition in the MPS-specific linear kernel.
Control Collectives Removal (#188617)
Removed the experimental control collectives dispatch mechanism. This simplifies the dispatch key space by removing keys that were used for prototype distributed control flow. Demonstrates that dispatch keys can be added and removed as the system evolves.
Files in This Module
| File | Description | Lines |
|---|---|---|
README.md | This guide — dispatcher internals explained | 400+ |
dispatch_keys.py | Explore dispatch keys, priority chains, tables | 250+ |
custom_dispatch.py | Register custom ops, autograd, compile integration | 250+ |
Key Takeaways
- Every op goes through the dispatcher — it's the central nervous system of PyTorch
- Dispatch keys are a bitset on each tensor, representing active features
- Priority chain determines which kernel runs — highest priority with a registered kernel wins
- Redispatch is how features compose — Autocast → Autograd → Backend
- CompositeImplicit ops decompose into primitives — autograd comes free
- torch.library and
@custom_oplet you register ops that work with all PyTorch features - torch.compile bypasses most dispatch — it captures a graph, then generates code that calls kernels directly
- The dispatcher is extensible — new backends/features just register new keys
Understanding the dispatcher transforms PyTorch from a "magic box" into a transparent, debuggable system.
Further Resources
- PyTorch Dispatcher Deep Dive (Edward Yang) — the definitive blog post
- torch.library documentation — official custom ops guide
- Module 19 — Tensor Dispatch —
__torch_function__and__torch_dispatch__ - Module 08 — torch.compile — how the compiler interacts with dispatch
Notebook: 35_dispatcher.ipynb
Source Files
dispatch_keys.py— Explore dispatch keys, priority chains, tablescustom_dispatch.py— Register custom ops, autograd, compile integration
Module 36: Custom C++ Extensions
Notebook: 36_cpp_extensions.ipynb
Prerequisites: Module 04 — Neural Networks, Module 35 — The Dispatcher
Time: ~3 hours
Files:cpp_extension_basics.py,cuda_extension_guide.py
Table of Contents
- Why C++ Extensions?
- Two Ways to Build
- JIT Compilation with load()
- Writing a C++ Extension (CPU)
- Accessing Tensor Data in C++
- Writing a CUDA Extension
- CppExtension vs CUDAExtension in setup.py
- Integrating with Autograd
- Error Handling
- Performance Tips
- Packaging and Distribution
- Modern Alternative: triton_op
- Upstream Updates (July 1-3, 2026)
1. Why C++ Extensions?
Python is wonderful for prototyping, but sometimes it's not fast enough:
When Python is too slow:
- Custom operators with tight inner loops that Python's overhead kills
- Operations that need to iterate over individual tensor elements
- Complex control flow that can't be expressed as PyTorch ops
When you need CUDA kernels:
- Novel GPU algorithms not covered by existing PyTorch ops
- Fused operations that eliminate memory round-trips
- Hardware-specific optimizations (shared memory, warp-level primitives)
When wrapping existing libraries:
- Integrating C/C++ numerical libraries (BLAS variants, custom solvers)
- Using vendor-specific GPU libraries alongside PyTorch
- Porting existing research code to the PyTorch ecosystem
PyTorch makes this easy with torch.utils.cpp_extension, which handles:
- Compiler invocation (gcc/g++, nvcc for CUDA)
- Include path management (Python, PyTorch, pybind11 headers)
- ABI compatibility across PyTorch versions
- Caching compiled shared objects
The extension mechanism builds on pybind11 for Python-C++ bindings and integrates with PyTorch's tensor library, autograd engine, and dispatcher.
2. Two Ways to Build
Option A: JIT Compilation with load()
from torch.utils.cpp_extension import load
module = load(
name="my_extension",
sources=["my_extension.cpp"],
)
| Pros | Cons |
|---|---|
| No setup.py needed | Compiles on first import (slow) |
| Great for development/iteration | Must have compiler on target machine |
| Automatic caching | Can't pip install |
| Verbose mode for debugging | Not suitable for distribution |
Option B: Setuptools with setup.py
from setuptools import setup
from torch.utils.cpp_extension import CppExtension, BuildExtension
setup(
name="my_extension",
ext_modules=[CppExtension("my_extension", ["my_extension.cpp"])],
cmdclass={"build_ext": BuildExtension},
)
| Pros | Cons |
|---|---|
| Standard Python packaging | Requires setup.py boilerplate |
pip install . works | Must rebuild after changes |
| Can build wheels for distribution | ABI compatibility concerns |
| Integrates with conda/pip ecosystem | More complex build configuration |
Rule of thumb: Use load() during development, switch to setup.py for distribution.
3. JIT Compilation with load()
The load() function is the fastest way to get C++ code running from Python:
from torch.utils.cpp_extension import load
module = load(
name="my_ext", # Name of the compiled module
sources=["my_ext.cpp"], # Source files
verbose=True, # Print compilation commands
)
How it works
- First call: Compiles sources → shared library (
.soon Linux,.pydon Windows) - Subsequent calls: Loads cached
.sofrom~/.cache/torch_extensions/ - Recompiles only when source files change (timestamp-based)
Common parameters
module = load(
name="fused_ops",
sources=["fused_ops.cpp", "fused_ops_kernel.cu"],
extra_include_paths=["/path/to/headers"],
extra_cflags=["-O3", "-march=native"],
extra_cuda_cflags=["-O3", "--use_fast_math"],
extra_ldflags=["-L/path/to/lib", "-lmylib"],
verbose=True,
with_cuda=True, # Auto-detected from .cu files
)
load_inline() for quick experiments
For small extensions, you can skip writing files entirely:
from torch.utils.cpp_extension import load_inline
cpp_source = """
torch::Tensor my_add(torch::Tensor a, torch::Tensor b) {
return a + b;
}
"""
module = load_inline(
name="inline_ext",
cpp_sources=cpp_source,
functions=["my_add"],
verbose=True,
)
Cache management
import torch.utils.cpp_extension as ext
# Default cache location
print(ext._get_build_directory("my_ext", verbose=False))
# Clear cache to force recompilation
import shutil
shutil.rmtree(ext._get_build_directory("my_ext", verbose=False))
4. Writing a C++ Extension (CPU)
Anatomy of an extension file
// my_add.cpp
#include <torch/extension.h>
// The operation itself — uses PyTorch's C++ tensor API
torch::Tensor my_add(torch::Tensor a, torch::Tensor b) {
TORCH_CHECK(a.sizes() == b.sizes(), "Size mismatch");
return a + b;
}
// Python bindings via pybind11
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("my_add", &my_add, "Element-wise addition");
}
Key components explained
#include <torch/extension.h> — The single header that includes:
<torch/torch.h>— Full PyTorch C++ API (tensors, autograd, nn)<pybind11/pybind11.h>— Python binding utilities- Macro definitions for
TORCH_EXTENSION_NAME,TORCH_CHECK, etc.
torch::Tensor — C++ equivalent of Python's torch.Tensor. Same underlying storage, same operations:
torch::Tensor result = torch::zeros({3, 4});
result = result + 1; // Broadcasting works
auto sliced = result.index({Slice(), 0}); // Indexing
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) — Creates the Python module. TORCH_EXTENSION_NAME is automatically set to the name argument passed to load() or CppExtension().
m.def("name", &function, "docstring") — Registers a C++ function as a Python callable. Pybind11 automatically converts between Python and C++ types (tensors, scalars, strings, lists, etc.).
A more realistic example: fused linear + ReLU
#include <torch/extension.h>
torch::Tensor fused_linear_relu(
torch::Tensor input,
torch::Tensor weight,
torch::Tensor bias
) {
TORCH_CHECK(input.dim() == 2, "Input must be 2D");
TORCH_CHECK(weight.dim() == 2, "Weight must be 2D");
auto output = torch::mm(input, weight.t());
if (bias.defined()) {
output = output + bias.unsqueeze(0);
}
return torch::relu(output);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fused_linear_relu", &fused_linear_relu,
"Fused Linear + ReLU (CPU)");
}
From Python:
from torch.utils.cpp_extension import load
ext = load(name="fused_ops", sources=["fused_linear_relu.cpp"])
# Use it like any Python function
output = ext.fused_linear_relu(x, weight, bias)
5. Accessing Tensor Data in C++
Raw pointer access with data_ptr<T>()
The fastest but most dangerous — no bounds checking:
float* data = tensor.data_ptr<float>();
for (int i = 0; i < tensor.numel(); i++) {
data[i] *= 2.0f;
}
Critical: Always ensure contiguity first:
auto t = tensor.contiguous(); // Copy if not contiguous
float* data = t.data_ptr<float>();
Safe strided access with accessor<T, N>()
Handles non-contiguous tensors correctly:
// For CPU tensors — includes bounds checking in debug mode
auto accessor = tensor.accessor<float, 2>(); // 2D float tensor
for (int i = 0; i < accessor.size(0); i++) {
for (int j = 0; j < accessor.size(1); j++) {
accessor[i][j] += 1.0f;
}
}
For CUDA tensors, use packed_accessor32<T, N>() or packed_accessor64<T, N>():
// In CUDA kernel — uses 32-bit indexing (faster, limits to ~2B elements)
auto acc = tensor.packed_accessor32<float, 2, torch::RestrictPtrTraits>();
Dtype dispatch with AT_DISPATCH_FLOATING_TYPES
Your C++ code needs to handle multiple dtypes. The AT_DISPATCH_* macros generate code for each supported type:
torch::Tensor scale_tensor(torch::Tensor input, double factor) {
auto output = torch::empty_like(input);
AT_DISPATCH_FLOATING_TYPES(input.scalar_type(), "scale_tensor", [&] {
// 'scalar_t' is the C++ type (float, double)
auto inp_a = input.accessor<scalar_t, 1>();
auto out_a = output.accessor<scalar_t, 1>();
for (int64_t i = 0; i < input.size(0); i++) {
out_a[i] = inp_a[i] * static_cast<scalar_t>(factor);
}
});
return output;
}
Available dispatch macros:
| Macro | Types |
|---|---|
AT_DISPATCH_FLOATING_TYPES | float, double |
AT_DISPATCH_FLOATING_TYPES_AND_HALF | float, double, Half |
AT_DISPATCH_ALL_TYPES | all integer + float + double |
AT_DISPATCH_ALL_TYPES_AND(ScalarType::Half, ...) | all + specified extras |
AT_DISPATCH_FLOATING_AND_COMPLEX_TYPES | float, double, complex |
Shape and stride access
auto sizes = tensor.sizes(); // IntArrayRef — shape
auto strides = tensor.strides(); // IntArrayRef — strides
int64_t dim = tensor.dim(); // Number of dimensions
int64_t n = tensor.numel(); // Total elements
bool contig = tensor.is_contiguous();
auto device = tensor.device(); // Device (cpu, cuda:0, ...)
auto dtype = tensor.scalar_type();
6. Writing a CUDA Extension
CUDA extensions have two parts:
.cufile — CUDA kernels (compiled by nvcc).cppfile — Python bindings and dispatch logic (compiled by gcc)
Example: fused add + ReLU CUDA kernel
fused_add_relu_kernel.cu:
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
// CUDA kernel — runs on GPU, one thread per element
template <typename scalar_t>
__global__ void fused_add_relu_kernel(
const scalar_t* __restrict__ a,
const scalar_t* __restrict__ b,
scalar_t* __restrict__ output,
int64_t size
) {
// Global thread index
const int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < size) {
scalar_t val = a[idx] + b[idx];
output[idx] = val > 0 ? val : 0; // ReLU
}
}
// Host function — called from C++, launches kernel
torch::Tensor fused_add_relu_cuda(torch::Tensor a, torch::Tensor b) {
TORCH_CHECK(a.device().is_cuda(), "Input a must be on CUDA");
TORCH_CHECK(b.device().is_cuda(), "Input b must be on CUDA");
TORCH_CHECK(a.sizes() == b.sizes(), "Size mismatch");
auto output = torch::empty_like(a);
const int64_t size = a.numel();
// Grid/block configuration
const int threads = 256;
const int blocks = (size + threads - 1) / threads;
AT_DISPATCH_FLOATING_TYPES(a.scalar_type(), "fused_add_relu", [&] {
fused_add_relu_kernel<scalar_t><<<blocks, threads>>>(
a.data_ptr<scalar_t>(),
b.data_ptr<scalar_t>(),
output.data_ptr<scalar_t>(),
size
);
});
return output;
}
fused_add_relu.cpp:
#include <torch/extension.h>
// Forward declaration of CUDA function
torch::Tensor fused_add_relu_cuda(torch::Tensor a, torch::Tensor b);
// Dispatch based on device
torch::Tensor fused_add_relu(torch::Tensor a, torch::Tensor b) {
TORCH_CHECK(a.device() == b.device(), "Tensors must be on same device");
if (a.is_cuda()) {
return fused_add_relu_cuda(a, b);
}
// CPU fallback
return torch::relu(a + b);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fused_add_relu", &fused_add_relu, "Fused add + ReLU");
}
Key CUDA concepts
__global__ — Marks a function as a CUDA kernel (called from host, runs on device).
<<<blocks, threads>>> — Kernel launch configuration:
blocks= number of thread blocks in the gridthreads= number of threads per block (max 1024)- Total threads = blocks × threads
Thread indexing:
int idx = blockIdx.x * blockDim.x + threadIdx.x; // 1D
int row = blockIdx.y * blockDim.y + threadIdx.y; // 2D
int col = blockIdx.x * blockDim.x + threadIdx.x;
Grid sizing — ensure enough threads to cover all elements:
const int threads = 256; // Common choice (multiple of warp size 32)
const int blocks = (num_elements + threads - 1) / threads;
__restrict__ — Tells the compiler pointers don't alias (enables optimizations).
Building the CUDA extension
# JIT
module = load(
name="fused_ops",
sources=["fused_add_relu.cpp", "fused_add_relu_kernel.cu"],
verbose=True,
)
# Or in setup.py
from torch.utils.cpp_extension import CUDAExtension
ext_modules = [CUDAExtension("fused_ops", [
"fused_add_relu.cpp",
"fused_add_relu_kernel.cu",
])]
7. CppExtension vs CUDAExtension in setup.py
CPU-only extension
from setuptools import setup
from torch.utils.cpp_extension import CppExtension, BuildExtension
setup(
name="my_cpu_ext",
ext_modules=[
CppExtension(
name="my_cpu_ext",
sources=["my_ext.cpp"],
extra_compile_args=["-O3", "-march=native"],
),
],
cmdclass={"build_ext": BuildExtension},
)
CUDA extension
from setuptools import setup
from torch.utils.cpp_extension import CUDAExtension, BuildExtension
setup(
name="my_cuda_ext",
ext_modules=[
CUDAExtension(
name="my_cuda_ext",
sources=[
"my_ext.cpp", # Host code (gcc)
"my_ext_kernel.cu", # Device code (nvcc)
],
extra_compile_args={
"cxx": ["-O3"],
"nvcc": ["-O3", "--use_fast_math"],
},
),
],
cmdclass={"build_ext": BuildExtension},
)
Key differences
| Feature | CppExtension | CUDAExtension |
|---|---|---|
| Compiler | gcc/g++ only | gcc + nvcc |
| Source files | .cpp, .c | .cpp, .cu |
| CUDA headers | Not included | Auto-included |
| GPU support | No | Yes |
| Requires CUDA toolkit | No | Yes |
BuildExtension
BuildExtension is a custom setuptools command class that:
- Detects the C++ compiler and its capabilities
- Manages mixed compilation (C++ and CUDA)
- Handles ABI compatibility flags
- Passes the correct PyTorch include paths
Install and test:
pip install .
python -c "import my_ext; print(my_ext.my_function(torch.ones(3)))"
8. Integrating with Autograd
To support backward(), you need to write a custom autograd function in C++.
C++ autograd function
#include <torch/extension.h>
using namespace torch::autograd;
class FusedLinearReLUFunction : public Function<FusedLinearReLUFunction> {
public:
static torch::Tensor forward(
AutogradContext* ctx,
torch::Tensor input,
torch::Tensor weight,
torch::Tensor bias
) {
auto output = torch::mm(input, weight.t()) + bias;
auto relu_output = torch::relu(output);
// Save tensors needed for backward
ctx->save_for_backward({input, weight, relu_output});
return relu_output;
}
static tensor_list backward(
AutogradContext* ctx,
tensor_list grad_outputs
) {
auto saved = ctx->get_saved_variables();
auto input = saved[0];
auto weight = saved[1];
auto relu_output = saved[2];
auto grad_output = grad_outputs[0];
// ReLU backward: zero gradient where output was zero
auto grad_relu = grad_output * (relu_output > 0).to(grad_output.dtype());
// Linear backward
auto grad_input = torch::mm(grad_relu, weight);
auto grad_weight = torch::mm(grad_relu.t(), input);
auto grad_bias = grad_relu.sum(0);
return {grad_input, grad_weight, grad_bias};
}
};
torch::Tensor fused_linear_relu(
torch::Tensor input,
torch::Tensor weight,
torch::Tensor bias
) {
return FusedLinearReLUFunction::apply(input, weight, bias);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fused_linear_relu", &fused_linear_relu);
}
Using the custom op with autograd
import torch
from torch.utils.cpp_extension import load
ext = load(name="fused_ops", sources=["fused_linear_relu.cpp"])
x = torch.randn(4, 8, requires_grad=True)
w = torch.randn(16, 8, requires_grad=True)
b = torch.randn(16, requires_grad=True)
out = ext.fused_linear_relu(x, w, b)
loss = out.sum()
loss.backward() # Calls C++ backward()
print(x.grad.shape) # torch.Size([4, 8])
Registering with the dispatcher (modern approach)
For better integration with torch.compile and other transforms, register ops via torch.library:
import torch
torch.library.define("myops::fused_linear_relu", "(Tensor x, Tensor w, Tensor b) -> Tensor")
@torch.library.impl("myops::fused_linear_relu", "cpu")
def fused_linear_relu_cpu(x, w, b):
return torch.relu(x @ w.t() + b)
@torch.library.register_fake("myops::fused_linear_relu")
def fused_linear_relu_fake(x, w, b):
return x.new_empty(x.shape[0], w.shape[0])
9. Error Handling
TORCH_CHECK — user-facing errors
Use for input validation and expected error conditions:
TORCH_CHECK(input.dim() == 2,
"Expected 2D input, got ", input.dim(), "D");
TORCH_CHECK(input.device().is_cuda(),
"Input must be a CUDA tensor, got ", input.device());
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"Expected float32, got ", input.scalar_type());
TORCH_CHECK(weight.size(1) == input.size(1),
"Weight columns (", weight.size(1),
") must match input columns (", input.size(1), ")");
TORCH_CHECK throws a c10::Error which becomes a Python RuntimeError.
TORCH_INTERNAL_ASSERT — developer-facing errors
Use for invariants that should never be violated (bugs in your code):
TORCH_INTERNAL_ASSERT(output.numel() == input.numel(),
"Output size mismatch — this is a bug");
In release builds, TORCH_INTERNAL_ASSERT can be compiled away for performance.
CUDA error checking
#define CUDA_CHECK(call) \
do { \
cudaError_t err = call; \
TORCH_CHECK(err == cudaSuccess, \
"CUDA error: ", cudaGetErrorString(err)); \
} while (0)
// Usage
CUDA_CHECK(cudaMemcpy(dst, src, size, cudaMemcpyDeviceToDevice));
Best practices
- Use
TORCH_CHECKfor all user-facing validation - Include the actual values in error messages (not just "size mismatch")
- Check device, dtype, shape, and contiguity at the entry point
- Use
TORCH_INTERNAL_ASSERTsparingly for true invariants
10. Performance Tips
1. Use AT_DISPATCH_* macros for dtype support
Don't write separate functions per dtype — the dispatch macros generate templatized code:
AT_DISPATCH_FLOATING_TYPES_AND_HALF(
input.scalar_type(), "my_kernel", [&] {
my_kernel<scalar_t><<<blocks, threads>>>(
input.data_ptr<scalar_t>(),
output.data_ptr<scalar_t>(),
size
);
}
);
2. Avoid unnecessary copies
// Bad — copies data
auto input_contig = input.contiguous(); // Might copy
// Better — check first
if (!input.is_contiguous()) {
input = input.contiguous();
}
// Best for read-only access — use accessor (handles strides)
auto acc = input.accessor<float, 2>();
3. Use torch::NoGradGuard for inference-only code
torch::Tensor my_inference_op(torch::Tensor input) {
torch::NoGradGuard no_grad; // Disables autograd tracking
return torch::relu(input);
}
4. Ensure contiguity before raw pointer access
auto t = tensor.contiguous(); // MUST do this first
float* ptr = t.data_ptr<float>(); // Now safe
5. OpenMP for CPU parallelism
#include <omp.h>
AT_DISPATCH_FLOATING_TYPES(input.scalar_type(), "parallel_op", [&] {
auto data = input.data_ptr<scalar_t>();
#pragma omp parallel for
for (int64_t i = 0; i < input.numel(); i++) {
data[i] = std::exp(data[i]);
}
});
Compile with OpenMP:
load(name="ext", sources=["ext.cpp"], extra_cflags=["-fopenmp"],
extra_ldflags=["-lgomp"])
6. Minimize kernel launches (CUDA)
Fuse operations into a single kernel instead of launching multiple kernels:
// Bad: 3 kernel launches
auto temp = a + b;
auto temp2 = temp * c;
auto output = torch::relu(temp2);
// Good: 1 kernel launch (fused)
fused_add_mul_relu_kernel<<<blocks, threads>>>(a, b, c, output, size);
7. Memory coalescing (CUDA)
Adjacent threads should access adjacent memory locations:
// Good: coalesced (threads access consecutive elements)
output[idx] = input[idx] * 2;
// Bad: strided access (threads access non-consecutive elements)
output[idx] = input[idx * stride] * 2;
11. Packaging and Distribution
Building wheels
# Build a wheel
python setup.py bdist_wheel
# Install the wheel
pip install dist/my_ext-0.1-cp310-cp310-linux_x86_64.whl
Conda packages
# meta.yaml
package:
name: my-pytorch-ext
version: "0.1.0"
requirements:
build:
- python
- setuptools
- pytorch
run:
- python
- pytorch
ABI compatibility
PyTorch extensions must match the ABI of the PyTorch installation:
import torch
print(torch._C._GLIBCXX_USE_CXX11_ABI) # 0 or 1
BuildExtension handles this automatically, but pre-built wheels must match:
- CXX11 ABI (most pip installs):
_GLIBCXX_USE_CXX11_ABI=1 - Pre-CXX11 ABI (some conda installs):
_GLIBCXX_USE_CXX11_ABI=0
Version compatibility
setup(
name="my_ext",
install_requires=[
"torch>=2.0",
],
# ...
)
For CUDA extensions, also check CUDA version compatibility:
import torch
print(torch.version.cuda) # e.g., "12.4"
Directory structure for a distributable extension
my_ext/
├── setup.py
├── my_ext/
│ ├── __init__.py
│ ├── _C.cpp # C++ bindings
│ ├── _C_cuda.cu # CUDA kernels (optional)
│ └── ops.py # Python wrappers
├── tests/
│ └── test_ops.py
└── README.md
The __init__.py loads the compiled module:
from torch.utils.cpp_extension import load
import os
_dir = os.path.dirname(os.path.abspath(__file__))
_C = load(
name="my_ext_C",
sources=[os.path.join(_dir, "_C.cpp")],
)
def my_op(x, y):
return _C.my_op(x, y)
12. Modern Alternative: triton_op
For many GPU kernels, Triton (see Module 25) is easier than writing CUDA C++:
| Aspect | CUDA C++ Extension | Triton Kernel |
|---|---|---|
| Language | C++/CUDA | Python |
| Compilation | nvcc (complex setup) | JIT (automatic) |
| Autotuning | Manual | Built-in @triton.autotune |
| torch.compile | Needs dispatcher registration | Native support |
| Debugging | gdb, cuda-gdb | Python debugger |
| Portability | NVIDIA only | Multi-backend (experimental) |
Registering Triton kernels as native ops
torch._native.triton provides triton_op for registering Triton kernels so they integrate with the dispatcher, autograd, and torch.compile:
import torch
import triton
import triton.language as tl
@triton.jit
def add_relu_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < n
x = tl.load(x_ptr + offs, mask=mask)
y = tl.load(y_ptr + offs, mask=mask)
out = tl.maximum(x + y, 0.0)
tl.store(out_ptr + offs, out, mask=mask)
@torch.library.custom_op("myops::add_relu", mutates_args=())
def add_relu(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
n = x.numel()
grid = lambda meta: (triton.cdiv(n, meta["BLOCK"]),)
add_relu_kernel[grid](x, y, out, n, BLOCK=1024)
return out
@add_relu.register_fake
def _(x, y):
return torch.empty_like(x)
When to choose which
Use CUDA C++ when:
- You need shared memory, warp-level primitives, or hardware intrinsics
- The algorithm requires complex thread synchronization
- You're wrapping an existing CUDA library
- Maximum performance is critical and you need full hardware control
Use Triton when:
- Writing elementwise, reduction, or matmul-like kernels
- You want autotuning without manual grid search
- Rapid iteration matters more than squeezing the last 5% of performance
- You need
torch.compileintegration
13. Upstream Updates (July 1-3, 2026)
Recent PyTorch changes relevant to C++ extensions and custom ops:
OpaqueBase → CustomClassBase rename (#188455)
The OpaqueBase class used for registering custom C++ classes has been renamed to CustomClassBase for clarity. If you have code using torch::OpaqueBase, update to torch::CustomClassBase:
// Before
class MyClass : public torch::OpaqueBase { ... };
// After
class MyClass : public torch::CustomClassBase { ... };
register_opaque_type → register_custom_class (#188456)
The function for registering custom C++ types with the dispatcher has been renamed:
// Before
torch::register_opaque_type<MyClass>("MyClass");
// After
torch::register_custom_class<MyClass>("MyClass");
triton_op schema support (#188722)
Enhanced schema validation for triton_op registrations. Triton ops now support richer type annotations in their schemas, improving integration with torch.compile and the dispatcher.
Native instrumentation module
New torch/_native/instrumentation.py provides a standardized instrumentation API for native ops, enabling profiling and tracing of custom kernels registered through torch.library.
Dynamo hasattr unification (#187226)
torch.compile (Dynamo) now handles hasattr checks uniformly across Python objects, improving compatibility when custom C++ extensions use hasattr-based feature detection in Python wrappers.
Inductor graph partition naming (#188700)
Improved naming for graph partitions in TorchInductor, making it easier to identify which partition corresponds to which custom op when debugging compiled code that includes C++ extensions.
Key Takeaways
torch.utils.cpp_extensionprovides two build paths:load()for development,setup.pyfor distributiontorch/extension.his the single include for all PyTorch C++ functionality- *
AT_DISPATCH_macros** handle dtype-generic code without manual template instantiation - CUDA extensions split into
.cpp(host) and.cu(device) files - Autograd integration works through
torch::autograd::Functionin C++ ortorch.libraryin Python TORCH_CHECKprovides clean error messages that surface as Python exceptions- Always ensure contiguity before calling
data_ptr<T>() - Triton is often simpler than CUDA C++ for standard GPU patterns — consider it first
- ABI compatibility matters for distribution — use
BuildExtensionto handle it
Understanding the dispatcher (Module 35) is essential — C++ extensions ultimately register kernels in the same dispatch system that powers all of PyTorch.
Further Resources
- PyTorch C++ Extension Tutorial — official tutorial
- pybind11 Documentation — Python-C++ binding details
- Module 25 — Triton Kernels — the modern GPU alternative
- Module 35 — The Dispatcher — how ops integrate with PyTorch internals
Notebook: 36_cpp_extensions.ipynb
Module 37: torch.export Deep Dive
Prerequisites: Module 08 — torch.compile, Module 11 — Export & Deployment
Time: ~3 hours
Files:export_advanced.py,control_flow_export.py
Table of Contents
- Beyond Basic Export
- ExportedProgram Anatomy
- Graph Signature
- torch.cond — Conditional Control Flow
- torch.while_loop — Loops in Export
- map — Applying a Function Over a Dimension
- Dynamic Shapes Advanced
- draft_export
- Pre-Dispatch vs Post-Dispatch IR
- Custom Ops in Export
- Retraceability
- Strict vs Non-Strict Export
- ExportBackwardSignature
- Serialization Format
- Practical Debugging Workflow
- Upstream Updates (July 3-6, 2026)
1. Beyond Basic Export
Module 11 covered the basics of torch.export: capturing a model into a graph, specifying dynamic shapes, and deploying via AOTInductor or ONNX. This module goes deeper.
Here we explore what's actually inside an ExportedProgram, how to handle control flow that export can't automatically trace, how to register custom ops for export, and how the two IR levels (pre-dispatch and post-dispatch) relate to each other. By the end, you'll be able to export models that most tutorials would call "unexportable."
The key insight: torch.export produces a complete, self-contained graph with no Python dependency. Unlike torch.jit.trace (which silently drops control flow) or torch.jit.script (which requires a subset of Python), torch.export is strict — if something can't be represented in the graph, it fails loudly. This is a feature, not a bug. The control flow primitives (torch.cond, torch.while_loop, torch.map) let you make branching and loops explicit so they survive export.
2. ExportedProgram Anatomy
When you call torch.export.export(model, args), you get back an ExportedProgram. It contains everything needed to run the model without the original Python source:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(10, 5)
self.register_buffer("scale", torch.tensor(2.0))
def forward(self, x):
return self.linear(x) * self.scale
model = MyModel()
ep = torch.export.export(model, (torch.randn(3, 10),))
The ExportedProgram object has these key attributes:
graph_module
The core FX graph that represents the computation. This is an fx.GraphModule whose nodes correspond to operations:
print(ep.graph_module.graph)
# Shows nodes: placeholder -> linear -> mul -> output
Each node has an op (call_function, placeholder, output, etc.), a target (the actual function being called), and args/kwargs.
graph_signature
Metadata describing what each graph input and output represents — is it a parameter, a buffer, a user input, or a gradient? This is how PyTorch knows which tensors in the flattened input list are weights vs data.
state_dict
The model's parameters and buffers, keyed by their fully qualified names (linear.weight, linear.bias, scale).
range_constraints
Constraints on symbolic integers. If you specify Dim("batch", min=1, max=128), the constraint 1 <= batch <= 128 appears here. These are checked at runtime.
module_call_graph
Records which submodules were called and in what order. Useful for understanding the call hierarchy in complex models.
constants
Non-parameter, non-buffer constants that appear in the graph (e.g., tensors created inside forward). These are lifted out and stored separately.
To inspect all of these:
print(f"Graph inputs: {len(ep.graph_signature.input_specs)}")
print(f"Graph outputs: {len(ep.graph_signature.output_specs)}")
print(f"State dict: {list(ep.state_dict.keys())}")
print(f"Constraints: {ep.range_constraints}")
print(f"Constants: {ep.constants}")
3. Graph Signature
The GraphSignature is critical for understanding how the exported graph maps to the original model. Every input and output has a spec:
InputSpec
Each graph input is categorized:
| Kind | Description |
|---|---|
InputKind.PARAMETER | Learnable parameter (e.g., linear.weight) |
InputKind.BUFFER | Registered buffer (e.g., scale) |
InputKind.CONSTANT_TENSOR | Lifted constant tensor |
InputKind.USER_INPUT | The actual user-provided data |
InputKind.TOKEN | Control flow token (for ordering side effects) |
OutputSpec
Each graph output is categorized:
| Kind | Description |
|---|---|
OutputKind.USER_OUTPUT | The actual return value |
OutputKind.LOSS_OUTPUT | Loss value (for training export) |
OutputKind.BUFFER_MUTATION | Buffer that was mutated in-place |
OutputKind.USER_INPUT_MUTATION | User input that was mutated |
OutputKind.GRADIENT_TO_PARAMETER | Gradient (training export) |
OutputKind.GRADIENT_TO_USER_INPUT | Gradient w.r.t. user input |
The ordering matters: in the flattened graph inputs, parameters come first, then buffers, then constants, then user inputs. The graph signature tells you where each section starts and ends.
for spec in ep.graph_signature.input_specs:
print(f" {spec.kind}: {spec.arg.name} -> {spec.target}")
This produces output like:
InputKind.PARAMETER: p_linear_weight -> linear.weight
InputKind.PARAMETER: p_linear_bias -> linear.bias
InputKind.BUFFER: b_scale -> scale
InputKind.USER_INPUT: x -> None
The target field maps back to the original state_dict key for parameters/buffers, or None for user inputs.
4. torch.cond — Conditional Control Flow
Python if/else statements are evaluated during tracing. Export only sees the branch that was taken for the example input. This means the other branch is silently dropped — a correctness bug.
torch.cond makes both branches explicit in the graph:
def f(x):
return torch.cond(
x.sum() > 0, # predicate (scalar bool tensor)
lambda x: x * 2, # true branch
lambda x: x * -1, # false branch
(x,), # operands passed to both branches
)
Rules for torch.cond
- Predicate must be a scalar boolean tensor — not a Python bool. Use tensor comparisons:
x.sum() > 0,x.shape[0] > 5won't work (that's a Python int comparison).
- Both branches must return the same structure — same number of tensors, same shapes, same dtypes. Export traces both branches and checks they match.
- Branches can't have side effects — no in-place ops on captured state, no mutation of external variables. The branches are pure functions of the operands.
- Operands are explicit — everything the branches need must be passed via the
operandstuple. Closures over external tensors are allowed but must follow the same rules.
Multiple Outputs
Branches can return tuples:
def true_fn(x, y):
return x + 1, y * 2
def false_fn(x, y):
return x - 1, y * 0.5
result_a, result_b = torch.cond(pred, true_fn, false_fn, (x, y))
Nesting
torch.cond calls can be nested — one branch can itself contain another torch.cond:
def outer_true(x):
return torch.cond(x[0] > 0, lambda x: x + 10, lambda x: x + 20, (x,))
def outer_false(x):
return x - 1
result = torch.cond(x.sum() > 0, outer_true, outer_false, (x,))
Graph Representation
In the exported graph, torch.cond appears as a higher_order_op node with two subgraph attributes — one for each branch. Both subgraphs are fully traced and available for inspection, optimization, and code generation.
5. torch.while_loop — Loops in Export
For data-dependent loops (where the number of iterations depends on runtime values), torch.while_loop provides exportable iteration:
def cond_fn(x, count):
return count < 10 # loop while this is True
def body_fn(x, count):
return x + 1, count + 1 # return updated carried inputs
init_x = torch.zeros(5)
init_count = torch.tensor(0)
result_x, result_count = torch.while_loop(cond_fn, body_fn, (init_x, init_count))
Rules
- cond_fn takes the carried inputs and returns a scalar boolean tensor.
- body_fn takes the carried inputs and returns updated values with the same structure, shapes, and dtypes.
- Carried inputs are the loop state — they're passed to both
cond_fnandbody_fnand updated each iteration. - No dynamic shape changes across iterations — the shapes of carried inputs are fixed.
Comparison with Python Loops
| Feature | Python while | torch.while_loop |
|---|---|---|
| Export | Only last iteration traced | Both branches traced |
| Iteration count | Can be dynamic | Can be dynamic |
| In graph | Unrolled (if static) or fails | Single loop node |
| Side effects | Allowed | Not allowed |
6. map — Applying a Function Over a Dimension
torch.map applies a function to each element along the first dimension of input tensors:
def double(x):
return x * 2
xs = torch.randn(5, 3) # 5 elements, each of shape (3,)
result = torch.map(double, xs) # applies double to each row
# result.shape == (5, 3)
This is useful when you want to express per-element operations that are more complex than what broadcasting handles. In the exported graph, torch.map appears as a higher-order op with a subgraph for the mapped function.
Rules
- The function is applied to slices along dimension 0.
- The function must return tensors with consistent shapes.
- Multiple input tensors can be passed — they must have the same size along dimension 0.
7. Dynamic Shapes Advanced
Module 11 introduced basic dynamic shapes. Here we cover the full Dim API and advanced constraint patterns.
The Dim API
from torch.export import Dim
batch = Dim("batch", min=1, max=128)
seq = Dim("seq", min=1, max=2048)
Dim creates a symbolic integer with optional bounds. Use it to tell export which dimensions can vary at runtime:
ep = torch.export.export(
model,
(torch.randn(4, 16),),
dynamic_shapes={"x": {0: batch, 1: seq}},
)
Shared Dims Across Inputs
When two inputs must have the same dynamic dimension, use the same Dim object:
batch = Dim("batch", min=1, max=64)
ep = torch.export.export(
model,
(torch.randn(4, 10), torch.randn(4, 10)),
dynamic_shapes={"x": {0: batch}, "y": {0: batch}},
)
This tells export that x.shape[0] == y.shape[0] at all times.
Dim.AUTO
When you don't want to manually specify every dimension, Dim.AUTO infers dynamism automatically:
ep = torch.export.export(
model,
(torch.randn(4, 10),),
dynamic_shapes={"x": {0: Dim.AUTO}},
)
Dim.AUTO examines the graph and infers appropriate constraints. It's convenient for exploratory work but may produce overly conservative or overly permissive constraints.
The dims() Helper
For models with many inputs, dims() creates multiple Dim objects at once:
from torch.export import dims
batch, seq, hidden = dims("batch", "seq", "hidden")
Runtime Assertions with torch._check
Add runtime constraints that export verifies symbolically:
def forward(self, x):
torch._check(x.shape[0] > 0)
torch._check(x.shape[1] % 2 == 0)
return self.linear(x)
These become range constraints in the exported program and are checked when the model is loaded and run with new inputs.
ShapesCollection
For complex models with many inputs, ShapesCollection provides a cleaner API:
from torch.export import ShapesCollection
shapes = ShapesCollection()
shapes[model.forward]["x"] = {0: batch}
shapes[model.forward]["y"] = {0: batch, 1: seq}
8. draft_export
When torch.export.export() fails, the error message can be cryptic. draft_export provides a more forgiving mode that returns a report explaining what went wrong:
from torch.export import draft_export
ep, report = draft_export(model, (example_input,))
Instead of raising on the first issue, draft_export continues tracing and collects all problems. The report includes:
- Missing fake implementations: Custom ops without
register_fake - Data-dependent control flow: Python
if/elsethat depends on tensor values - Unsupported Python constructs: Features that can't be represented in the graph
- Dynamic shape issues: Constraints that couldn't be satisfied
The returned ExportedProgram may be incomplete or incorrect — it's a debugging aid, not a production artifact. The workflow is:
- Try
export()— if it succeeds, you're done - If it fails, run
draft_export()to get the full picture - Fix issues one by one (add
torch.cond, register fakes, add constraints) - Try
export()again
ep, report = draft_export(model, args)
if report:
for issue in report:
print(f"Issue: {issue}")
9. Pre-Dispatch vs Post-Dispatch IR
Export produces graphs at two levels of abstraction:
Post-Dispatch (Default)
The default export() decomposes high-level ops into lower-level primitives. For example, nn.Linear becomes aten.mm + aten.add, and F.relu becomes aten.clamp_min:
ep = torch.export.export(model, args)
# Graph contains: aten.mm, aten.add, aten.clamp_min, etc.
This is closer to what hardware backends need. It's the right choice for deployment, AOTInductor, and ONNX export.
Pre-Dispatch
With pre_dispatch=True, export captures higher-level ops that are closer to the user's code:
ep = torch.export.export(model, args, pre_dispatch=True)
# Graph contains: aten.linear, aten.relu, etc.
Pre-dispatch preserves composite ops — you see linear instead of mm + add. This is useful for:
- Graph analysis: Understanding what the model does at a high level
- Custom transformations: Rewriting ops before decomposition
- Debugging: Matching graph nodes back to source code
Converting Between IRs
You can convert from pre-dispatch to post-dispatch using run_decompositions():
ep_pre = torch.export.export(model, args, pre_dispatch=True)
ep_post = ep_pre.run_decompositions()
You cannot go in the other direction — decomposition is one-way. The typical workflow is to export at pre-dispatch for analysis, then decompose for deployment.
Choosing the Right IR
| Use Case | IR Level | Why |
|---|---|---|
| AOTInductor | Post-dispatch | Backend needs decomposed ops |
| ONNX export | Post-dispatch | Maps directly to ONNX ops |
| Graph analysis | Pre-dispatch | Higher-level, easier to read |
| Custom passes | Pre-dispatch | Operate on meaningful ops |
| Training export | Pre-dispatch | Preserve autograd-relevant ops |
10. Custom Ops in Export
Custom ops (defined via torch.library) need a fake implementation (also called a Meta implementation) for export to work. Export doesn't run the real computation — it traces symbolically, so it needs to know the output shapes and dtypes without executing the kernel.
The Pattern
# Step 1: Define the op
@torch.library.custom_op("mylib::relu_squared", mutates_args=())
def relu_squared(x: torch.Tensor) -> torch.Tensor:
return torch.relu(x) ** 2
# Step 2: Register the fake implementation
@relu_squared.register_fake
def relu_squared_fake(x):
return torch.empty_like(x)
The fake implementation:
- Receives
FakeTensorinputs (tensors with shapes/dtypes but no data) - Must return tensors with the correct shapes and dtypes
- Should NOT do actual computation — just allocate empty tensors with the right metadata
Without register_fake
If you skip register_fake, export fails:
@torch.library.custom_op("mylib::bad_op", mutates_args=())
def bad_op(x: torch.Tensor) -> torch.Tensor:
return x * 2
# This will fail:
# torch.export.export(model_using_bad_op, args)
# RuntimeError: mylib::bad_op does not have a fake impl
Complex Output Shapes
When the output shape depends on input values (not just shapes), you need torch.library.FakeTensorMode:
@torch.library.custom_op("mylib::nonzero_count", mutates_args=())
def nonzero_count(x: torch.Tensor) -> torch.Tensor:
return torch.tensor(x.nonzero().shape[0])
@nonzero_count.register_fake
def nonzero_count_fake(x):
ctx = torch.library.get_ctx()
# Data-dependent output: we don't know the count at trace time
u = ctx.new_dynamic_size()
return torch.empty(u, dtype=torch.long, device=x.device)
11. Retraceability
An exported program can be re-exported (re-traced). This enables several workflows:
Adding Decompositions
Export a model, then re-export with additional decompositions to break down composite ops:
ep1 = torch.export.export(model, args)
# ep1 has high-level ops
ep2 = torch.export.export(ep1.module(), args)
# ep2 may have different decompositions
Changing Dynamic Shape Constraints
Re-export with different dynamic shapes to tighten or loosen constraints:
batch = Dim("batch", min=1, max=64)
ep1 = torch.export.export(model, args, dynamic_shapes={"x": {0: batch}})
# Later, tighten the constraint
batch_tight = Dim("batch", min=1, max=32)
ep2 = torch.export.export(ep1.module(), args, dynamic_shapes={"x": {0: batch_tight}})
Composing Exported Programs
Re-tracing lets you compose multiple exported programs into a pipeline:
class Pipeline(nn.Module):
def __init__(self, encoder_ep, decoder_ep):
super().__init__()
self.encoder = encoder_ep.module()
self.decoder = decoder_ep.module()
def forward(self, x):
z = self.encoder(x)
return self.decoder(z)
pipeline_ep = torch.export.export(Pipeline(ep_enc, ep_dec), args)
Caveats
- Re-tracing may produce different graphs if the model has Python-level logic
- Constraints from the original export are not automatically carried forward
- Custom ops need their fake implementations registered for the re-trace too
12. Strict vs Non-Strict Export
Strict Mode (Default)
strict=True uses full symbolic tracing via Dynamo. Every Python operation is symbolically evaluated:
ep = torch.export.export(model, args, strict=True)
Strict mode catches issues early:
- Data-dependent control flow → error
- Graph breaks → error
- Unsupported Python → error
This produces the most reliable exported programs.
Non-Strict Mode
strict=False is more permissive. It allows some Python constructs that strict mode rejects:
ep = torch.export.export(model, args, strict=False)
Non-strict mode:
- Uses a less aggressive tracing approach
- May allow some Python control flow (if it doesn't depend on tensor data)
- Can handle some patterns that cause graph breaks in strict mode
- May miss some issues that strict mode catches
When to Use Each
| Scenario | Mode |
|---|---|
| Production deployment | strict=True (catches all issues) |
| Iterative development | strict=False (get something working first) |
| Complex Python logic | strict=False (then migrate to strict) |
| Custom frameworks | strict=False (framework code may not trace) |
| Maximum reliability | strict=True (gold standard) |
The recommended workflow is to start with strict=False to get a working export, then switch to strict=True and fix any issues. Strict mode produces more optimizable graphs.
13. ExportBackwardSignature
For training-time export (experimental), torch.export can capture both forward and backward:
# Experimental API
ep = torch.export.export_for_training(model, args)
The backward signature specifies:
- Which parameters have gradients in the output
- The mapping from output gradients back to parameter gradients
- Loss outputs that drive the backward pass
This is used by frameworks that need to export the training loop itself, not just inference. The API is still evolving — check the PyTorch docs for the latest status.
Key considerations:
- Autograd graphs are more complex than forward-only graphs
- In-place operations and aliasing add complications
- Not all models support training export yet
14. Serialization Format
PT2 Archive Format
torch.export.save() produces a PT2 archive — a zip file containing:
model.pt2
├── data/ # Serialized tensor data
│ ├── 0 # Weight tensor 0
│ ├── 1 # Weight tensor 1
│ └── ...
├── constants/ # Non-parameter constants
├── export_graph.json # The FX graph (serialized)
├── version # Format version
└── extra/ # Sample inputs, metadata
└── sample_inputs.json
Save and Load
# Save
torch.export.save(ep, "model.pt2")
# Load
ep_loaded = torch.export.load("model.pt2")
# Run
result = ep_loaded.module()(input_tensor)
Backward Compatibility
The serialization format is versioned. PyTorch maintains backward compatibility:
- Newer PyTorch can load older archives
- Older PyTorch may not load newer archives (forward compatibility is not guaranteed)
- The graph serialization uses a stable schema based on the ATen operator set
Comparison with Other Formats
| Format | Use Case | Preserves Graph | Portable |
|---|---|---|---|
torch.save (pickle) | Checkpointing | No | No (needs source) |
torch.export.save (PT2) | Deployment | Yes | Yes |
torch.package | Hermetic archive | Source code | Yes |
| ONNX | Cross-framework | Yes | Yes |
15. Practical Debugging Workflow
When exporting a complex model, issues are common. Here's a systematic approach:
Step 1: Try Export
try:
ep = torch.export.export(model, args)
print("Export succeeded!")
except Exception as e:
print(f"Export failed: {e}")
Step 2: Use draft_export
ep, report = torch.export.draft_export(model, args)
for issue in report:
print(f"Issue: {issue}")
Step 3: Fix Issues One by One
# Before (fails)
def forward(self, x):
if x.sum() > 0: # data-dependent
return x * 2
return x * -1
# After (works)
def forward(self, x):
return torch.cond(
x.sum() > 0,
lambda x: x * 2,
lambda x: x * -1,
(x,),
)
@my_custom_op.register_fake
def my_custom_op_fake(x):
return torch.empty_like(x)
batch = Dim("batch", min=1, max=256)
ep = torch.export.export(
model, args,
dynamic_shapes={"x": {0: batch}},
)
# Try non-strict first
ep = torch.export.export(model, args, strict=False)
# Then work toward strict=True
Step 4: Validate
# Test with original inputs
original_result = model(*args)
exported_result = ep.module()(*args)
assert torch.allclose(original_result, exported_result)
# Test with different shapes (if dynamic)
new_args = (torch.randn(8, 10),) # different batch size
exported_result = ep.module()(*new_args)
Step 5: Save
torch.export.save(ep, "model.pt2")
ep_loaded = torch.export.load("model.pt2")
16. Upstream Updates (July 3-6, 2026)
Recent PyTorch changes relevant to export and compilation:
Fix Strict Export of Unregistered Parameters (#185728)
Previously, strict=True export could fail when a model referenced parameters that weren't registered via register_parameter. This fix handles the case where tensors used as parameters are regular attributes (not nn.Parameter), ensuring they're correctly captured as constants rather than causing a tracing failure.
Inductor FFT f16/bf16 Support (#180766)
TorchInductor now supports FFT operations in float16 and bfloat16. Previously, FFT ops were silently upcasted to float32, causing unexpected memory usage and performance degradation. This is relevant to export because the post-dispatch IR may contain decomposed FFT ops that Inductor needs to handle.
max_autotune Under CUDA Graph (#179246)
The max_autotune mode in Inductor now works correctly when CUDA Graphs are enabled. This fixes a class of bugs where autotuning would select a kernel configuration during tracing that was incompatible with CUDA Graph capture. For exported models deployed via AOTInductor, this means more reliable performance tuning.
Inductor TF32 Advisory Suppression (#185541)
Inductor no longer emits spurious TF32 advisory warnings during compilation. Previously, every compilation would warn about TF32 matmul precision even when the user had explicitly set the precision policy. This reduces noise in export/compile logs.
ShardedTensor Device-Agnostic Transfers (#187939)
ShardedTensor now supports device-agnostic transfers, enabling export of models that use tensor parallelism. Previously, exporting a model with sharded parameters required all shards to be on the same device type. This unblocks export workflows for distributed models.
Key Takeaways
ExportedProgramis fully self-contained — graph, weights, constraints, and metadata in one object- Graph signature tells you exactly what each input/output represents (parameter, buffer, user data)
torch.condandtorch.while_loopmake control flow explicit — the graph captures both branchesDimAPI with shared dims andDim.AUTOgives fine-grained control over dynamic shapesdraft_exportis your best debugging friend — it reports all issues instead of failing on the first- Pre-dispatch IR preserves high-level ops; post-dispatch IR decomposes for backends
- Custom ops need
register_fake— without it, export can't trace through your op - Strict mode is the gold standard — use non-strict for development, strict for production
- PT2 archive is the serialization format — versioned, portable, self-contained
Further Resources
- torch.export Documentation — official API reference
- Export Tutorial — step-by-step walkthrough
- Module 08 — torch.compile — the compilation pipeline that export feeds into
- Module 11 — Export & Deployment — basics of export and deployment paths
- Module 35 — The Dispatcher — how ops are dispatched (relevant to custom ops)
Notebook: 37_export_deep_dive.ipynb
Module 38: Compiled Autograd & AOTAutograd
Prerequisites: Module 03 — Autograd, Module 04 — Neural Networks, Module 08 — torch.compile, Module 35 — The Dispatcher
Time: ~3 hours
Files:aot_autograd_explained.py,compiled_backward.py
Table of Contents
- The Problem: Eager Backward is Slow
- AOTAutograd — Ahead-of-Time Autograd
- How AOTAutograd Works
- The Joint Graph
- min_cut_rematerialization_partition
- Compiled Autograd
- Viewing the Forward and Backward Graphs
- What AOTAutograd Enables
- Activation Memory in AOTAutograd
- AOTAutograd vs Standard Autograd
- Functorch and AOTAutograd
- Debugging AOTAutograd
- Upstream Updates (July 5-7, 2026)
1. The Problem: Eager Backward is Slow
In eager mode, PyTorch's backward pass executes operations one-by-one through the C++ autograd engine. Each op launches independently: a kernel for the matmul gradient, another for the activation gradient, another for the bias gradient, and so on. Every launch carries overhead — kernel dispatch, memory allocation, synchronization.
torch.compile solves this for the forward pass. Dynamo captures the forward graph, Inductor fuses and optimizes it, and the result runs as an efficient kernel sequence. But by default, the backward pass is still eager. The autograd engine builds the backward graph at runtime by walking the chain of grad_fn objects, and then executes each node sequentially.
This means a compiled model is only half-optimized:
Forward: [Compiled — fused, optimized, fast]
Backward: [Eager — one op at a time, dispatch overhead, no fusion]
For training workloads, the backward pass typically takes 2-3x longer than the forward (more ops, more memory traffic). Leaving it unoptimized is a major missed opportunity. AOTAutograd and Compiled Autograd fix this by bringing the backward pass into the compilation pipeline.
The fundamental tension:
- Standard autograd must be general — it handles arbitrary dynamic graphs, in-place ops, hooks, and complex control flow
- Compilation wants static graphs — known shapes, known ops, no Python callbacks
- AOTAutograd bridges this by trading generality for performance: trace the backward at compile time, producing a static graph that Inductor can optimize
2. AOTAutograd — Ahead-of-Time Autograd
AOTAutograd's core idea is simple: instead of building the backward graph at runtime, trace both forward and backward at compile time. This produces two FX graphs — a forward graph and a backward graph — and both get independently optimized by the backend (typically Inductor).
User Model → [AOTAutograd] → Forward Graph (FX) + Backward Graph (FX)
↓ ↓
[Inductor] [Inductor]
↓ ↓
Optimized Fwd Optimized Bwd
This is what happens under the hood when you call torch.compile on a model that requires gradients. Dynamo captures the forward operations, then hands them to AOTAutograd, which:
- Traces the forward to get an FX graph
- Runs
torch.autograd.grad()on the traced forward to produce the backward - Combines them into a joint graph
- Partitions the joint graph into separate forward and backward graphs
- Passes each graph to the backend compiler
The result: both forward and backward run at compiled speed, with cross-op fusion, memory planning, and kernel optimization.
AOTAutograd is not a user-facing API in most workflows — it's a component inside the torch.compile pipeline. But understanding it is essential for debugging compilation issues, understanding memory behavior, and using advanced features like custom partitioning.
3. How AOTAutograd Works
Here's the step-by-step process AOTAutograd follows:
Step 1: Functionalization
The input model may contain mutations (in-place ops like x.add_(1)) and views (ops like x.view(-1) that share memory). These are problematic for tracing because they create implicit dependencies.
Functionalization rewrites the model to eliminate mutations and views:
x.add_(1)becomesx = x + 1(out-of-place)x.view(-1)becomesx.reshape(-1)with explicit copy semantics
This produces a pure functional model — no side effects, no aliasing. Every operation takes inputs and produces new outputs.
Step 2: Trace Forward with make_fx
Using functorch's make_fx, AOTAutograd traces the functionalized forward pass. This produces an FX graph where every node is an ATen operation:
# Conceptually:
from torch.fx.experimental.proxy_tensor import make_fx
fx_forward = make_fx(functionalized_model)(*example_inputs)
The trace uses FakeTensor mode — no actual computation happens. Instead, tensor metadata (shape, dtype, device) flows through the graph, recording every operation.
Step 3: Derive Backward with torch.autograd.grad
With the traced forward graph, AOTAutograd calls torch.autograd.grad() on it to produce the backward operations. This is possible because the forward graph is itself a differentiable program — every ATen op has a registered derivative formula.
# Conceptually:
outputs = fx_forward(*inputs)
grad_inputs = torch.autograd.grad(outputs, inputs, grad_outputs)
This produces the backward ops as additional nodes in the graph.
Step 4: Build the Joint Graph
The forward ops and backward ops are combined into a single joint graph. This graph takes the original inputs and grad_outputs, and produces both the forward outputs and the gradients:
Joint Graph:
Inputs: [x, weight, bias, grad_output]
Forward ops: linear, relu, ...
Backward ops: relu_backward, linear_backward, ...
Outputs: [forward_output, grad_x, grad_weight, grad_bias]
Having everything in one graph is critical — it allows the partitioner to reason about which forward activations are needed by which backward ops.
Step 5: Partition into Forward and Backward
The joint graph is split into two separate graphs. The key decision: which intermediate tensors from the forward need to be saved for the backward?
The partitioner inserts "save" nodes at the boundary — the forward graph's extra outputs become the backward graph's extra inputs. These are the saved tensors (equivalent to what ctx.save_for_backward() does in a manual autograd.Function).
Step 6: Compile Each Graph
Both graphs are independently passed to the backend compiler (Inductor). Each gets the full optimization treatment: operator fusion, memory planning, kernel generation, and code generation.
4. The Joint Graph
Before partitioning, there's one unified graph containing both forward and backward operations. Understanding this graph is key to understanding AOTAutograd's behavior.
For a simple linear model y = relu(Wx + b):
Joint Graph:
%x : input
%weight : parameter
%bias : parameter
%grad_out : gradient of loss w.r.t. output
# Forward
%mm = aten.mm(%x, %weight.t())
%add = aten.add(%mm, %bias)
%relu = aten.relu(%add)
# Backward
%relu_bwd = aten.threshold_backward(%grad_out, %relu, 0)
%grad_b = aten.sum(%relu_bwd, dim=0)
%grad_w = aten.mm(%relu_bwd.t(), %x)
%grad_x = aten.mm(%relu_bwd, %weight)
return (%relu, %grad_x, %grad_w, %grad_b)
Notice that the backward ops reference forward tensors:
threshold_backwardneeds%relu(to know which elements were zeroed)grad_wcomputation needs%x(the original input)grad_xcomputation needs%weight
The partitioner must decide: should %relu be saved from forward, or recomputed during backward? Should %x be saved? These decisions directly control memory usage.
5. min_cut_rematerialization_partition
The default partitioner in AOTAutograd uses a min-cut algorithm to decide which activations to save vs recompute. This is the same concept as activation checkpointing, but applied automatically at the operator level.
The Tradeoff
Every forward activation used by the backward pass presents a choice:
- Save it: use memory to store it from forward until backward needs it
- Recompute it: don't save it, but recompute it during backward (uses extra FLOPs)
The min-cut partitioner formulates this as a graph cut problem:
- Nodes have costs (memory for saving, FLOPs for recomputing)
- The algorithm finds the cut that minimizes total memory while respecting a compute budget
What Gets Saved vs Recomputed
The partitioner uses heuristics about operation costs:
Typically saved (expensive to recompute):
- Matrix multiplications (
aten.mm,aten.bmm) - Convolutions (
aten.convolution) - Attention scores
- Any op with high FLOP count
Typically recomputed (cheap to recompute):
- Element-wise ops:
relu,add,mul,sigmoid - Reductions:
sum,mean - Type conversions:
to,float - Shape ops:
view,reshape,transpose
Example
For a model with y = relu(linear(x)):
Without rematerialization (save everything):
Forward saves: [mm_result, add_result, relu_result, x, weight]
Memory: 5 tensors
With min-cut rematerialization:
Forward saves: [mm_result, x, weight] # relu/add are recomputed
Memory: 3 tensors
Backward recomputes: add = mm_result + bias; relu = clamp(add, 0)
The relu and add are cheap to recompute (element-wise), so the partitioner drops them from saved tensors. The matmul result is expensive to recompute, so it's saved.
Controlling the Partitioner
You can influence partitioner behavior:
# Force specific ops to be saved (not recomputed)
torch._functorch.config.ban_recompute_ops = ["aten.mm"]
# Force specific ops to be recomputed (not saved)
torch._functorch.config.force_recompute_ops = ["aten.relu"]
For debugging:
torch._functorch.config.debug_partitioner = True
This prints which ops are saved, which are recomputed, and why.
6. Compiled Autograd
Compiled Autograd goes a step further than AOTAutograd. While AOTAutograd compiles the forward and backward graphs, the autograd engine itself is still in C++ and dispatches backward ops one-by-one. Compiled Autograd compiles the engine's execution — the entire backward pass becomes one compiled unit.
Enabling Compiled Autograd
torch._dynamo.config.compiled_autograd = True
model = torch.compile(model)
# Both forward and backward are compiled
loss = model(x).sum()
loss.backward() # The backward is ALSO compiled
What Changes
Without Compiled Autograd:
loss.backward()
→ C++ autograd engine walks grad_fn chain
→ Dispatches MulBackward0 → eager kernel
→ Dispatches AddmmBackward0 → eager kernel
→ Dispatches ReluBackward0 → eager kernel
→ ... (one dispatch per op)
With Compiled Autograd:
loss.backward()
→ Dynamo captures the entire backward execution
→ Produces one FX graph for the full backward
→ Inductor compiles it into fused kernels
→ Runs as optimized kernel sequence
The key difference from AOTAutograd: AOTAutograd traces the backward at the op level using torch.autograd.grad(). Compiled Autograd captures the actual autograd engine's execution, including any hooks, accumulation logic, and multi-output handling.
When to Use Compiled Autograd
Compiled Autograd is most beneficial when:
- The backward pass has many small operations that can be fused
- You're training on GPU and kernel launch overhead is significant
- The model structure is static (same operations every iteration)
It's less beneficial when:
- The model uses complex autograd hooks
- The backward graph changes dynamically
- You're on CPU where kernel launch overhead is minimal
Interaction with torch.compile
Compiled Autograd works in conjunction with torch.compile. When both are enabled:
torch.compilecaptures the forward via Dynamo- AOTAutograd traces forward+backward and partitions them
- Compiled Autograd captures the autograd engine's backward execution
- Inductor compiles everything
The result is that both the forward and backward passes run as optimized, fused kernel sequences.
7. Viewing the Forward and Backward Graphs
For debugging and understanding, you can inspect the graphs that AOTAutograd produces using the low-level aot_function API:
from torch._functorch.aot_autograd import aot_function
def inspect_compiler(gm, example_inputs):
print("=" * 60)
print("Graph:")
gm.graph.print_tabular()
print(f"Number of nodes: {len(list(gm.graph.nodes))}")
return gm # return the graph module as-is (no optimization)
def my_fn(x, weight):
return torch.relu(x @ weight)
compiled = aot_function(
my_fn,
fw_compiler=inspect_compiler, # called with forward graph
bw_compiler=inspect_compiler, # called with backward graph
)
x = torch.randn(4, 8, requires_grad=True)
w = torch.randn(8, 4, requires_grad=True)
out = compiled(x, w)
out.sum().backward()
This prints both the forward and backward FX graphs, showing every operation and the data flow between them.
Using TORCH_LOGS
For torch.compile, use environment variables:
# See AOTAutograd forward/backward graphs
TORCH_LOGS="aot" python train.py
# See generated Inductor code
TORCH_LOGS="output_code" python train.py
# See everything
TORCH_LOGS="aot,output_code,graph_breaks" python train.py
Programmatic Logging
import logging
torch._logging.set_logs(aot=logging.DEBUG)
This produces verbose output showing the joint graph, the partitioning decisions, and the final forward/backward graphs.
8. What AOTAutograd Enables
By compiling the backward pass alongside the forward, AOTAutograd unlocks optimizations that are impossible in eager mode:
Operator Fusion Across Forward/Backward Boundary
In eager mode, the forward and backward are separate execution phases. The compiler can't fuse a forward activation with its corresponding backward derivative. With AOTAutograd, both are visible in a single compilation unit, enabling cross-boundary fusion.
Dead Code Elimination in Backward
If a gradient is unused (e.g., a parameter's requires_grad=False), AOTAutograd can eliminate the entire backward subgraph for that parameter. In eager mode, the autograd engine would still compute it.
Constant Folding in Backward
Operations that depend only on constants (like weight shapes) can be folded at compile time. The backward graph is simplified before code generation.
Memory Planning
Inductor can plan memory allocation for the entire backward pass upfront. In eager mode, each backward op allocates its output independently. With compilation, Inductor can reuse memory buffers across non-overlapping operations.
Kernel Fusion
Small backward operations (element-wise gradients, accumulations) are fused into larger kernels. Instead of launching 20 separate kernels for 20 backward ops, Inductor might generate 3-4 fused kernels that do the same work with much less launch overhead.
9. Activation Memory in AOTAutograd
The partitioner controls what "saved tensors" are passed from forward to backward. This directly determines training memory consumption — fewer saved tensors means less memory, but potentially more recomputation.
Viewing Saved Tensors
from torch._functorch.aot_autograd import aot_function
saved_tensors_count = []
def counting_compiler(gm, example_inputs):
# Count extra outputs in forward = saved tensors
output_node = [n for n in gm.graph.nodes if n.op == "output"][0]
n_outputs = len(output_node.args[0])
saved_tensors_count.append(n_outputs)
return gm
compiled = aot_function(fn, fw_compiler=counting_compiler, bw_compiler=lambda gm, _: gm)
Memory Impact
For a transformer layer with hidden_size=1024, batch_size=32, seq_len=512:
| Activation | Size | Saved? (min-cut) |
|---|---|---|
| QKV projection output | 48 MB | Yes (matmul) |
| Attention scores | 32 MB | Yes (matmul) |
| Post-softmax attention | 32 MB | Yes (expensive) |
| ReLU mask | 2 MB | No (recomputed) |
| LayerNorm intermediate | 4 MB | No (recomputed) |
| Residual add result | 8 MB | No (recomputed) |
The min-cut partitioner saves roughly 112 MB instead of 126 MB — a 11% reduction for this layer. Across a 24-layer model, that's over 300 MB saved.
Interaction with Activation Checkpointing
AOTAutograd's rematerialization is complementary to user-level activation checkpointing (torch.utils.checkpoint). Checkpointing operates at the module level (recompute an entire layer), while AOTAutograd's min-cut operates at the operator level (recompute individual cheap ops).
You can use both:
# Module-level: recompute entire transformer layers
model = checkpoint_wrapper(model)
# Op-level: within each layer, AOTAutograd's min-cut
# further reduces saved tensors
model = torch.compile(model)
10. AOTAutograd vs Standard Autograd
| Feature | Standard Autograd | AOTAutograd |
|---|---|---|
| Graph built | Runtime (during forward) | Compile time (ahead-of-time) |
| Backward optimized | No (eager dispatch) | Yes (Inductor compiles backward) |
| Memory planning | Manual (user calls checkpoint) | Automatic (min-cut partitioner) |
| Kernel fusion | None (one kernel per op) | Yes (Inductor fuses backward ops) |
| Works with compile | Forward only | Forward + backward |
| Dynamic graphs | Fully supported | Requires recompilation on change |
| Autograd hooks | Fully supported | Limited support |
| In-place ops | Fully supported | Functionalized (no in-place) |
| Debugging | Easy (Python stack traces) | Harder (compiled code) |
| Overhead | None | Compilation cost (amortized) |
When Standard Autograd is Better
- Debugging: When you need Python-level stack traces and step-through debugging
- Dynamic models: Models where the computation graph changes every iteration (e.g., tree-RNNs)
- Complex hooks: Models that rely heavily on autograd hooks for gradient manipulation
- One-off computations: When compilation cost isn't amortized (few iterations)
When AOTAutograd is Better
- Training throughput: When you need maximum training speed
- Large models: Where kernel launch overhead is a bottleneck
- Static models: Where the graph doesn't change between iterations (transformers, CNNs)
- Memory-constrained: Where automatic rematerialization helps fit larger batches
11. Functorch and AOTAutograd
AOTAutograd is built on top of functorch's primitives. Understanding this connection clarifies how the system works.
make_fx
make_fx is functorch's functional tracing tool. It runs a function with FakeTensor inputs and records every ATen operation into an FX graph:
from torch.fx.experimental.proxy_tensor import make_fx
def fn(x):
return torch.relu(x) + 1
gm = make_fx(fn)(torch.randn(4))
print(gm.graph)
AOTAutograd uses make_fx to trace the forward pass.
torch.autograd.grad on Traced Graphs
The traced forward graph is itself differentiable — every ATen op has registered derivatives. AOTAutograd calls torch.autograd.grad() on the traced graph to produce backward operations:
def joint_fn(primals, tangents):
# Forward
output = traced_forward(*primals)
# Backward (using autograd on the trace)
grads = torch.autograd.grad(output, primals, tangents)
return output, grads
joint_graph = make_fx(joint_fn)(primals, tangents)
This is the joint graph — forward and backward combined.
grad and vmap
Functorch's grad transform is closely related. While torch.autograd.grad computes gradients imperatively, functorch's grad wraps a function to return its gradient:
from torch.func import grad
def loss_fn(x):
return (x ** 2).sum()
grad_fn = grad(loss_fn)
gradient = grad_fn(torch.randn(4))
AOTAutograd uses the imperative torch.autograd.grad rather than functorch's grad transform, but the mathematical operation is the same.
12. Debugging AOTAutograd
Environment Variables
# See forward and backward graphs
TORCH_LOGS="aot" python script.py
# See generated Inductor code for both graphs
TORCH_LOGS="output_code" python script.py
# See partitioner decisions
TORCH_LOGS="aot" python script.py
# Look for lines containing "partition" in the output
# Combined: full picture
TORCH_LOGS="aot,output_code,graph_breaks" python script.py
Partitioner Debugging
torch._functorch.config.debug_partitioner = True
This prints which tensors the partitioner decided to save vs recompute, and the cost estimates that drove those decisions.
Common Issues
Graph breaks in backward: If Dynamo encounters an unsupported operation during backward tracing, it inserts a graph break. This fragments the backward into multiple compiled regions with eager transitions between them.
# Check for graph breaks
torch._dynamo.config.verbose = True
# Look for "graph break" in output
Shape mismatch errors: These occur when the backward graph expects a different shape than what the forward produces. Usually caused by dynamic shapes or data-dependent operations.
# Use TORCH_LOGS to see the shapes at each node
TORCH_LOGS="aot,dynamic" python script.py
Functionalization errors: Some in-place operations can't be functionalized. The error message will mention functionalize or FunctionalTensorWrapper.
Recompilation storms: If the model's shapes change frequently, AOTAutograd recompiles the forward and backward graphs each time. Use torch._dynamo.config.cache_size_limit and dynamic shapes to mitigate.
Comparing Eager vs Compiled Results
model_eager = MyModel()
model_compiled = torch.compile(MyModel())
# Same weights
model_compiled.load_state_dict(model_eager.state_dict())
x = torch.randn(4, 8, requires_grad=True)
x_copy = x.clone().detach().requires_grad_(True)
# Forward
y_eager = model_eager(x)
y_compiled = model_compiled(x_copy)
# Backward
y_eager.sum().backward()
y_compiled.sum().backward()
# Compare gradients
print(torch.allclose(x.grad, x_copy.grad)) # Should be True
13. Upstream Updates (July 5-7, 2026)
Recent changes relevant to compiled autograd and AOTAutograd:
Inductor Accumulator addmm Preservation (#184296)
The Inductor backend now preserves accumulator precision for addmm operations during backward pass compilation. Previously, intermediate accumulations could lose precision when Inductor fused multiple matmul backward ops. This fix ensures that gradient accumulation in compiled backward passes matches eager autograd numerics, particularly important for large-scale training where small numerical differences compound across many steps.
standalone_compile Fake Mode Fix (#185638)
Fixed an issue where standalone_compile would fail when entering fake mode for AOTAutograd tracing. The bug manifested as shape inference errors during joint graph construction — fake tensors would lose their symbolic shape information when passed through certain decomposition rules. This fix ensures that standalone compilation (used for ahead-of-time compilation workflows) correctly maintains shape metadata throughout the AOTAutograd pipeline.
CPU Outer-Loop Buffer Reuse (#185855)
Inductor's CPU backend now supports buffer reuse across outer loop iterations in compiled backward graphs. When the backward pass contains reduction operations that produce intermediate buffers, those buffers are now recycled rather than reallocated each iteration. This reduces memory allocation pressure during compiled backward passes on CPU, particularly for models with many small reduction ops in their gradient computation.
Avoid CUDA Init in CPU Compile (#186403)
AOTAutograd and Inductor no longer trigger CUDA initialization when compiling models on CPU. Previously, importing certain compilation modules would call torch.cuda.is_available() or torch.cuda.device_count(), which initializes the CUDA runtime. This was wasteful for CPU-only workloads and caused failures in environments without GPU drivers. The fix lazily gates CUDA-specific code paths.
Stateless RNG Clone Fix (#188495)
Fixed a bug in AOTAutograd's functionalization pass where stateless RNG operations (used by dropout and similar stochastic layers) were not correctly cloned during joint graph construction. The symptom was that compiled training would produce different random masks in forward vs backward, leading to incorrect gradients for models with dropout. The fix ensures that the RNG state is properly snapshotted at the partition boundary.
NativeRT Warp Size Query for Triton (#188881)
The NativeRT inference engine now correctly queries the GPU's warp size when running Triton-generated kernels from AOTInductor. This is relevant to AOTAutograd because compiled backward passes may be deployed via AOTInductor for inference-time gradient computation (e.g., in differentiable rendering or physics simulation). Previously, NativeRT assumed a warp size of 32, which fails on non-NVIDIA hardware.
Key Takeaways
- Eager backward is the bottleneck —
torch.compilealone only optimizes the forward; the backward is still eager dispatch with per-op overhead - AOTAutograd traces both passes at compile time — it uses
make_fx+torch.autograd.gradto produce forward and backward FX graphs that Inductor can optimize - The joint graph is central — forward and backward ops start in one graph, then the partitioner splits them, deciding what to save vs recompute
- min-cut partitioner automates checkpointing — instead of manually wrapping layers in
checkpoint(), the partitioner does operator-level save/recompute decisions based on cost - Compiled Autograd goes further — it compiles the autograd engine's execution itself, capturing hooks, accumulation, and multi-output handling
- Use
aot_functionto inspect graphs — pass customfw_compilerandbw_compilercallbacks to see exactly what AOTAutograd produces - Debugging uses TORCH_LOGS —
TORCH_LOGS="aot"shows graphs,debug_partitioner=Trueshows save/recompute decisions - Both forward and backward run at Inductor speed — operator fusion, memory planning, and kernel optimization apply to the backward pass too
- Complementary to activation checkpointing — AOTAutograd's min-cut is op-level; user checkpointing is module-level; they compose
Further Resources
- AOTAutograd Documentation — official API reference
- Compiled Autograd Tutorial — step-by-step walkthrough
- Module 03 — Autograd — foundational autograd concepts
- Module 08 — torch.compile — the compilation pipeline
- Module 16 — Activation Checkpointing — manual memory optimization
- Module 35 — The Dispatcher — how ops are dispatched
Notebook: 38_compiled_autograd.ipynb
Module 39: Building a Text Classifier
Prerequisites: Module 04 — Neural Networks, Module 05 — Optimizers, Module 06 — Data Loading, Module 07 — Training Pipelines, Module 09 — Attention Mechanisms
Time: ~4 hours
Files:tokenizer.py,text_classifier.py,train_and_evaluate.py
Table of Contents
- Project Overview
- Tokenization
- Embedding Layer
- Model Architecture
- Dataset & DataLoader
- Training Loop
- Evaluation Metrics
- Inference Pipeline
- torch.compile for Serving
- Improvements
- Upstream Updates (July 7-15, 2026)
1. Project Overview
This module builds a sentiment classifier entirely from scratch. No pretrained models, no Hugging Face tokenizers, no external NLP libraries. You will implement every piece yourself:
Raw Text → Tokenizer → Token IDs → Embedding → Transformer Encoder → Classification Head → Prediction
↓
positive / negative / neutral
Why build from scratch? Using pretrained models is the right production choice, but it hides the mechanics. Building each component teaches you:
- How text becomes numbers (tokenization)
- How numbers become vectors (embeddings)
- How vectors become predictions (transformer + classifier)
- How to train the whole system end-to-end
What We'll Build
| Component | What It Does | File |
|---|---|---|
| Character Tokenizer | Splits text into characters, maps to integers | tokenizer.py |
| Word Tokenizer | Splits on whitespace, builds frequency-based vocabulary | tokenizer.py |
| TransformerTextClassifier | Embedding + positional encoding + transformer encoder + classification head | text_classifier.py |
| Training Pipeline | Synthetic data, training loop, evaluation, inference | train_and_evaluate.py |
Input/Output
Input: "This movie was absolutely fantastic and I loved every moment"
Output: {"label": "positive", "confidence": 0.94}
Input: "Terrible acting, awful script, waste of time"
Output: {"label": "negative", "confidence": 0.91}
Input: "The movie was okay, nothing special"
Output: {"label": "neutral", "confidence": 0.72}
Everything runs on CPU — no GPU required.
2. Tokenization
Tokenization converts raw text into a sequence of integers that a neural network can process. This is the first — and often most impactful — design decision in any NLP pipeline.
Why Tokenize?
Neural networks operate on numbers, not strings. We need a mapping:
"hello world" → [7, 4, 11, 11, 14, 0, 22, 14, 17, 11, 3] # character-level
"hello world" → [42, 103] # word-level
"hello world" → [7592, 1079] # subword (BPE)
Each approach trades off vocabulary size, sequence length, and expressiveness.
Character-Level Tokenizer
The simplest approach: each character is a token.
text = "hello"
tokens = ['h', 'e', 'l', 'l', 'o']
ids = [7, 4, 11, 11, 14]
Advantages:
- Tiny vocabulary (26 letters + punctuation + digits ≈ 100 tokens)
- No out-of-vocabulary (OOV) words — every string can be tokenized
- Can handle typos, neologisms, any language with the same alphabet
Disadvantages:
- Long sequences — "transformer" is 11 tokens, not 1
- Hard for the model to learn word-level semantics from characters
- Self-attention cost is O(n²) in sequence length
Word-Level Tokenizer
Split on whitespace and punctuation. Each word is a token.
text = "Hello, world!"
tokens = ['hello', ',', 'world', '!']
ids = [42, 3, 103, 4]
Advantages:
- Short sequences — each word is one token
- Semantically meaningful units
- Easy to implement
Disadvantages:
- Large vocabulary (English has 170,000+ words)
- OOV problem — unseen words map to
[UNK] - Can't handle morphology ("running" and "runs" are separate tokens)
Subword Tokenization (BPE Concept)
Byte-Pair Encoding (BPE) finds a middle ground. It starts with characters and iteratively merges the most frequent pairs:
Iteration 0: ['l', 'o', 'w', 'e', 'r', 'l', 'o', 'w', 'e', 's', 't']
Merge ('l', 'o') → 'lo': ['lo', 'w', 'e', 'r', 'lo', 'w', 'e', 's', 't']
Merge ('lo', 'w') → 'low': ['low', 'e', 'r', 'low', 'e', 's', 't']
Merge ('low', 'e') → 'lowe': ['lowe', 'r', 'lowe', 's', 't']
BPE handles rare words by splitting them into known subwords:
"unhappiness" → ["un", "happiness"]
"transformers" → ["transform", "ers"]
We won't implement BPE in this module (it's complex), but understanding the concept explains why production systems use it. Our word-level tokenizer with [UNK] handling is sufficient for learning the full pipeline.
Special Tokens
Every tokenizer needs special tokens that carry structural information:
| Token | Purpose | Typical ID |
|---|---|---|
[PAD] | Fills sequences to equal length for batching | 0 |
[UNK] | Replaces out-of-vocabulary words | 1 |
[CLS] | Classification token — its embedding becomes the sequence representation | 2 |
[SEP] | Separator between segments (for pair classification) | 3 |
Building a Vocabulary
The vocabulary is a bidirectional mapping between tokens and integer IDs:
# Forward: token → id
vocab = {"[PAD]": 0, "[UNK]": 1, "[CLS]": 2, "the": 3, "movie": 4, ...}
# Reverse: id → token
id_to_token = {0: "[PAD]", 1: "[UNK]", 2: "[CLS]", 3: "the", 4: "movie", ...}
Building from data:
- Tokenize all training texts
- Count token frequencies
- Keep the top N most frequent tokens (vocabulary size)
- Assign integer IDs (special tokens first)
Encoding and Decoding
text = "great movie"
# 1. Tokenize: ["great", "movie"]
# 2. Prepend [CLS]: ["[CLS]", "great", "movie"]
# 3. Map to IDs: [2, 57, 4]
# 4. Pad to max_length=8: [2, 57, 4, 0, 0, 0, 0, 0]
ids = [2, 57, 4, 0, 0, 0, 0, 0]
# 1. Map to tokens: ["[CLS]", "great", "movie", "[PAD]", "[PAD]", ...]
# 2. Remove special tokens: ["great", "movie"]
# 3. Join: "great movie"
Padding and Truncation
Batching requires all sequences to have the same length. Two operations handle this:
Padding — add [PAD] tokens to short sequences:
"good" → [CLS, good, PAD, PAD, PAD] # length 5
"very good" → [CLS, very, good, PAD, PAD] # length 5
Truncation — cut long sequences to max_length:
"this is a very long sentence that exceeds the limit"
→ [CLS, this, is, a, very] # max_length=5, truncated
The padding mask tells the model which positions are real tokens vs padding:
tokens: [CLS, good, PAD, PAD, PAD]
mask: [ 1, 1, 0, 0, 0] # 1=attend, 0=ignore
See tokenizer.py for the complete implementation.
3. Embedding Layer
Tokenization gives us integers. The embedding layer converts each integer into a dense vector that the model can work with.
nn.Embedding: The Lookup Table
nn.Embedding is simply a matrix of shape (vocab_size, embed_dim). Looking up token ID i returns row i:
embedding = nn.Embedding(num_embeddings=1000, embedding_dim=128)
# embedding.weight.shape = (1000, 128)
token_ids = torch.tensor([42, 7, 103])
vectors = embedding(token_ids) # shape: (3, 128)
# vectors[0] = embedding.weight[42]
# vectors[1] = embedding.weight[7]
# vectors[2] = embedding.weight[103]
Random Init vs Learned Embeddings
At initialization, embedding vectors are random — "dog" and "cat" have no special relationship. During training, backpropagation adjusts the vectors so that semantically similar words end up near each other in embedding space:
Before training:
"good" = [0.23, -0.81, 0.45, ...] (random)
"great" = [-0.12, 0.67, -0.33, ...] (random)
"terrible" = [0.56, 0.09, 0.71, ...] (random)
After training:
"good" = [0.82, 0.41, -0.15, ...]
"great" = [0.79, 0.38, -0.12, ...] (close to "good")
"terrible" = [-0.71, -0.33, 0.65, ...] (far from "good")
padding_idx
Padding tokens should not contribute to the model's computation. Setting padding_idx=0 ensures:
- The embedding vector for token 0 is always zeros
- Gradients don't flow through padding positions
embedding = nn.Embedding(1000, 128, padding_idx=0)
# embedding.weight[0] is always [0, 0, 0, ..., 0]
# No gradient updates for token 0
Positional Embeddings
Self-attention is permutation-invariant — it doesn't know token order. "The cat sat" and "sat cat the" produce identical attention outputs without positional information.
Sinusoidal positional encoding (from "Attention Is All You Need"):
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
This produces a unique vector for each position. Properties:
- Each position gets a distinct encoding
- The model can learn to attend to relative positions
- Generalizes to longer sequences than seen during training
Learned positional embeddings are an alternative — just another nn.Embedding(max_length, embed_dim) that's added to the token embeddings. We use sinusoidal in our implementation.
The final input to the transformer is:
input = token_embedding(token_ids) + positional_encoding(positions)
Both are (batch, seq_len, embed_dim) tensors.
4. Model Architecture
Our classifier uses a Transformer encoder architecture. This is the same encoder used in BERT, but built from scratch with modern best practices.
Architecture Diagram
Input Token IDs: (batch, seq_len)
│
▼
┌─────────────────────┐
│ nn.Embedding │ (vocab_size, d_model)
│ + PositionalEnc │ (max_len, d_model)
│ + Dropout │
└─────────┬───────────┘
│
▼ (batch, seq_len, d_model)
┌─────────────────────┐
│ TransformerEncoder │ N layers ×:
│ ┌─────────────────┐ │ - LayerNorm (pre-norm)
│ │ Self-Attention │ │ - Multi-Head Attention (SDPA)
│ │ (SDPA) │ │ - Residual connection
│ ├─────────────────┤ │ - LayerNorm (pre-norm)
│ │ Feed-Forward │ │ - FFN: Linear → GELU → Linear
│ │ Network │ │ - Residual connection
│ └─────────────────┘ │
│ (× N layers) │
└─────────┬───────────┘
│
▼ (batch, seq_len, d_model)
┌─────────────────────┐
│ Pooling │ CLS token → (batch, d_model)
│ (CLS or Mean) │ or Mean pool → (batch, d_model)
└─────────┬───────────┘
│
▼ (batch, d_model)
┌─────────────────────┐
│ Classification Head │ LayerNorm → Linear → num_classes
└─────────┬───────────┘
│
▼ (batch, num_classes)
Logits
Component Details
Embedding + Positional Encoding:
# Token embedding: (batch, seq_len) → (batch, seq_len, d_model)
x = self.embedding(token_ids) # shape: (B, S, D)
x = x + self.pos_encoder(positions) # shape: (B, S, D)
x = self.dropout(x) # shape: (B, S, D)
TransformerEncoderLayer (pre-norm variant):
Pre-norm applies LayerNorm before (not after) each sublayer. This improves training stability — gradients flow more smoothly through the residual connections.
# Pre-norm self-attention
residual = x
x = self.norm1(x) # shape: (B, S, D)
x = self.self_attn(x, x, x, mask) # shape: (B, S, D)
x = residual + self.dropout(x) # shape: (B, S, D)
# Pre-norm feed-forward
residual = x
x = self.norm2(x) # shape: (B, S, D)
x = self.ffn(x) # shape: (B, S, D)
x = residual + self.dropout(x) # shape: (B, S, D)
SDPA (Scaled Dot-Product Attention):
PyTorch's F.scaled_dot_product_attention fuses Q·K^T/√d, mask, softmax, and V multiplication into one efficient operation:
attn_output = F.scaled_dot_product_attention(
query, key, value,
attn_mask=padding_mask,
dropout_p=self.dropout_p if self.training else 0.0,
)
Pooling Strategies:
Two approaches to get a fixed-size vector from variable-length sequences:
- CLS token pooling: Take the hidden state of the
[CLS]token (position 0). The[CLS]token attends to all other tokens, so its representation captures the whole sequence.
pooled = encoder_output[:, 0, :] # (batch, d_model) — CLS position
- Mean pooling: Average all non-padding token representations. More robust — doesn't rely on a single token learning to summarize everything.
mask = padding_mask.unsqueeze(-1) # (batch, seq_len, 1)
pooled = (encoder_output * mask).sum(1) # (batch, d_model)
pooled = pooled / mask.sum(1).clamp(min=1e-9) # normalize by actual length
Classification Head:
A simple linear projection from d_model to num_classes:
logits = self.classifier(pooled) # (batch, d_model) → (batch, num_classes)
Shape Annotations
Following the data through the model for batch=32, seq_len=64, d_model=128, num_heads=4, num_layers=2, num_classes=3:
token_ids: (32, 64) — input
embedded: (32, 64, 128) — after embedding + positional
encoder_out: (32, 64, 128) — after transformer encoder
pooled: (32, 128) — after CLS/mean pooling
logits: (32, 3) — final output
See text_classifier.py for the complete implementation.
5. Dataset & DataLoader
Custom TextDataset
Our dataset stores tokenized sequences and their labels:
class TextDataset(Dataset):
def __init__(self, texts, labels, tokenizer, max_length=128):
self.encodings = [tokenizer.encode(t, max_length) for t in texts]
self.labels = labels
def __len__(self):
return len(self.labels)
def __getitem__(self, idx):
return {
"input_ids": torch.tensor(self.encodings[idx], dtype=torch.long),
"label": torch.tensor(self.labels[idx], dtype=torch.long),
}
Custom collate_fn
Texts have different lengths. Our collate_fn pads each batch to the length of the longest sequence in that batch (dynamic padding):
def collate_fn(batch):
input_ids = [item["input_ids"] for item in batch]
labels = torch.stack([item["label"] for item in batch])
# Pad to max length in this batch
input_ids = nn.utils.rnn.pad_sequence(input_ids, batch_first=True, padding_value=0)
# Create attention mask: 1 for real tokens, 0 for padding
attention_mask = (input_ids != 0).float()
return {"input_ids": input_ids, "attention_mask": attention_mask, "labels": labels}
Dynamic padding is more efficient than padding all sequences to the global max_length — shorter batches waste less compute.
Train/Val/Test Split
Standard practice: 80% train, 10% validation, 10% test.
from torch.utils.data import random_split
dataset = TextDataset(texts, labels, tokenizer)
train_size = int(0.8 * len(dataset))
val_size = int(0.1 * len(dataset))
test_size = len(dataset) - train_size - val_size
train_set, val_set, test_set = random_split(dataset, [train_size, val_size, test_size])
Why three splits?
- Train: model learns from this data
- Validation: tune hyperparameters, decide when to stop training
- Test: final evaluation — never used for any decisions during training
6. Training Loop
Our training loop includes every best practice for a production-quality training run.
Optimizer: AdamW
AdamW decouples weight decay from the gradient update. This is the standard optimizer for transformer training:
optimizer = torch.optim.AdamW(
model.parameters(),
lr=1e-3,
weight_decay=0.01,
betas=(0.9, 0.999),
)
Learning Rate Schedule: OneCycleLR
OneCycleLR implements warmup + cosine decay in one scheduler:
LR
↑ ╱╲
│ ╱ ╲
│ ╱ ╲
│ ╱ ╲
│╱ ╲________
└─────────────────────────→ step
warmup decay final
scheduler = torch.optim.lr_scheduler.OneCycleLR(
optimizer,
max_lr=1e-3,
epochs=num_epochs,
steps_per_epoch=len(train_loader),
)
Loss Function: CrossEntropyLoss
For multi-class classification (positive/negative/neutral), CrossEntropyLoss combines LogSoftmax + NLLLoss:
criterion = nn.CrossEntropyLoss()
# Input: logits (batch, num_classes) — NOT probabilities
# Target: class indices (batch,) — integers 0, 1, 2
loss = criterion(logits, labels)
Mixed Precision (BF16)
BFloat16 reduces memory and speeds up training on modern hardware without the numerical instability of FP16:
with torch.autocast(device_type="cpu", dtype=torch.bfloat16):
logits = model(input_ids, attention_mask)
loss = criterion(logits, labels)
We use device_type="cpu" since this module runs without GPU. On GPU, change to "cuda".
Gradient Clipping
Prevents exploding gradients that can destabilize training:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
This clips the global gradient norm to 1.0. If the total gradient norm exceeds 1.0, all gradients are scaled down proportionally.
Validation After Each Epoch
After each training epoch, evaluate on the validation set:
model.eval()
with torch.no_grad():
for batch in val_loader:
logits = model(batch["input_ids"], batch["attention_mask"])
loss = criterion(logits, batch["labels"])
# Accumulate metrics...
Early Stopping
Stop training when validation loss stops improving:
patience = 3
best_val_loss = float("inf")
epochs_without_improvement = 0
for epoch in range(max_epochs):
# ... train and validate ...
if val_loss < best_val_loss:
best_val_loss = val_loss
epochs_without_improvement = 0
torch.save(model.state_dict(), "best_model.pt")
else:
epochs_without_improvement += 1
if epochs_without_improvement >= patience:
print("Early stopping!")
break
Best Model Checkpointing
Save the model whenever validation loss improves:
if val_loss < best_val_loss:
best_val_loss = val_loss
checkpoint = {
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"epoch": epoch,
"val_loss": val_loss,
}
torch.save(checkpoint, "best_model.pt")
See train_and_evaluate.py for the complete training loop.
7. Evaluation Metrics
Accuracy alone is misleading for imbalanced datasets. If 90% of reviews are positive, a model that always predicts "positive" gets 90% accuracy while being useless.
Precision, Recall, F1
Precision: Of all predictions for class C, how many were correct?
Precision(C) = True Positives / (True Positives + False Positives)
Recall: Of all actual instances of class C, how many did we find?
Recall(C) = True Positives / (True Positives + False Negatives)
F1 Score: Harmonic mean of precision and recall:
F1(C) = 2 × (Precision × Recall) / (Precision + Recall)
Macro F1: Average F1 across all classes (treats all classes equally):
Macro-F1 = mean(F1(positive), F1(negative), F1(neutral))
Per-Class Metrics Example
Precision Recall F1-Score Support
positive 0.91 0.93 0.92 120
negative 0.88 0.85 0.86 95
neutral 0.79 0.82 0.80 85
macro avg 0.86 0.87 0.86 300
Confusion Matrix
A confusion matrix shows where the model confuses classes:
Predicted
pos neg neu
Actual pos [ 112 3 5 ]
neg [ 5 81 9 ]
neu [ 6 9 70 ]
Reading: Row = actual class, Column = predicted class. Diagonal = correct predictions. Off-diagonal = errors.
Key insights from this matrix:
- The model sometimes confuses neutral for negative (9 cases)
- Positive is the easiest class to predict
- Neutral has the lowest recall (most often misclassified)
Implementation
We compute these metrics without scikit-learn — just PyTorch:
def compute_metrics(all_preds, all_labels, num_classes):
confusion = torch.zeros(num_classes, num_classes, dtype=torch.long)
for pred, label in zip(all_preds, all_labels):
confusion[label][pred] += 1
per_class = {}
for c in range(num_classes):
tp = confusion[c][c].item()
fp = confusion[:, c].sum().item() - tp
fn = confusion[c, :].sum().item() - tp
precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0
recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0
f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0.0
per_class[c] = {"precision": precision, "recall": recall, "f1": f1}
return confusion, per_class
8. Inference Pipeline
Single Text Inference
def predict(text, model, tokenizer, label_names, max_length=128):
model.eval()
input_ids = torch.tensor([tokenizer.encode(text, max_length)])
attention_mask = (input_ids != 0).float()
with torch.no_grad():
logits = model(input_ids, attention_mask)
probs = torch.softmax(logits, dim=-1)
confidence, predicted = probs.max(dim=-1)
return {
"text": text,
"label": label_names[predicted.item()],
"confidence": confidence.item(),
"all_probs": {name: probs[0][i].item() for i, name in enumerate(label_names)},
}
Batch Inference
For multiple texts, batch them together for efficiency:
def predict_batch(texts, model, tokenizer, label_names, max_length=128):
model.eval()
encoded = [tokenizer.encode(t, max_length) for t in texts]
input_ids = nn.utils.rnn.pad_sequence(
[torch.tensor(e) for e in encoded], batch_first=True, padding_value=0,
)
attention_mask = (input_ids != 0).float()
with torch.no_grad():
logits = model(input_ids, attention_mask)
probs = torch.softmax(logits, dim=-1)
confidences, predictions = probs.max(dim=-1)
return [
{"text": t, "label": label_names[p.item()], "confidence": c.item()}
for t, p, c in zip(texts, predictions, confidences)
]
Confidence Scores
Softmax converts logits to probabilities:
logits: [2.1, -0.5, 0.3]
softmax: [0.72, 0.05, 0.23] (sums to 1.0)
Use confidence scores to set a threshold:
if result["confidence"] < 0.5:
result["label"] = "uncertain" # don't trust low-confidence predictions
9. torch.compile for Serving
Once trained, compile the model for faster inference:
model.eval()
compiled_model = torch.compile(model, mode="reduce-overhead")
# Warmup (first call triggers compilation)
dummy = torch.randint(0, vocab_size, (1, 64))
mask = torch.ones(1, 64)
_ = compiled_model(dummy, mask)
# Now inference is faster
with torch.no_grad():
logits = compiled_model(input_ids, attention_mask)
Benchmarking Compiled vs Eager
import time
def benchmark(model, input_ids, mask, n_runs=100):
# Warmup
for _ in range(10):
model(input_ids, mask)
start = time.perf_counter()
for _ in range(n_runs):
model(input_ids, mask)
elapsed = time.perf_counter() - start
return elapsed / n_runs * 1000 # ms per inference
eager_ms = benchmark(model, input_ids, mask)
compiled_ms = benchmark(compiled_model, input_ids, mask)
print(f"Eager: {eager_ms:.2f} ms")
print(f"Compiled: {compiled_ms:.2f} ms")
print(f"Speedup: {eager_ms / compiled_ms:.2f}x")
Typical results on CPU for a small model:
Eager: 3.45 ms
Compiled: 2.10 ms
Speedup: 1.64x
On GPU with mode="reduce-overhead", speedups can reach 2-4x for small models due to CUDA graph capture.
Compile Modes for Serving
| Mode | Compile Time | Inference Speed | Use Case |
|---|---|---|---|
default | Fast | Good | General purpose |
reduce-overhead | Medium | Best for small models | Low-latency serving |
max-autotune | Slow | Best for large models | Throughput-optimized |
10. Improvements
Pretrained Embeddings
Instead of learning embeddings from scratch, initialize with pretrained vectors (GloVe, FastText):
pretrained = load_glove("glove.6B.100d.txt") # word → 100d vector
for word, idx in vocab.items():
if word in pretrained:
embedding.weight.data[idx] = torch.tensor(pretrained[word])
This gives the model a head start — it already knows that "good" and "great" are similar before seeing any training data.
Data Augmentation
Synonym replacement: Replace words with synonyms to create new training examples:
"This movie was great" → "This film was excellent"
Random deletion: Randomly remove words (the model should still classify correctly):
"This movie was absolutely great" → "This movie great"
Back-translation concept: Translate to another language and back to get paraphrases:
"Great movie" → (French) "Super film" → (English) "Awesome film"
Attention Visualization
Extract attention weights to see what the model focuses on:
# Register hooks to capture attention weights
attention_weights = []
def hook_fn(module, input, output):
if hasattr(module, 'attn_weights'):
attention_weights.append(module.attn_weights)
for layer in model.encoder.layers:
layer.self_attn.register_forward_hook(hook_fn)
High attention on "terrible" for a negative prediction confirms the model learned meaningful patterns.
Model Distillation
Train a smaller, faster model (student) to mimic the larger model (teacher):
teacher_logits = teacher_model(input_ids, mask)
student_logits = student_model(input_ids, mask)
# Soft label loss (KL divergence between teacher and student distributions)
T = 4.0 # temperature
soft_loss = F.kl_div(
F.log_softmax(student_logits / T, dim=-1),
F.softmax(teacher_logits / T, dim=-1),
reduction="batchmean",
) * (T * T)
# Hard label loss (standard cross-entropy)
hard_loss = F.cross_entropy(student_logits, labels)
# Combined
loss = 0.7 * soft_loss + 0.3 * hard_loss
The student learns from both the teacher's soft probability distributions and the hard labels. The temperature T softens the teacher's outputs, exposing more information about inter-class relationships ("this review is 60% positive, 30% neutral, 10% negative" is more informative than just "positive").
11. Upstream Updates (July 7-15, 2026)
Recent upstream changes relevant to the text classifier workflow covered in this module:
Inductor Padding Fusion for Ragged Sequences (#189112)
The Inductor backend improved its handling of padded tensor operations common in NLP pipelines. When torch.compile encounters masked reductions over padded sequences (exactly the pattern used in mean pooling over variable-length texts), the fused kernel now avoids reading padding positions entirely rather than multiplying by zero. For short-sequence batches with heavy padding, this reduces unnecessary memory bandwidth by up to 30%. This directly benefits the compiled inference pipeline in Section 9.
nn.TransformerEncoderLayer Pre-Norm Default (#189445)
nn.TransformerEncoderLayer now defaults to norm_first=True (pre-norm) instead of norm_first=False (post-norm). Pre-norm is the standard in modern transformers and improves training stability. Our model already uses norm_first=True explicitly, so no code changes are needed, but new code no longer needs to specify it. The post-norm default dated back to the original "Attention Is All You Need" paper but caused training instability without careful learning rate warmup.
SDPA Nested Tensor Padding Mask (#190023)
F.scaled_dot_product_attention can now accept a NestedTensor as input and automatically handles the padding mask. Instead of constructing an explicit (batch, 1, 1, seq_len) mask and passing it as attn_mask, you can pack variable-length sequences into a NestedTensor and SDPA derives the mask internally. This eliminates the mask broadcasting overhead and enables Flash Attention on padded NLP batches where it previously fell back to the math kernel.
CrossEntropyLoss label_smoothing BF16 Fix (#190187)
Fixed a numerical issue in nn.CrossEntropyLoss with label_smoothing > 0 when running in BF16 autocast. The smoothing computation accumulated small probabilities in BF16, which underflowed for large vocabularies. The fix promotes the smoothing arithmetic to FP32 internally. This matters for our training loop if users enable label smoothing as an improvement (a common regularization technique for classification).
OneCycleLR step() Warning Suppression (#190542)
OneCycleLR no longer emits a deprecation warning when step() is called per-batch (its intended usage pattern). Previously, calling scheduler.step() inside the batch loop printed a warning suggesting epoch-level stepping, which was incorrect for OneCycleLR. Our training loop calls scheduler.step() per-batch, which is the correct pattern, and this warning no longer appears.
torch.compile Dynamic Sequence Length Cache (#191003)
torch.compile improved its recompilation behavior for models with dynamic sequence lengths. Previously, each new sequence length triggered a full recompilation. With this change, the compiler generates specialized code for a set of "buckets" (powers of 2, common lengths) and falls back to a generic dynamic-shape kernel for other lengths. For our text classifier with variable-length inputs, this means the compiled model handles different batch configurations without excessive recompilation after the initial warmup.
Key Takeaways
- Tokenization is the foundation — the choice of character-level, word-level, or subword tokenization fundamentally shapes model capacity and sequence length
- Embeddings are learned —
nn.Embeddingstarts random and learns semantic relationships during training;padding_idxprevents the padding token from contributing - Positional encoding adds order — sinusoidal or learned position embeddings tell the transformer where each token sits in the sequence
- Pre-norm transformers are stabler — applying LayerNorm before (not after) attention and FFN sublayers improves gradient flow
- Dynamic padding saves compute — pad to the longest sequence in each batch, not the global maximum, using a custom
collate_fn - OneCycleLR handles warmup + decay — one scheduler replaces manual warmup + cosine decay configuration
- Metrics beyond accuracy matter — precision, recall, and F1 per class reveal where the model struggles; the confusion matrix shows which classes are confused
- Confidence scores enable thresholding — softmax probabilities let you reject uncertain predictions at inference time
- torch.compile speeds up serving — compile the trained model for 1.5-4x faster inference with no accuracy change
Further Resources
- Module 04 — Neural Networks —
nn.Module, layers, losses - Module 05 — Optimizers — AdamW, learning rate schedulers
- Module 06 — Data Loading — Dataset, DataLoader, custom collate
- Module 07 — Training Pipelines — full training loops, mixed precision
- Module 09 — Attention Mechanisms — SDPA, multi-head attention, transformers
- Module 08 — torch.compile — compilation for inference
- PyTorch Text Classification Tutorial
Notebook: 39_text_classifier.ipynb
Operational guide
Cross-Repository CI Relay (CRCR)
PyTorch's Cross-Repository CI Relay enables downstream backends (Intel XPU, AMD ROCm, Apple MPS, Red Hat RHEL, etc.) to run their own CI against upstream PyTorch PRs and report results back to the PyTorch HUD.
Architecture
pytorch/pytorch (PR merged or updated)
│
▼ repository_dispatch
downstream/backend-ci
│
├─ Build PyTorch from source at dispatched SHA
├─ Run backend-specific tests
└─ POST callback → CRCR Lambda → ClickHouse → HUD
Registration (Allowlist)
Downstream repos register in pytorch/pytorch/.github/allowlist.yml:
L2:
- intel/torch-xpu-ops
- TorchedHat/pytorch-redhat-ci
L3:
- some-org/experimental-backend
Levels:
- L1: Receive dispatches, no HUD reporting
- L2: Report to HUD, non-blocking
- L3: Report to HUD, visible but non-blocking with distinct treatment
- L4: Report to HUD, blocking (viable/strict)
Receiving Dispatches
# .github/workflows/crcr-ci.yml
on:
repository_dispatch:
types: [pull_request]
jobs:
test:
runs-on: self-hosted
steps:
- name: Get PR info
run: |
PR_NUM="${{ github.event.client_payload.pr_number }}"
SHA="${{ github.event.client_payload.head_sha }}"
echo "Testing PR #${PR_NUM} at ${SHA}"
- name: Build PyTorch
run: |
git clone https://github.com/pytorch/pytorch --depth 1
cd pytorch && git fetch origin ${SHA} && git checkout ${SHA}
git submodule update --init --recursive
pip install -e . -v --no-build-isolation
Reporting Results (Callback Action)
- name: Report to CRCR
if: always()
uses: pytorch/test-infra/.github/actions/cross-repo-ci-relay-callback@main
with:
conclusion: ${{ job.status }}
The callback action:
- Mints an OIDC token (proves identity)
- Builds a JSON payload (job name, conclusion, URL, timing)
- POSTs to the CRCR Lambda endpoint
- Lambda validates JWT, writes to ClickHouse
Nightly & Periodic CI
For nightly builds, downstream repos use schedule triggers and build from the pytorch/pytorch nightly branch:
on:
schedule:
- cron: "0 4 * * *"
jobs:
nightly:
steps:
- name: Get nightly SHA
run: |
# Nightly commits reference the source main SHA in their message
NIGHTLY_SHA=$(curl -fsSL \
"https://api.github.com/repos/pytorch/pytorch/commits?sha=nightly&per_page=1" \
| jq -r '.[0].sha')
COMMIT_MSG=$(curl -fsSL \
"https://api.github.com/repos/pytorch/pytorch/commits/${NIGHTLY_SHA}" \
| jq -r '.commit.message')
SOURCE_SHA=$(echo "$COMMIT_MSG" | grep -oP '\(([a-f0-9]{40})\)' | tr -d '()')
echo "Building from main@${SOURCE_SHA}"
HUD Integration
CRCR results appear on:
- PR pages: As a distinct "CRCR" section showing downstream results
- Commit pages: Same CRCR section for commits with associated PRs
- CRCR Summary page: Aggregated health metrics across all backends
- Main HUD grid: CRCR columns alongside in-tree CI (grouped by level)
Quick Start: Adding a New Downstream Backend
- Fork the template: Start from the
- Register: Open a PR to add your repo to
.github/allowlist.ymlunder L2. - Configure dispatch handler: Add a
repository_dispatchworkflow that
builds PyTorch from the dispatched SHA and runs your backend tests.
- Add the callback step: Include
pytorch/test-infra/.github/actions/cross-repo-ci-relay-callback@main
as the final step with if: always().
- Verify on HUD: After merge, trigger a test dispatch and check
hud.pytorch.org for your results.
Troubleshooting
| Symptom | Cause | Fix |
|---|---|---|
| Dispatch never received | Repo not in allowlist | Add to .github/allowlist.yml |
| OIDC token mint fails | Missing id-token: write permission | Add permissions: id-token: write to workflow |
| Results not on HUD | Callback URL wrong or Lambda down | Check callback action logs; verify endpoint |
| Build fails at dispatched SHA | Submodules out of sync | Run git submodule update --init --recursive |
| Nightly SHA resolution fails | Commit message format changed | Update grep pattern for SHA extraction |
Environment Variables
The dispatch payload sets these environment variables for your workflow:
| Variable | Description |
|---|---|
github.event.client_payload.pr_number | PR number that triggered the dispatch |
github.event.client_payload.head_sha | Git SHA to build and test against |
github.event.client_payload.base_sha | Base branch SHA for diff context |
github.event.client_payload.sender | GitHub user who authored the PR |
Key Resources
Module 40: Building an Image Classifier — End-to-End Computer Vision Project
Build a complete image classification system from scratch: data pipeline with augmentation, CNN and ResNet models, transfer learning, training with mixed precision, evaluation metrics, test-time augmentation, and Grad-CAM visualization.
No pretrained weights required — every component is implemented from scratch so you understand the full pipeline.
| Input | Output |
|---|---|
32x32 RGB image of a circle | circle (0.95) |
32x32 RGB image of a star | star (0.91) |
32x32 RGB image of a triangle | triangle (0.88) |
Table of Contents
- Overview
- Data Pipeline
- Data Augmentation
- MixUp and CutMix
- CNN Architecture
- ResNet Architecture
- Transfer Learning
- Training Pipeline
- Evaluation Metrics
- Test-Time Augmentation
- Grad-CAM Visualization
- Confidence Calibration
- Inference Pipeline
- Key Takeaways
Overview
Image classification is the canonical computer vision task: given an image, assign it to one of N categories. This module builds the full pipeline:
Raw Images → Augmentation → CNN/ResNet → Training (AMP + MixUp) → Evaluation → TTA → Grad-CAM
Architecture
┌─────────────────────────────────────────────────────────────────────────┐
│ IMAGE CLASSIFICATION PIPELINE │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌───────────┐ ┌──────────────┐ ┌────────────┐ ┌───────────┐ │
│ │ Data │ │ Augmentation │ │ Model │ │ Training │ │
│ │ Pipeline │───▶│ Transform │───▶│ (CNN / │───▶│ Loop │ │
│ │ (Synthetic)│ │ MixUp/Cut │ │ ResNet) │ │ AMP + EMA │ │
│ └───────────┘ └──────────────┘ └────────────┘ └─────┬─────┘ │
│ │ │
│ ┌───────────┐ ┌──────────────┐ ┌────────────┐ ┌─────▼─────┐ │
│ │ Inference │ │ Grad-CAM │ │ TTA │ │ Evaluation│ │
│ │ Pipeline │◀───│ Visualization│◀───│ (5-aug) │◀───│ Metrics │ │
│ └───────────┘ └──────────────┘ └────────────┘ └───────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────┘
Files
| File | Lines | Description |
|---|---|---|
data_pipeline.py | 270+ | Synthetic dataset, augmentation, MixUp, CutMix |
model_and_training.py | 310+ | SimpleCNN, MiniResNet, transfer learning, training loop |
evaluation.py | 270+ | Metrics, TTA, Grad-CAM, confidence analysis, inference |
Data Pipeline
Synthetic Shape Dataset
We generate synthetic images with 10 geometric shape classes. This lets us train and demonstrate the full pipeline without downloading large datasets.
CLASS_NAMES = [
"circle", "square", "triangle", "cross", "diamond",
"star", "ring", "arrow", "pentagon", "hexagon",
]
Each image is a 3-channel (RGB) tensor with:
- A random background intensity (0.0–0.3)
- A shape drawn with a foreground intensity (0.5–1.0)
- Gaussian noise for realism
class SyntheticShapeDataset(Dataset):
def __init__(self, num_samples=5000, img_size=32, channels=3, transform=None, seed=42):
...
def __getitem__(self, idx):
# 1. Create blank image with random background
img = torch.full((self.channels, self.img_size, self.img_size), bg)
# 2. Draw the shape (circle, square, etc.)
_DRAW_FNS[label](img, cx, cy, r, fg)
# 3. Add noise
img = (img + torch.randn_like(img) * 0.05).clamp(0, 1)
# 4. Apply transforms
if self.transform:
img = self.transform(img)
return img, label
Data Splits
Standard three-way split with separate seeds for reproducibility:
| Split | Samples | Purpose | Augmentation |
|---|---|---|---|
| Train | 4,000 | Weight updates | Full augmentation |
| Val | 500 | Hyperparameter tuning, early stopping | Normalize only |
| Test | 500 | Final evaluation | Normalize only |
train_loader, val_loader, test_loader = build_dataloaders(
num_train=4000, num_val=500, num_test=500,
img_size=32, batch_size=64,
)
Data Augmentation
Data augmentation artificially increases training set diversity by applying random transformations. This reduces overfitting and improves generalization.
Why Augmentation Works
A model that sees many variations of the same shape learns invariances:
- Flip invariance → recognizes shapes regardless of orientation
- Color invariance → focuses on shape, not brightness
- Occlusion robustness → handles partially visible objects
Implemented Transforms
| Transform | Parameters | Effect |
|---|---|---|
RandomHorizontalFlip | p=0.5 | Mirror left-right |
RandomVerticalFlip | p=0.3 | Mirror top-bottom |
RandomRotation90 | — | Rotate 0°/90°/180°/270° |
ColorJitter | brightness=0.2, contrast=0.2 | Vary brightness and contrast |
RandomErasing | p=0.3, scale=(0.02, 0.15) | Cutout-style occlusion |
Normalize | mean=0.5, std=0.5 | Center to [-1, 1] range |
Augmentation Pipeline
train_transform = Compose([
RandomHorizontalFlip(p=0.5),
RandomVerticalFlip(p=0.3),
RandomRotation90(),
ColorJitter(brightness=0.2, contrast=0.2),
RandomErasing(p=0.3),
Normalize(),
])
val_transform = Compose([
Normalize(), # No augmentation for validation/test
])
Pure PyTorch Transforms
All transforms are implemented using pure PyTorch tensor operations — no torchvision dependency:
class RandomHorizontalFlip:
def __init__(self, p=0.5):
self.p = p
def __call__(self, img):
if random.random() < self.p:
return img.flip(-1) # Flip width dimension
return img
class ColorJitter:
def __init__(self, brightness=0.2, contrast=0.2):
self.brightness = brightness
self.contrast = contrast
def __call__(self, img):
b_factor = 1.0 + (random.random() * 2 - 1) * self.brightness
c_factor = 1.0 + (random.random() * 2 - 1) * self.contrast
mean = img.mean()
img = (img - mean) * c_factor + mean # Contrast around mean
img = img * b_factor # Brightness scaling
return img.clamp(0, 1)
Random Erasing (Cutout)
Random erasing occludes a rectangular region, forcing the model to use all spatial regions rather than relying on a single discriminative patch:
class RandomErasing:
def __init__(self, p=0.3, scale=(0.02, 0.15)):
self.p = p
self.scale = scale
def __call__(self, img):
if random.random() > self.p:
return img
C, H, W = img.shape
area = H * W
erase_area = random.uniform(*self.scale) * area
eh = int(math.sqrt(erase_area * aspect))
ew = int(math.sqrt(erase_area / aspect))
img[:, y0:y0+eh, x0:x0+ew] = torch.rand(C, eh, ew)
return img
MixUp and CutMix
MixUp and CutMix are regularization techniques that blend training samples, creating soft labels that improve generalization.
MixUp
MixUp creates virtual training examples by linearly interpolating between two samples:
x_mixed = λ · x_a + (1 - λ) · x_b
loss = λ · L(pred, y_a) + (1 - λ) · L(pred, y_b)
Where λ ~ Beta(α, α), typically α = 0.2.
def mixup(images, labels, alpha=0.2):
lam = Beta(alpha, alpha).sample().item()
perm = torch.randperm(images.size(0))
mixed = lam * images + (1 - lam) * images[perm]
return mixed, labels, labels[perm], lam
CutMix
CutMix pastes a rectangular patch from one image onto another, combining spatial regions:
x_mixed[:, :, cy:cy+ch, cx:cx+cw] = x_b[:, :, cy:cy+ch, cx:cx+cw]
λ_actual = 1 - (ch · cw) / (H · W)
def cutmix(images, labels, alpha=1.0):
lam = Beta(alpha, alpha).sample().item()
cut_ratio = sqrt(1 - lam)
ch, cw = int(H * cut_ratio), int(W * cut_ratio)
mixed = images.clone()
mixed[:, :, cy:cy+ch, cx:cx+cw] = images[perm, :, cy:cy+ch, cx:cx+cw]
return mixed, labels, labels[perm], actual_lam
MixUp vs CutMix
| Aspect | MixUp | CutMix |
|---|---|---|
| Blending | Global pixel-wise | Local rectangular patch |
| Label mixing | Based on λ | Based on patch area ratio |
| Effect | Smoother decision boundaries | Better localization |
| Best α | 0.2 | 1.0 |
| Use case | General regularization | When spatial features matter |
Training with Mixed Labels
def mixup_criterion(criterion, pred, y_a, y_b, lam):
return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)
CNN Architecture
SimpleCNN
A lightweight 3-layer CNN with batch normalization and adaptive pooling:
Input (3, 32, 32)
│
├─ Conv2d(3→32, 3x3) + BN + ReLU + MaxPool(2) → (32, 16, 16)
├─ Conv2d(32→64, 3x3) + BN + ReLU + MaxPool(2) → (64, 8, 8)
├─ Conv2d(64→128, 3x3) + BN + ReLU + AdaptPool(4) → (128, 4, 4)
│
├─ Flatten → (2048,)
├─ Dropout(0.3) + Linear(2048→256) + ReLU
├─ Dropout(0.2) + Linear(256→10)
│
└─ Output: 10 class logits
class SimpleCNN(nn.Module):
def __init__(self, in_channels=3, num_classes=10, img_size=32):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(in_channels, 32, 3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, 3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.AdaptiveAvgPool2d(4),
)
self.classifier = nn.Sequential(
nn.Dropout(0.3),
nn.Linear(128 * 4 * 4, 256),
nn.ReLU(inplace=True),
nn.Dropout(0.2),
nn.Linear(256, num_classes),
)
def forward(self, x):
x = self.features(x)
x = x.flatten(1)
return self.classifier(x)
Design Decisions
| Choice | Reason |
|---|---|
| BatchNorm after Conv | Normalizes activations, enables higher learning rates |
| ReLU (inplace) | Saves memory, standard non-linearity |
| AdaptiveAvgPool(4) | Works with any input resolution |
| Dropout before Linear | Regularization in classifier head |
| 3x3 kernels only | Modern best practice (VGG insight) |
Parameter Count
SimpleCNN: ~570K parameters
Conv layers: ~110K
Classifier: ~460K (dominated by the 2048→256 linear)
ResNet Architecture
Residual Connections
The key insight: instead of learning H(x), learn the residual F(x) = H(x) - x:
output = F(x) + x
This solves the degradation problem — deeper networks can always learn the identity by setting F(x) = 0.
BasicBlock
x ──┬── Conv(3x3) → BN → ReLU → Conv(3x3) → BN ──┐
│ │
└──────────── Shortcut (identity or 1x1) ───────┘
│
ReLU(sum)
class BasicBlock(nn.Module):
def __init__(self, in_planes, planes, stride=1):
super().__init__()
self.conv1 = nn.Conv2d(in_planes, planes, 3, stride=stride, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(planes)
self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(planes)
# Shortcut: identity when dimensions match, 1x1 conv otherwise
self.shortcut = nn.Identity()
if stride != 1 or in_planes != planes:
self.shortcut = nn.Sequential(
nn.Conv2d(in_planes, planes, 1, stride=stride, bias=False),
nn.BatchNorm2d(planes),
)
def forward(self, x):
out = F.relu(self.bn1(self.conv1(x)), inplace=True)
out = self.bn2(self.conv2(out))
out = F.relu(out + self.shortcut(x), inplace=True) # Residual add
return out
MiniResNet
Adapted for 32x32 images (CIFAR-style, no initial downsampling):
Input (3, 32, 32)
├─ Conv(3→64, 3x3) + BN + ReLU → (64, 32, 32)
├─ Layer1: 2 × BasicBlock(64→64) → (64, 32, 32)
├─ Layer2: 2 × BasicBlock(64→128, ↓2) → (128, 16, 16)
├─ Layer3: 2 × BasicBlock(128→256, ↓2) → (256, 8, 8)
├─ AdaptiveAvgPool(1) → (256, 1, 1)
└─ Linear(256→10) → (10,)
Weight Initialization
Kaiming initialization for convolutional layers ensures variance is preserved through ReLU networks:
def _init_weights(self):
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, nn.BatchNorm2d):
nn.init.ones_(m.weight)
nn.init.zeros_(m.bias)
Transfer Learning
Transfer learning reuses features learned on a large dataset (e.g., ImageNet) and fine-tunes for a new task.
The Pattern
┌──────────────────────────────────────────────┐
│ Pretrained Backbone (frozen) │
│ Conv layers → learned features │
│ [Parameters: requires_grad = False] │
├──────────────────────────────────────────────┤
│ New Classification Head (trainable) │
│ Dropout → Linear → ReLU → Linear → Output │
│ [Parameters: requires_grad = True] │
└──────────────────────────────────────────────┘
Two-Phase Training
Phase 1: Train head only (backbone frozen)
model = TransferModel(backbone, feature_dim=256, num_classes=10, freeze_backbone=True)
optimizer = AdamW(model.head.parameters(), lr=1e-3) # Only head params
Phase 2: Fine-tune everything (backbone unfrozen with lower LR)
param_groups = model.unfreeze_backbone(lr_factor=0.1)
optimizer = AdamW([
{"params": model.backbone.parameters(), "lr": 1e-4}, # Lower LR
{"params": model.head.parameters(), "lr": 1e-3}, # Normal LR
])
Implementation
class TransferModel(nn.Module):
def __init__(self, backbone, feature_dim, num_classes=10, freeze_backbone=True):
super().__init__()
self.backbone = backbone
if freeze_backbone:
for p in self.backbone.parameters():
p.requires_grad = False
self.head = nn.Sequential(
nn.Dropout(0.3),
nn.Linear(feature_dim, 128),
nn.ReLU(inplace=True),
nn.Linear(128, num_classes),
)
def unfreeze_backbone(self, lr_factor=0.1):
for p in self.backbone.parameters():
p.requires_grad = True
return [
{"params": self.backbone.parameters(), "lr": lr_factor},
{"params": self.head.parameters()},
]
Training Pipeline
Components
| Component | Implementation | Purpose |
|---|---|---|
| Optimizer | AdamW | Weight decay decoupled from gradient |
| Loss | Label Smoothing CE | Prevents overconfident predictions |
| Scheduler | Cosine + Warmup | Smooth LR decay with warmup |
| AMP | torch.autocast | Mixed precision for speed |
| EMA | Exponential moving average | Smoother, more stable weights |
| MixUp/CutMix | Random per batch | Regularization |
| Early Stopping | Patience-based | Prevents overfitting |
| Gradient Clipping | Max norm = 1.0 | Training stability |
Label Smoothing
Instead of hard targets [0, 0, 1, 0, ...], use soft targets [ε/K, ε/K, 1-ε+ε/K, ε/K, ...]:
L = (1 - ε) · NLL(pred, target) + ε · mean(-log_probs)
With ε = 0.1, the model is penalized less for being "not 100% sure," which improves calibration.
class LabelSmoothingCrossEntropy(nn.Module):
def __init__(self, smoothing=0.1):
super().__init__()
self.smoothing = smoothing
def forward(self, pred, target):
log_probs = F.log_softmax(pred, dim=-1)
nll_loss = F.nll_loss(log_probs, target, reduction="none")
smooth_loss = -log_probs.mean(dim=-1)
return ((1 - self.smoothing) * nll_loss + self.smoothing * smooth_loss).mean()
Cosine Warmup Scheduler
LR
│ ╱──╲
│ ╱ ╲
│╱ ╲
│ warmup ╲ cosine decay
│ ╲
└────────────╲──────▶ epoch
0 3 20
class CosineWarmupScheduler(LRScheduler):
def __init__(self, optimizer, warmup_epochs, total_epochs, min_lr=1e-6):
...
def get_lr(self):
if self.last_epoch < self.warmup_epochs:
factor = self.last_epoch / max(1, self.warmup_epochs)
return [base_lr * factor for base_lr in self.base_lrs]
progress = (self.last_epoch - self.warmup_epochs) / (self.total_epochs - self.warmup_epochs)
cosine = 0.5 * (1 + cos(pi * progress))
return [self.min_lr + (base_lr - self.min_lr) * cosine for base_lr in self.base_lrs]
EMA (Exponential Moving Average)
Maintains a shadow copy of parameters that updates slowly:
shadow = decay · shadow + (1 - decay) · current_params
With decay = 0.999, the EMA model is a smoothed version of the training model — often generalizes better.
class EMA:
def __init__(self, model, decay=0.999):
self.decay = decay
self.shadow = {name: p.clone() for name, p in model.named_parameters() if p.requires_grad}
@torch.no_grad()
def update(self, model):
for name, p in model.named_parameters():
if p.requires_grad and name in self.shadow:
self.shadow[name].mul_(self.decay).add_(p.data, alpha=1 - self.decay)
Training Loop with AMP
for images, labels in loader:
images, labels = images.to(device), labels.to(device)
# Optional MixUp / CutMix
if use_mixup:
images, y_a, y_b, lam = mixup(images, labels, alpha=0.2)
# Forward with AMP
with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=use_amp):
logits = model(images)
loss = mixup_criterion(criterion, logits, y_a, y_b, lam)
# Backward with GradScaler
optimizer.zero_grad(set_to_none=True)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
# EMA update
ema.update(model)
Early Stopping
if val_acc > best_val_acc:
best_val_acc = val_acc
best_state = model.state_dict()
no_improve = 0
else:
no_improve += 1
if no_improve >= patience:
print(f"Early stopping at epoch {epoch}")
break
model.load_state_dict(best_state) # Restore best model
Evaluation Metrics
Beyond Accuracy
Accuracy alone can be misleading, especially with class imbalance. We compute:
| Metric | Formula | What It Tells You |
|---|---|---|
| Precision | TP / (TP + FP) | Of predicted positives, how many are correct |
| Recall | TP / (TP + FN) | Of actual positives, how many were found |
| F1 Score | 2 · P · R / (P + R) | Harmonic mean of precision and recall |
| Macro F1 | avg(F1_per_class) | Equal weight to each class |
| Weighted F1 | weighted avg(F1_per_class) | Weight by class support |
| Top-k Acc | correct in top-k / total | Useful when k>1 classes are reasonable |
Per-Class Report
metrics = ClassificationMetrics(num_classes=10, class_names=CLASS_NAMES)
for images, labels in test_loader:
logits = model(images)
metrics.update(logits, labels)
metrics.print_report()
Output:
Class Precision Recall F1 Support
--------------------------------------------------
circle 0.9200 0.9200 0.9200 50
square 0.8800 0.8800 0.8800 50
triangle 0.8600 0.8600 0.8600 50
...
--------------------------------------------------
Accuracy 0.8900
Macro F1 0.8900
Weighted F1 0.8900
Top-3 Acc 0.9700
Confusion Matrix
metrics.print_confusion_matrix()
The confusion matrix reveals which classes are most often confused with each other — a circle might be confused with a ring, for example.
Test-Time Augmentation
How TTA Works
At test time, apply multiple augmentations to the same image and average the predictions:
┌─── original ──── pred_1 ───┐
├─── h-flip ────── pred_2 ───┤
image ────┼─── v-flip ────── pred_3 ───┼── average → final prediction
├─── rot90 ─────── pred_4 ───┤
└─── rot180 ────── pred_5 ───┘
Implementation
class TTAAugmentation:
def __init__(self, num_augmentations=5):
self.augmentations = [
lambda x: x, # Original
lambda x: x.flip(-1), # Horizontal flip
lambda x: x.flip(-2), # Vertical flip
lambda x: torch.rot90(x, 1, [-2, -1]), # 90° rotation
lambda x: torch.rot90(x, 2, [-2, -1]), # 180° rotation
]
@torch.no_grad()
def predict(self, model, images):
all_probs = []
for aug_fn in self.augmentations[:self.num_augmentations]:
probs = F.softmax(model(aug_fn(images)), dim=-1)
all_probs.append(probs)
return torch.stack(all_probs).mean(dim=0)
When to Use TTA
| Scenario | Use TTA? |
|---|---|
| Competition / final submission | Yes — free accuracy boost |
| Real-time inference | No — multiplies latency by N |
| Medical imaging / safety-critical | Yes — reliability matters |
| Development / prototyping | No — slower iteration |
Typical improvement: +0.5–2% accuracy at the cost of N× inference time.
Grad-CAM Visualization
What Is Grad-CAM?
Gradient-weighted Class Activation Mapping produces a heatmap showing which spatial regions of the input image most influenced the model's prediction.
How It Works
- Forward pass: record activations at the target conv layer
- Backward pass: record gradients flowing into that layer
- Weight: global-average-pool the gradients → per-channel importance weights
- Combine: weighted sum of activation maps → heatmap
- ReLU: keep only positive contributions (features that increase the score)
Activations A (C, H, W) Gradients G (C, H, W)
│ │
│ GAP over (H,W)
│ │
│ Weights α (C,)
│ │
└─── Σ(αc · Ac) ──── ReLU ─── Upsample ─── Heatmap
Implementation
class GradCAM:
def __init__(self, model, target_layer):
self.model = model
target_layer.register_forward_hook(self._forward_hook)
target_layer.register_full_backward_hook(self._backward_hook)
def _forward_hook(self, module, input, output):
self.activations = output.detach()
def _backward_hook(self, module, grad_input, grad_output):
self.gradients = grad_output[0].detach()
def generate(self, input_tensor, target_class=None):
output = self.model(input_tensor)
if target_class is None:
target_class = output.argmax(dim=-1)
# Backward for target class
one_hot = torch.zeros_like(output)
one_hot[range(len(target_class)), target_class] = 1.0
output.backward(gradient=one_hot)
# Compute CAM
weights = self.gradients.mean(dim=(-2, -1), keepdim=True)
cam = (weights * self.activations).sum(dim=1, keepdim=True)
cam = F.relu(cam)
cam = F.interpolate(cam, size=input_tensor.shape[-2:], mode="bilinear")
# Normalize to [0, 1]
cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)
return cam.squeeze(1)
Using Grad-CAM
grad_cam = GradCAM(model, target_layer=model.features[-3])
heatmaps = grad_cam.generate(images)
# heatmaps shape: (B, H, W), values in [0, 1]
# High values = important regions
Confidence Calibration
Expected Calibration Error (ECE)
A well-calibrated model's confidence should match its accuracy:
- If the model says "90% confident" for a set of predictions, 90% should be correct.
ECE = Σ (|Bm| / N) · |accuracy(Bm) - confidence(Bm)|
Where Bm are bins of predictions grouped by confidence level.
def confidence_analysis(model, loader, device):
# Compute per-prediction confidence
probs = F.softmax(logits, dim=-1)
max_probs, preds = probs.max(dim=-1)
# Bin predictions by confidence
for bin in confidence_bins:
bin_acc = accuracy of predictions in this bin
bin_conf = mean confidence in this bin
ece += |bin_acc - bin_conf| * bin_size / total
return {
"mean_correct_confidence": ...,
"mean_incorrect_confidence": ...,
"ece": ece,
}
Interpreting ECE
| ECE | Calibration Quality |
|---|---|
| < 0.02 | Excellent |
| 0.02–0.05 | Good |
| 0.05–0.10 | Fair |
| > 0.10 | Poor — consider temperature scaling |
Inference Pipeline
Production-Ready Wrapper
class ImageClassifier:
def __init__(self, model, class_names, device, use_tta=False):
self.model = model.to(device).eval()
self.class_names = class_names
self.normalize = Normalize()
self.tta = TTAAugmentation(5) if use_tta else None
@torch.no_grad()
def predict(self, images, top_k=3):
images = self.normalize(images).to(self.device)
if self.tta:
probs = self.tta.predict(self.model, images)
else:
probs = F.softmax(self.model(images), dim=-1)
results = []
for i in range(images.size(0)):
top_probs, top_indices = probs[i].topk(top_k)
results.append({
"top_class": self.class_names[top_indices[0]],
"confidence": top_probs[0].item(),
"predictions": [...],
})
return results
Usage
classifier = ImageClassifier(model, CLASS_NAMES, device, use_tta=False)
results = classifier.predict(image_tensor)
print(results[0]["top_class"]) # "circle"
print(results[0]["confidence"]) # 0.95
Saving and Loading
# Save
checkpoint = {
"model_state_dict": model.state_dict(),
"class_names": CLASS_NAMES,
"num_classes": NUM_CLASSES,
"history": history,
}
torch.save(checkpoint, "classifier.pt")
# Load
checkpoint = torch.load("classifier.pt", weights_only=True)
model.load_state_dict(checkpoint["model_state_dict"])
torch.compile for Inference
compiled_model = torch.compile(model, mode="reduce-overhead")
# First call triggers compilation, subsequent calls are faster
with torch.no_grad():
output = compiled_model(images)
Running the Scripts
cd 40_image_classifier
# Step 1: Data pipeline and augmentation demo
python data_pipeline.py
# Step 2: Model training (trains SimpleCNN with full pipeline)
python model_and_training.py
# Step 3: Evaluation, TTA, Grad-CAM, confidence analysis
python evaluation.py
Key Takeaways
- Data augmentation is essential — random flips, color jitter, and random erasing significantly reduce overfitting on small datasets
- MixUp and CutMix create virtual samples — by blending images and labels, they smooth decision boundaries and improve generalization
- Residual connections enable depth — the skip connection in BasicBlock lets gradients flow through deep networks without degradation
- Transfer learning saves compute — freeze a pretrained backbone, train only the head, then optionally fine-tune with differential learning rates
- Label smoothing improves calibration — soft targets prevent overconfident predictions and reduce the gap between confidence and accuracy
- EMA stabilizes training — maintaining a moving average of parameters often yields a model that generalizes better than any single checkpoint
- TTA boosts accuracy for free — averaging predictions over augmented views of the same image typically adds +0.5–2% accuracy
- Grad-CAM reveals what the model sees — heatmaps show which spatial regions drive predictions, which is crucial for debugging and trust
- ECE measures calibration quality — a model's confidence should match its accuracy; ECE quantifies the gap
Further Resources
- Module 04 — Neural Networks —
nn.Module, layers, losses - Module 06 — Data Loading — Dataset, DataLoader, custom collate
- Module 07 — Training Pipelines — full training loops, mixed precision
- Module 12 — Model Architectures — ResNet, ViT complete implementations
- Module 29 — Mixed Precision — AMP, GradScaler deep dive
- Module 33 — Interpretability — Grad-CAM, saliency maps, hooks
- Module 39 — Text Classifier — End-to-end NLP classification project
- PyTorch Vision Transfer Learning Tutorial
Notebook: 40_image_classifier.ipynb
Source Files
data_pipeline.py— 270+model_and_training.py— 310+evaluation.py— 270+
Module 41: Building a Diffusion Model — From Noise to Data
Build a complete diffusion model from scratch: noise schedules, UNet architecture with time conditioning, DDPM training, DDPM and DDIM sampling, and classifier-free guidance — all demonstrated on 2D distributions for intuitive visualization.
No images required — we train on 2D point distributions (Swiss roll, moons, circles) so you can visualize the entire diffusion process in 2D.
| Input | Output |
|---|---|
Pure Gaussian noise (1000 points × 2D) | Swiss roll distribution |
Pure Gaussian noise (1000 points × 2D) | Two moons distribution |
50 DDIM steps (vs 1000 DDPM) | Same quality, 20× faster |
Table of Contents
- Overview
- What Are Diffusion Models?
- Forward Process — Adding Noise
- Noise Schedules
- Reverse Process — Denoising
- UNet Architecture
- DDPM Training
- DDPM Sampling
- DDIM Sampling
- Classifier-Free Guidance
- Training on 2D Distributions
- Key Takeaways
Overview
Diffusion models learn to generate data by reversing a gradual noising process. They are the backbone of modern image generators (Stable Diffusion, DALL-E 3, Imagen) and have achieved state-of-the-art results in image, audio, and video generation.
Forward Process (fixed): x_0 ──→ x_1 ──→ x_2 ──→ ··· ──→ x_T ~ N(0, I)
data slightly noisy pure noise
Reverse Process (learned): x_T ──→ x_{T-1} ──→ ··· ──→ x_1 ──→ x_0
noise slightly denoised generated data
Architecture
┌─────────────────────────────────────────────────────────────────────────┐
│ DIFFUSION MODEL PIPELINE │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌───────────┐ ┌──────────────┐ ┌────────────┐ ┌───────────┐ │
│ │ Noise │ │ Forward │ │ UNet │ │ Training │ │
│ │ Schedule │───▶│ Process │───▶│ (denoise) │───▶│ Loop │ │
│ │ (beta_t) │ │ q(x_t|x_0) │ │ eps_theta │ │ MSE loss │ │
│ └───────────┘ └──────────────┘ └────────────┘ └─────┬─────┘ │
│ │ │
│ ┌───────────┐ ┌──────────────┐ ┌────────────┐ ┌─────▼─────┐ │
│ │ Generated │ │ DDIM │ │ DDPM │ │ Trained │ │
│ │ Samples │◀───│ Sampling │◀───│ Sampling │◀───│ Model │ │
│ └───────────┘ └──────────────┘ └────────────┘ └───────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────┘
Files
| File | Lines | Description |
|---|---|---|
noise_schedule.py | 200+ | Linear/cosine beta schedules, forward diffusion, alpha_cumprod |
unet_model.py | 300+ | Sinusoidal embeddings, ResBlocks, UNet with skip connections |
train_diffusion.py | 300+ | 2D data generation, DDPM/DDIM training and sampling, visualization |
What Are Diffusion Models?
Diffusion models are a class of generative models that learn a data distribution by:
- Defining a forward process that gradually destroys data by adding Gaussian noise over T steps
- Learning a reverse process (a neural network) that removes noise one step at a time
The key insight: destroying data is easy (just add noise); learning to reverse this destruction teaches the model what real data looks like.
Comparison with Other Generative Models
| Model | Training | Sampling | Mode Coverage | Quality |
|---|---|---|---|---|
| GANs | Adversarial (unstable) | Single pass (fast) | Mode collapse risk | High |
| VAEs | ELBO (stable) | Single pass (fast) | Good coverage | Blurry |
| Diffusion | Denoising (stable) | Iterative (slow) | Excellent coverage | Highest |
| Flow | Exact likelihood | Single pass | Good coverage | High |
Mathematical Foundation
A diffusion model defines a Markov chain of latent variables x_1, ..., x_T:
q(x_{1:T} | x_0) = prod_{t=1}^{T} q(x_t | x_{t-1})
The forward process adds Gaussian noise at each step:
q(x_t | x_{t-1}) = N(x_t; sqrt(1 - beta_t) * x_{t-1}, beta_t * I)
Where beta_t is a variance schedule that controls how much noise is added at step t.
Forward Process — Adding Noise
Step-by-Step Noising
At each timestep t, we add a small amount of Gaussian noise:
x_t = sqrt(1 - beta_t) * x_{t-1} + sqrt(beta_t) * epsilon
Where epsilon ~ N(0, I).
Closed-Form Sampling at Arbitrary Timestep
A key property: we can sample x_t directly from x_0 without iterating through all intermediate steps:
alpha_t = 1 - beta_t
alpha_bar_t = prod_{s=1}^{t} alpha_s (cumulative product)
q(x_t | x_0) = N(x_t; sqrt(alpha_bar_t) * x_0, (1 - alpha_bar_t) * I)
Or equivalently:
x_t = sqrt(alpha_bar_t) * x_0 + sqrt(1 - alpha_bar_t) * epsilon
This is crucial for efficient training — we can jump to any timestep t directly.
Implementation
def q_sample(x_0, t, noise, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod):
"""Sample x_t from q(x_t | x_0) — the forward process."""
sqrt_alpha = sqrt_alphas_cumprod[t].unsqueeze(-1)
sqrt_one_minus = sqrt_one_minus_alphas_cumprod[t].unsqueeze(-1)
return sqrt_alpha * x_0 + sqrt_one_minus * noise
Noise Level Progression
t=0: x_0 (clean data) alpha_bar ≈ 1.0
t=250: x_250 (slightly noisy) alpha_bar ≈ 0.7
t=500: x_500 (noisy) alpha_bar ≈ 0.3
t=750: x_750 (very noisy) alpha_bar ≈ 0.05
t=1000: x_T (pure noise) alpha_bar ≈ 0.0
Noise Schedules
The noise schedule {beta_1, ..., beta_T} controls how quickly data is destroyed. The choice of schedule significantly affects training and generation quality.
Linear Schedule
The simplest schedule — linearly interpolate between beta_start and beta_end:
beta_t = beta_start + (beta_end - beta_start) * t / T
Typical values: beta_start = 0.0001, beta_end = 0.02, T = 1000.
def linear_beta_schedule(timesteps, beta_start=1e-4, beta_end=0.02):
return torch.linspace(beta_start, beta_end, timesteps)
Problem: the linear schedule destroys information too quickly at the end. By t=600, most signal is already gone, wasting the remaining 400 steps.
Cosine Schedule
Proposed in "Improved DDPM" (Nichol & Dhariwal, 2021). Designs alpha_bar_t to follow a cosine curve:
alpha_bar_t = f(t) / f(0)
where f(t) = cos((t/T + s) / (1 + s) * pi/2)^2
The offset s = 0.008 prevents beta_t from being too small near t = 0.
def cosine_beta_schedule(timesteps, s=0.008):
steps = torch.linspace(0, timesteps, timesteps + 1)
alphas_cumprod = torch.cos((steps / timesteps + s) / (1 + s) * math.pi * 0.5) ** 2
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
return torch.clamp(betas, 0.0001, 0.9999)
Linear vs Cosine Comparison
alpha_bar_t
1.0 ┤
│ ╲ cosine (gradual)
0.8 ┤ ╲
│ ╲╲
0.6 ┤ ╲╲
│ ╲ ╲
0.4 ┤ ╲ ╲ linear (aggressive)
│ ╲ ╲
0.2 ┤ ╲ ╲
│ ╲ ╲
0.0 ┤ ╲──╲──
└───────────────────▶ t
0 200 400 600 800 1000
The cosine schedule distributes information destruction more evenly across timesteps, leading to better sample quality.
Derived Quantities
From the beta schedule, we precompute all quantities needed for training and sampling:
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value=1.0)
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
sqrt_recip_alphas = torch.sqrt(1.0 / alphas)
posterior_variance = betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)
Reverse Process — Denoising
The Goal
Learn a model p_theta(x_{t-1} | x_t) that reverses the forward process:
p_theta(x_{t-1} | x_t) = N(x_{t-1}; mu_theta(x_t, t), sigma_t^2 * I)
Predicting Noise vs Predicting x_0
The model can parameterize the reverse process in several ways:
| Parameterization | Model predicts | Used by |
|---|---|---|
| Epsilon (noise) | eps_theta(x_t, t) | DDPM (default) |
| x_0 (clean data) | x_0_theta(x_t, t) | Some methods |
| v (velocity) | v_theta(x_t, t) | Progressive distillation |
| Score | score_theta(x_t, t) | Score matching |
We use epsilon prediction (the DDPM default). Given the predicted noise eps_theta, the mean of p_theta is:
mu_theta(x_t, t) = 1/sqrt(alpha_t) * (x_t - beta_t/sqrt(1 - alpha_bar_t) * eps_theta(x_t, t))
Why Noise Prediction Works
The model's task is simple: "given a noisy version of the data and the noise level t, predict what noise was added." This is equivalent to learning the score function (gradient of log-density), which is theoretically well-justified.
UNet Architecture
The denoising network takes a noisy sample x_t and a timestep t, and predicts the noise epsilon. For 2D point diffusion, we use a simplified UNet-style architecture.
Sinusoidal Time Embedding
Timesteps are embedded using sinusoidal positional encoding (from the Transformer paper), which gives the model a smooth, continuous representation of the noise level:
PE(t, 2i) = sin(t / 10000^(2i/d))
PE(t, 2i+1) = cos(t / 10000^(2i/d))
class SinusoidalTimeEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, t):
half = self.dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half)
args = t[:, None].float() * freqs[None, :]
return torch.cat([torch.sin(args), torch.cos(args)], dim=-1)
ResBlock with Time Conditioning
Each residual block receives the time embedding, allowing the network to adapt its behavior based on the noise level:
x ──┬── Linear → SiLU → Linear ──┐
│ + time_emb │
└──────── Shortcut ───────────┘
│
Output
class ResBlock(nn.Module):
def __init__(self, dim, time_dim):
super().__init__()
self.net = nn.Sequential(
nn.Linear(dim, dim),
nn.SiLU(),
nn.Linear(dim, dim),
)
self.time_mlp = nn.Sequential(
nn.SiLU(),
nn.Linear(time_dim, dim),
)
self.shortcut = nn.Identity()
def forward(self, x, t_emb):
return self.shortcut(x) + self.net(x) + self.time_mlp(t_emb)
UNet Structure (for 2D Points)
For 2D point data, our UNet operates on feature vectors rather than spatial grids:
Input (B, 2)
│
├─ Project to hidden dim → (B, 256)
│
├─ DownBlock 1: ResBlock(256) + ResBlock(256) → (B, 256) ─── skip_1
├─ Down: Linear(256 → 128) → (B, 128)
├─ DownBlock 2: ResBlock(128) + ResBlock(128) → (B, 128) ─── skip_2
├─ Down: Linear(128 → 64) → (B, 64)
│
├─ MidBlock: ResBlock(64) + ResBlock(64) → (B, 64)
│
├─ Up: Linear(64 → 128) → (B, 128)
├─ UpBlock 2: ResBlock(256) + ResBlock(128) → (B, 128) ← cat(skip_2)
├─ Up: Linear(128 → 256) → (B, 256)
├─ UpBlock 1: ResBlock(512) + ResBlock(256) → (B, 256) ← cat(skip_1)
│
├─ Project to output dim → (B, 2)
│
Output: predicted noise (B, 2)
Skip Connections
Skip connections from the downsampling path to the upsampling path are essential — they preserve fine-grained information that would be lost through the bottleneck:
class PointUNet(nn.Module):
def forward(self, x, t):
t_emb = self.time_mlp(self.time_embed(t))
x = self.input_proj(x)
# Down path (save skips)
skips = []
for down_block, downsample in self.downs:
x = down_block(x, t_emb)
skips.append(x)
x = downsample(x)
# Mid
x = self.mid(x, t_emb)
# Up path (use skips)
for up_block, upsample in self.ups:
x = upsample(x)
x = torch.cat([x, skips.pop()], dim=-1) # Skip connection
x = up_block(x, t_emb)
return self.output_proj(x)
Simple 1D UNet for Understanding
To build intuition, we also provide a minimal 1D UNet that processes scalar inputs — the simplest possible diffusion setup:
class SimpleUNet1D(nn.Module):
"""Minimal UNet for 1D data — useful for understanding the core idea."""
def __init__(self, time_dim=32):
super().__init__()
self.time_embed = SinusoidalTimeEmbedding(time_dim)
self.net = nn.Sequential(
nn.Linear(1 + time_dim, 128),
nn.SiLU(),
nn.Linear(128, 128),
nn.SiLU(),
nn.Linear(128, 1),
)
def forward(self, x, t):
t_emb = self.time_embed(t)
return self.net(torch.cat([x, t_emb], dim=-1))
DDPM Training
Training Algorithm
DDPM (Denoising Diffusion Probabilistic Models) training is remarkably simple:
repeat:
1. Sample x_0 ~ data distribution
2. Sample t ~ Uniform{1, ..., T}
3. Sample epsilon ~ N(0, I)
4. Compute x_t = sqrt(alpha_bar_t) * x_0 + sqrt(1 - alpha_bar_t) * epsilon
5. Predict eps_theta = model(x_t, t)
6. Loss = MSE(epsilon, eps_theta)
7. Gradient step
Implementation
def train_step(model, optimizer, x_0, schedule):
batch_size = x_0.shape[0]
# Sample random timesteps
t = torch.randint(0, schedule.num_timesteps, (batch_size,), device=x_0.device)
# Sample noise
noise = torch.randn_like(x_0)
# Forward process: add noise to get x_t
x_t = schedule.q_sample(x_0, t, noise)
# Predict noise
predicted_noise = model(x_t, t)
# MSE loss between true and predicted noise
loss = F.mse_loss(predicted_noise, noise)
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss.item()
Why MSE on Noise?
The simplified DDPM objective is:
L_simple = E_{t, x_0, eps} [ ||eps - eps_theta(x_t, t)||^2 ]
This is a simplified version of the variational lower bound (VLB). It works because:
- Predicting noise is equivalent to predicting the score (gradient of log probability)
- MSE on noise is equivalent to a reweighted VLB where all timesteps contribute equally
- Empirically, this simple objective produces better samples than the full VLB
Training Tips
| Tip | Reason |
|---|---|
| Use cosine schedule | Better noise distribution across timesteps |
| AdamW optimizer | Stable training with weight decay |
| Learning rate ~1e-3 for 2D, ~2e-4 for images | 2D data is simpler |
| Gradient clipping (max_norm=1.0) | Prevents training instability |
| EMA of model weights | Smoother, better-quality samples |
DDPM Sampling
Algorithm
DDPM sampling iterates from pure noise x_T back to clean data x_0:
x_T ~ N(0, I)
for t = T, T-1, ..., 1:
z ~ N(0, I) if t > 1, else z = 0
x_{t-1} = 1/sqrt(alpha_t) * (x_t - beta_t/sqrt(1-alpha_bar_t) * eps_theta(x_t, t)) + sigma_t * z
Where sigma_t = sqrt(beta_t) (the posterior standard deviation).
Implementation
@torch.no_grad()
def ddpm_sample(model, schedule, shape, device):
x = torch.randn(shape, device=device) # Start from pure noise
for t in reversed(range(schedule.num_timesteps)):
t_batch = torch.full((shape[0],), t, device=device, dtype=torch.long)
# Predict noise
predicted_noise = model(x, t_batch)
# Compute mean
alpha = schedule.alphas[t]
alpha_bar = schedule.alphas_cumprod[t]
beta = schedule.betas[t]
mean = (1 / alpha.sqrt()) * (x - (beta / (1 - alpha_bar).sqrt()) * predicted_noise)
# Add noise (except at t=0)
if t > 0:
noise = torch.randn_like(x)
sigma = beta.sqrt()
x = mean + sigma * noise
else:
x = mean
return x
DDPM Sampling Properties
| Property | Value |
|---|---|
| Steps | T (typically 1000) |
| Stochastic? | Yes (random noise at each step) |
| Quality | Excellent |
| Speed | Slow (1000 forward passes) |
| Deterministic? | No (different noise = different samples) |
DDIM Sampling
Motivation
DDPM requires T steps (e.g., 1000) for sampling, which is slow. DDIM (Denoising Diffusion Implicit Models, Song et al. 2021) enables sampling with far fewer steps by using a non-Markovian process.
Key Idea
DDIM defines a family of non-Markovian forward processes that all share the same marginal q(x_t | x_0) as DDPM. The reverse process can skip timesteps:
x_{t-1} = sqrt(alpha_bar_{t-1}) * predicted_x_0 + sqrt(1 - alpha_bar_{t-1} - sigma_t^2) * predicted_direction + sigma_t * noise
Where:
predicted_x_0 = (x_t - sqrt(1 - alpha_bar_t) * eps_theta) / sqrt(alpha_bar_t)predicted_direction = eps_theta(pointing toward x_t)sigma_t = 0for deterministic sampling (eta = 0)
Implementation
@torch.no_grad()
def ddim_sample(model, schedule, shape, device, num_steps=50, eta=0.0):
# Create subsequence of timesteps
step_size = schedule.num_timesteps // num_steps
timesteps = list(range(0, schedule.num_timesteps, step_size))[::-1]
x = torch.randn(shape, device=device)
for i, t in enumerate(timesteps):
t_batch = torch.full((shape[0],), t, device=device, dtype=torch.long)
predicted_noise = model(x, t_batch)
alpha_bar_t = schedule.alphas_cumprod[t]
alpha_bar_prev = schedule.alphas_cumprod[timesteps[i+1]] if i < len(timesteps)-1 else torch.tensor(1.0)
# Predict x_0
pred_x0 = (x - (1 - alpha_bar_t).sqrt() * predicted_noise) / alpha_bar_t.sqrt()
pred_x0 = pred_x0.clamp(-3, 3) # Clip for stability
# Direction pointing to x_t
direction = (1 - alpha_bar_prev).sqrt() * predicted_noise
# Noise (eta=0 for deterministic)
sigma = eta * ((1 - alpha_bar_prev) / (1 - alpha_bar_t) * (1 - alpha_bar_t / alpha_bar_prev)).sqrt()
x = alpha_bar_prev.sqrt() * pred_x0 + direction
if eta > 0 and i < len(timesteps) - 1:
x = x + sigma * torch.randn_like(x)
return x
DDPM vs DDIM Comparison
| Aspect | DDPM | DDIM |
|---|---|---|
| Steps | 1000 | 50-100 (tunable) |
| Stochastic | Yes | Configurable (eta) |
| Deterministic mode | No | Yes (eta=0) |
| Sample quality at 50 steps | Poor | Good |
| Same noise → same output | No | Yes (when eta=0) |
| Speed | ~1000 forward passes | ~50 forward passes |
Choosing the Number of Steps
Steps Quality Speed
10 Poor Very fast (20ms)
25 Decent Fast (50ms)
50 Good Moderate (100ms)
100 Very good Slow (200ms)
250 Excellent Very slow (500ms)
1000 Best DDPM speed (2s)
Classifier-Free Guidance
Concept
Classifier-free guidance (Ho & Salimans, 2022) improves sample quality by steering generation toward a condition (e.g., class label) without needing a separate classifier.
How It Works
During training, randomly drop the condition with some probability (e.g., 10%):
# Training: randomly use unconditional or conditional
if random() < 0.1:
eps_theta = model(x_t, t, condition=None) # Unconditional
else:
eps_theta = model(x_t, t, condition=class_label) # Conditional
During sampling, interpolate between conditional and unconditional predictions:
eps_guided = eps_unconditional + w * (eps_conditional - eps_unconditional)
Where w is the guidance scale:
- w = 1.0: no guidance (normal conditional generation)
- w > 1.0: stronger guidance (higher quality, less diversity)
- w = 7.5: common default for image generation
Intuition
guidance strength w
│
│ w=1: balanced (diverse but sometimes off-target)
│ w=3: moderate guidance (good quality, good diversity)
│ w=7: strong guidance (high quality, less diversity)
│ w=20: extreme (very sharp but repetitive)
│
The model learns both what the data looks like (unconditional) and what data with a specific condition looks like (conditional). Guidance amplifies the difference, pushing samples more strongly toward the conditioned distribution.
Training on 2D Distributions
Why 2D?
Training on 2D point distributions lets you:
- Visualize the entire forward/reverse process — scatter plots show noise being added and removed
- See mode coverage — verify the model generates all modes of the distribution
- Fast iteration — training takes seconds, not hours
- Understand failure modes — easy to spot mode collapse, poor mixing, etc.
Available Distributions
| Distribution | Shape | Characteristics |
|---|---|---|
| Swiss Roll | Spiral | Tests ability to learn curved manifolds |
| Two Moons | Two crescents | Tests multi-modal generation |
| Circles | Concentric rings | Tests ring-shaped distributions |
Data Generation
def make_swiss_roll(n_samples=1000):
t = 1.5 * math.pi * (1 + 2 * torch.rand(n_samples))
x = t * torch.cos(t)
y = t * torch.sin(t)
data = torch.stack([x, y], dim=-1)
data = data / data.std() # Normalize
return data
def make_moons(n_samples=1000):
n = n_samples // 2
# Upper moon
theta1 = torch.linspace(0, math.pi, n)
x1, y1 = torch.cos(theta1), torch.sin(theta1)
# Lower moon
theta2 = torch.linspace(0, math.pi, n_samples - n)
x2, y2 = 1 - torch.cos(theta2), 1 - torch.sin(theta2) - 0.5
x = torch.cat([x1, x2]) + torch.randn(n_samples) * 0.05
y = torch.cat([y1, y2]) + torch.randn(n_samples) * 0.05
data = torch.stack([x, y], dim=-1)
data = (data - data.mean(0)) / data.std()
return data
Visualizing the Forward Process
t=0 (clean) t=250 t=500 t=750 t=1000 (noise)
Swiss roll Blurred Fuzzy blob Nearly noise Pure Gaussian
·····••• ··· ··· · · · · · · · · · · · ·
· •• ·· ·· · · · · · · · · ·
· ••• • ·· · · · · · · · · · · ·
· • • • · · · · · · · · · · · · ·
·••••• ·· ··· · · · · ·· · · · · · ·
Visualizing the Reverse Process (Sampling)
t=1000 (start) t=750 t=500 t=250 t=0 (generated)
Random noise Structure Rough shape Refined Swiss roll
· · · · · · · · ·····• ·····•• ·····•••
· · · · · · · · • · •• · ••
· · · · · · · · ••• • · ••• • · ••• •
· · · · · · · • • • · • • • · • • •
· · · · · · · · ·••••• ·••••• ·•••••
Key Takeaways
- Diffusion models learn by denoising — the forward process adds noise; the model learns to reverse it, implicitly learning the data distribution
- Closed-form forward process is key — q(x_t | x_0) lets us jump to any timestep directly, making training efficient (just MSE on predicted noise)
- The noise schedule matters — cosine schedule distributes information destruction more evenly than linear, leading to better samples
- UNet with skip connections preserves detail — the encoder-decoder structure with skip connections lets the model process at multiple resolutions
- Time embedding conditions the network — sinusoidal embeddings give the model continuous knowledge of the current noise level
- DDPM is slow but high quality — 1000 iterative denoising steps produce excellent samples but take time
- DDIM enables fast sampling — by using a non-Markovian process, DDIM generates comparable samples in 50 steps (20x faster) and supports deterministic generation
- Classifier-free guidance trades diversity for quality — by interpolating between conditional and unconditional predictions, guidance produces sharper, more targeted samples
- 2D distributions build intuition — visualizing diffusion on points before scaling to images reveals the core mechanics without GPU overhead
Further Resources
- Module 04 — Neural Networks —
nn.Module, layers, losses - Module 07 — Training Pipelines — complete training loops, mixed precision
- Module 12 — Model Architectures — ResNet, VAE implementations
- Module 40 — Image Classifier — End-to-end CNN/ResNet project
- DDPM Paper — Ho et al. 2020
- Improved DDPM — Nichol & Dhariwal 2021
- DDIM — Song et al. 2021
- Classifier-Free Guidance — Ho & Salimans 2022
Upstream Updates (PyTorch 2.14+)
| Feature | Impact on Diffusion Models |
|---|---|
torch.compile | Compiles the denoising UNet for faster training and sampling |
| FlexAttention | Custom attention patterns for UNet self-attention layers |
torch.float8 | FP8 training for larger diffusion models |
| FSDP2 | Distributed training of billion-parameter diffusion models |
torch.export | Export trained diffusion models for deployment |
Notebook: 41_diffusion_model.ipynb
Source Files
noise_schedule.py— 200+unet_model.py— 300+train_diffusion.py— 300+
Operational guide
Targeted Test Selection
Running the full PyTorch test suite takes hours. Targeted test selection identifies which tests are affected by a code change and runs only those, reducing CI time from hours to minutes.
Approaches
1. File-Path Heuristic (targeted_tests.py)
Maps changed source files to test files using submodule-level rules:
SUBMODULE_MAP = {
"torch/nn/": [("test/test_nn.py", None)],
"torch/optim/": [("test/test_optim.py", None)],
"torch/cuda/": [("test/test_cuda.py", None)],
"torch/fx/": [("test/test_fx.py", None)],
"aten/src/ATen/": [("test/test_torch.py", "test_type")],
}
Usage:
python targeted_tests.py OLD_SHA NEW_SHA --pytorch-dir /pytorch --category cpu --commands-only
# Output:
# python test/run_test.py -i test_nn
# python test/run_test.py -i test_optim -k "test_adam"
2. Structural Analysis (TorchTalk)
For C++ changes, file-path heuristics miss transitive callers. TorchTalk uses libclang to:
- Parse
compile_commands.jsonfor accurate compilation flags - Extract changed C++ symbols
- Walk the call graph to find all callers
- Map callers through pybind11/TORCH_LIBRARY bindings
- Resolve to Python test files and classes
python torchtalk_tests.py OLD_SHA NEW_SHA --pytorch-dir /pytorch --commands-only
3. Unified Merger
Combines both approaches by taking their union:
python merge_test_results.py OLD_SHA NEW_SHA --pytorch-dir /pytorch --category cpu --commands-only
If TorchTalk is unavailable, gracefully falls back to heuristic only.
Test Categories
Tests are classified into 4 categories for parallel execution:
| Category | Prefix patterns | Runner requirement |
|---|---|---|
| cpu | test/test_*.py (default) | CPU only |
| inductor | test/inductor/, test/dynamo/, test/export/ | CPU (some GPU) |
| sgpu | test/test_cuda* | 1 GPU |
| mgpu | test/distributed/ | 2+ GPUs |
run_test.py Format
PyTorch's test runner (test/run_test.py) handles path resolution, environment setup, and timeout management:
# Run entire test file
python test/run_test.py -i test_torch
# Run with keyword filter
python test/run_test.py -i test_nn -k "test_linear or test_conv2d"
# Subdirectory tests use / separator
python test/run_test.py -i nn/test_multihead_attention
# Inductor tests
python test/run_test.py -i inductor/test_torchinductor
Full-Suite Triggers
Some changes are too broad for targeted selection and trigger the full test suite:
FULL_SUITE_TRIGGERS = [
"setup.py",
"CMakeLists.txt",
"torch/__init__.py",
"torch/csrc/", # Core C++ runtime
"c10/", # Core library
"caffe2/", # Legacy but affects build
".ci/", # CI infrastructure
]
CI Integration Example
determine-tests:
steps:
- name: Resolve tests
run: |
for cat in cpu inductor sgpu mgpu; do
CMDS=$(python merge_test_results.py $PREV_SHA $HEAD_SHA \
--pytorch-dir /pytorch --category $cat --commands-only)
echo "${cat}_tests=$(echo "$CMDS" | base64 -w0)" >> $GITHUB_OUTPUT
done
cpu-tests:
needs: determine-tests
steps:
- name: Run
run: |
COMMANDS=$(echo "${{ needs.determine-tests.outputs.cpu_tests }}" | base64 -d)
while IFS= read -r cmd; do
eval "$cmd"
done <<< "$COMMANDS"
Key Repositories
- pytorch-targeted-tests — Heuristic engine
- TorchTalk — Structural C++ analysis
- pytorch-redhat-ci — Integration example
Performance Metrics
| Approach | Avg. tests selected | Time saved vs full suite |
|---|---|---|
| File-path heuristic | ~5-15% of suite | 70-85% wall time |
| Structural (TorchTalk) | ~3-10% of suite | 80-90% wall time |
| Merged (union) | ~8-20% of suite | 65-80% wall time |
| Full suite | 100% | Baseline (~4-6 hours) |
Extending the Mapping
To add a new test mapping for a source directory:
# In targeted_tests.py, add to SUBMODULE_MAP:
SUBMODULE_MAP = {
# existing mappings...
"torch/my_feature/": [
("test/test_my_feature.py", None), # run entire file
("test/test_related.py", "test_my_func"), # run specific test
],
}
The tuple format is (test_file, optional_keyword_filter). When the keyword is None, the entire test file runs. Otherwise, it becomes a -k filter passed to run_test.py.
FAQ
Q: What if targeted tests miss a regression? Nightly full-suite CI catches regressions that slip through targeted selection. The merge of heuristic + structural analysis minimizes false negatives.
Q: Can I run targeted tests locally? Yes: python targeted_tests.py HEAD~1 HEAD --pytorch-dir . --category cpu --commands-only prints the commands you can paste into your terminal.
Q: How does TorchTalk handle header-only changes? Header changes in c10/ or aten/src/ATen/core/ trigger the full suite via FULL_SUITE_TRIGGERS, since their call graph is too broad to scope.
Module 42: Building a RAG Pipeline with PyTorch
Overview
Retrieval-Augmented Generation (RAG) combines a retrieval system with a generative language model. Instead of relying solely on parametric knowledge, the model retrieves relevant context from an external knowledge base before generating a response. This module builds a complete RAG pipeline using only PyTorch and standard libraries — no LangChain, no vector database services.
Architecture
┌─────────────────────────────────────────────────────────────┐
│ RAG Pipeline │
│ │
│ Query ──► Encoder ──► Similarity Search ──► Top-K Docs │
│ │ │
│ ▼ │
│ Context + Query ──► Generator ──► Response │
└─────────────────────────────────────────────────────────────┘
Components
1. Document Encoder
Encodes documents and queries into dense vector representations using a pre-trained transformer. We use mean-pooling over token embeddings to produce fixed-size vectors.
import torch
import torch.nn.functional as F
from transformers import AutoTokenizer, AutoModel
def mean_pooling(model_output, attention_mask):
token_embeddings = model_output.last_hidden_state
input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(
input_mask_expanded.sum(1), min=1e-9
)
class DocumentEncoder:
def __init__(self, model_name="sentence-transformers/all-MiniLM-L6-v2", device="cpu"):
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
self.model = AutoModel.from_pretrained(model_name).to(device)
self.device = device
@torch.no_grad()
def encode(self, texts: list[str]) -> torch.Tensor:
encoded = self.tokenizer(
texts, padding=True, truncation=True, max_length=512, return_tensors="pt"
).to(self.device)
output = self.model(**encoded)
embeddings = mean_pooling(output, encoded["attention_mask"])
return F.normalize(embeddings, p=2, dim=1)
2. Vector Index (FAISS-free, Pure PyTorch)
A simple but effective in-memory vector store using cosine similarity:
class VectorIndex:
def __init__(self):
self.embeddings: torch.Tensor | None = None
self.documents: list[str] = []
def add(self, documents: list[str], embeddings: torch.Tensor):
if self.embeddings is None:
self.embeddings = embeddings
else:
self.embeddings = torch.cat([self.embeddings, embeddings], dim=0)
self.documents.extend(documents)
def search(self, query_embedding: torch.Tensor, top_k: int = 5) -> list[tuple[str, float]]:
similarities = torch.mm(query_embedding, self.embeddings.T).squeeze(0)
scores, indices = torch.topk(similarities, min(top_k, len(self.documents)))
return [(self.documents[idx], scores[i].item()) for i, idx in enumerate(indices)]
3. Context Assembly
Assembles retrieved documents into a prompt for the generator:
def build_prompt(query: str, retrieved_docs: list[tuple[str, float]], max_context_tokens: int = 1024) -> str:
context_parts = []
for doc, score in retrieved_docs:
context_parts.append(f"[Relevance: {score:.3f}] {doc}")
context = "\n\n".join(context_parts)
return f"""Answer the question based on the provided context.
Context:
{context}
Question: {query}
Answer:"""
4. Generator
Uses a causal language model to generate answers conditioned on the retrieved context:
class RAGGenerator:
def __init__(self, model_name="gpt2", device="cpu"):
from transformers import AutoModelForCausalLM, AutoTokenizer
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
self.model = AutoModelForCausalLM.from_pretrained(model_name).to(device)
self.device = device
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
@torch.no_grad()
def generate(self, prompt: str, max_new_tokens: int = 200, temperature: float = 0.7) -> str:
inputs = self.tokenizer(prompt, return_tensors="pt", truncation=True, max_length=1024).to(
self.device
)
outputs = self.model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
do_sample=True,
top_p=0.9,
pad_token_id=self.tokenizer.pad_token_id,
)
generated = outputs[0][inputs["input_ids"].shape[1] :]
return self.tokenizer.decode(generated, skip_special_tokens=True)
Complete Pipeline
class RAGPipeline:
def __init__(self, encoder: DocumentEncoder, index: VectorIndex, generator: RAGGenerator):
self.encoder = encoder
self.index = index
self.generator = generator
def ingest(self, documents: list[str], batch_size: int = 32):
for i in range(0, len(documents), batch_size):
batch = documents[i : i + batch_size]
embeddings = self.encoder.encode(batch)
self.index.add(batch, embeddings)
def query(self, question: str, top_k: int = 5) -> str:
query_emb = self.encoder.encode([question])
retrieved = self.index.search(query_emb, top_k=top_k)
prompt = build_prompt(question, retrieved)
return self.generator.generate(prompt)
Key Concepts
Why RAG?
| Approach | Pros | Cons |
|---|---|---|
| Fine-tuning | Fast inference, no retrieval latency | Expensive, stale knowledge, hallucinations |
| RAG | Up-to-date knowledge, verifiable sources | Retrieval latency, context window limits |
| RAG + Fine-tuning | Best of both worlds | Most complex to build and maintain |
Embedding Quality Matters
The quality of your retrieval depends entirely on the embedding model. Key considerations:
- Domain adaptation: General-purpose encoders may not capture domain-specific semantics
- Chunking strategy: Document splitting affects what gets retrieved
- Normalization: Always L2-normalize embeddings for cosine similarity
Chunking Strategies
def chunk_by_sentences(text: str, chunk_size: int = 3) -> list[str]:
sentences = text.split(". ")
chunks = []
for i in range(0, len(sentences), chunk_size):
chunk = ". ".join(sentences[i : i + chunk_size])
if not chunk.endswith("."):
chunk += "."
chunks.append(chunk)
return chunks
def chunk_by_tokens(text: str, tokenizer, max_tokens: int = 256, overlap: int = 50) -> list[str]:
tokens = tokenizer.encode(text)
chunks = []
start = 0
while start < len(tokens):
end = start + max_tokens
chunk_tokens = tokens[start:end]
chunks.append(tokenizer.decode(chunk_tokens))
start += max_tokens - overlap
return chunks
Performance Optimization
Batch Encoding with torch.compile
@torch.compile
def batch_encode_optimized(model, input_ids, attention_mask):
output = model(input_ids=input_ids, attention_mask=attention_mask)
return mean_pooling(output, attention_mask)
GPU-Accelerated Similarity Search
For large indices (>100K documents), move embeddings to GPU:
class GPUVectorIndex(VectorIndex):
def __init__(self, device="cuda"):
super().__init__()
self.search_device = device
def search(self, query_embedding: torch.Tensor, top_k: int = 5) -> list[tuple[str, float]]:
query_gpu = query_embedding.to(self.search_device)
embeddings_gpu = self.embeddings.to(self.search_device)
similarities = torch.mm(query_gpu, embeddings_gpu.T).squeeze(0)
scores, indices = torch.topk(similarities, min(top_k, len(self.documents)))
return [(self.documents[idx.item()], scores[i].item()) for i, idx in enumerate(indices)]
Files in This Module
| File | Description |
|---|---|
README.md | This guide |
rag_pipeline.py | Complete RAG pipeline implementation |
chunking_strategies.py | Document chunking utilities |
evaluation.py | RAG evaluation metrics (retrieval recall, answer quality) |
References
- Retrieval-Augmented Generation for Knowledge-Intensive NLP Tasks (Lewis et al., 2020)
- Dense Passage Retrieval (Karpukhin et al., 2020)
- Sentence-BERT (Reimers & Gurevych, 2019)
Source Files
rag_pipeline.py— Complete RAG pipeline implementationchunking_strategies.py— Document chunking utilitiesevaluation.py— RAG evaluation metrics (retrieval recall, answer quality)
Module 43: Production Serving Patterns
Production-grade inference serving with PyTorch: from single-request latency optimization to high-throughput batched serving with monitoring.
Topics Covered
- Batched Inference — Vectorized forward passes over variable-length inputs with padding/masking
- Dynamic Batching — Accumulate requests over a time window, dispatch as a single batch
- torch.compile + CUDA Graphs — Eliminate kernel launch overhead for fixed-shape workloads
- Model Warmup — Pre-fill CUDA caches and JIT traces before serving live traffic
- Health & Metrics — Latency histograms, throughput counters, queue depth monitoring
Architecture
Clients ──► Request Queue ──► Dynamic Batcher ──► Model (compiled) ──► Response Fan-out
│
Timeout / Max-batch
triggers dispatch
Key Concepts
- Padding & Attention Masks: Variable-length sequences padded to batch max, masked during attention
- CUDA Graphs: Record a static computation graph once, replay with near-zero launch overhead
- torch.compile: Fuses kernels, reduces memory bandwidth; combine with
mode="reduce-overhead"for serving - Queue Discipline: FIFO with configurable max wait time and max batch size
- Graceful Degradation: Shed load via queue depth limits; return 503 before OOM
Files
| File | Description |
|---|---|
batched_inference.py | Padding, masking, and batched forward pass utilities |
dynamic_batcher.py | Async request queue with time/size-triggered dispatch |
compiled_serving.py | torch.compile + CUDA Graphs integration for serving |
monitoring.py | Metrics collection, health checks, latency tracking |
server.py | End-to-end serving example tying all components together |
Running
# Single-file demos
python batched_inference.py
python dynamic_batcher.py
python compiled_serving.py
python monitoring.py
# Full server example
python server.py
Performance Tips
- Use
torch.compile(model, mode="reduce-overhead")for serving workloads - Pre-allocate output tensors to avoid allocation jitter
- Pin memory for CPU→GPU transfers in the request path
- Profile with
torch.profilerto find launch-bound vs compute-bound phases - Set
torch.set_float32_matmul_precision('high')for TF32 on Ampere+
Source Files
batched_inference.py— Padding, masking, and batched forward pass utilitiesdynamic_batcher.py— Async request queue with time/size-triggered dispatchcompiled_serving.py— torch.compile + CUDA Graphs integration for servingmonitoring.py— Metrics collection, health checks, latency trackingserver.py— End-to-end serving example tying all components together
Module 44: Performance Case Studies
Case Studies
| # | Case Study | Bottleneck | Fix | Speedup |
|---|---|---|---|---|
| 1 | Memory-bound DataLoader | CPU→GPU copy on main thread | Pin memory + non-blocking transfers | 2–3× |
| 2 | Naive attention scaling | O(n²) memory, repeated allocation | Flash attention pattern + in-place ops | 4–8× |
| 3 | Training loop overhead | Python dispatch + autograd bookkeeping | torch.compile with graph breaks analysis | 1.5–3× |
| 4 | Inference memory bloat | Gradients + BN running stats retained | Freezing, inference_mode, weight-only quantization | 60–75% memory reduction |
| 5 | Multi-GPU communication | All-reduce blocking compute | Gradient bucketing + overlap comm/compute | 1.3–1.8× at scale |
Files
| File | Description |
|---|---|
case1_dataloader.py | DataLoader pinned memory and prefetch optimization |
case2_attention.py | Attention implementation: naive → memory-efficient → flash |
case3_compile.py | torch.compile graph breaks analysis and fix |
case4_inference_memory.py | Inference memory reduction techniques |
case5_distributed.py | Multi-GPU communication overlap patterns |
How to Read These
Each script follows the same pattern:
# 1. BEFORE: naive implementation
def slow_version(...): ...
# 2. PROFILING: identify the bottleneck
profile_and_report(slow_version)
# 3. AFTER: optimized implementation
def fast_version(...): ...
# 4. COMPARISON: measure speedup
benchmark_comparison(slow_version, fast_version)
Running
python case1_dataloader.py
python case2_attention.py
python case3_compile.py
python case4_inference_memory.py
python case5_distributed.py # requires 2+ GPUs or use NCCL CPU backend
Key Takeaways
- Profile first, optimize second —
torch.profilerandtorch.cuda.memory_summary()are your friends - Memory bandwidth is usually the bottleneck, not compute (especially for inference)
torch.compilewins come from kernel fusion and reduced memory traffic- Overlapping communication with compute is critical for multi-GPU scaling
- Small changes (pin_memory, non_blocking, inference_mode) compound into large gains
Source Files
case1_dataloader.py— DataLoader pinned memory and prefetch optimizationcase2_attention.py— Attention implementation: naive → memory-efficient → flashcase3_compile.py— torch.compile graph breaks analysis and fixcase4_inference_memory.py— Inference memory reduction techniquescase5_distributed.py— Multi-GPU communication overlap patterns
Module 45: PyTorch Profiler Deep Dive
Overview
The PyTorch Profiler (torch.profiler) provides fine-grained visibility into CPU and GPU execution, memory allocation, and operator-level timing. It replaces the legacy torch.autograd.profiler with a unified API that integrates with Chrome Trace Viewer, TensorBoard, and HTA (Holistic Trace Analysis).
Key Concepts
Profiler Context Manager
The primary entry point is torch.profiler.profile(), which records events within a with block. Configure it with activities (CPU, CUDA), schedule for warm-up/active/repeat cycles, and on_trace_ready callbacks for automatic export.
Scheduling
torch.profiler.schedule(wait, warmup, active, repeat) controls when the profiler records. The wait phase skips steps, warmup discards initial noisy steps, and active captures the trace. This avoids profiling cold-start overhead.
Trace Export
- Chrome Trace:
prof.export_chrome_trace("trace.json")produces a JSON file
viewable at chrome://tracing or Perfetto UI.
- TensorBoard:
torch.profiler.tensorboard_trace_handler("./logs")writes
traces consumable by the TensorBoard PyTorch Profiler plugin.
- Stacks:
prof.export_stacks("stacks.txt")exports flame-graph-ready data.
Key Averages
prof.key_averages() aggregates events by operator name, returning a table with self_cpu_time_total, self_cuda_time_total, cpu_memory_usage, and call counts. Sort by any column to find hotspots.
Memory Profiling
Enable profile_memory=True to track tensor allocations. Combine with torch.cuda.memory._record_memory_history() for allocation-level snapshots viewable in the Memory Visualizer.
Examples
import torch
from torch.profiler import profile, ProfilerActivity, schedule, tensorboard_trace_handler
with profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
schedule=schedule(wait=1, warmup=1, active=3, repeat=1),
on_trace_ready=tensorboard_trace_handler("./tb_logs"),
record_shapes=True,
profile_memory=True,
with_stack=True,
) as prof:
for step in range(10):
train_step(model, data)
prof.step()
print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=20))
Files in This Module
| File | Description |
|---|---|
profiler_basics.py | Context manager usage, scheduling, key averages |
trace_analysis.py | Advanced trace analysis, memory snapshots, CUDA activity |
References
Source Files
profiler_basics.py— Context manager usage, scheduling, key averagestrace_analysis.py— Advanced trace analysis, memory snapshots, CUDA activity
Module 46: Quantization Recipes
Overview
Quantization reduces model size and inference latency by representing weights and activations with lower-precision data types (INT8, INT4, FP8). PyTorch provides three main approaches: dynamic quantization, static quantization, and quantization-aware training (QAT).
Key Concepts
Dynamic Quantization
Weights are quantized ahead of time; activations are quantized on-the-fly during inference. Best for models dominated by nn.Linear (e.g., LSTMs, Transformers). No calibration data required.
Static Quantization
Both weights and activations are quantized using calibration data to determine activation ranges. Requires inserting QuantStub/DeQuantStub and running a representative dataset through the model. Produces faster inference than dynamic.
Quantization-Aware Training (QAT)
Simulates quantization during training using fake-quantize operators. The model learns to compensate for quantization error, yielding higher accuracy than post-training quantization—especially for aggressive quantization (INT4).
PyTorch 2 Export Quantization (pt2e)
The modern path uses torch.export + torchao for quantization. Define a Quantizer that annotates the FX graph, then lower to a backend (XNNPack, Executorch, etc.). This replaces the legacy torch.quantization eager-mode API.
Data Types
| Type | Bits | Use Case |
|---|---|---|
| INT8 | 8 | General-purpose server/edge inference |
| INT4 | 4 | LLM weight-only quantization |
| FP8 (E4M3/E5M2) | 8 | Training and inference on Hopper+ GPUs |
| UINT4 | 4 | Asymmetric weight packing (torchao) |
Files in This Module
| File | Description |
|---|---|
dynamic_quantization.py | Dynamic quantization with torch.ao.quantization |
static_quantization.py | Static quantization with calibration and QAT |
Examples
# Dynamic quantization (simplest path)
import torch.ao.quantization as quant
quantized_model = torch.ao.quantization.quantize_dynamic(
model, {torch.nn.Linear}, dtype=torch.qint8
)
References
- Quantization docs
- PT2E Quantization tutorial
- torchao — Modern quantization and sparsity
- Quantization-Aware Training
- LLM.int8() paper
Source Files
dynamic_quantization.py— Dynamic quantization with `torch.ao.quantization`static_quantization.py— Static quantization with calibration and QAT
Module 47: DistributedDataParallel Patterns
Overview
DistributedDataParallel (DDP) is PyTorch's primary API for multi-GPU and multi-node training. Unlike DataParallel, DDP uses one process per GPU with NCCL/Gloo backends for efficient gradient all-reduce, avoiding the GIL bottleneck.
Key Concepts
Process Groups
DDP communicates via process groups. init_process_group() establishes a default group; custom sub-groups enable overlapping communication patterns (e.g., intra-node all-reduce + inter-node reduce-scatter).
Gradient Synchronization
DDP hooks into autograd to overlap gradient all-reduce with backward computation. Gradients are bucketed by parameter order (reverse of model.parameters()) for coalesced communication. Bucket size is tunable via bucket_cap_mb.
Launch Methods
- torchrun: Recommended launcher (
torchrun --nproc_per_node=4 train.py) - mp.spawn: Programmatic launch for testing
- SLURM: Multi-node via
srunwithMASTER_ADDR/MASTER_PORTenv vars
Gradient Accumulation
To accumulate gradients across micro-batches without synchronizing each step, use model.no_sync() as a context manager for all but the last micro-batch.
Mixed Precision + DDP
Combine DDP with torch.amp.autocast and GradScaler for FP16/BF16 training. The scaler handles gradient unscaling before the all-reduce.
Files in This Module
| File | Description |
|---|---|
ddp_training.py | Full DDP training loop with process groups, gradient sync, mixed precision |
Examples
# Single-node, 4 GPUs
torchrun --nproc_per_node=4 ddp_training.py
# Multi-node (2 nodes, 4 GPUs each)
torchrun --nnodes=2 --nproc_per_node=4 \
--rdzv_backend=c10d --rdzv_endpoint=$MASTER_ADDR:29500 \
ddp_training.py
Common Pitfalls
- Unused parameters: Set
find_unused_parameters=Trueif not all params
participate in every forward pass (but this adds overhead).
- Non-deterministic reductions: Floating-point all-reduce is not bitwise
reproducible across runs. Use torch.use_deterministic_algorithms(True) cautiously.
- State dict loading: Save/load the
moduleattribute (model.module.state_dict())
to get a clean checkpoint without DDP wrapper keys.
References
Source Files
ddp_training.py— Full DDP training loop with process groups, gradient sync, mixed precision
Module 48: Custom Autograd Functions
Overview
torch.autograd.Function lets you define custom forward and backward behavior when built-in ops are not enough: fused kernels, numerically stable formulas, or ops that need a hand-written gradient. This module goes beyond the intro in Module 03 and focuses on production patterns: ctx usage, @staticmethod, gradcheck, and preparing for higher-order gradients.
Key Concepts
Subclassing torch.autograd.Function
Implement two static methods:
- *
forward(ctx,args)** — compute the output; save tensors needed for backward
with ctx.save_for_backward(...) (preferred) or set attributes on ctx.
- *
backward(ctx,grad_outputs)** — return gradients for each forward input
that requires grad (use None for non-tensor / non-differentiable args).
Call the function with MyOp.apply(x, y) — never instantiate the class yourself.
What to Save on ctx
- Prefer
ctx.save_for_backward(*tensors)so autograd can free memory when safe. - Store non-tensor metadata as plain
ctxattributes (ctx.dim,ctx.eps). - Avoid saving huge intermediates if you can recompute them cheaply in backward.
Correctness: gradcheck / gradgradcheck
Always verify analytical gradients against finite differences:
from torch.autograd import gradcheck
assert gradcheck(MyOp.apply, (x.double(),), eps=1e-6, atol=1e-4)
Use gradgradcheck when you need second-order correctness (see double_backward.py).
Common Patterns
| Pattern | When |
|---|---|
| Elementwise fused op | Custom CUDA/Triton forward + matching backward |
| Numerically stable log-sum-exp | Forward uses max-trick; backward uses softmax |
| Straight-through estimator (STE) | Forward discrete; backward identity |
| Non-differentiable arg | Return None in that backward slot |
Examples
import torch
from torch.autograd import Function
class Square(Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, grad_output):
(x,) = ctx.saved_tensors
return grad_output * 2 * x
x = torch.tensor(3.0, requires_grad=True)
y = Square.apply(x)
y.backward()
print(x.grad) # 6.0
When to Use
- You need a custom gradient (STE, clipped grads, surrogate losses).
- You fuse ops for performance and must teach autograd the fused backward.
- Built-in ops lack the formula you need, or you want a stable formulation.
- Prefer functorch / existing ops when they already cover the math — less code, fewer bugs.
Files in This Module
| File | Description |
|---|---|
autograd_function_basics.py | Forward/backward patterns, STE, gradcheck |
double_backward.py | Higher-order grads and create_graph |
References
- Extending PyTorch — Custom Functions
- torch.autograd.Function
- gradcheck
- Module 03: Autograd · Module 38: Compiled Autograd
Source Files
autograd_function_basics.py— Forward/backward patterns, STE, gradcheckdouble_backward.py— Higher-order grads and `create_graph`
Module 49: Gradient Checkpointing — Advanced
Overview
Module 16 covers basic activation checkpointing. This module goes deeper: selective policies, reentrant vs non-reentrant implementations, activation offload to CPU, and how checkpointing interacts with torch.compile and DDP/FSDP. Use these tools when naive checkpoint(layer, x) is not enough for memory or correctness.
Key Concepts
Reentrant vs Non-Reentrant
| Mode | Flag | Notes |
|---|---|---|
| Reentrant (legacy) | use_reentrant=True | Re-enters autograd; needed for some older patterns; more edge cases |
| Non-reentrant (preferred) | use_reentrant=False | Cleaner saved-tensor handling; recommended in modern PyTorch |
Prefer use_reentrant=False unless you hit a known compatibility issue.
Selective Activation Checkpointing (SAC)
Instead of recomputing everything in a region, apply a policy that saves cheap ops (e.g. pointwise) and recomputes expensive ones (matmul/attention), or the reverse — depending on memory vs compute priorities.
from torch.utils.checkpoint import checkpoint, create_selective_checkpoint_contexts
Policies typically inspect ops / op types during the forward pack phase.
Offloading Activations
When GPU memory is still tight after SAC, offload saved activations to CPU (pinned) memory during forward and prefetch them for backward. This adds PCIe traffic but can unlock larger batch sizes.
Interaction with Distributed & Compile
- DDP / FSDP: Checkpoint inside each transformer block so recompute stays
local; avoid wrapping the entire model in one checkpoint.
torch.compile: Non-reentrant checkpointing composes better; graph breaks
can appear at checkpoint boundaries — profile before/after.
Examples
import torch
from torch.utils.checkpoint import checkpoint
def block(x, weight):
return torch.nn.functional.gelu(x @ weight)
x = torch.randn(2, 128, requires_grad=True)
w = torch.randn(128, 128, requires_grad=True)
y = checkpoint(block, x, w, use_reentrant=False)
y.sum().backward()
When to Use
- Model OOMs with full activation storage (LLM / ViT training).
- You need a finer memory/compute tradeoff than “checkpoint every layer”.
- GPU memory is scarce but host RAM / PCIe bandwidth is available (offload).
- Skip if the model already fits — checkpointing adds ~20–40% compute.
Files in This Module
| File | Description |
|---|---|
selective_checkpoint.py | Reentrant vs non-reentrant, SAC-style policy, offload sketch |
References
- torch.utils.checkpoint
- Activation Checkpointing tutorial
- Module 16: Activation Checkpointing
- Module 10: Distributed · Module 26: Memory Profiling
Source Files
selective_checkpoint.py— Reentrant vs non-reentrant, SAC-style policy, offload sketch
Module 50: Sparse Tensors & Sparse Ops
Overview
PyTorch sparse tensors store only nonzero values (plus indices), saving memory and compute when data is highly sparse — recommendation matrices, graph adjacency, embeddings with large vocabularies. This module surveys layouts (coo, csr, csc, csc), conversion, and the sparse operator surface.
Key Concepts
Layouts
| Layout | Best for | Structure |
|---|---|---|
| COO | Construction, irregular sparsity | indices [ndim, nnz] + values |
| CSR | Fast row slices, SpMM on CPU/CUDA | crow_indices, col_indices, values |
| CSC | Fast column slices | ccol_indices, row_indices, values |
| BSR / BSC | Block-sparse (structured) | block indices + dense blocks |
Create with torch.sparse_coo_tensor / torch.sparse_csr_tensor, or dense.to_sparse() / dense.to_sparse_csr().
Hybrid Tensors
A tensor can be sparse in some leading dims and dense in others (batched sparse). Example: batched graph adjacency [B, N, N] stored as sparse batches.
Sparse Ops
Common ops with sparse support (version-dependent):
- Arithmetic:
+,*,sparse.addmm/torch.sparse.mm - Reductions:
sum,softmax(layout-specific) - Conversions:
to_dense(),coalesce()(merge duplicate COO indices)
Always coalesce() COO tensors before relying on unique indices.
Gradients
Sparse tensors can participate in autograd when ops support it. Gradients w.r.t. sparse parameters may be sparse or dense depending on the op — check docs and test with small examples.
Examples
import torch
i = torch.tensor([[0, 1, 1],
[2, 0, 2]])
v = torch.tensor([3.0, 4.0, 5.0])
s = torch.sparse_coo_tensor(i, v, (2, 3)).coalesce()
print(s.to_dense())
# CSR SpMM-style matmul with a dense vector
csr = s.to_sparse_csr()
x = torch.randn(3)
y = torch.sparse.mm(csr, x.unsqueeze(1)).squeeze(1)
When to Use
- Density below ~1–5% and ops you need are sparse-aware.
- Graph NNs, sparse attention masks, large embedding bags.
- Prefer dense + masking when sparsity is moderate or ops lack sparse kernels.
- Structured 2:4 sparsity for Ampere+ Tensor Cores — see Module 31 (torchao).
Files in This Module
| File | Description |
|---|---|
sparse_basics.py | COO/CSR creation, coalesce, matmul, autograd sketch |
References
- torch.sparse docs
- Sparse semi-structured
- Module 13: Advanced · Module 31: torchao
Source Files
sparse_basics.py— COO/CSR creation, coalesce, matmul, autograd sketch