../

PyTorch

Tensors, autograd, nn.Module, the training loop, data loading, devices, mixed precision, torch.compile and checkpoints in PyTorch 2.14 (Python 3.14). Array basics carry over from NumPy; for classical models on tables reach for scikit-learn first.

Setup & mental model

uv add torch torchvision   # or: pip install torch ...
uv pip install torch --torch-backend=auto  # picks CUDA/CPU
uv run python -c "import torch; print(torch.__version__)"

The default PyPI wheels include MPS on macOS and CUDA on Linux. For a specific CUDA or CPU-only build, point uv at the PyTorch index in pyproject.toml:

pyproject.toml
[tool.uv.sources]
torch = [{ index = "pytorch-cu130" }]
 
[[tool.uv.index]]
name = "pytorch-cu130"      # or "pytorch-cpu" + /whl/cpu
url = "https://download.pytorch.org/whl/cu130"
explicit = true             # only torch comes from here
PyTorchClosest TS ideaWhere the analogy breaks
Tensora Float32Array plus a shapeit lives on a device and can record how it was computed
dtypewhich typed array (Float32Array, Int32Array)ops need matching dtypes; no silent number
devicea GPUBuffer in WebGPU: memory owned elsewheremoving data is an explicit copy (.to(device))
nn.Modulea class with state and a forward(x) methodcall model(x), not model.forward(x), so hooks run
state_dict()a Record<string, Tensor> snapshot you can serializeholds weights only, never the class or code
DataLoaderan iterator over batches (for...of)synchronous; parallelism comes from worker processes
torch.compileV8's JIT optimizing hot functions"graph breaks" and recompiles are its deopts

Tensors

tensors.py
import numpy as np
import torch
 
a = torch.tensor([[1.0, 2.0], [3.0, 4.0]])  # copies data
z = torch.zeros(2, 3)                        # float32
g = torch.Generator().manual_seed(0)
r = torch.randn(4, 3, generator=g)           # seeded
i = torch.arange(0, 10, 2)                   # int64
o = torch.ones_like(a)                       # same shape
n = torch.from_numpy(np.ones(3))             # float64!
print(a.shape, a.dtype, a.device)  # (2, 2) float32 cpu
 
a @ a            # matrix multiply
a * a            # elementwise
a.sum(dim=0)     # reduce over rows -> shape (2,)
a.mean(dim=1, keepdim=True)  # shape (2, 1)
a[a > 2]         # boolean mask -> tensor([3., 4.])
a[0, 0].item()   # Python float 1.0
CreateGives
torch.tensor(data)copy of a list / array; infers dtype
torch.as_tensor(arr) / torch.from_numpy(arr)shares memory with a NumPy array when it can
zeros, ones, full(shape, v), emptyfilled (or uninitialised) tensors
arange(start, end, step), linspace(a, b, n)ranges
rand, randn, randint(lo, hi, shape)uniform, normal, integers
eye(n)identity matrix
zeros_like(t), randn_like(t)same shape, dtype and device as t
dtypeUse
float32default for weights and inputs
bfloat16 / float16mixed precision; bf16 has float32's range
float64NumPy's default: cast with .float() before feeding a model; MPS lacks it
int64 (long)default integer; class labels for CrossEntropyLoss, indices
boolmasks
castt.float(), t.long(), t.to(torch.bfloat16)

Shape operations

OpDoesNote
t.shape, t.size(0), t.ndim, t.numel()inspectshape is a torch.Size (a tuple)
t.view(2, -1)reshape without copyingneeds contiguous memory
t.reshape(2, -1)reshape, copies if neededthe safe default
t.unsqueeze(0) / t.squeeze(0)add / drop a size-1 dimadd a batch dim: x.unsqueeze(0)
t.permute(0, 2, 1), t.transpose(1, 2), t.mTreorder dimsresult is non-contiguous
t.flatten(1)flatten all but the batch dimbefore nn.Linear
torch.cat([a, b], dim=0)join along an existing dimshapes match except dim
torch.stack([a, b])join along a new dimshapes must match exactly
t.split(2), t.chunk(3)cut into piecesviews, not copies
t.expand(4, -1) / t.repeat(4, 1)broadcast view / real copyprefer expand
t.contiguous()copy into standard layoutfixes view errors
torch.einsum("bij,bjk->bik", a, b)named-index contractionsreadable batched math

Broadcasting

Shapes are aligned from the right. Each pair of dims must be equal, or one of them is 1 (or missing), and a size-1 dim is stretched. These are the same rules as in NumPy.

