PyTorch: The Complete Reference Guide

From Absolute Beginner to Advanced Practitioner

PyTorch 2.14+ 50 Curriculum Modules 150 Code Examples 50 Notebooks 28 Reference Cards

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?

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

AspectPyTorchTensorFlow
ExecutionEager by default (define-by-run)Historically graph-based (define-then-run), now eager via tf.function
DebuggingStandard Python debugger worksHarder to debug graph mode
Research adoptionDominant in academia (~80%+ of papers)Strong in industry/production
DeploymentTorchServe, ONNX, torch.exportTF Serving, TFLite, TF.js
API StylePythonic, object-orientedKeras-based high-level API
Compilationtorch.compile (TorchDynamo + Inductor)XLA compiler
CommunityMassive 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

Module 02: Tensors — The Complete Guide

Table of Contents


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:

DimensionsMath NamePyTorch ShapeExample
0Scalartorch.Size([])A single loss value: 0.543
1Vectortorch.Size([n])A word embedding: [0.2, -0.1, ...]
2Matrixtorch.Size([m, n])A linear layer's weights
33-tensortorch.Size([a, b, c])A batch of sequences (batch, seq, embed)
44-tensortorch.Size([a, b, c, d])A batch of images (batch, channels, H, W)
NN-tensortorch.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:

dtypeBitsRange/PrecisionUse Case
torch.float32 (default)32~7 decimal digitsStandard training
torch.float6464~15 decimal digitsNumerical verification, scientific computing
torch.float1616~3 decimal digitsMixed-precision training (with loss scaling)
torch.bfloat1616~3 decimal digits, wider rangeLLM training (Transformer-preferred)
torch.int88[-128, 127]Quantized inference
torch.int1616[-32768, 32767]Rarely used
torch.int3232[-2^31, 2^31-1]Indices, counts
torch.int64 (default for ints)64[-2^63, 2^63-1]Default integer type
torch.bool8True/FalseMasks, conditions
torch.complex6464Two float32Signal processing, FFT
torch.complex128128Two float64High-precision complex math
torch.float8_e4m3fn8~2 digits, narrowTransformer engine inference
torch.float8_e5m28~1 digit, wider rangeTransformer 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 inspection
  • operations.py — element-wise and reduction operations
  • indexing_and_slicing.py — all forms of indexing
  • broadcasting.py — broadcasting rules with examples
  • views_strides_memory.py — views, strides, and memory layout

📓 Open Notebook — Interactive version of this module

Source Files

Module 03: Autograd — Automatic Differentiation

Table of Contents


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

SituationUse
Validation loop during trainingtorch.no_grad()
Updating parameters manuallytorch.no_grad()
Production inferencetorch.inference_mode()
Need to use result in autograd laterNeither (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

  • forward receives the actual tensor values and returns output tensors.
  • backward receives 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 backward must 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): Use log(x + epsilon) or torch.clamp(x, min=1e-8)
  • sqrt(0): derivative of sqrt at 0 is infinity. Use sqrt(x + epsilon)
  • 0/0: Can occur in normalization layers when variance is zero
  • Division by a very small number: use torch.clamp on 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 and inference_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 loop
  • computation_graph.py — visualizing and understanding the computation graph
  • custom_functions.py — writing custom autograd functions
  • higher_order_gradients.py — second derivatives, Jacobians, and Hessians

📓 Open Notebook — Interactive version of this module

Source Files

Module 04: Neural Networks in PyTorch

Table of Contents


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 forward method)
  • 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 borders
  • dilation: 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

ConceptKey Takeaway
nn.ModuleBase class; register layers in __init__, compute in forward
ParametersLearnable tensors (weights), updated by optimizer
BuffersState tensors without gradients (running stats, masks)
SequentialSimple chain of layers
Conv2dSpatial feature extraction with kernel sliding
BatchNormNormalize across batch; different train/eval behavior
LayerNormNormalize across features; batch-independent
DropoutRandom zeroing during training only
LSTMRecurrent with forget/input/output gates
TransformerSelf-attention + feedforward
CrossEntropyLossMulti-class classification standard
Kaiming initDefault for ReLU networks
HooksInspect/modify without changing model code
state_dictStandard save/load mechanism

📓 Open Notebook — Interactive version of this module

Source Files

Module 05: Optimizers and Learning Rate Schedulers

Table of Contents


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.001
  • beta1 = 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?

TaskRecommendedLR RangeNotes
Vision (CNNs)SGD+momentum or AdamW0.01-0.1SGD often generalizes better
NLP/TransformersAdamW1e-5 to 5e-4With cosine schedule + warmup
Fine-tuningAdamW1e-5 to 3e-5Lower LR for pretrained weights
GANsAdam (beta1=0.0)1e-4 to 2e-4Two separate optimizers
RLAdam3e-4Simpler schedules
Small datasetsSGD+momentum0.01Better 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

OptimizerKey FeatureBest For
SGD+momentumSimple, good generalizationVision, large-scale
AdamAdaptive per-param LRGeneral, fast convergence
AdamWProper weight decayTransformers, modern default
AdagradAdapts to sparse featuresNLP with sparse embeddings
RMSpropFixes Adagrad decayRNNs (historically)
LBFGSSecond-orderSmall problems, fine-tuning
SchedulerPatternBest For
CosineAnnealingSmooth decay to 0Most tasks
OneCycleLRWarmup + decayFastest convergence
ReduceLROnPlateauAdaptive decayWhen you have val metric
Sequential(Linear+Cosine)Warmup + cosineTransformers
CosineWarmRestartsPeriodic resetsLong training, exploration

📓 Open Notebook — Interactive version of this module

Source Files

Module 06: Data Loading in PyTorch

Table of Contents


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 uint8 for 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

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

Propertyfloat32float16bfloat16
Total bits321616
Exponent bits858
Mantissa bits23107
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?

StrategyData AmountDomain SimilarityRisk
Linear probeVery smallAnyLow
Freeze + new headSmallSimilarLow
Progressive unfreezeMediumDifferentMedium
Full fine-tuneLargeDifferentHigher

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 = False can slow down training
  • DataLoader with num_workers > 0 needs 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

TechniquePurposeMemory ImpactSpeed Impact
Mixed precisionFaster math, less memoryReduces 50%2-3x faster
Gradient accumulationSimulate large batchesNo changeSlight slow
Gradient checkpointingReduce activation memorySaves 60-80%~33% slower
Gradient clippingPrevent exploding gradientsNoneNegligible
EMASmoother final model2x paramsNegligible
Label smoothingPrevent overconfidenceNoneNone
Early stoppingPrevent overfittingNoneSaves time

📓 Open Notebook — Interactive version of this module

Source Files

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

ModeCompile TimeRuntime SpeedMemoryBest For
defaultFastGoodNormalDevelopment
reduce-overheadMediumBetterHigherSmall ops (GPU)
max-autotuneSlowBestNormalProduction

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() or fullgraph=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 inputs
  • call_function — calls to functions (like torch.add)
  • call_method — method calls on tensors (.relu(), .view(), etc.)
  • call_module — calls to nn.Module submodules
  • output — return value

Understanding FX graphs helps when:

  • Writing custom backends
  • Debugging compilation issues
  • Understanding what optimizations are applied

Summary

FeatureWhat It DoesWhen to Use
torch.compileCompiles model for speedAlways (production)
fullgraph=TrueErrors on graph breaksEnsuring no breaks
dynamic=TrueHandles varying shapesVariable batch sizes
max-autotuneMaximum optimizationDeployment
reduce-overheadMinimizes launch overheadMany small GPU ops
explain()Shows compilation infoDebugging
custom backendCustom graph processingResearch/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

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^T has 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

ConceptPurposeKey Shape
Scaled dot-productCore attention computation[B, N, N] scores
Causal maskPrevent attending to futureLower triangular
Multi-headMultiple attention patternsh heads, d_k = d/h
SDPAOptimized PyTorch attentionSame as manual
Flash AttentionO(N) memory attentionTiled computation
FlexAttentionCustom patterns, fusedscore_mod/mask_mod
RoPERelative position encoding2D rotations
KV CacheFast autoregressive generationCached K, V tensors

📓 Open Notebook — Interactive version of this module

Source Files

Module 10: Distributed Training in PyTorch

Table of Contents


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

StrategyWhat 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 parametersModel barely fits or doesn't fit on 1 GPU
Tensor Parallel (TP)Individual layers/tensorsVery large layers (e.g., huge linear layers)
Pipeline Parallel (PP)Model stages (groups of layers)Very deep models, many GPUs
Context Parallel (CP)Sequence dimensionVery long sequences
3D ParallelismCombination of DP + TP + PPLarge-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

BackendDevicesUse Case
NCCLGPU (NVIDIA)Default for GPU training. Highly optimized for NVIDIA hardware
GlooCPU, GPUCPU training, or as a fallback. Also used for CPU collectives in GPU training
UCCGPUAlternative 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

VariableDescription
RANKGlobal rank of this process
LOCAL_RANKLocal rank on this node
WORLD_SIZETotal number of processes
MASTER_ADDRAddress of the master node
MASTER_PORTPort for master node communication
LOCAL_WORLD_SIZENumber 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=True flag 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:

PlacementDescriptionExample
Shard(dim)Tensor is sharded along dimension dimA [4, 8] tensor Shard(1) across 2 GPUs → each gets [4, 4]
Replicate()Tensor is fully replicated on each deviceA [4, 8] tensor Replicate() → each GPU has [4, 8]
Partial()Each device has a partial result; needs reductionIntermediate 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 FSDP module

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-gather reconstructs 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-gather again to get full parameters for gradient

