PyTorch Cheatsheet

Training Loop

Use this PyTorch reference while you build software engineering projects, review code, or refresh the syntax you reach for most.

Minimal Training Loop

import torch
import torch.nn as nn

model = MyModel().to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
criterion = nn.CrossEntropyLoss()

for epoch in range(num_epochs):
    model.train()
    for x, y in train_loader:
        x, y = x.to(device), y.to(device)

        optimizer.zero_grad()
        logits = model(x)
        loss = criterion(logits, y)
        loss.backward()
        optimizer.step()

Full Production Loop

import torch
from torch.amp import GradScaler, autocast

model = MyModel().to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
scaler = GradScaler('cuda')               # for mixed-precision training

best_val_loss = float('inf')

for epoch in range(num_epochs):
    # ── Training ──────────────────────────────
    model.train()
    train_loss = 0.0

    for batch_idx, (x, y) in enumerate(train_loader):
        x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)

        optimizer.zero_grad(set_to_none=True)

        with autocast('cuda', dtype=torch.float16):
            logits = model(x)
            loss = criterion(logits, y)

        scaler.scale(loss).backward()
        scaler.unscale_(optimizer)
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        scaler.step(optimizer)
        scaler.update()

        train_loss += loss.item()

    avg_train_loss = train_loss / len(train_loader)

    # ── Validation ────────────────────────────
    model.eval()
    val_loss = 0.0
    correct = 0
    total = 0

    with torch.inference_mode():
        for x, y in val_loader:
            x, y = x.to(device), y.to(device)
            logits = model(x)
            val_loss += criterion(logits, y).item()
            preds = logits.argmax(dim=1)
            correct += (preds == y).sum().item()
            total += y.size(0)

    avg_val_loss = val_loss / len(val_loader)
    accuracy = correct / total

    scheduler.step()

    print(f"Epoch {epoch+1}/{num_epochs}  "
          f"train_loss={avg_train_loss:.4f}  "
          f"val_loss={avg_val_loss:.4f}  "
          f"acc={accuracy:.4f}  "
          f"lr={scheduler.get_last_lr()[0]:.2e}")

    # ── Checkpoint best model ─────────────────
    if avg_val_loss < best_val_loss:
        best_val_loss = avg_val_loss
        torch.save({'epoch': epoch,
                    'model_state_dict': model.state_dict(),
                    'optimizer_state_dict': optimizer.state_dict(),
                    'val_loss': best_val_loss}, 'best.pt')

Mixed-Precision Training (AMP)

from torch.amp import autocast, GradScaler

scaler = GradScaler('cuda')

with autocast('cuda', dtype=torch.float16):   # or bfloat16
    output = model(x)
    loss = criterion(output, y)

scaler.scale(loss).backward()

# Unscale BEFORE gradient clipping
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)

scaler.step(optimizer)   # only steps if no NaN/Inf in grads
scaler.update()          # adjust scale factor

Use bfloat16 on Ampere+ GPUs (A100, H100) — it has the same exponent range as float32, so no loss scaling needed: autocast('cuda', dtype=torch.bfloat16).

Gradient Accumulation

accumulate = 4     # effective batch = batch_size × accumulate

optimizer.zero_grad(set_to_none=True)

for step, (x, y) in enumerate(loader):
    x, y = x.to(device), y.to(device)

    with autocast('cuda'):
        loss = model(x, y) / accumulate   # scale loss

    scaler.scale(loss).backward()

    if (step + 1) % accumulate == 0:
        scaler.unscale_(optimizer)
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        scaler.step(optimizer)
        scaler.update()
        optimizer.zero_grad(set_to_none=True)
        scheduler.step()

Tracking and Logging

# TensorBoard
from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter('runs/experiment_1')

writer.add_scalar('Loss/train', avg_train_loss, epoch)
writer.add_scalar('Loss/val', avg_val_loss, epoch)
writer.add_scalar('Accuracy/val', accuracy, epoch)
writer.add_scalar('LR', scheduler.get_last_lr()[0], epoch)

# Log histograms of weights/grads
for name, param in model.named_parameters():
    writer.add_histogram(name, param, epoch)
    if param.grad is not None:
        writer.add_histogram(f'{name}.grad', param.grad, epoch)