x = torch.randn(32, 10)       # batch of 32 rows
mu = x.mean(dim=0)            # (10,)
std = x.std(dim=0)            # (10,)
xn = (x - mu) / std           # (32, 10) op (10,) -> (32, 10)
col = torch.randn(32, 1)
col + mu                      # (32, 1) + (10,) -> (32, 10)

NumPy interop

DirectionCodeMemory
NumPy to tensortorch.from_numpy(arr)shared (CPU)
NumPy to tensortorch.tensor(arr)copied
Tensor to NumPyt.numpy()shared; CPU only, no grad
Any tensor to NumPyt.detach().cpu().numpy()the always-works form
Scalart.item()Python number; syncs with the GPU
Listt.tolist()nested Python lists

Autograd

Operations on tensors with requires_grad=True are recorded; backward() walks the record and adds gradients into each leaf's .grad.

w = torch.tensor([2.0, -1.0], requires_grad=True)
x = torch.tensor([3.0, 4.0])
loss = ((w * x).sum() - 1.0) ** 2   # (6 - 4 - 1)^2 = 1
loss.backward()                     # fills w.grad
print(w.grad)                       # tensor([6., 8.])
 
with torch.no_grad():               # manual SGD step
    w -= 0.01 * w.grad
w.grad = None                       # grads accumulate!
APIDoes
requires_grad=True / t.requires_grad_()track operations on this leaf tensor
loss.backward()compute gradients of a scalar; frees the graph afterward
t.gradaccumulated gradient: adds on every backward until zeroed
opt.zero_grad()reset grads (sets them to None by default)
with torch.no_grad():no recording; for weight updates and eval
with torch.inference_mode():like no_grad but faster; its tensors can never join autograd later
t.detach()same data, cut from the graph
torch.autograd.grad(out, [x])gradients returned instead of stored
backward(retain_graph=True)keep the graph for a second backward
t.grad_fnthe op that produced t (None for leaves)

nn.Module

model.py
import torch
from torch import Tensor, nn
 
 
class MLP(nn.Module):
    def __init__(
        self, d_in: int, d_hidden: int, n_classes: int
    ) -> None:
        super().__init__()          # required first
        self.net = nn.Sequential(
            nn.Linear(d_in, d_hidden),
            nn.ReLU(),
            nn.Dropout(0.1),
            nn.Linear(d_hidden, n_classes),  # logits
        )
 
    def forward(self, x: Tensor) -> Tensor:
        return self.net(x)
 
 
model = MLP(20, 64, 3)
logits = model(torch.randn(8, 20))   # (8, 3)
n_params = sum(p.numel() for p in model.parameters())
LayerInput to outputUse
nn.Linear(in, out)(N, *, in) to (N, *, out)dense layer
nn.Conv2d(c_in, c_out, k, padding=)(N, C, H, W)images
nn.MaxPool2d(k), nn.AdaptiveAvgPool2d(1)shrink H and Wdownsampling, global pooling
nn.BatchNorm1d/2d(c)same shapeCNNs; behaves differently in train/eval
nn.LayerNorm(d), nn.RMSNorm(d)same shapetransformers, MLPs
nn.Dropout(p)same shaperegularization; off in eval()
nn.Embedding(vocab, d)int64 (N, L) to (N, L, d)token or category ids
nn.LSTM, nn.GRU (batch_first=True)(N, L, d)sequences
nn.MultiheadAttention(d, heads, batch_first=True)(N, L, d)attention
nn.TransformerEncoderLayer(d, heads, batch_first=True)(N, L, d)stock encoder block
nn.ReLU, nn.GELU, nn.SiLUsame shapeactivations
nn.Flatten()(N, ...) to (N, -1)between conv and linear
nn.Sequential, nn.ModuleList, nn.ModuleDictcontainersregister sub-modules (a plain list does not)
MemberDoes
model.parameters(), named_parameters()trainable tensors, for the optimizer
model.train() / model.eval()toggle dropout and batch-norm behavior; nothing to do with grads
model.to(device)moves parameters and buffers, in place
self.register_buffer("mask", t)state that moves and saves with the model but is not trained
p.requires_grad_(False)freeze a parameter
model.apply(fn)run fn on every sub-module (custom init)
print(model)the layer tree

Losses, optimizers & schedulers

