A training loop is four lines of Python: forward, loss, backward, optimizer step. Underneath, those lines launch thousands of GPU kernels, allocate and free tens of gigabytes, and pass through several points where the CPU and GPU wait for each other. Understanding what happens inside a single step, in order, is what turns out-of-memory errors, slow steps and mysterious first-step crashes from guesswork into arithmetic.

This article walks through one step on one GPU as a timeline. It follows the batch from host memory to the loss, through the backward pass and the optimizer, and tracks what occupies memory at each moment. The per-bucket byte accounting, how many bytes a parameter costs under different optimizers and precisions, is covered in Memory Math for Training. The focus here is on when each bucket appears and disappears, and what that implies for peak memory and step time.

Advertisement

The step as a stream of kernels

PyTorch executes eagerly, but the GPU does not run in lockstep with Python. Each operation enqueues one or more kernels on a CUDA stream and returns immediately; the GPU executes the queue in order while the host races ahead to enqueue more. In a healthy loop the host is several operations ahead of the device, and the GPU never waits for Python. A step is therefore best thought of as a long sequence of kernels on one stream: copy, embedding lookup, a few hundred GEMMs and element-wise kernels for the forward, a larger number for the backward, then the optimizer's kernels.

This asynchrony has two consequences. Wall-clock timing of Python lines is meaningless without synchronization, because a line returns before its kernels run. And anything that requires a value on the host, such as loss.item(), printing a tensor, .cpu(), or a data-dependent shape from nonzero(), forces the host to wait until the queue drains. One such sync per step is harmless; one per layer can leave the GPU idle for a large fraction of the step.

Allocated memory across one step (illustrative shape, not to scale)GBweights + optimizer state (resident across steps)activations savedfreed layer by layergradients appearoptimizer tempspeakload batchforwardbackwardoptimizerzero_gradHost (Python)enqueues kernelsCUDA streamkernels run in orderSync points.item(), print, .cpu()
Resident weights and optimizer state form the floor. Activations accumulate through the forward pass, peak around the start of backward, and are released layer by layer as gradients are produced. The optimizer adds short-lived temporaries; zero_grad releases gradients.

Phase 0: getting the batch onto the device

The step begins with a host-to-device copy of the input ids, labels and masks. For language models this is small, a few megabytes, but it still matters when it is synchronous. Pinned (page-locked) host memory lets the copy run asynchronously with non_blocking=True, overlapping with the previous step's compute. With pageable memory the driver stages through a bounce buffer and the copy behaves synchronously. Vision and audio pipelines move much more data per step, and a slow loader shows up as idle gaps at the start of each step in the profiler timeline.

Advertisement

Phase 1: the forward pass saves what backward will need

The forward pass computes the loss, but for training its more important side effect is saving tensors. Autograd records, for each operation, the inputs or intermediates needed to compute its gradient: a linear layer saves its input to compute the weight gradient, a GELU saves its input, dropout saves its mask, layer norm saves its mean and reciprocal standard deviation, attention saves its inputs and, with fused attention kernels, small softmax statistics rather than the full score matrix. These saved tensors are the activation memory, and they accumulate layer by layer through the forward pass.

Activation memory scales with batch size times sequence length times hidden size times depth, which is why it, rather than the weights, usually decides whether a configuration fits. A widely used estimate for a transformer layer in 16-bit precision without recomputation is about 34 x s x b x h bytes, where s is sequence length, b micro-batch size and h hidden size, plus an attention-score term that fused attention kernels largely remove. The derivation is in Activation Memory Math.

Under autocast the GEMMs run in BF16 on tensor cores while master weights stay in FP32. Autocast casts the weights to BF16 for each matmul, and those casts are cached for the duration of the forward pass, which adds a copy of the weights in BF16 to memory. It is a small bucket next to activations, but not zero.

Phase 2: backward frees activations and creates gradients

loss.backward() walks the recorded graph in reverse. For each linear layer it launches two GEMMs of roughly the forward's size, one for the gradient with respect to the input, which is passed down to the previous layer, and one for the gradient with respect to the weight, which is accumulated into .grad. That is why the backward pass costs about twice the forward's FLOPs, and a full step about three times. Element-wise ops contribute their own small kernels.

