PyTorch Cheatsheet

Training Loop

Use this PyTorch reference while you build software engineering projects, review code for technical interview prep, or polish examples for a software engineer resume.

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()