LossModel outputTarget
nn.CrossEntropyLoss()raw logits (N, C), no softmaxint64 class ids (N,) or float probabilities (N, C)
nn.BCEWithLogitsLoss(pos_weight=)raw logits, any shapefloat 0/1, same shape (multi-label too)
nn.MSELoss() / nn.L1Loss()valuesvalues, same shape
nn.HuberLoss(delta=)valuesrobust regression
nn.NLLLoss()log_softmax outputint64 class ids
nn.KLDivLoss(reduction="batchmean")log-probabilitiesprobabilities (distillation)
CrossEntropyLoss(weight=w, label_smoothing=0.1)logitsper-class weights for imbalance; smoothing curbs over-confidence
OptimizerNotes
AdamW(params, lr=1e-3, weight_decay=0.01)the default choice; decoupled weight decay
AdamAdamW without decoupled decay
SGD(lr, momentum=0.9, nesterov=True)CNNs with a schedule; needs more LR tuning
RMSpropRNNs, reinforcement learning
Muon2-D hidden weight matrices only; pair with AdamW for the rest
LBFGSsmall full-batch problems; step(closure)
Param groupsAdamW([{"params": a, "lr": 1e-4}, {"params": b}], lr=1e-3)
Schedulerstep() everyShape
StepLR(opt, step_size, gamma)epochdrop by gamma every step_size
CosineAnnealingLR(opt, T_max)epoch or batchcosine decay to eta_min
OneCycleLR(opt, max_lr, total_steps)batchwarm up then anneal; fast convergence
LinearLR + SequentialLRepoch or batchlinear warm-up, then another schedule
ReduceLROnPlateau(opt, patience=)epoch, step(val_loss)cut LR when the metric stalls
LambdaLR(opt, fn)anycustom multiplier

Call scheduler.step() after optimizer.step(), and save both state dicts in checkpoints.

opt = torch.optim.AdamW(model.parameters(), lr=3e-3)
warmup = torch.optim.lr_scheduler.LinearLR(
    opt, start_factor=0.1, total_iters=5
)
cosine = torch.optim.lr_scheduler.CosineAnnealingLR(
    opt, T_max=45
)
sched = torch.optim.lr_scheduler.SequentialLR(
    opt, [warmup, cosine], milestones=[5]
)

Training loop

loop.py
from torch.utils.data import DataLoader
 
Batch = tuple[Tensor, Tensor]
 
 
def train_epoch(
    model: nn.Module,
    loader: DataLoader[Batch],
    loss_fn: nn.Module,
    opt: torch.optim.Optimizer,
    device: torch.device,
) -> float:
    model.train()                     # dropout/BN on
    total, seen = 0.0, 0
    for xb, yb in loader:
        xb, yb = xb.to(device), yb.to(device)
        opt.zero_grad()               # 1. clear old grads
        loss = loss_fn(model(xb), yb) # 2. forward
        loss.backward()               # 3. backprop
        opt.step()                    # 4. update weights
        total += loss.item() * len(xb)
        seen += len(xb)
    return total / seen
@torch.inference_mode()               # no graph, faster
def evaluate(
    model: nn.Module,
    loader: DataLoader[Batch],
    loss_fn: nn.Module,
    device: torch.device,
) -> tuple[float, float]:
    model.eval()                      # dropout/BN off
    loss_sum, correct, seen = 0.0, 0, 0
    for xb, yb in loader:
        xb, yb = xb.to(device), yb.to(device)
        logits = model(xb)
        loss_sum += loss_fn(logits, yb).item() * len(xb)
        correct += (logits.argmax(1) == yb).sum().item()
        seen += len(xb)
    return loss_sum / seen, correct / seen

Dataset & DataLoader

data.py
from torch.utils.data import TensorDataset, random_split
 
g = torch.Generator().manual_seed(0)
X = torch.randn(1_000, 20, generator=g)
y = (X[:, 0] + X[:, 1] > 0).long() + (X[:, 2] > 1).long()
full = TensorDataset(X, y)            # rows of (x, y)
train_ds, val_ds = random_split(
    full, [0.8, 0.2], generator=g
)
 
train_loader = DataLoader(
    train_ds, batch_size=64, shuffle=True, generator=g
)
val_loader = DataLoader(val_ds, batch_size=256)
xb, yb = next(iter(train_loader))     # (64, 20), (64,)
PieceDoes
Dataset (map-style)implement __len__ and __getitem__(i), returning one sample
IterableDatasetimplement __iter__; streams, no random access
TensorDataset(X, y)wrap tensors that are already in memory
random_split(ds, [0.8, 0.2])fractions or counts; pass a generator
Subset(ds, indices)a view over some indices (e.g. from sklearn splits)
batch_size, shuffle=Trueshuffle training data only
num_workers=4load in subprocesses; guard the script with if __name__ == "__main__": on macOS/Windows
pin_memory=Truefaster host-to-GPU copies with .to(device, non_blocking=True)
persistent_workers=Truekeep workers alive between epochs
drop_last=Trueskip a ragged final batch (helps BatchNorm, torch.compile)
collate_fn=fncustom batching: padding variable-length sequences

