One number, several memory roles
During training, a parameter can appear as a weight, a gradient and in two Adam moment states. With all four arrays in float32, that is roughly 16 bytes per parameter. For 15 million parameters, that comes to about 240 million bytes, around 229 MiB. Activations and temporary workspaces are not yet included.
Classic mixed-precision training with Adam also needs about 16 bytes per parameter: 2 for 16-bit weights, 2 for gradients, 4 for FP32 master weights and 8 for the two moments (ZeRO, Rajbhandari et al.). It mainly saves activation memory and compute time, not optimizer state. In every calculation, state explicitly which states have which dtype.
Why activations are expensive
In the backward pass, autograd needs certain intermediate results from the forward pass. These depend on batch size, sequence length, width and depth. A naive attention matrix adds a term that is quadratic in T. An efficient kernel can avoid certain large intermediate matrices, but it can't make all activations disappear.
If memory is tight, first reduce the microbatch, then check the context length and model size. Gradient accumulation can keep the effective batch size, but it doesn't give the same runtime as a larger batch computed all at once.
FP32, FP16 and BF16
FP32 has 32 bits and is a simple, robust starting point. FP16 has a smaller exponent range; large or very small values can become problematic sooner. BF16 has a wider exponent range than FP16, but a coarser mantissa. None of these dtypes is "always accurate enough".
Large runs now also use FP8 (8 bits) for many matrix multiplications, with fine-grained scaling factors, higher-precision accumulators and master weights (DeepSeek-V3); for inference, 4-bit formats are also common. That is engineering work for special hardware, not a switch for your course project.
The reference trainer deliberately starts in float32. This reduces the number of mechanisms you have to learn at the same time and supports CPU, MPS and CUDA with one shared, clear path. Mixed precision is a later extension, not a silent promise of the download.
Using autocast deliberately
On suitable CUDA hardware, training with autocast and BF16 can run certain operations in lower precision. Other operations stay in float32, depending on the rules. One possible extension is:
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
logits = model(x)
loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1))
loss.backward()Check what your hardware supports. For FP16, a gradient scaler is often used; it scales the loss and correctly scales it back before the optimizer step. Clipping then belongs after the unscale. Don't just use the same snippet untested on every backend. PyTorch AMP documentation.
Quantization is a different intervention
A 4-bit weight file reduces the memory of the represented weights. It doesn't automatically make every training array 4-bit, nor all operations four times faster. Scales, grouping metadata, activations, cache and dequantization are part of the real bill.
In quantization, you group weights, for example, choose a scale and approximate floating-point numbers with a limited set of integer values. This introduces errors. How much the output suffers depends on the method, data, model and sensitive layers. "Four bits" alone doesn't describe the overall quality.
Activation checkpointing
With checkpointing, you store fewer intermediate values from the forward pass and recompute them in the backward pass. This trades activation memory for extra compute time. It is not the same as a model checkpoint on disk.
from torch.utils.checkpoint import checkpoint
for block in model.blocks:
x = checkpoint(block, x, use_reentrant=False)For a small model, the overhead can be larger than the benefit. Only add this extension once you have a measured memory bottleneck. Also check random operations, dtypes and compatibility with your backend.
An honest memory measurement
On CUDA, you can measure PyTorch's peak allocation after a warmup. This doesn't exactly match the device's total memory usage: libraries and other programs can claim additional memory. On Apple Silicon, shared memory needs to be looked at differently again.
if device.type == 'cuda':
torch.cuda.reset_peak_memory_stats()
# A real forward/backward/optimizer step
print(torch.cuda.max_memory_allocated() / 2**20, 'MiB')A model that fits during the forward pass can fail at the first Adam step, because the moment states are only created then. So measure a complete training step.
Your optimization log
Compare a float32 baseline, smaller microbatches plus accumulation, and then possibly mixed precision. Record validation loss, throughput and peak usage. A speedup that changes the loss or becomes unstable calls for a decision about quality. Optimization is not just the search for the smallest number in the memory monitor.