writer.flush()
writer.close()

# Weights & Biases
import wandb
wandb.init(project='my-project', config={'lr': 1e-3, 'epochs': 20})
wandb.log({'train_loss': avg_train_loss, 'val_loss': avg_val_loss, 'epoch': epoch})
wandb.finish()

Early Stopping (Manual)

class EarlyStopping:
    def __init__(self, patience=7, min_delta=0.0):
        self.patience = patience
        self.min_delta = min_delta
        self.counter = 0
        self.best = float('inf')

    def __call__(self, val_loss) -> bool:
        if val_loss < self.best - self.min_delta:
            self.best = val_loss
            self.counter = 0
        else:
            self.counter += 1
        return self.counter >= self.patience   # True → stop

stopper = EarlyStopping(patience=5)
for epoch in range(max_epochs):
    ...
    if stopper(val_loss):
        print(f'Early stopping at epoch {epoch}')
        break

Reproducibility

import torch, random, numpy as np, os

def set_seed(seed: int = 42):
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    np.random.seed(seed)
    random.seed(seed)
    # Deterministic algorithms (may be slower)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False
    os.environ['PYTHONHASHSEED'] = str(seed)

set_seed(42)

torch.backends.cudnn.benchmark = True (the opposite) speeds up training when input sizes are fixed — it finds the optimal convolution algorithm. Use it when reproducibility is not required.

Multi-GPU: DataParallel

# Single-machine, multiple GPUs (easiest but less efficient)
model = nn.DataParallel(model, device_ids=[0, 1, 2, 3])
model = model.to('cuda:0')

# DataParallel splits batches along dim 0, runs forward on each GPU,
# gathers on device_ids[0].
# Access original module: model.module.state_dict()

Multi-GPU: DistributedDataParallel

# DDP — preferred for all multi-GPU training
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

def train(rank, world_size):
    dist.init_process_group('nccl', rank=rank, world_size=world_size)
    torch.cuda.set_device(rank)

    model = MyModel().to(rank)
    model = DDP(model, device_ids=[rank])

    # ... training loop as normal ...

    dist.destroy_process_group()

# Launch with torchrun
# torchrun --nproc_per_node=4 train.py

FSDP (Fully Sharded Data Parallel)

# FSDP2 (recommended, PyTorch >= 2.6): fully_shard — composable, per-module
from torch.distributed.fsdp import fully_shard

for block in model.layers:        # shard each transformer block
    fully_shard(block)
fully_shard(model)                # shard the root module last
# model stays your own nn.Module — train as normal (DTensor parameters)
# FSDP1 (legacy wrapper API)
import functools
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy

wrap_policy = functools.partial(size_based_auto_wrap_policy,
                                min_num_params=1_000_000)
model = FSDP(model, auto_wrap_policy=wrap_policy, device_id=rank)

Profiling the Loop

with torch.profiler.profile(
    activities=[
        torch.profiler.ProfilerActivity.CPU,
        torch.profiler.ProfilerActivity.CUDA,
    ],
    schedule=torch.profiler.schedule(wait=1, warmup=1, active=3),
    on_trace_ready=torch.profiler.tensorboard_trace_handler('./log'),
    record_shapes=True,
    with_stack=True,
) as prof:
    for step, (x, y) in enumerate(loader):
        train_step(x, y)
        prof.step()
        if step >= 5:
            break

Metrics Patterns

# Accuracy (multi-class)
preds = logits.argmax(dim=1)
acc = (preds == targets).float().mean()

# Top-k accuracy
_, topk = logits.topk(k=5, dim=1)
correct_topk = topk.eq(targets.view(-1, 1).expand_as(topk))
top5_acc = correct_topk.float().sum(1).mean()

# Using torchmetrics (recommended for correct distributed reduction)
from torchmetrics import Accuracy, F1Score
acc_metric = Accuracy(task='multiclass', num_classes=10).to(device)
f1_metric = F1Score(task='multiclass', num_classes=10, average='macro').to(device)

for x, y in loader:
    preds = model(x)
    acc_metric.update(preds, y)
    f1_metric.update(preds, y)

print(acc_metric.compute())   # aggregated over all batches
acc_metric.reset()