Devices

def pick_device() -> torch.device:
    if torch.cuda.is_available():
        return torch.device("cuda")
    if torch.backends.mps.is_available():   # Apple GPU
        return torch.device("mps")
    return torch.device("cpu")
 
# device-agnostic API; None when no accelerator
acc = torch.accelerator.current_accelerator(
    check_available=True
)
device = acc if acc is not None else torch.device("cpu")
model = model.to(device)            # before the optimizer
RuleWhy
Model and batch on the same deviceotherwise "Expected all tensors to be on the same device"
Create the optimizer after model.to(device)it should hold the moved parameters
Create tensors on the device: torch.zeros(3, device=device)avoids a CPU allocation plus a copy
Avoid .item(), .cpu(), print(t) inside the hot loopeach one waits for the GPU
MPS: no float64cast to float32 first
CUDA_VISIBLE_DEVICES=1pick the GPU from the shell
torch.cuda.max_memory_allocated()peak memory, for batch-size tuning

Mixed precision

Run matmuls and convolutions in 16-bit inside torch.autocast; keep weights and the optimizer in float32. float16 needs a GradScaler to stop small gradients underflowing to zero; bfloat16 does not.

DevicedtypeScaler
CUDA (Ampere+)torch.bfloat16not needed
CUDA (older)torch.float16torch.amp.GradScaler("cuda")
CPUtorch.bfloat16not needed
MPStorch.float16 or torch.bfloat16fp16: scaler recommended
use_fp16 = device.type == "cuda"
amp_dtype = torch.float16 if use_fp16 else torch.bfloat16
scaler = torch.amp.GradScaler(device.type, enabled=use_fp16)
loss_fn = nn.CrossEntropyLoss()
 
for xb, yb in train_loader:
    xb, yb = xb.to(device), yb.to(device)
    opt.zero_grad()
    with torch.autocast(device.type, dtype=amp_dtype):
        loss = loss_fn(model(xb), yb)   # forward in 16-bit
    scaler.scale(loss).backward()       # no-op if disabled
    scaler.step(opt)                    # skips inf/NaN steps
    scaler.update()

torch.cuda.amp.* and torch.cpu.amp.* are deprecated: use torch.amp.GradScaler(device) and torch.autocast(device_type).

torch.compile

Traces Python into graphs and generates fused kernels. The first call is slow while it compiles; later calls are faster, often 1.3–2× for training on GPU.

compiled = torch.compile(model)       # returns a wrapper
model.compile()                       # or compile in place
 
@torch.compile(fullgraph=True)        # error on graph breaks
def gelu_mlp(x: Tensor, w: Tensor) -> Tensor:
    return torch.nn.functional.gelu(x @ w)
OptionEffect
mode="default"good balance, quick compile
mode="reduce-overhead"CUDA graphs; small batches, less Python overhead
mode="max-autotune"benchmarks kernel choices; slow compile, fastest run
fullgraph=Trueraise instead of silently splitting at unsupported code
dynamic=Trueexpect varying shapes and avoid recompiles
TORCH_LOGS="graph_breaks,recompiles"see what broke the graph and why it recompiled
torch.compile(model) state dictkeys gain an _orig_mod. prefix; model.compile() keeps them clean

Compile once, then train. Varying batch shapes, .item() and data-dependent Python branches in forward cause graph breaks or recompiles.

Saving & loading

torch.save(model.state_dict(), "model.pt")  # weights only
 
restored = MLP(20, 64, 3)                   # rebuild class
state = torch.load(
    "model.pt", map_location="cpu", weights_only=True
)
missing, unexpected = restored.load_state_dict(state)
restored.eval()
TopicRule
What to savestate_dict(), not the module object (torch.save(model) pickles code paths)
weights_onlydefaults to True since 2.6: only tensors and plain containers load
Custom classes in a checkpointtorch.serialization.add_safe_globals([Cls]), or weights_only=False for files you trust
map_location="cpu"load GPU checkpoints on a CPU-only machine
load_state_dict(strict=False)partial loads; returns the missing and unexpected keys
Sharing weightssafetensors (safetensors.torch.save_file): no pickle at all
Deploymenttorch.export.export(model, (x,)), or ONNX via torch.onnx.export(..., dynamo=True)