Memory moves in two directions during backward. Each layer's saved activations are released as soon as its backward kernels have consumed them, so activation memory falls steadily from the top of the network downward. Meanwhile gradient buffers appear, one per parameter, the first time each parameter receives a gradient. The peak therefore usually occurs early in backward: nearly all activations are still alive, the first gradients exist, and the last layer's backward needs temporary workspace. In models with a large vocabulary the peak often sits right at the loss, because the logits tensor of shape batch x sequence x vocabulary, and its gradient, can be several gigabytes by themselves.

Activation checkpointing trades compute for this peak: selected layers save only their inputs, and backward recomputes their internals just before use. See activation checkpointing for where to place checkpoints.

# What one step does, written out. Model has layers L1..Ln, loss function f.
x = batch.to(device, non_blocking=True)          # H2D copy (pinned memory)

# forward: compute, and save what backward will need
saved = []
h = x
for layer in layers:
    h, ctx = layer.forward(h)                    # ctx: inputs / intermediates
    saved.append(ctx)                            # activation memory grows here
loss = f(h, targets)

# backward: walk layers in reverse, consuming saved tensors
g = dloss_dh = f.backward()
for layer, ctx in reversed(list(zip(layers, saved))):
    g, dW = layer.backward(g, ctx)               # two GEMMs per linear: dX and dW
    layer.W.grad = dW if layer.W.grad is None else layer.W.grad + dW
    free(ctx)                                    # activation memory shrinks here

# optimizer: element-wise, bandwidth-bound, touches every parameter's state
for W in params:
    m, v = state[W]                              # allocated lazily on first step
    m = b1*m + (1-b1)*W.grad
    v = b2*v + (1-b2)*W.grad**2
    W -= lr * (m/(1-b1**t)) / (sqrt(v/(1-b2**t)) + eps) + lr*wd*W
    state[W] = (m, v)
for W in params:
    W.grad = None                                # zero_grad(set_to_none=True)

Phase 3: gradient clipping and the optimizer update

Gradient clipping computes a global norm over all gradients, which means one reduction per parameter tensor and a final combine; the clip factor is then applied on the device without a host sync in current PyTorch implementations. The optimizer step follows. For AdamW each parameter needs its two moment buffers updated and the weight rewritten, a handful of element-wise operations per element. The arithmetic is trivial; the cost is bandwidth. AdamW with FP32 state reads about 16 bytes per parameter (weight, gradient, two moments) and writes about 12 back, roughly 28 bytes of traffic; for a 1.3B-parameter model that is roughly 36 GB, about ten milliseconds at 3.35 TB/s.

How those operations are launched matters a great deal. A naive per-parameter loop launches several kernels for each of hundreds of tensors. The foreach implementation, PyTorch's default when available, groups tensors into multi-tensor kernels, and fused=True performs the whole update in one kernel pass per group. The grouped variants may allocate intermediate buffers, which is the short-lived bump at the right of the memory curve.

One detail catches many people: Adam's state is allocated lazily, on the first call to step(). The first forward and backward pass fit comfortably, and then the first optimizer step allocates two full-size FP32 buffers per parameter and runs out of memory. Memory estimates that were measured before the first step are missing the largest bucket. See AdamW math for the update rule itself.

Phase 4: zero_grad and the floor between steps

optimizer.zero_grad() sets gradients to None by default in PyTorch 2.x, which returns their memory to the allocator instead of writing zeros into it. Between steps, the floor is therefore weights plus optimizer state. With gradient accumulation, zero_grad is called only every few micro-batches, gradients stay resident across the accumulation window, and each micro-batch adds and releases its own activations on top; see gradient accumulation.

Worked example: a 1.3B-parameter model on an 80 GB GPU

Take a decoder with 24 layers, hidden size 2,048 and a 50,000-token vocabulary: about 1.2 billion parameters in the blocks from 12 x L x h2 plus about 0.1 billion in the embedding, so roughly 1.3 billion. Train with BF16 autocast, FP32 master weights and AdamW. The resident floor is FP32 weights (4 bytes), two FP32 moments (8 bytes), and FP32 gradients during backward (4 bytes), about 16 bytes per parameter or roughly 21 GB.

Activations at sequence length 2,048 and micro-batch 8, using the 34 x s x b x h estimate: 34 x 2,048 x 8 x 2,048 bytes, about 1.14 GB per layer, or roughly 27 GB for 24 layers. Logits add 8 x 2,048 x 50,000 elements, about 0.8 billion values; in FP32 for the loss that is over 3 GB, and its gradient as much again. The estimated peak is around 21 + 27 + 6 plus workspace and allocator slack, somewhere in the 55 to 60 GB range, which fits on an 80 GB card. Doubling the micro-batch to 16 doubles activations and logits and does not fit without checkpointing. These are estimates to plan with; measure the real peak with the code below before committing to a configuration.