computation.

  • After backward: reduce-scatter synchronizes 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:

ScheduleBubble RatioMemoryDescription
GPipe(p-1)/mHigh (all activations)All forwards, then all backwards
1F1B(p-1)/mLow (1 activation)Alternating forward-backward in steady state
Interleaved 1F1B(p-1)/(m×v)LowMultiple virtual stages per rank, smaller bubbles
Zero Bubble~0ModerateOverlaps weight gradient with next forward
DualPipeV~0ModerateBidirectional 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

FileDescriptionRun Command
concepts_and_collectives.pyCollective operations on CPU with Glootorchrun --nproc_per_node=3 concepts_and_collectives.py
ddp_example.pyComplete DDP training with synthetic datatorchrun --nproc_per_node=2 ddp_example.py
fsdp2_example.pyFSDP2 API patterns and setuptorchrun --nproc_per_node=2 fsdp2_example.py
device_mesh_example.pyDeviceMesh creation and sub-mesh accesstorchrun --nproc_per_node=4 device_mesh_example.py
parallelism_overview.pyAPI patterns for TP and PPReference script showing code patterns

📓 Open Notebook — Interactive version of this module

Source Files

Module 11: Export and Deployment

Table of Contents


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
PathOutputTargetKey Advantage
AOTInductor.so shared libraryServer (C++ or Python)Maximum performance, no Python needed
NativeRTSerialized modelServer (C++)C++ inference engine, easy deployment
ONNX.onnx fileAny ONNX RuntimeCross-framework interop
TorchServeModel archiveServer (Python)Full serving stack (batching, scaling)
ExecuTorch.pte fileMobile/edge devicesSmall 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.GraphModule containing 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.cond instead)

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 > 0 with torch.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 value
  • get_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

FeatureAOTInductorNativeRT
CompilationAhead of time (slow)No compilation
Inference speedFastest (native code)Fast (C++ runtime)
Startup timeFast (pre-compiled)Fast (load + interpret)
FlexibilityFixed graphMore flexible
Python neededNoNo

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

PlatformDelegate
iOSCore ML, Metal
AndroidXNNPACK, Qualcomm QNN
MicrocontrollersCustom delegates
WebWebAssembly

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

TypePrecisionSpeedAccuracyUse Case
FP3232-bitBaselineBestTraining
FP1616-bit~2×Negligible lossGPU inference
BF1616-bit~2×Negligible lossGPU inference
INT88-bit~2-4×Small lossServer inference
INT44-bit~4-8×Moderate lossEdge/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 Dim for 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 PathRelative LatencySetup Effort
Eager (Python)1.0× (baseline)None
torch.compile0.5-0.8×One line
AOTInductor0.3-0.6×Moderate
Quantized (INT8)0.2-0.4×Significant
ExecuTorch (mobile)Varies by hardwareSignificant

Files in This Module

FileDescriptionRun Command
export_basics.pyBasic export, run exported modelpython export_basics.py
dynamic_shapes.pyDim API, multiple dynamic dimspython dynamic_shapes.py
export_inspection.pyInspect the graph, list opspython export_inspection.py
save_and_load.pySave/load PT2 archivespython save_and_load.py

📓 Open Notebook — Interactive version of this module

Source Files

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:

PaperCore Innovation
ResNetSkip (residual) connections
TransformerScaled dot-product self-attention
GPTDecoder-only Transformer + autoregressive pretraining
ViTTreat image patches as token sequences
VAEReparameterization trick for differentiable sampling
U-NetEncoder-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.Module subclass (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

ModelBlockLayers per stageTotal layersParameters
ResNet-18BasicBlock[2, 2, 2, 2]18~11M
ResNet-34BasicBlock[3, 4, 6, 3]34~21M
ResNet-50Bottleneck[3, 4, 6, 3]50~25M
ResNet-101Bottleneck[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

ArchitectureYearInnovationKey Pattern
ResNet2015Residual connectionsIdentity shortcut
U-Net2015Encoder-decoder + skipConcatenation skip
VAE2013Reparameterization trickStochastic latent
Transformer2017Self-attentionQ/K/V attention
GPT2018Decoder-only + pretrainCausal masking
ViT2020Patch tokenizationImage as sequence

Files in This Module

  • resnet.py — Complete ResNet with BasicBlock, Bottleneck, configs for 18/34/50/101
  • transformer.py — Full encoder-decoder Transformer from scratch
  • gpt.py — GPT with generation: greedy, temperature, top-k, top-p sampling
  • vit.py — Vision Transformer for image classification
  • vae.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 Bottleneck
  • transformer.py — Transformer — complete encoder-decoder implementation
  • gpt.py — GPT (Generative Pre-trained Transformer) — decoder-only implementation
  • vit.py — Vision Transformer (ViT) — complete implementation for image classification
  • vae.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 jacrev when output dimension is smaller than input dimension
  • Use jacfwd when 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() or with 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 without retain_graph=True
  • Fix: either use retain_graph=True or 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

FeatureUse CaseKey API
vmapBatch any functiontorch.func.vmap
gradFunctional gradientstorch.func.grad
jacrev/jacfwdJacobian matricestorch.func.jacrev
Per-sample gradsDP-SGD, influence functionsvmap(grad(...))
Sparse tensorsGraphs, sparse datatorch.sparse_coo_tensor
Complex numbersFFT, signal processingtorch.complex, torch.fft
Custom opsExtending PyTorchtorch.library
ProfilingPerformance optimizationtorch.profiler
Anomaly detectionDebugging NaN/Inftorch.autograd.detect_anomaly
torch.fxGraph transformstorch.fx.symbolic_trace
Meta deviceShape analysistorch.device("meta")

Files in This Module

  • functorch_transforms.py — vmap, grad, jacrev, hessian demonstrations
  • per_sample_gradients.py — Per-sample gradient computation with vmap+grad
  • custom_operators.py — Defining custom ops with torch.library
  • profiling.py — Profiler usage, timing, and analysis
  • sparse_and_complex.py — Sparse tensors, complex numbers, and FFT
  • debugging_tips.py — Anomaly detection, gradient flow checking, and common fixes

📓 Open Notebook — Interactive version of this module

Source Files

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

TopicKey ToolWhen to Use
Basic testingTestCase, assertEqualEvery project
Parametrized tests@parametrizeTesting across configs
Device testsinstantiate_device_type_testsCross-device testing
Reproducibilityset_seed(), deterministic modeDebugging, CI
Benchmarkingtorch.utils.benchmark.TimerPerformance comparison
Gradient checkstorch.autograd.gradcheckCustom autograd

Files in This Module

  • test_example.py — Complete test file using PyTorch's TestCase
  • reproducibility.py — Full reproducibility setup and verification
  • benchmarking.py — Benchmarking with torch.utils.benchmark

📓 Open Notebook — Interactive version of this module

Source Files

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


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

MethodWhat It Does
random_unstructuredZero out random individual weights
l1_unstructuredZero out weights with smallest L1 magnitude
random_structuredZero out entire channels/neurons randomly
ln_structuredZero out channels with smallest Ln norm
global_unstructuredPrune across all layers by global ranking

How Pruning Works in PyTorch

  • The original weight is moved to weight_orig
  • A binary mask weight_mask is created
  • A forward hook computes weight = weight_orig * weight_mask before 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

UtilityWhat It DoesWhen to Use
parametrizeConstrain weights (orthogonal, symmetric, positive)Research, stable training
pruneZero out weights by magnitude/structureModel compression, edge deploy
spectral_normNormalize by largest singular valueGAN training stability
weight_normDecouple magnitude and directionFaster convergence
pack_padded_sequenceEfficient variable-length RNN processingAny RNN with variable lengths
fuse_conv_bn_evalMerge Conv+BN for inferenceInference optimization
nested_tensorVariable-length batches without paddingAttention, NLP, Flash Attention
clip_grad_norm_Prevent exploding gradientsAny training with transformers
skip_initSkip weight initializationLoading large pretrained models

Further Reading


📓 Open Notebook — Interactive version

Source Files

  • parametrization.py — Weight parametrization — enforcing constraints like symmetry, orthogonality, and positivity on parameters
  • pruning.py — Model pruning — making neural networks smaller by removing weights
  • sequence_packing_and_nested.py — Sequence packing & nested tensors — efficient variable-length processing
  • conv_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


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 ops
  • use_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

PolicyBehavior
MUST_SAVEAlways save this op's output (never recompute)
PREFER_SAVESave unless torch.compile decides otherwise
MUST_RECOMPUTEAlways recompute (never save)
PREFER_RECOMPUTERecompute unless torch.compile decides otherwise
MUST_CPU_OFFLOADSave to CPU during forward, reload to GPU during backward
PREFER_CPU_OFFLOADOffload 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

ScenarioRecommendation
Model fits in GPU memoryDon't checkpoint (unnecessary overhead)
OOM with desired batch sizeCheckpoint every Transformer layer
Still OOMAdd selective checkpointing (save matmuls, recompute rest)
Still OOMCombine with FSDP2, gradient accumulation, or CPU offload

Rules of Thumb

  • Checkpoint at the Transformer layer granularity — each TransformerEncoderLayer or 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


📓 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


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
StanceBehaviorUse Case
defaultNormal compilationProduction
force_eagerSkip all compilationDebugging, profiling eager
eager_on_recompileCompile once, eager on recompileAvoid compile-time storms
fail_on_recompileError on recompilationCI, catch shape issues
eager_then_compileEager first call, compile on secondWarmup 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

APIWhat It Does
set_stance()Global compilation behavior
@disableSkip compilation for a function
@allow_in_graphTreat as opaque graph node
substitute_in_graph()Replace with traceable version
mark_dynamic()Declare dynamic dimension
mark_static()Declare static dimension
fullgraph=TrueError on any graph break
graph_break()Force a graph break
explain()Get compilation report
comptime.breakpoint()Debug during compilation
CompileCounterCount compilations in tests
EagerAndRecordGraphsInspect 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


1. What is torch.package?

torch.package creates a hermetic zip archive containing:

  • Your model's Python source code (the actual .py files)
  • 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

ActionWhat It DoesWhen to Use
intern(pattern)Bundle the module's source code into the packageYour own code, custom modules
extern(pattern)Expect the module to be installed on the target machinePyTorch, NumPy, standard libs
mock(pattern)Replace with a stub that returns MockedObjectUnused optional dependencies
deny(pattern)Error if this module is encounteredKnown-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

Featuretorch.savetorch.packagetorch.export
Saves weightsYesYesYes
Saves codeNoYes (source)Yes (graph)
Hermetic loadingNoYesYes
Python control flowN/AFullLimited
Works cross-versionFragileRobustRobust
Deployment targetPythonPythonC++, mobile, ONNX
File formatpicklezip (with source)PT2 archive
SpeedFastMediumCompile required

When to use each:

  • torch.save: Quick checkpoints during training
  • torch.package: Ship Python models with all dependencies, research sharing
  • torch.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_NAME generation (#186402)
  • Dynamo operator support — Added divmod, remainder, true_divide, floor_divide operators (#185652-#185655)
  • Deterministic topk — torch.topk now respects torch.use_deterministic_algorithms() (#186653)
  • XPU oneDNN LSTM — Intel GPU LSTM inference via oneDNN primitives (#185531)
  • Stable ABI generator — New torch/csrc/stable/generator.h for stable C API

Further Reading


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


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 arguments
  • kwargs — 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__
LevelPython APIATen operators
Inputtorch.add, torch.nn.functional.reluaten.add.Tensor, aten.relu.default
DecompositionBeforeAfter (sees primitive ops)
Used byCustom wrappers, loggingDTensor, FakeTensor, torch.compile
Subclass requiredNo (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

ScenarioUse
Log/trace all torch function callsTorchFunctionMode
Custom tensor wrapper (non-subclass)__torch_function__
Override ATen ops for a tensor subclass__torch_dispatch__
Count/profile ops at ATen levelTorchDispatchMode
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.stream context 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


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

FileDescription
backends_tuning.pyRunnable 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=FalseDeterministic=True
SpeedFasterSlower (sometimes 2–3x)
ReproducibilityRun-to-run varianceBit-exact results
Use caseNormal trainingDebugging, 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:

StrategySpeedQualityUse case
"greedy"FastGoodDefault for most cases
"optimal"SlowBestSmall expressions (<10 indices)
"dp"MediumGoodBalanced for larger expressions
"auto"AdaptiveBest tradeoffRecommended

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

OperationSettingSpeedup (A100)Precision Loss
Conv2dcudnn.allow_tf32=True~2x~0.1% relative
matmulcuda.matmul.allow_tf32=True~2–3x~0.1% relative
LinearVia 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
LevelMatmul PrecisionSpeedUse Case
"highest"Full FP32BaselineDebugging, validation
"high"TF32 on Ampere+~2–3x fasterStandard training
"medium"TF32 + BF16 reductionsFastestLarge-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

SettingTrainingInferenceDebug
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

Module 21: CUDA Graphs — Eliminating CPU Launch Overhead

Day 7 of the incremental learning series


Table of Contents


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

ScenarioTypical 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 InitializationWhy It Matters
cuDNN algorithm selectionFirst conv/GEMM benchmarks multiple algorithms
CUDA context creationFirst CUDA call initializes the driver context
Memory allocator warmupCaching allocator builds its pool
JIT kernel compilationSome ops compile PTX on first use
cuBLAS handle creationFirst 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 CUDAGraphtorch.compile(mode="reduce-overhead")
You manage static inputsInputs handled automatically
You do warmupWarmup is automatic
Entire model must be graphablePer-region graphs (partial capture)
No fusionKernel fusion + graphs combined
Fixed shapes onlyMultiple 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:

OperationWhy It Fails
print(tensor) inside graphRequires 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 allocationGraph can't record variable-size allocs

Silent Failures (Wrong Results)

These won't crash but will produce incorrect output:

PatternProblem
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 tensorsWrites to wrong addresses
Random ops without manual seedSame 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

ScenarioRecommendation
Inference, fixed shapes, max throughputManual CUDAGraph or reduce-overhead
Inference, variable shapestorch.compile(mode="default") — dynamic shape support
Training, single GPUtorch.compile(mode="reduce-overhead") — handles optimizer
Training, multi-GPU (DDP/FSDP)torch.compile(mode="default") — NCCL compatibility
Model has data-dependent control flowtorch.compile with graph breaks
Quick prototypingmake_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


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


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 signal
  • W1 (dim → hidden): produces the value
  • W2 (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?

Propertyfp16bf16
Exponent bits58
Mantissa bits107
Max value~65504~3.4×10³⁸
Min normal~6×10⁻⁵~1.2×10⁻³⁸
PrecisionHigherLower
Overflow riskHighVery low
Loss scaling neededYesNo

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_steps to 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.permutations now supported)
  • Bug fixes for argmin/argmax on 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

TechniqueWhat It DoesMemory ImpactSpeed Impact
RoPERelative position via rotationNoneSlight compute
KV CacheCache K/V for generation+cache memory~1000× faster decode
GQAFewer KV heads4× less KV cacheFaster decode
Sliding WindowLimit attention spanO(n·W) vs O(n²)Faster for long seq
RMSNormFaster normalizationSame~10% faster norm
SwiGLUGated FFN with SiLU+50% FFN paramsBetter convergence
Weight TyingShare embed/output-30% for vocab-heavySame
bf16Wider exponent rangeHalf vs fp322× throughput
Grad AccumulationLarge effective batchConstantLinear in steps

Further Reading


Notebook: 22_llm_recipes.ipynb

Source Files

  • rope_embeddings.py — RoPE — precompute freqs, apply rotary embeddings, position encoding visualization
  • kv_cache.py — KV Cache — pre-allocated cache, prefill/decode phases, GQA repeat_kv, benchmarking
  • llm_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


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 — the Graph object (node-level IR)
  • traced.code — auto-generated Python source code
  • traced.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:

LimitationExampleWhy It Fails
Data-dependent control flowif x.sum() > 0:Proxy doesn't have a real value to branch on
Dynamic shapesx[:, :n] where n is runtime-determinedProxy doesn't have concrete shapes
Non-torch Python opsprint(x.shape), list comprehensions over tensorsProxies don't support arbitrary Python
Non-forward methodsself.helper(x) not called from forwardOnly 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.opMeaningnode.targetExample
placeholderFunction inputParameter namex
get_attrAccess self.attrAttribute path (string)self.weight
call_functionFree function callThe function itselftorch.relu, operator.add
call_methodMethod on a valueMethod name (string).view(), .relu()
call_moduleCall a submoduleModule path (string)self.linear1
outputReturn value"output"The return node

Node Anatomy

Every node has:

  • op — one of the six operations above
  • name — 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 match x + y (which becomes operator.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 args and kwargs reference valid nodes in the same graph
  • There is exactly one output node
  • Placeholder nodes come before all other nodes
  • No cycles exist
  • All call_module targets 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

MethodCalled WhenUse Case
run_node(node)Every nodeProfiling, logging, error handling
call_function(target, args, kwargs)call_function nodesMock functions, replace ops
call_method(target, args, kwargs)call_method nodesIntercept method calls
call_module(target, args, kwargs)call_module nodesSwap modules, add hooks
placeholder(target, args, kwargs)Input nodesModify inputs
get_attr(target, args, kwargs)Attribute accessIntercept param loads
output(target, args, kwargs)Return nodePost-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

AspectTransformerManual (graph.nodes iteration)
Creates new graphYesNo (in-place)
Node remappingAutomaticManual
Easier for per-node transformsYesNo
Better for structural changesNoYes (inserting/removing)
Risk of dangling referencesLowHigher

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

PassWhat It Does
decompositions.pyBreak complex ops into primitives
fuse_attention.pyPattern-match and fuse attention
group_batch_fusion.pyBatch small ops together
joint_graph.pyOptimizations on the joint fwd+bwd graph
post_grad.pyPost-autograd optimizations
pre_grad.pyPre-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 performance
  • logical_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

GoalAPI
Inspect model structuresymbolic_trace + iterate graph.nodes
Simple op replacementManual node iteration, change node.target
Structural transforms (fuse, split)Manual inserting_before/after + erase_node
Pattern-based replacementreplace_pattern
Per-node behavior (profiling, logging)Interpreter subclass
Clean per-node transformsTransformer subclass
Production optimization passesCustom Inductor passes

Key Rules

  • Always call graph.lint() after modifying a graph
  • Always call gm.recompile() after modifying the graph of a GraphModule
  • Erase nodes bottom-up — a node can only be erased when it has zero users
  • replace_all_uses_with before erasing a node with users
  • symbolic_trace ≠ torch.compile — symbolic trace is simpler but less powerful; torch.compile (Dynamo) handles control flow and dynamic shapes

Further Reading


Notebook: 23_fx_transforms.ipynb

Source Files

torch.masked — First-Class Missing Data in PyTorch

Table of Contents


1. The Problem: Missing Data & Masking

Missing or invalid data appears in virtually every domain of deep learning:

DomainScenarioWhat's "Missing"
NLPPadded sequences in a batchPositions beyond each sequence's true length
VisionIrregular shapes, masked regionsPixels outside the region of interest
TabularIncomplete recordsColumns with no observed value
AttentionCausal masks, padding masksFuture 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 tensor
  • mt.get_mask() — returns the boolean mask
  • True means valid, False means 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

FunctionWhat it does
torch.masked._ops.sumSum of valid elements
torch.masked._ops.meanMean of valid elements (divides by valid count)
torch.masked._ops.amaxMaximum of valid elements
torch.masked._ops.aminMinimum of valid elements
torch.masked._ops.prodProduct of valid elements
torch.masked._ops.normNorm over valid elements
torch.masked._ops.varVariance over valid elements
torch.masked._ops.stdStandard 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:

OperationManual Maskingtorch.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)
Maxdata.masked_fill(~mask, -inf).max(dim)torch.masked._ops.amax(data, dim, mask=mask)
Mindata.masked_fill(~mask, inf).min(dim)torch.masked._ops.amin(data, dim, mask=mask)
Softmaxsoftmax(data.masked_fill(~mask, -inf), dim)torch.masked.softmax(data, dim, mask=mask)
NormalizeCompute norm manually, dividetorch.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.compile in 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

ConceptDescription
The ProblemMissing data is everywhere: padding, irregular shapes, missing values. Manual masking is verbose and error-prone.
MaskedTensorA tensor subclass pairing data + boolean mask. Operations respect the mask automatically.
torch.masked.softmaxSoftmax that correctly ignores masked positions and normalizes over valid elements only.
Masked Reductionstorch.masked._ops.sum/mean/amax/amin — correct reductions that ignore masked elements.
Mask ConventionTrue = valid, False = masked/missing.
Unary OpsPreserve the mask.
Binary OpsAND the masks (both must be valid).
Prototype StatusNot 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


Notebook: 24_masked_tensor.ipynb

Source Files

Custom Triton Kernels — GPU Programming in Python

Table of Contents


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.compile backend) generates Triton code for fused operations. When you torch.compile a 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.compile compatibility.

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 CaseExample
Fuse operationsCombine elementwise ops, reductions, and activations into one kernel
Custom opsImplement operations that don't exist in PyTorch (novel attention variants, custom normalizations)
Eliminate overheadRemove Python/dispatch overhead by running everything in a single GPU launch
PrototypingIterate on GPU kernel ideas 10x faster than CUDA C++
Match InductorWrite 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

ConceptDescription
@triton.jitDecorator 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_SIZENumber 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
GridTotal 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 = pid BLOCK_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 stability
  • x = x - max_val — shift
  • x = exp(x) — exponentiate
  • sum_val = sum(x) — normalize
  • out = 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, so exp(-inf) = 0 and 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

FactorGuidance
Too few blocksGPU SMs sit idle. Aim for at least num_SMs * 4 blocks
Too many blocksMinor overhead from scheduling. Generally harmless
Block size too smallInstruction overhead dominates. Use 256+ for elementwise
Block size too largeRegister 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

AspectTritonCUDA C++
LanguagePythonC++
Iteration speedFast (Python workflow, auto-compile)Slow (compile, link, debug cycle)
Shared memoryAutomatic (compiler-managed)Manual (__shared__, bank conflict avoidance)
Thread-level controlBlock-level onlyFull warp/thread control
Performance~80-95% of hand-tuned CUDA100% (by definition)
Warp primitivesLimited (tl.atomic_*, basic reductions)Full (__shfl_*, warp vote, cooperative groups)
Tensor CoresVia tl.dot (automatic)Via wmma or mma.sync (manual)
PortabilityNVIDIA GPUs (AMD ROCm support WIP)NVIDIA GPUs
DebuggingPrint + assertCUDA-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):

PRAreaSummary
#187402OptimizersMuon optimizer: spectral_unclamped scaling — new scaling strategy for the Muon optimizer that avoids clamping spectral norms, improving convergence for certain architectures
#186300Distributedc10d abort hooks and pre/post collective hooks — new extensibility points for distributed collectives: register callbacks before and after collectives, and abort hooks for cleanup
#187387DistributedPublic torch.distributed.set_timeout — exposes a public API for setting distributed operation timeouts, replacing internal-only mechanisms
#183838InductorUnbacked FlexAttention predicates — Inductor now supports FlexAttention score_mod/mask_mod with unbacked SymInt predicates, enabling more dynamic attention patterns
#187406Testingtorchfuzz ~190 ops coverage expansion — the torchfuzz fuzzing framework now covers approximately 190 PyTorch operators, up from the initial set
#186976Dynamoobject() 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

TopicKey Takeaway
TritonWrite GPU kernels in Python with near-CUDA performance
Programming modelGrid of blocks, program_id, BLOCK_SIZE, load/store with masks
FusionCombine operations to eliminate memory bandwidth waste
SoftmaxPractical kernel: load row, compute in registers, write once
MatmulTiled approach with tl.dot for Tensor Core utilization
PyTorch integrationcustom_op → register_fake → register_autograd pipeline
Autotuning@triton.autotune automatically finds the best config
TorchInductorGenerates 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


Notebook: 25_triton_kernels.ipynb

Source Files

GPU Memory Profiling & Optimization — Every Byte Accounted For

Table of Contents


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:

OptimizerState per ParameterTotal for N params (fp32)
SGD (no momentum)00
SGD + momentum1× (momentum buffer)4N bytes
Adam / AdamW2× (first moment m, second moment v)8N bytes
Adam + master weights (bf16 training)2× moments + 1× master copy12N 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 you del a tensor, this number drops.
  • memory_reserved() — total memory PyTorch has claimed from CUDA. This only drops when you call empty_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

RowMeaning
Allocated BytesMemory occupied by live tensors
Active BytesSame as allocated (non-released blocks)
Reserved BytesTotal memory held by the caching allocator
Inactive Split BytesFreed memory within split blocks (fragmentation indicator)
Allocation countHow many cudaMalloc calls were made
Active allocsNumber 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 PatternDescription
allocated_bytes.all.currentCurrent allocated memory (bytes)
allocated_bytes.all.peakPeak allocated memory
reserved_bytes.all.currentCurrent reserved memory
reserved_bytes.all.peakPeak reserved memory
active.all.currentNumber of currently active allocations
active.all.peakPeak number of simultaneous allocations
inactive_split_bytes.all.currentCurrent fragmentation (inactive splits)
num_alloc_retriesHow many times allocator retried after cudaMalloc failure
num_oomsNumber 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_entries to 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/cudaFree calls
# 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

ApproachMemory vs. AdamTrade-off
SGD + momentum50% lessMay need different hyperparameters
8-bit Adam (bitsandbytes)75% lessSlight accuracy impact
Adafactor~50% lessUses row/column factorization
LoRA (low-rank adaptation)90%+ lessOnly trains adapter weights
GaLore~65% lessProjects 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

TechniqueSavesCostTypical Reduction
Gradient checkpointingActivations~33% more compute60-70% activation memory
Mixed precision (bf16)Parameters + activationsNone (sometimes better)~50%
Gradient accumulationActivationsNoneProportional to accum steps
In-place operationsIntermediate tensorsAutograd limitations5-15%
del + gc.collect()Named intermediatesManual managementVariable
CPU offloadingParameters + optimizerTransfer overheadUp to 90% GPU memory
8-bit optimizersOptimizer stateSlight accuracy impact75% 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 LengthStandard AttentionFlash AttentionSavings
51216 MB0.25 MB64×
2048256 MB1 MB256×
40961 GB2 MB512×
1638416 GB8 MB2048×
65536256 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_attention is 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


Notebook: 26_memory_profiling.ipynb

Source Files

Multi-GPU Inference Patterns — Serving Large Models at Scale

Table of Contents


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_heads and num_kv_heads evenly
  • 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

AspectTensor ParallelPipeline Parallel
CommunicationAll-reduce per layerPoint-to-point between stages
LatencyLower (all GPUs active per token)Higher (pipeline bubble)
ThroughputLimited by TP commScales with micro-batches
Best interconnectNVLink (intra-node)Works on PCIe/InfiniBand
Memory balanceEven (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

AspectPython (torch.compile)AOTInductor (.so)
Startup timeCompile on first callPre-compiled, instant
Python GILYes, limits concurrencyNo Python needed
DeploymentNeeds Python + PyTorchJust the .so + libtorch
DebuggingFull Python stack tracesC++ debugging
Use caseDevelopment, prototypingProduction serving

10. Benchmarking Inference

10.1 Key Metrics

MetricDefinitionTarget
TTFT (Time to First Token)Time from request to first generated token< 500ms
ITL (Inter-Token Latency)Time between consecutive generated tokens< 30ms
ThroughputTotal tokens generated per second across all requestsMaximize
p50 / p95 / p99 latencyPercentile latency across requestsp99 < 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

ScenarioRecommended StrategyWhy
7B model, 1 GPUtorch.compile(mode="reduce-overhead")Fits easily, maximize latency
7B model, high QPSContinuous batching on 1 GPUMaximize throughput
70B model, 2 GPUsTP=2 + INT4 quantizationFits with INT4, TP for latency
70B model, 4 GPUsTP=4 (FP16) or TP=2 (INT8)Full precision or quant + TP
70B model, 8 GPUsTP=4 + PP=2TP within node, PP across
405B model, 16 GPUsTP=8 + PP=2 + INT8Multi-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-overhead for latency, max-autotune for 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


Notebook: 27_multi_gpu_inference.ipynb

Source Files

torch.utils.benchmark Deep Dive — Measuring Performance Correctly

Table of Contents


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

FactorEffectFix
No CUDA syncMeasures launch time, not compute timetorch.cuda.synchronize()
Cold startFirst call initializes CUDA context (~1-3s)Warmup iterations
JIT compilationtorch.compile first call is slowSeparate warmup phase
cuDNN autotuningFirst convolution triggers autotunertorch.backends.cudnn.benchmark = True before warmup
Garbage collectionGC pauses inject random latency spikesDisable GC during measurement
CPU frequency scalingDynamic clocks cause variancePin CPU frequency or use instruction counts
Memory cachingCUDA caching allocator reuses memoryConsistent 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

ParameterTypeDescription
stmtstrThe code to benchmark (can be multi-line)
setupstrCode run once before measurement (imports, tensor creation)
globalsdictVariables accessible in stmt and setup
num_threadsintCPU threads to use (controls torch.set_num_threads)
labelstrRow label for Compare tables
sub_labelstrSub-row label for Compare tables
descriptionstrColumn label for Compare tables
envstrEnvironment name (for cross-environment comparison)
timercallableCustom 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_time seconds
  • 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

ParameterDefaultDescription
min_run_time2.0Minimum total wall time in seconds
callbackNoneCalled after each block with intermediate results

4. Measurement Object

Both timeit() and blocked_autorange() return a Measurement object:

result = t.blocked_autorange()

Key Attributes

AttributeTypeDescription
result.meanfloatMean time per execution (seconds)
result.medianfloatMedian time per execution (seconds)
result.timesList[float]All measured times (per execution)
result.number_per_runintNumber of stmt executions per block
result.raw_timesList[float]Raw block times (total, not per execution)
result.iqrfloatInterquartile range
result.significant_figuresintStable 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 section
  • description: 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

ParameterDescription
minval, maxvalRange for generated values
distribution"uniform" or "loguniform"

FuzzedTensor Options

ParameterDescription
sizeTuple of parameter names or ints
probability_contiguousProbability the tensor is contiguous (0.0-1.0)
min_elementsMinimum total elements
max_elementsMaximum total elements
dtypeTensor 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 CaseWall ClockCallgrind
Comparing two implementations✓✓
Micro-benchmarks (< 1μs)Noisy✓
CI regression testingNoisy✓
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:

PRFeatureImpact
#187218FlexGEMM BMM supportNew batched matmul paths to benchmark
#187605Dynamo RangeVariable symbolic specializationChanged compile behavior for range-based loops
#187494Distributed backend accessorsCleaner backend switching for distributed benchmarks
#187602ShapesSpec in non-strict exportBetter shape control for exported model benchmarks
#186398DTensor logspaceNew distributed tensor op to benchmark
N/ACUPTI profiler refactoring into _cupti/ packageCleaner 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

ToolPurposeWhen to Use
Timer.timeit(N)Fixed N iterationsQuick checks, known iteration count
Timer.blocked_autorange()Auto iterations, robust statsMost benchmarks (recommended default)
CompareSide-by-side formatted tableComparing implementations, shapes, configs
FuzzerRandom test configurationsThorough coverage, avoiding bias
CallgrindDeterministic instruction countsMicro-benchmarks, CI regression tests

Key Rules

  • Always use torch.utils.benchmark — never time.time() or raw timeit
  • Warmup torch.compile before measuring — compilation time is not runtime
  • Use blocked_autorange() — it handles warmup, GC, iteration count
  • Pin num_threads for CPU benchmarks — reproducibility requires it
  • Use Compare tables — organized comparison beats ad-hoc prints
  • Report median, not mean — outlier resistance matters

Further Resources


Notebook: 28_benchmarking.ipynb

Source Files

Mixed Precision Deep Dive — FP32, FP16, BF16, and FP8

Table of Contents


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

FormatBitsExponentMantissaRangePrecisionPyTorch dtype
FP3232823±3.4e38Hightorch.float32
TF3219810±3.4e38Medium(internal)
BF161687±3.4e38Lowtorch.bfloat16
FP1616510±65504Medium-lowtorch.float16
FP8 E4M3843±448Very lowtorch.float8_e4m3fn
FP8 E5M2852±57344Lowesttorch.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

ComponentFP32Mixed (BF16 compute)Savings
Model parameters (compute copy)4B/param2B/param2×
Activations4B/element2B/element2×
Gradients4B/param2B/param2×
Optimizer states (Adam)8B/param8B/param1× (kept in FP32)
Master weights—4B/paramOverhead

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.bmm
  • torch.nn.functional.linear
  • torch.nn.functional.conv1d/2d/3d
  • torch.baddbmm

Keep in FP32 (numerically sensitive):

  • torch.nn.functional.softmax
  • torch.nn.functional.cross_entropy, all loss functions
  • torch.nn.functional.layer_norm, batch_norm, group_norm
  • torch.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_interval consecutive steps have no inf/nan → multiply scale by growth_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

PropertyFP16BF16
Max value655043.4×10^38
Min positive normal6.1×10^-51.2×10^-38
Precision (decimal digits)~3.3~2.1
GradScaler neededYesNo
Overflow riskHighNone (same as FP32)
Hardware supportAll GPUs with tensor coresAmpere+ (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

SettingEffectRecommendation
param_dtypeCast parameters to this dtype for forward/backwardtorch.bfloat16
reduce_dtypeDtype for gradient all-reduce communicationtorch.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)

PrecisionMatmul ThroughputMemoryUse Case
FP321× (baseline)4B/paramDebugging, validation
TF32~2×4B/paramDefault (transparent)
FP16 + scaler2-3×2B/paramOlder GPUs (V100, T4)
BF162-3×2B/paramLLM training (standard)
FP8 (H100)4-6× vs FP321B/paramLarge-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

ConceptKey Takeaway
BF16Default for training on Ampere+ GPUs — same range as FP32, no scaler needed
FP16 + GradScalerRequired for older GPUs — GradScaler prevents gradient underflow
FP8Cutting edge — 2× over BF16 on H100, requires careful scaling
autocastAutomatically handles per-op precision — just wrap forward pass
GradScalerOnly needed for FP16 — dynamically scales loss to prevent underflow
FSDP2 MixedPrecisionPolicyCompute in BF16, reduce in FP32 for distributed stability
torch.compileFuses casts, eliminates redundant precision changes

Further Resources


Notebook: 29_mixed_precision.ipynb

Source Files

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


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

ProblemWithout detect_anomalyWith detect_anomaly
NaN in backwardSilent propagationRuntimeError with traceback
In-place op on grad tensorCryptic error laterImmediate error at the op
Double backward without retain_graphConfusing errorClear 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

CauseExampleFix
Learning rate too highGradients explode → weights overflowReduce LR, use gradient clipping
log(0)torch.log(probabilities) where some are 0torch.log(x + 1e-8) or torch.clamp(x, min=1e-8)
Division by zerox / norm where norm is 0x / (norm + 1e-8)
Softmax overflowVery large logits → exp overflowUse log_softmax instead of log(softmax(x))
sqrt(0) gradienttorch.sqrt(x) at x=0 has infinite gradienttorch.sqrt(x + 1e-8)
Unstable lossCross-entropy with raw probabilitiesUse 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

ErrorCauseFix
mat1 and mat2 shapes cannot be multipliedLinear layer input size wrongCheck 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 inputCalculate 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) not torch.zeros(...)
  • Loss target on wrong device: target = target.to(device)
  • Buffer not registered: Use self.register_buffer('name', tensor) not self.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

CauseExampleFix
print() in compiled codeprint(x.shape)Remove or guard with if not torch.compiler.is_compiling()
Data-dependent control flowif x.sum() > 0:Use torch.where or torch.cond
Unsupported Python builtinsorted(list)Rewrite with torch ops
Non-tensor data structuresBuilding a list in a loopUse tensor operations
Calling uncompiled functionsExternal library callsWrap 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

#ErrorCauseFix
1RuntimeError: CUDA out of memoryGPU memory exhaustedReduce batch size, use gradient checkpointing, use mixed precision, call torch.cuda.empty_cache()
2RuntimeError: Expected all tensors to be on the same deviceMixing CPU and CUDA tensorsMove all tensors to same device with .to(device)
3RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operationIn-place op on a tensor needed for backwardReplace x.add_(1) with x = x + 1, avoid in-place ops on leaf tensors
4RuntimeError: element 0 of tensors does not require grad and does not have a grad_fnCalling .backward() on a non-grad tensorEnsure inputs have requires_grad=True, check model parameters
5torch._dynamo.exc.UnsupportedGraph break in torch.compileSee Section 8 — remove unsupported ops or use torch.compiler.is_compiling() guard
6CUDA error: device-side assert triggeredIndex out of bounds in CUDA kernelRun with CUDA_LAUNCH_BLOCKING=1, check label indices < num_classes
7RuntimeError: Trying to backward through the graph a second timeCalling .backward() twice without retain_graph=TrueAdd retain_graph=True or restructure to avoid double backward
8RuntimeError: expected scalar type Float but found HalfDtype mismatch between FP32 and FP16Use autocast or explicit .float() / .half() conversion
9RuntimeError: mat1 and mat2 shapes cannot be multipliedLinear layer shape mismatchCheck in_features matches the flattened input dimension
10ValueError: optimizer got an empty parameter listNo parameters passed to optimizerCheck model.parameters() is not empty, ensure modules are registered as attributes
11RuntimeError: Input type and weight type should be the sameMixed dtypes (e.g., double input, float weights)Use x = x.float() or model.double()
12RuntimeError: expected stride to be a single integer or a list of integersWrong argument type to operationCheck 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:

PRTitleImpact
#187768MPS FlexAttention lse returnFlexAttention on MPS now correctly returns log-sum-exp alongside attention output
#187758Sequential.__getitem__ type overloadsBetter type checking when indexing nn.Sequential — clearer errors for invalid indexing
#187776SymmMem copy optimizationOptimized symmetric memory copy for distributed training — reduced latency
#187702vmap batching rule for repeat_interleavetorch.vmap now supports repeat_interleave — no more manual unbatching workaround
#184653Dynamo globals fix for unregistered modulesFixed graph break when accessing global modules not registered as submodules — helps torch.compile debugging
#187778all_to_all_nd narrow-row throughput fixImproved 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 lse values — 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

PrincipleImplementation
Always create a minimal reproStrip to smallest code that reproduces the bug
Use anomaly detectiontorch.autograd.set_detect_anomaly(True) — but only when debugging
Check NaN/Inf earlyRegister forward hooks on all modules
Monitor gradient normsLog per-layer gradient norms each step
Print shapes systematicallyHooks > manual prints > torchinfo
Fix device mismatches at data boundary.to(device) right after data loading
Use environment variables for C++ issuesTORCH_SHOW_CPP_STACKTRACES=1, CUDA_LAUNCH_BLOCKING=1
Use explain() for compile issuestorch._dynamo.explain(fn)(inputs)
Store scalars not tensorsloss.item() not loss
Set all seeds for reprotorch.manual_seed, np.random.seed, random.seed

Further Resources


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


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.compile for fused, optimized kernels
  • Tensor subclass-based: Uses PyTorch's tensor subclass system — no graph rewrites needed

Why torchao matters

Quantization and sparsity can provide:

TechniqueMemory ReductionSpeedupUse 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_mapping boilerplate
  • 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's nn.Module tree
  • 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_point maps 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):

PrecisionMemoryRelative
FP3228 GB1.0×
FP16/BF1614 GB0.5×
INT87 GB0.25×
INT43.5 GB0.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-OnlyDynamic
WeightsINT8/INT4INT8
ActivationsFP16/BF16INT8 (computed at runtime)
Matmul precisionFP16INT8
Best forMemory-bound (batch=1)Compute-bound (batch>1)
Accuracy impactLowerSlightly higher
Extra overheadNonePer-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

FeatureMinimum GPU
FP8 inferenceH100, L40S, MI300
FP8 trainingH100, MI300
FP8 with per-row scalingH100

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

QuantizerTarget HardwareTypical Use
XNNPACKQuantizerARM CPU (mobile)Android/iOS inference
X86InductorQuantizerx86 CPUServer-side CPU inference
QNNPackQuantizerARM CPULegacy mobile path

When to use PT2E vs quantize_()

quantize_() (torchao)PT2E
Ease of useOne-linerMulti-step pipeline
Needs calibrationNo (weight-only/dynamic)Yes (static quant)
Backend-specificNo (generic)Yes (XNNPack, x86, etc.)
Works with compileYes (primary use case)Yes (through export)
Best forGPU inference, LLM servingMobile/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, then torch.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

ScenarioMethodMemorySpeedupAccuracy
LLM serving (batch=1)int4_weight_only(group_size=128)4× less2–4×Good
LLM serving (batch=1, quality)int8_weight_only()2× less1.5–2×Very good
Batch inference (batch>8)int8_dynamic_activation_int8_weight()2× less2–3×Good
H100 inferencefloat8_dynamic_activation_float8_weight()2× less1.5–2×Excellent
H100 training (large models)Float8Linear2× less1.3–1.5×Excellent
Mobile deploymentPT2E + XNNPack4× less2–4×Good
Prunable model2:4 sparsity2× 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

ConceptDescription
QuantizationReducing numerical precision (FP16→INT8/INT4) to save memory and increase speed
Weight-onlyOnly weights are quantized; activations stay in higher precision
DynamicBoth weights and activations are quantized; activation scales computed at runtime
Tensor subclasstorchao's approach: quantized weights are special tensor objects that handle dequant transparently
Semi-structured sparsity2:4 pattern — hardware-accelerated on NVIDIA Ampere+
PT2EExport-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.compile works 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


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


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

MetricTargetProblem if missed
GPU utilization>90%Data starvation
Data loading time< compute timeGPU idle cycles
Memory usageStable over timeOOM from leaks
Worker utilizationBalancedStragglers 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 caseMap-styleIterable
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_factor batches 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=True requires 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 %DiagnosisFix
< 10%Compute-bound (good)Focus on model optimization
10-30%Mild bottleneckMore workers, prefetching
30-50%Significant bottleneckPre-process data, faster storage
> 50%Severe bottleneckRestructure pipeline entirely

Solutions by Severity

  • Quick wins: Increase num_workers, set pin_memory=True, increase prefetch_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

PatternI/O costCPU costFlexibilityMemory
Offline pre-processLowLowLowHigh (disk)
On-the-flyHighHighHighLow
Cached transformsLow (after warmup)LowMediumHigh (disk)
Multi-stageLowDistributedHighMedium

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


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


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 TypeRegistrationSignatureWhen It Runs
Forward hookmodule.register_forward_hook(fn)fn(module, input, output)After forward() returns
Forward pre-hookmodule.register_forward_pre_hook(fn)fn(module, input)Before forward() executes
Backward hookmodule.register_full_backward_hook(fn)fn(module, grad_input, grad_output)During backward()

Plus one tensor-level hook:

Hook TypeRegistrationSignatureWhen It Runs
Tensor hooktensor.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 L to 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

MethodGranularitySharpnessSpeedComplexity
Saliency mapsPixelLowFastTrivial
Grad-CAMRegionMediumFastLow
Guided backpropPixelHighFastLow
Guided Grad-CAMPixelHighFastMedium

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.compile in 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

GoalMethod
"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


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


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

AspectFull Fine-TuningParameter-Efficient (PEFT)
Parameters updatedAll (billions)Small subset (millions)
MemoryVery high (full optimizer state)Low (only adapter state)
Training speedSlowFast
Risk of forgettingHigherLower
Multiple tasksOne model per taskOne base + multiple adapters
GPU requirementMulti-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-projection
  • A ∈ R^(r×k) — up-projection
  • r << 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 × k trainable 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

  • A is initialized from N(0, 1/r) so the initial magnitude is controlled
  • B is initialized to zeros so that B @ A = 0 at 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)

MethodBase WeightsAdaptersOptimizerTotal
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:

LayerBenefitTypically Adapted
Q projectionHighYes
K projectionMediumYes
V projectionHighYes
O projectionMediumSometimes
FFN up-projectionMediumYes
FFN down-projectionMediumYes
EmbeddingsLowNo
LayerNorm/RMSNormLowNo

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

StrategyDescriptionUse Case
Greedy (temperature=0)Always pick highest probabilityFactual Q&A
TemperatureScale logits before softmaxControl randomness
Top-kKeep only k highest-probability tokensModerate diversity
Top-p (nucleus)Keep smallest set with cumulative prob >= pDynamic 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:

Parameter1B7B13B70B
LoRA rank (r)8161632
LoRA alpha16323264
LoRA target layersQKV + FFNQKV + FFNQKV + FFNQKV + FFN
Learning rate3e-42e-41e-45e-5
Batch size (effective)3264128128
Grad accumulation steps481616
Max sequence length512102420482048
Warmup ratio0.030.030.030.03
Epochs3321
Weight decay0.010.010.010.01
Gradient clip1.01.01.01.0
PrecisionBF16BF16BF16BF16
MethodLoRALoRA/QLoRAQLoRAQLoRA
GPUs needed11 (QLoRA) / 2 (LoRA)2-44-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

FileDescriptionLines
README.mdThis guide — complete theory and workflow450+
lora_adapter.pyLoRA implementation, apply/merge, QLoRA concept250+
finetuning_pipeline.pyMini-LLM + LoRA + full training loop300+
evaluation_and_export.pyPerplexity, generation, merge, export200+

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


Notebook: 34_llm_finetuning.ipynb

Source Files

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


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

PriorityKeyPurpose
HighestPythonTLSSnapshotThread-local state snapshot
HighPythonDispatcherPython-level dispatch (torch.compile)
FuncTorchDynamicLayerFrontFront guard for functorch
FunctionalizeConvert mutations to functional ops
AutocastMixed precision dtype casting
AutogradCPURecord op for backward (CPU tensors)
AutogradCUDARecord op for backward (CUDA tensors)
AutogradMPSRecord op for backward (MPS tensors)
AutogradXPURecord op for backward (XPU tensors)
ADInplaceOrViewTrack in-place ops and views for autograd
FuncTorchBatchedvmap batching rules
FuncTorchVmapModevmap mode (outer)
BackendSelectRoute to correct backend for factory ops
LowCPUActual computation on CPU
LowCUDAActual computation on CUDA
LowMPSActual computation on MPS
LowXPUActual computation on XPU
LowMetaShape/dtype computation (no data)
LowestCompositeImplicitAutogradDefault decompositions (autograd-aware)
LowestCompositeExplicitAutogradDecompositions 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

PropertyKey Added
device='cpu'CPU
device='cuda'CUDA
device='mps'MPS
device='meta'Meta
requires_grad=TrueAutogradCPU/CUDA/... (matches device)
Is a view or in-place resultADInplaceOrView

From Thread-Local State

ContextKey Added
Inside torch.autocast(...)Autocast
Inside torch.vmap(...)FuncTorchBatched
Inside torch._dynamoPythonDispatcher
Custom TorchDispatchMode activePython

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

FeatureLibrary API@custom_op
BoilerplateHighLow
Schema inferenceManualAutomatic from type hints
torch.compileManual Meta regregister_fake
AutogradManual Function classregister_autograd
ComposabilityManualBuilt-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.add calls 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_fake work seamlessly — Dynamo uses the fake implementation for tracing

Dispatch Keys Relevant to Compile

  • PythonDispatcher — active when Dynamo is tracing
  • FakeTensor — uses Meta kernels to track shapes during tracing
  • ProxyTorchDispatchMode — 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

FileDescriptionLines
README.mdThis guide — dispatcher internals explained400+
dispatch_keys.pyExplore dispatch keys, priority chains, tables250+
custom_dispatch.pyRegister custom ops, autograd, compile integration250+

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_op let 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


Notebook: 35_dispatcher.ipynb

Source Files

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


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"],
)
ProsCons
No setup.py neededCompiles on first import (slow)
Great for development/iterationMust have compiler on target machine
Automatic cachingCan't pip install
Verbose mode for debuggingNot 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},
)
ProsCons
Standard Python packagingRequires setup.py boilerplate
pip install . worksMust rebuild after changes
Can build wheels for distributionABI compatibility concerns
Integrates with conda/pip ecosystemMore 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 (.so on Linux, .pyd on Windows)
  • Subsequent calls: Loads cached .so from ~/.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:

MacroTypes
AT_DISPATCH_FLOATING_TYPESfloat, double
AT_DISPATCH_FLOATING_TYPES_AND_HALFfloat, double, Half
AT_DISPATCH_ALL_TYPESall integer + float + double
AT_DISPATCH_ALL_TYPES_AND(ScalarType::Half, ...)all + specified extras
AT_DISPATCH_FLOATING_AND_COMPLEX_TYPESfloat, 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:

  • .cu file — CUDA kernels (compiled by nvcc)
  • .cpp file — 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 grid
  • threads = 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

FeatureCppExtensionCUDAExtension
Compilergcc/g++ onlygcc + nvcc
Source files.cpp, .c.cpp, .cu
CUDA headersNot includedAuto-included
GPU supportNoYes
Requires CUDA toolkitNoYes

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_CHECK for 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_ASSERT sparingly 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++:

AspectCUDA C++ ExtensionTriton Kernel
LanguageC++/CUDAPython
Compilationnvcc (complex setup)JIT (automatic)
AutotuningManualBuilt-in @triton.autotune
torch.compileNeeds dispatcher registrationNative support
Debugginggdb, cuda-gdbPython debugger
PortabilityNVIDIA onlyMulti-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.compile integration

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_extension provides two build paths: load() for development, setup.py for distribution
  • torch/extension.h is 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::Function in C++ or torch.library in Python
  • TORCH_CHECK provides 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 BuildExtension to 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


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


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:

KindDescription
InputKind.PARAMETERLearnable parameter (e.g., linear.weight)
InputKind.BUFFERRegistered buffer (e.g., scale)
InputKind.CONSTANT_TENSORLifted constant tensor
InputKind.USER_INPUTThe actual user-provided data
InputKind.TOKENControl flow token (for ordering side effects)

OutputSpec

Each graph output is categorized:

KindDescription
OutputKind.USER_OUTPUTThe actual return value
OutputKind.LOSS_OUTPUTLoss value (for training export)
OutputKind.BUFFER_MUTATIONBuffer that was mutated in-place
OutputKind.USER_INPUT_MUTATIONUser input that was mutated
OutputKind.GRADIENT_TO_PARAMETERGradient (training export)
OutputKind.GRADIENT_TO_USER_INPUTGradient 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] > 5 won'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 operands tuple. 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_fn and body_fn and updated each iteration.
  • No dynamic shape changes across iterations — the shapes of carried inputs are fixed.

Comparison with Python Loops

FeaturePython whiletorch.while_loop
ExportOnly last iteration tracedBoth branches traced
Iteration countCan be dynamicCan be dynamic
In graphUnrolled (if static) or failsSingle loop node
Side effectsAllowedNot 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/else that 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 CaseIR LevelWhy
AOTInductorPost-dispatchBackend needs decomposed ops
ONNX exportPost-dispatchMaps directly to ONNX ops
Graph analysisPre-dispatchHigher-level, easier to read
Custom passesPre-dispatchOperate on meaningful ops
Training exportPre-dispatchPreserve 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 FakeTensor inputs (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

ScenarioMode
Production deploymentstrict=True (catches all issues)
Iterative developmentstrict=False (get something working first)
Complex Python logicstrict=False (then migrate to strict)
Custom frameworksstrict=False (framework code may not trace)
Maximum reliabilitystrict=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

FormatUse CasePreserves GraphPortable
torch.save (pickle)CheckpointingNoNo (needs source)
torch.export.save (PT2)DeploymentYesYes
torch.packageHermetic archiveSource codeYes
ONNXCross-frameworkYesYes

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

  • ExportedProgram is 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.cond and torch.while_loop make control flow explicit — the graph captures both branches
  • Dim API with shared dims and Dim.AUTO gives fine-grained control over dynamic shapes
  • draft_export is 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


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


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) becomes x = x + 1 (out-of-place)
  • x.view(-1) becomes x.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_backward needs %relu (to know which elements were zeroed)
  • grad_w computation needs %x (the original input)
  • grad_x computation 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.compile captures 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:

ActivationSizeSaved? (min-cut)
QKV projection output48 MBYes (matmul)
Attention scores32 MBYes (matmul)
Post-softmax attention32 MBYes (expensive)
ReLU mask2 MBNo (recomputed)
LayerNorm intermediate4 MBNo (recomputed)
Residual add result8 MBNo (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

FeatureStandard AutogradAOTAutograd
Graph builtRuntime (during forward)Compile time (ahead-of-time)
Backward optimizedNo (eager dispatch)Yes (Inductor compiles backward)
Memory planningManual (user calls checkpoint)Automatic (min-cut partitioner)
Kernel fusionNone (one kernel per op)Yes (Inductor fuses backward ops)
Works with compileForward onlyForward + backward
Dynamic graphsFully supportedRequires recompilation on change
Autograd hooksFully supportedLimited support
In-place opsFully supportedFunctionalized (no in-place)
DebuggingEasy (Python stack traces)Harder (compiled code)
OverheadNoneCompilation 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.compile alone 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.grad to 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_function to inspect graphs — pass custom fw_compiler and bw_compiler callbacks to see exactly what AOTAutograd produces
  • Debugging uses TORCH_LOGS — TORCH_LOGS="aot" shows graphs, debug_partitioner=True shows 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


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


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

ComponentWhat It DoesFile
Character TokenizerSplits text into characters, maps to integerstokenizer.py
Word TokenizerSplits on whitespace, builds frequency-based vocabularytokenizer.py
TransformerTextClassifierEmbedding + positional encoding + transformer encoder + classification headtext_classifier.py
Training PipelineSynthetic data, training loop, evaluation, inferencetrain_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:

TokenPurposeTypical ID
[PAD]Fills sequences to equal length for batching0
[UNK]Replaces out-of-vocabulary words1
[CLS]Classification token — its embedding becomes the sequence representation2
[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

ModeCompile TimeInference SpeedUse Case
defaultFastGoodGeneral purpose
reduce-overheadMediumBest for small modelsLow-latency serving
max-autotuneSlowBest for large modelsThroughput-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.Embedding starts random and learns semantic relationships during training; padding_idx prevents 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


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

CRCR starter workflow.

  • Register: Open a PR to add your repo to .github/allowlist.yml under L2.
  • Configure dispatch handler: Add a repository_dispatch workflow 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

SymptomCauseFix
Dispatch never receivedRepo not in allowlistAdd to .github/allowlist.yml
OIDC token mint failsMissing id-token: write permissionAdd permissions: id-token: write to workflow
Results not on HUDCallback URL wrong or Lambda downCheck callback action logs; verify endpoint
Build fails at dispatched SHASubmodules out of syncRun git submodule update --init --recursive
Nightly SHA resolution failsCommit message format changedUpdate grep pattern for SHA extraction

Environment Variables

The dispatch payload sets these environment variables for your workflow:

VariableDescription
github.event.client_payload.pr_numberPR number that triggered the dispatch
github.event.client_payload.head_shaGit SHA to build and test against
github.event.client_payload.base_shaBase branch SHA for diff context
github.event.client_payload.senderGitHub 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.

InputOutput
32x32 RGB image of a circlecircle (0.95)
32x32 RGB image of a starstar (0.91)
32x32 RGB image of a triangletriangle (0.88)

Table of Contents


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

FileLinesDescription
data_pipeline.py270+Synthetic dataset, augmentation, MixUp, CutMix
model_and_training.py310+SimpleCNN, MiniResNet, transfer learning, training loop
evaluation.py270+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:

SplitSamplesPurposeAugmentation
Train4,000Weight updatesFull augmentation
Val500Hyperparameter tuning, early stoppingNormalize only
Test500Final evaluationNormalize 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

TransformParametersEffect
RandomHorizontalFlipp=0.5Mirror left-right
RandomVerticalFlipp=0.3Mirror top-bottom
RandomRotation90—Rotate 0°/90°/180°/270°
ColorJitterbrightness=0.2, contrast=0.2Vary brightness and contrast
RandomErasingp=0.3, scale=(0.02, 0.15)Cutout-style occlusion
Normalizemean=0.5, std=0.5Center 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

AspectMixUpCutMix
BlendingGlobal pixel-wiseLocal rectangular patch
Label mixingBased on λBased on patch area ratio
EffectSmoother decision boundariesBetter localization
Best α0.21.0
Use caseGeneral regularizationWhen 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

ChoiceReason
BatchNorm after ConvNormalizes activations, enables higher learning rates
ReLU (inplace)Saves memory, standard non-linearity
AdaptiveAvgPool(4)Works with any input resolution
Dropout before LinearRegularization in classifier head
3x3 kernels onlyModern 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

ComponentImplementationPurpose
OptimizerAdamWWeight decay decoupled from gradient
LossLabel Smoothing CEPrevents overconfident predictions
SchedulerCosine + WarmupSmooth LR decay with warmup
AMPtorch.autocastMixed precision for speed
EMAExponential moving averageSmoother, more stable weights
MixUp/CutMixRandom per batchRegularization
Early StoppingPatience-basedPrevents overfitting
Gradient ClippingMax norm = 1.0Training 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:

MetricFormulaWhat It Tells You
PrecisionTP / (TP + FP)Of predicted positives, how many are correct
RecallTP / (TP + FN)Of actual positives, how many were found
F1 Score2 · P · R / (P + R)Harmonic mean of precision and recall
Macro F1avg(F1_per_class)Equal weight to each class
Weighted F1weighted avg(F1_per_class)Weight by class support
Top-k Acccorrect in top-k / totalUseful 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

ScenarioUse TTA?
Competition / final submissionYes — free accuracy boost
Real-time inferenceNo — multiplies latency by N
Medical imaging / safety-criticalYes — reliability matters
Development / prototypingNo — 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

ECECalibration Quality
< 0.02Excellent
0.02–0.05Good
0.05–0.10Fair
> 0.10Poor — 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


Notebook: 40_image_classifier.ipynb

Source Files

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.

InputOutput
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

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

FileLinesDescription
noise_schedule.py200+Linear/cosine beta schedules, forward diffusion, alpha_cumprod
unet_model.py300+Sinusoidal embeddings, ResBlocks, UNet with skip connections
train_diffusion.py300+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

ModelTrainingSamplingMode CoverageQuality
GANsAdversarial (unstable)Single pass (fast)Mode collapse riskHigh
VAEsELBO (stable)Single pass (fast)Good coverageBlurry
DiffusionDenoising (stable)Iterative (slow)Excellent coverageHighest
FlowExact likelihoodSingle passGood coverageHigh

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:

ParameterizationModel predictsUsed 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
Scorescore_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

TipReason
Use cosine scheduleBetter noise distribution across timesteps
AdamW optimizerStable training with weight decay
Learning rate ~1e-3 for 2D, ~2e-4 for images2D data is simpler
Gradient clipping (max_norm=1.0)Prevents training instability
EMA of model weightsSmoother, 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

PropertyValue
StepsT (typically 1000)
Stochastic?Yes (random noise at each step)
QualityExcellent
SpeedSlow (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 = 0 for 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

AspectDDPMDDIM
Steps100050-100 (tunable)
StochasticYesConfigurable (eta)
Deterministic modeNoYes (eta=0)
Sample quality at 50 stepsPoorGood
Same noise → same outputNoYes (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

DistributionShapeCharacteristics
Swiss RollSpiralTests ability to learn curved manifolds
Two MoonsTwo crescentsTests multi-modal generation
CirclesConcentric ringsTests 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


Upstream Updates (PyTorch 2.14+)

FeatureImpact on Diffusion Models
torch.compileCompiles the denoising UNet for faster training and sampling
FlexAttentionCustom attention patterns for UNet self-attention layers
torch.float8FP8 training for larger diffusion models
FSDP2Distributed training of billion-parameter diffusion models
torch.exportExport trained diffusion models for deployment

Notebook: 41_diffusion_model.ipynb

Source Files

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.json for 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:

CategoryPrefix patternsRunner requirement
cputest/test_*.py (default)CPU only
inductortest/inductor/, test/dynamo/, test/export/CPU (some GPU)
sgputest/test_cuda*1 GPU
mgputest/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

Performance Metrics

ApproachAvg. tests selectedTime saved vs full suite
File-path heuristic~5-15% of suite70-85% wall time
Structural (TorchTalk)~3-10% of suite80-90% wall time
Merged (union)~8-20% of suite65-80% wall time
Full suite100%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?

ApproachProsCons
Fine-tuningFast inference, no retrieval latencyExpensive, stale knowledge, hallucinations
RAGUp-to-date knowledge, verifiable sourcesRetrieval latency, context window limits
RAG + Fine-tuningBest of both worldsMost 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

FileDescription
README.mdThis guide
rag_pipeline.pyComplete RAG pipeline implementation
chunking_strategies.pyDocument chunking utilities
evaluation.pyRAG evaluation metrics (retrieval recall, answer quality)

References


Source Files

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

FileDescription
batched_inference.pyPadding, masking, and batched forward pass utilities
dynamic_batcher.pyAsync request queue with time/size-triggered dispatch
compiled_serving.pytorch.compile + CUDA Graphs integration for serving
monitoring.pyMetrics collection, health checks, latency tracking
server.pyEnd-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.profiler to find launch-bound vs compute-bound phases
  • Set torch.set_float32_matmul_precision('high') for TF32 on Ampere+

Source Files

Module 44: Performance Case Studies

Case Studies

#Case StudyBottleneckFixSpeedup
1Memory-bound DataLoaderCPU→GPU copy on main threadPin memory + non-blocking transfers2–3×
2Naive attention scalingO(n²) memory, repeated allocationFlash attention pattern + in-place ops4–8×
3Training loop overheadPython dispatch + autograd bookkeepingtorch.compile with graph breaks analysis1.5–3×
4Inference memory bloatGradients + BN running stats retainedFreezing, inference_mode, weight-only quantization60–75% memory reduction
5Multi-GPU communicationAll-reduce blocking computeGradient bucketing + overlap comm/compute1.3–1.8× at scale

Files

FileDescription
case1_dataloader.pyDataLoader pinned memory and prefetch optimization
case2_attention.pyAttention implementation: naive → memory-efficient → flash
case3_compile.pytorch.compile graph breaks analysis and fix
case4_inference_memory.pyInference memory reduction techniques
case5_distributed.pyMulti-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.profiler and torch.cuda.memory_summary() are your friends
  • Memory bandwidth is usually the bottleneck, not compute (especially for inference)
  • torch.compile wins 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

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

FileDescription
profiler_basics.pyContext manager usage, scheduling, key averages
trace_analysis.pyAdvanced trace analysis, memory snapshots, CUDA activity

References

Source Files

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

TypeBitsUse Case
INT88General-purpose server/edge inference
INT44LLM weight-only quantization
FP8 (E4M3/E5M2)8Training and inference on Hopper+ GPUs
UINT44Asymmetric weight packing (torchao)

Files in This Module

FileDescription
dynamic_quantization.pyDynamic quantization with torch.ao.quantization
static_quantization.pyStatic 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

Source Files

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 srun with MASTER_ADDR/MASTER_PORT env 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

FileDescription
ddp_training.pyFull 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=True if 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 module attribute (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 ctx attributes (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

PatternWhen
Elementwise fused opCustom CUDA/Triton forward + matching backward
Numerically stable log-sum-expForward uses max-trick; backward uses softmax
Straight-through estimator (STE)Forward discrete; backward identity
Non-differentiable argReturn 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

FileDescription
autograd_function_basics.pyForward/backward patterns, STE, gradcheck
double_backward.pyHigher-order grads and create_graph

References

Source Files

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

ModeFlagNotes
Reentrant (legacy)use_reentrant=TrueRe-enters autograd; needed for some older patterns; more edge cases
Non-reentrant (preferred)use_reentrant=FalseCleaner 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

FileDescription
selective_checkpoint.pyReentrant vs non-reentrant, SAC-style policy, offload sketch

References

Source Files

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

LayoutBest forStructure
COOConstruction, irregular sparsityindices [ndim, nnz] + values
CSRFast row slices, SpMM on CPU/CUDAcrow_indices, col_indices, values
CSCFast column slicesccol_indices, row_indices, values
BSR / BSCBlock-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

FileDescription
sparse_basics.pyCOO/CSR creation, coalesce, matmul, autograd sketch

References

Source Files