PyTorch for LLMs
The first five notes cover PyTorch mechanics that apply to any model. Training and serving language models adds a specific toolkit: fused attention kernels, compilation, bf16, and - once a model outgrows one GPU - distributed data parallelism, parameter sharding and distributed checkpoints. This note maps that toolkit to the current PyTorch APIs.
- Use scaled_dot_product_attention and FlexAttention instead of hand-written attention, and explain why they are faster
- Decide when torch.compile helps, and recognize graph breaks and recompilation
- Choose bf16, fp16 or fp32 for training and explain why bf16 needs no GradScaler
- Launch multi-GPU training with torchrun and choose between DDP and FSDP2 (fully_shard)
- Save and load large-model checkpoints safely - safetensors, torch.load weights_only, and torch.distributed.checkpoint
The LLM Toolkit at a Glance
flowchart LR
subgraph ONE["๐ฅ๏ธ One GPU"]
direction TB
S["โก SDPA / FlexAttention<br/>fused attention"] --> C["๐ง torch.compile<br/>fused kernels"] --> BF["๐ข bf16 autocast"]
end
subgraph MANY["๐ฅ๏ธ๐ฅ๏ธ Many GPUs"]
direction TB
TR["๐ torchrun<br/>one process per GPU"] --> DDP["๐ DDP<br/>replicate model"]
TR --> FS["๐งฉ FSDP2 fully_shard<br/>shard params, grads, optimizer"]
FS --> DC["๐พ distributed.checkpoint<br/>sharded save / load"]
end
ONE --> MANY
style S fill:#d8dfe8,stroke:#b0bac8
style C fill:#d8dfe8,stroke:#b0bac8
style BF fill:#d8dfe8,stroke:#b0bac8
style DDP fill:#dde4dc,stroke:#b0c4b0
style FS fill:#dde4dc,stroke:#b0c4b0
style DC fill:#e8e0d4,stroke:#c8b89a
Fused Attention: SDPA and FlexAttention
Concept
A naive attention implementation materializes the full T ร T score matrix in GPU memory, then reads it back for softmax and again for the value product. torch.nn.functional.scaled_dot_product_attention (SDPA) dispatches to fused kernels - FlashAttention, memory-efficient attention, or cuDNN - that compute attention in tiles and never store the full matrix.
import torch.nn.functional as F
# q: (B, n_head, T, d_head); k, v: (B, n_kv_head, T, d_head)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True)
is_causal=Trueapplies the causal mask inside the kernel - don't build a mask tensor.enable_gqa=Truehandles grouped-query attention without repeating K/V in memory.- To force or debug a backend, use
torch.nn.attention.sdpa_kernel(SDPBackend.FLASH_ATTENTION)(the oldertorch.backends.cuda.sdp_kernelis deprecated).
FlexAttention (torch.nn.attention.flex_attention) covers the attention variants SDPA doesn't: sliding windows, document masking for packed sequences, prefix-LM masks, soft-capping. You write the mask or score modification as a small Python function; torch.compile turns it into a fused kernel.
from torch.nn.attention.flex_attention import flex_attention, create_block_mask
def sliding_window_causal(b, h, q_idx, kv_idx):
return (q_idx >= kv_idx) & (q_idx - kv_idx <= 1024)
block_mask = create_block_mask(sliding_window_causal, B=None, H=None, Q_LEN=T, KV_LEN=T)
flex = torch.compile(flex_attention) # uncompiled flex_attention is a slow reference path
y = flex(q, k, v, block_mask=block_mask)
torch.compile
Concept
torch.compile(model) traces the model's Python into a graph and generates fused kernels (via TorchInductor/Triton on GPUs). For transformer training it typically removes a large share of the small elementwise kernels around the matmuls - normalization, activation, residual adds - and speeds up training noticeably; the gain depends on the model and GPU, so measure it.
What to watch:
- Compile time. The first iterations are slow while kernels compile; this is amortized over a long run.
- Graph breaks. Python the compiler can't trace (data-dependent control flow,
.item()in the forward pass, some third-party calls) splits the graph and loses fusion.TORCH_LOGS="graph_breaks"lists them. - Recompilation. Changing input shapes can trigger recompiles. Keep sequence lengths fixed (or bucketed) during training, or mark dynamic dimensions.
Precision: bf16 by Default
| Format | Exponent / mantissa bits | Range | Needs loss scaling? | Use for |
|---|---|---|---|---|
| fp32 | 8 / 23 | Wide | No | Master weights, optimizer state, reductions |
| fp16 | 5 / 10 | Narrow (max ~65,504) | Yes (GradScaler) | Older GPUs without bf16 |
| bf16 | 8 / 7 | Same as fp32 | No | Default for LLM training on A100/H100-class GPUs |
| fp8 (E4M3 / E5M2) | 4/3 or 5/2 | Narrow | Per-tensor or per-block scaling | Large-scale training and inference on H100+ (see Pretraining at Scale) |
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
logits, loss = model(x, y)
loss.backward() # no GradScaler: bf16 has fp32's exponent range, so gradients don't underflow
optimizer.step()
Autocast keeps precision-sensitive operations (softmax, norms, losses) in fp32. Compute the cross-entropy on fp32 logits - the GPT lab calls logits.float() before the loss for this reason.
Multi-GPU: torchrun, DDP and FSDP2
torchrun
torchrun --nproc_per_node=8 train.py starts one process per GPU and sets RANK, LOCAL_RANK and WORLD_SIZE. Each process calls torch.distributed.init_process_group("nccl") and pins itself to LOCAL_RANK. Multi-node runs add --nnodes and a rendezvous endpoint.
DDP - replicate the model
DistributedDataParallel keeps a full copy of the model, gradients and optimizer state on every GPU and all-reduces gradients during backward(). It is the simplest option and the fastest when the whole training state fits on one GPU.
FSDP2 - shard everything
When it doesn't fit (mixed-precision Adam needs ~16 bytes per parameter - see GPU & Hardware), fully sharded data parallelism splits parameters, gradients and optimizer state across GPUs and gathers each layer's full weights only while that layer runs. PyTorch's current API is fully_shard (FSDP2), which represents each parameter as a sharded DTensor:
from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
mp = MixedPrecisionPolicy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32)
for block in model.blocks: # shard each transformer block as its own unit...
fully_shard(block, mp_policy=mp)
fully_shard(model, mp_policy=mp) # ...then the root (embeddings, head)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4) # created after sharding
| DDP | FSDP2 (fully_shard) | |
|---|---|---|
| Memory per GPU | Full model + grads + optimizer | ~1/N of each, plus one layer's full weights at a time |
| Communication | All-reduce gradients | All-gather params (forward/backward) + reduce-scatter grads |
| When | Model state fits on one GPU | It doesn't, or you want larger batches |
| Composes with | - | Tensor / context parallelism via DeviceMesh (see Pretraining at Scale) |
Checkpoints for Large Models
| Tool | What it does | Use when |
|---|---|---|
safetensors | Stores raw tensors plus a JSON header; loading can't execute code; memory-mapped | Publishing and loading model weights (the Hugging Face default) |
torch.load(..., weights_only=True) | Restricts unpickling to tensors and simple types. It is the default since PyTorch 2.6 | Loading any .pt file you didn't create |
torch.distributed.checkpoint (DCP) | Each rank saves its own shards in parallel; loading can reshard to a different GPU count; async_save overlaps saving with training | Training checkpoints for FSDP / multi-GPU runs |
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
model_sd, optim_sd = get_state_dict(model, optimizer)
dcp.save({"model": model_sd, "optim": optim_sd}, checkpoint_id=f"ckpt/step_{step}")
Why not pickle? A .pt file saved with torch.save is a pickle; unpickling can run arbitrary code. Downloading someone's checkpoint and calling torch.load with weights_only=False is equivalent to running their script.
Profiling
torch.profiler records CPU and CUDA activity per operator; export a trace and open it in Perfetto to see whether the GPU is waiting on the data loader, on communication, or on small kernels.
from torch.profiler import profile, ProfilerActivity, schedule
with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
schedule=schedule(wait=1, warmup=1, active=3)) as prof:
for step in range(5):
train_step()
prof.step()
prof.export_chrome_trace("trace.json")
For end-to-end efficiency, compute model FLOPs utilization (MFU): achieved training FLOPs per second (โ 6 ร parameters ร tokens per second for a dense transformer) divided by the GPU's peak. Well-tuned large runs report roughly 35-50% MFU on H100s; far lower usually means a data, communication or small-kernel bottleneck.
Check Yourself
- Why does bf16 training not need a GradScaler while fp16 does?
- A 13B model's full training state doesn't fit on one 80 GB GPU. What does FSDP2 change compared with DDP?
- You downloaded
model.ptfrom an unfamiliar repository. What is the risk intorch.load('model.pt', weights_only=False), and what should you do instead?
Exercises
Using the GPT From Scratch lab on a GPU, time 200 steps with and without --compile (discard the first 20 steps of each). Report tokens/second and the compile time. Then add a .item() call inside GPT.forward and use TORCH_LOGS="graph_breaks" to see what it does.
You will fine-tune an 8B dense model with full fine-tuning (not LoRA) on 8ร 80 GB GPUs with mixed-precision AdamW. Estimate the model-state memory per GPU under DDP and under FSDP2, and say which is feasible.
Solution
Model state โ 16 bytes ร 8B = 128 GB. DDP puts all of it on every GPU - 128 GB > 80 GB, so it does not fit. FSDP2 shards it: ~128 / 8 = 16 GB per GPU plus one layer's gathered weights and activations, which fits comfortably.
References
- PyTorch docs: scaled_dot_product_attention, FlexAttention, torch.compile, FSDP2
fully_shard, Distributed Checkpoint - Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel (2023)
- Dao et al., FlashAttention (2022); Dao, FlashAttention-2 (2023)
- Liang et al., TorchTitan (2024) - a reference implementation of FSDP2 + tensor/pipeline/context parallelism
- Hugging Face, safetensors
Last reviewed: 2026-09