The caching allocator: allocated versus reserved

PyTorch does not call cudaMalloc for each tensor. Its caching allocator requests large segments from the driver and carves tensors out of them, keeping freed blocks for reuse. Two numbers follow: allocated, the bytes held by live tensors, and reserved, the bytes the allocator holds from the driver. nvidia-smi shows reserved plus the CUDA context, not what your tensors use.

Because activation sizes vary, with variable sequence lengths for example, the cache can fragment: enough total free memory exists, but no single block is large enough, and an allocation fails with reserved far above allocated. Setting PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True, documented as experimental, lets segments grow instead of multiplying, which reduces this in many workloads. Padding sequences to a few fixed lengths also helps, and it lets you capture the step with CUDA graphs.

import torch

model = build_model().cuda()
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, fused=True)
torch.cuda.memory._record_memory_history(max_entries=100_000)   # allocation trace

for step, (x, y) in enumerate(loader):
    x = x.cuda(non_blocking=True); y = y.cuda(non_blocking=True)
    torch.cuda.reset_peak_memory_stats()
    with torch.autocast("cuda", dtype=torch.bfloat16):
        loss = model(x, y)                       # forward: saves activations
    loss.backward()                              # backward: grads, frees activations
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    opt.step()                                   # step 0 allocates Adam state
    opt.zero_grad(set_to_none=True)              # release gradient memory

    if step % 50 == 0:                           # sync deliberately, not every step
        print(step, loss.item(),
              torch.cuda.max_memory_allocated() / 2**30, "GiB peak",
              torch.cuda.memory_reserved() / 2**30, "GiB reserved")
    if step == 3:
        torch.cuda.memory._dump_snapshot("step_snapshot.pickle")  # open in pytorch.org/memory_viz

The snapshot records every allocation with a stack trace, and the viewer at pytorch.org/memory_viz, which runs locally in the browser, draws the memory curve over time with each block attributed to the code that allocated it. It is the most direct way to see the shape in the diagram above for your own model. LLM memory profiling covers the wider tooling.

Failure modes

  • OOM on the first optimizer step. Adam state is created lazily; budget for it before the first step, or create it early with a dry step.
  • OOM at the loss. Large-vocabulary logits in FP32 plus their gradient dominate the peak; chunk the loss computation or use a fused cross-entropy kernel.
  • Memory growing across steps. Keeping the loss tensor instead of loss.item() in a list retains its graph; appending tensors to logs is the usual leak.
  • Fragmentation OOMs with variable shapes. Reserved far above allocated at the failure point; enable expandable segments or bucket shapes.
  • Hidden syncs. A .item() inside the model, a host-side condition on a tensor, or torch.cuda.synchronize() in a logging helper leaves the GPU idle between kernels; the profiler timeline shows gaps.
  • Idle GPU at step start. Pageable memory and a slow data loader show up as a gap before the first forward kernel; use pinned memory and enough workers.

Operational guidance and trade-offs

Profile one step before tuning anything. Record peak allocated, reserved and step time for three or four steps after warm-up, and take a profiler trace to see the kernel stream. Then decide which bucket to attack. When activations dominate, reduce micro-batch and accumulate, checkpoint selected layers, or shorten sequences. When states dominate, shard them with ZeRO or FSDP, or use a lower-precision optimizer. When the step is slow but memory is fine, look for host syncs, small unfused kernels and GEMMs running below peak. Mixed precision affects both columns at once.

Every memory saving costs something. Checkpointing adds roughly one extra forward pass of compute for the checkpointed layers. Smaller micro-batches lower GEMM efficiency. Sharding adds communication. Offloading states to CPU memory trades GPU memory for PCIe bandwidth. The right choice depends on which resource is scarce, and the timeline view of the step is what tells you.

Key takeaway: One training step is a stream of kernels with a predictable memory shape: weights and optimizer state form a floor, activations pile up through the forward pass, the peak lands near the loss or early in backward, gradients replace activations as backward proceeds, and the optimizer adds brief temporaries before zero_grad releases gradients. Budget for lazy Adam state and large logits, avoid host syncs inside the step, read allocated and reserved separately, and use a memory snapshot to see the curve for your own model before choosing checkpointing, sharding or a smaller micro-batch.