Debugging errors

Error (abridged)Usual causeFix
mat1 and mat2 shapes cannot be multiplied (32x128 and 784x10)in_features doesn't match the inputprint x.shape before the layer; fix nn.Linear(128, ...) or flatten
mat1 and mat2 must have the same dtype, but got Double and FloatNumPy float64 inputtorch.from_numpy(a).float()
Expected all tensors to be on the same devicemodel or batch left on CPU.to(device) both
0D or 1D target tensor expected, multi-target not supportedCrossEntropyLoss target shaped (N, 1)y.squeeze(1)
expected target dtype to be Long or Byte, but got Floatclass ids stored as floatsy.long()
The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 1shapes don't broadcastcheck both shapes; unsqueeze the smaller
Trying to backward through the graph a second timereusing a graph (RNN hidden state across batches)hidden.detach() between batches
element 0 of tensors does not require gradloss built under no_grad/inference_mode or frozen paramscompute the loss outside those blocks
a leaf Variable that requires grad is being used in an in-place operationmanual update on a parameterwrap it in torch.no_grad()
view size is not compatible with input tensor's size and strideview after permute/transpose.reshape(...) or .contiguous().view(...)
Can't call numpy() on Tensor that requires gradconverting a live tensort.detach().cpu().numpy()
Weights only load failedcheckpoint contains non-tensor objectsadd_safe_globals, or trust the file and pass weights_only=False
CUDA out of memorybatch too big; eval keeping graphssmaller batch, AMP, inference_mode for eval, gradient accumulation
Loss is nanLR too high, log(0), fp16 overflowlower LR, clip grads, torch.autograd.set_detect_anomaly(True)
Loss flat, accuracy randomforgot zero_grad, softmax before CrossEntropyLoss, LR offcheck the four-step loop; overfit one batch first
Eval scores differ run to runstill in train() modemodel.eval()

Reproducibility

import os
import random
 
import numpy as np
 
 
def seed_everything(seed: int = 0) -> torch.Generator:
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)          # CPU, CUDA and MPS
    return torch.Generator().manual_seed(seed)
 
 
def strict_determinism() -> None:
    os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
    torch.use_deterministic_algorithms(True)  # raise if not
    torch.backends.cudnn.benchmark = False
Source of randomnessControl
Weight init, dropouttorch.manual_seed
Shuffling, random_splitpass generator= to DataLoader and random_split
Worker processesworker_init_fn that seeds random/numpy from torch.initial_seed()
Non-deterministic GPU kernelstorch.use_deterministic_algorithms(True) (slower)
Hardware, driver, library versionsonly identical setups give bitwise-identical results

Recipes

These build on MLP, train_epoch, evaluate, train_loader, val_loader and device from the sections above.

MLP classifier on synthetic data

When you want to check that a model, loss and loop can learn at all before touching real data.

torch.manual_seed(0)
device = pick_device()
model = MLP(20, 64, 3).to(device)
opt = torch.optim.AdamW(model.parameters(), lr=3e-3)
loss_fn = nn.CrossEntropyLoss()
 
for epoch in range(20):
    tr = train_epoch(
        model, train_loader, loss_fn, opt, device
    )
    va, acc = evaluate(model, val_loader, loss_fn, device)
    if epoch % 5 == 4:
        print(f"{epoch:2} train {tr:.3f} val {va:.3f} "
              f"acc {acc:.2%}")

Custom Dataset from a CSV

When features live in a CSV: parse once with pandas, keep tensors in memory, return one row per index.

import pandas as pd
from torch.utils.data import Dataset
 
 
class CsvDataset(Dataset[tuple[Tensor, Tensor]]):
    def __init__(self, path: str, target: str) -> None:
        df = pd.read_csv(path)
        feats = df.drop(columns=[target])
        self.x = torch.tensor(
            feats.to_numpy(), dtype=torch.float32
        )
        self.y = torch.tensor(
            df[target].to_numpy(), dtype=torch.long
        )
 
    def __len__(self) -> int:
        return len(self.y)
 
    def __getitem__(self, i: int) -> tuple[Tensor, Tensor]:
        return self.x[i], self.y[i]
 
 
toy = {"a": [0.1, 0.9], "b": [1.0, 2.0], "label": [0, 1]}
pd.DataFrame(toy).to_csv("toy.csv", index=False)
ds = CsvDataset("toy.csv", target="label")
loader = DataLoader(ds, batch_size=2, shuffle=True)

Early stopping

When validation loss stops improving: stop training and keep the best weights, not the last ones.

import copy
 
 
class EarlyStopping:
    def __init__(self, patience: int = 5) -> None:
        self.patience, self.bad = patience, 0
        self.best = float("inf")
        self.state: dict[str, Tensor] | None = None
 
    def step(self, loss: float, model: nn.Module) -> bool:
        """Return True when training should stop."""
        if loss < self.best:
            self.best, self.bad = loss, 0
            self.state = copy.deepcopy(model.state_dict())
        else:
            self.bad += 1
        return self.bad >= self.patience
 
 
stopper = EarlyStopping(patience=3)
for epoch in range(100):
    train_epoch(model, train_loader, loss_fn, opt, device)
    val, _ = evaluate(model, val_loader, loss_fn, device)
    if stopper.step(val, model):
        break
assert stopper.state is not None
model.load_state_dict(stopper.state)  # best weights

Checkpoint and resume

When a long run may be interrupted: save everything that changes during training, then continue where it stopped.

from pathlib import Path
 
ckpt = Path("ckpt.pt")
sched = torch.optim.lr_scheduler.StepLR(opt, step_size=10)
 
 
def save(epoch: int) -> None:
    torch.save({
        "epoch": epoch,
        "model": model.state_dict(),
        "opt": opt.state_dict(),
        "sched": sched.state_dict(),
        "rng": torch.get_rng_state(),
    }, ckpt)
 
 
start = 0
if ckpt.exists():
    c = torch.load(ckpt, map_location=device)  # weights_only
    model.load_state_dict(c["model"])
    opt.load_state_dict(c["opt"])
    sched.load_state_dict(c["sched"])
    torch.set_rng_state(c["rng"].cpu())
    start = c["epoch"] + 1
for epoch in range(start, start + 3):
    train_epoch(model, train_loader, loss_fn, opt, device)
    sched.step()
    save(epoch)

Fine-tune a pretrained torchvision model

When you have a small image dataset: reuse ImageNet features, freeze the backbone, train a new head.

from torchvision.models import ResNet18_Weights, resnet18
 
weights = ResNet18_Weights.DEFAULT      # downloads once
net = resnet18(weights=weights)
preprocess = weights.transforms()       # resize + normalize
 
for p in net.parameters():
    p.requires_grad_(False)             # freeze everything
net.fc = nn.Linear(net.fc.in_features, 5)  # new 5-class head
 
net = net.to(device)
head_opt = torch.optim.AdamW(
    (p for p in net.parameters() if p.requires_grad), lr=1e-3
)
imgs = torch.rand(4, 3, 256, 256)       # stand-in batch
logits = net(preprocess(imgs).to(device))  # (4, 5)
# later: unfreeze net.layer4 with a 10x lower LR group

In train() mode frozen BatchNorm layers still update their running statistics. Call net.eval() on the backbone, or leave it in eval mode, if that matters for your data.

Gradient clipping and accumulation

When loss spikes or goes nan (clip), or when the batch you want doesn't fit in memory (accumulate).

accum = 4                            # effective batch x4
model.train()
opt.zero_grad()
for step, (xb, yb) in enumerate(train_loader, start=1):
    xb, yb = xb.to(device), yb.to(device)
    loss = loss_fn(model(xb), yb) / accum  # mean over steps
    loss.backward()                        # grads add up
    if step % accum == 0:
        nn.utils.clip_grad_norm_(      # returns the norm
            model.parameters(), max_norm=1.0
        )
        opt.step()
        opt.zero_grad()

With a GradScaler, call scaler.unscale_(opt) before clip_grad_norm_ so the norm is measured on real gradients.

Accuracy under inference_mode

When you need predictions and accuracy on a held-out set quickly, with no autograd bookkeeping.

@torch.inference_mode()
def predict(
    model: nn.Module, loader: DataLoader[Batch]
) -> tuple[Tensor, Tensor]:
    model.eval()
    preds, targets = [], []
    for xb, yb in loader:
        logits = model(xb.to(device))
        preds.append(logits.argmax(dim=1).cpu())
        targets.append(yb)
    return torch.cat(preds), torch.cat(targets)
 
 
pred, true = predict(model, val_loader)
acc = (pred == true).float().mean().item()
print(f"accuracy {acc:.2%}")

References