Part of Language AI Handbook
Explains how mixed precision training uses FP16 and BF16 floating point formats to speed up LLM training and cut memory usage without sacrificing accuracy.
Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.
Article links
Make inline references clickable
Mixed Precision Training
Training large language models demands enormous computational resources. A single forward and backward pass through GPT-3 moves billions of numbers through GPU registers, each consuming memory and time. One of the most practical ways to reduce both is to change which number format those values are stored in. Mixed precision training does exactly this: it keeps some tensors in lower-precision floating point to save memory and speed up computation, while preserving enough numerical accuracy for the model to learn effectively.
The gain is substantial. Modern GPUs execute FP16 matrix multiplications roughly two to four times faster than FP32, and each tensor takes half the memory. For a training run that would otherwise require 80 GB of GPU memory in full precision, mixed precision may bring it down to 40 GB or less. That difference can mean fitting a larger batch, training a bigger model, or using fewer GPUs entirely.
This chapter builds the conceptual foundation for understanding mixed precision training. We start with how floating point numbers work and why precision matters, then walk through the FP16 format and its limitations, introduce BF16 as a more numerically stable alternative, explain the loss scaling technique that makes FP16 training stable, and show how to implement mixed precision training in PyTorch using the torch.amp module. By the end, you will understand how to toggle a setting in your training loop, why the whole system works, and when it might fail.
This chapter sits within the broader Part XXXI: Training Infrastructure. The chapters on GPU Architecture and Memory Management earlier in this part explained why memory is scarce and how the GPU operates. Mixed precision training is one of the most effective answers to that scarcity. Later chapters on Communication Optimization and Checkpointing and Recovery will address other bottlenecks that emerge as training scales further.
Floating Point Formats
Before you can understand mixed precision, you need a clear picture of how computers represent real numbers. The standard format used in most numerical computing is IEEE 754 floating point, a specification that governs how real numbers are approximated using a fixed number of bits. Understanding this representation is not just academic. The failure modes of mixed precision training, from gradient underflow to overflow-induced NaN values, all trace directly back to properties of this format.
The IEEE 754 Standard
IEEE 754 was standardized in 1985, emerging from a decade-long effort to bring consistency to floating point arithmetic across different computer architectures. Before IEEE 754, each hardware vendor used its own representation, making numerical code unreliable and non-portable. The standard specified the bit layout of numbers along with the behavior of arithmetic operations, rounding rules, and special values like infinity and NaN. For over thirty years, virtually every programming language and every CPU conformed to this standard, making FP32 (single precision) the de facto format for scientific computing.
A floating point number is stored as three fields: a sign bit, an exponent, and a fraction (also called the mantissa or significand). The value represented is:
where:
- : 0 for positive, 1 for negative
- : raw exponent bits interpreted as an unsigned integer, minus a bias that centers the range
- : a constant that shifts the exponent so both very small and very large numbers can be represented
- : the fractional part of the significand; the leading 1 is implied (for normalized numbers)
The bias plays a subtle but role. Without it, the exponent field would represent only non-negative powers of 2, limiting representation to large numbers. By subtracting the bias from the stored exponent, you can represent numbers much smaller than 1. For FP32, the bias is 127, so a stored exponent of 0 represents , and a stored exponent of 254 represents . The stored exponent 255 is reserved for special values (Inf and NaN), and 0 is reserved for subnormals and zero.
The leading-1 convention deserves explicit mention. In the formula above, the fraction contributes values from 0.0 to just under 1.0, and the term means the significand always lies in the interval . The leading 1 is implicit: it is not stored in the bits, but assumed. This is called a normalized number. By always storing values in this normalized form, you get one extra bit of precision for free. The special case of a stored exponent of 0 represents subnormal numbers, where the leading 1 is not assumed, allowing very small numbers close to zero to be represented with gradually decreasing precision.
The three standard formats used in deep learning are:
| Format | Total bits | Sign | Exponent | Fraction |
|---|---|---|---|---|
| FP32 (float) | 32 | 1 | 8 | 23 |
| FP16 (half) | 16 | 1 | 5 | 10 |
| BF16 (bfloat) | 16 | 1 | 8 | 7 |
Each format makes a different tradeoff between bit budget and numerical properties. FP32 is the general-purpose workhorse with high precision and wide range. FP16 compresses to 16 bits by narrowing both the exponent and fraction. BF16 also uses 16 bits but makes a different tradeoff: it preserves the full 8-bit exponent from FP32, sacrificing even more fraction bits than FP16 to achieve this.
Numerical Range and Precision
The exponent bits control the range of representable values, meaning the span from the smallest to the largest representable number. The fraction bits control precision, meaning how finely you can distinguish two numbers that are close together.
For FP32 with 8 exponent bits and a bias of 127, the maximum exponent is 127, so the largest representable value is approximately . For FP16 with only 5 exponent bits and a bias of 15, the maximum exponent is 15, giving a maximum value of . That is a dramatically narrower range.
The smallest normal FP16 value (the smallest number that can be represented with full precision) is approximately . Values smaller than this are handled by subnormal numbers, which sacrifice precision to extend the range further downward. The smallest subnormal FP16 value is about . Subnormals are important because gradient values during deep learning training often live in this range, and their representation becomes progressively less accurate as values approach zero.
To understand subnormals concretely: when the exponent field contains all zeros, the number is subnormal. The formula changes to , meaning the leading 1 is dropped and replaced with a leading 0. This allows the representation of values smaller than the smallest normal number, but at reduced precision. A subnormal FP16 value near might have only 3 or 4 bits of true precision rather than the full 10, since the leading bits of the fraction are zeros that carry no information.
Precision is governed by the fraction bits. FP32 provides roughly 7 decimal digits of precision. FP16 provides roughly 3 decimal digits. This means that if you have a value like 0.001234 in FP32 and store it in FP16, it may round to 0.001234 or 0.001235 or some other nearby value, with the rounding error on the order of relative to the magnitude.
This rounding error is called quantization error, and it is the root cause of training instability when using FP16 naively. The gradient values that flow back through a deep network can be extremely small, sometimes well below , and they can overflow FP16's narrow range when they are large. Both cases cause training to diverge. The tension between range and precision is precisely what the different 16-bit formats resolve in opposite directions.
A Worked Example: What FP16 Can and Cannot Hold
Let us trace through a concrete scenario to make these abstractions tangible. Suppose a weight in the embedding layer of a transformer has a value of in FP32. When stored in FP16, this rounds to , an error of about , which is well within FP16's approximately 3 decimal digits of precision near values around 1.
Now consider the gradient for that weight at some training step: . In FP32, this is represented with full precision. In FP16, the nearest representable value is approximately , still acceptable. But suppose the learning rate is and the gradient update is . The weight value has a unit in the last place (ULP) in FP16 of roughly . A gradient update of is about 16,000 times smaller than the ULP. In FP16, this update vanishes entirely: the weight does not change. In FP32, the ULP near 0.73 is about , so the update of is small but representable. After enough steps, those tiny updates compound into meaningful weight changes.
This is not a contrived edge case. During the later stages of training, when a model is near convergence, many gradient updates are precisely of this small magnitude. The model has learned the large-scale structure and is fine-tuning minor adjustments. In FP32, these fine adjustments accumulate correctly over thousands of steps. In FP16, they disappear. The model appears to have converged, but it has stalled at a solution that is slightly worse than the FP32 equivalent.

Why Precision Matters During Training
Consider what happens during a gradient update. The optimizer computes a gradient value, multiplies it by the learning rate (often or smaller), and adds the result to a weight parameter. If the weight is around 0.5 and the gradient update is , then in FP32 the operation is:
This small difference accumulates over millions of steps and represents a real learning signal. In FP16, the value 0.5 has a unit in the last place (ULP) of approximately . A gradient update of is far smaller than the ULP, so it gets rounded away completely. The weight does not update. This phenomenon is called vanishing updates, and it is one of the main failure modes of FP16 training.
The vanishing update problem is especially insidious because it is silent. When gradients overflow to NaN, training collapses visibly and immediately. When updates vanish, training simply makes no progress in certain parameters. If you are not monitoring gradient magnitudes closely, you may train for thousands of steps without realizing that a large fraction of your model's weights have stopped updating entirely.
Overflow is the other failure mode. Gradients can grow large, especially early in training or during instability. If a gradient value exceeds 65504 (the FP16 maximum), it becomes infinity (Inf) or not-a-number (NaN), and the entire training run collapses. Overflow propagates: a single NaN in one gradient contaminates the entire optimizer step, then the updated weights, then future forward passes. A training run that hits overflow is essentially unrecoverable without reverting to a checkpoint.
The ULP is the gap between a floating point number and the next representable value. For FP16, the ULP near 1.0 is approximately , meaning any difference smaller than that cannot be represented. For FP32, the ULP near 1.0 is approximately , more than three orders of magnitude smaller.
The following visualization shows how loss scaling shifts the gradient distribution relative to FP16's representable range. Without scaling, many values fall below the FP16 normal threshold and lose precision; the smaller tail below the subnormal floor becomes zero. A scale factor of shifts the distribution by decades, placing almost all values in the normal range while exposing a small upper tail to overflow risk.


FP16 Training and Its Challenges
The term "FP16 training" is slightly misleading. Truly training entirely in FP16 does not work reliably for most models. What practitioners mean is a mixed precision approach: use FP16 where it is safe and beneficial, and FP32 where it is necessary. The art is in knowing where to draw that line.
The foundational observation that made mixed precision training practical was established empirically by Micikevicius et al. at NVIDIA in their 2018 paper "Mixed Precision Training." They tested a broad range of model types, including convolutional image classifiers, RNNs for speech recognition, and language models, and found a consistent pattern: the operations dominating training time (large matrix multiplications) are less sensitive to precision than the operations affecting training stability (gradient accumulation, optimizer state updates). This asymmetry allows you to use the cheaper format where speed matters and the expensive format where accuracy matters.
The NVIDIA team also identified the specific mechanisms by which naive FP16 training fails and proposed targeted solutions for each. Their framework is essentially the one still in use today in PyTorch, JAX, and every other major deep learning library. Understanding the problem they were solving clarifies why the solution looks the way it does.
What Runs in FP16
The primary beneficiaries of FP16 are the large matrix multiplications that dominate compute time in transformer training: the attention score computations, the feed-forward layer projections, and the embedding lookups. These operations are memory-bandwidth-bound and compute-bound, and using FP16 halves the data movement and doubles the throughput on hardware with dedicated FP16 units.
When an operation runs in FP16, both the inputs and the outputs are in FP16. The CUDA kernels for these operations internally accumulate in FP32 to avoid precision loss in the accumulation itself, then convert the result back to FP16 before writing to memory. This is a hardware-level behavior for tensor cores, and you get it automatically. The intermediate FP32 accumulation prevents summing thousands of FP16 values from accumulating so much rounding error that the result becomes meaningless. The hardware designers anticipated this and built the accumulation into the tensor core itself.
Modern NVIDIA tensor cores operate in a specific way: they take FP16 inputs, compute the matrix multiply-accumulate in FP32, and write an FP16 result. This is called mixed precision multiplication at the hardware level, and it is the key reason why FP16 matrix multiplications are accurate despite the low-precision inputs. The precision loss occurs only when writing the result back to FP16, not during the arithmetic itself.
What Stays in FP32
Certain operations are sensitive to precision and should remain in FP32:
- The master copy of model weights: Stored in FP32, then cast to FP16 before each forward pass.
- Gradient accumulation: Gradients are accumulated in FP32 to avoid the vanishing update problem.
- The optimizer state: Momentum terms and Adam's second-moment estimates should stay in FP32.
- Batch normalization and layer normalization statistics: Mean and variance computations that reduce over large sets of values benefit from FP32 accumulation.
- Loss computation: The scalar loss itself, especially with cross-entropy over large vocabularies, can involve large intermediate sums.
This "keep a master copy in FP32" pattern is the defining feature of the NVIDIA mixed precision training recipe from their 2018 paper "Mixed Precision Training" by Micikevicius et al. The key insight in that paper was empirical: they tested a wide range of models (convolutional networks, RNNs, language models) and found that the FP32 master weight copy, combined with loss scaling, was sufficient to match FP32 training accuracy in virtually all cases.
The optimizer state represents a large chunk of the memory saved by not needing FP16 for these values. For AdamW, the optimizer state includes the first moment (momentum estimate) and second moment (variance estimate) for every parameter, each stored as FP32. This is already 8N bytes for N parameters. Storing them in FP16 would halve this to 4N bytes, but the precision loss would cause optimizer updates to be inaccurate, ultimately harming model quality. In practice, the optimizer state is kept in FP32 in essentially all serious mixed precision training setups.
To understand intuitively why the optimizer state must stay in FP32: Adam's second moment estimate tracks the running average of squared gradients. For a parameter that receives occasional large gradients and frequent small ones, might be a small positive number like . The Adam update divides the gradient by , which is sensitive to the exact value of . If rounds from to due to FP16 precision, the effective step size for that parameter changes by a factor of . Over millions of steps, this corrupts the adaptive learning rate signal that Adam depends on.
The Overflow Problem
Even with master weights in FP32, a persistent problem remains: the FP16 gradients can overflow. Because the FP16 maximum is only 65504, gradient values that legitimately arise during training can become Inf or NaN when cast to FP16. This causes the entire training step to be corrupted.
The key point is that gradient overflow is not a sign of a broken model or bad hyperparameters. Gradient values that are perfectly valid in FP32, and that would produce a correct optimizer step, can simply exceed FP16's representable range. The gradient is not intrinsically wrong; FP16 just cannot hold it. This distinction matters because it explains why scaling the gradient, rather than clipping or discarding it, is the right solution. The gradient carries accurate information about the loss surface; we just need to shift its magnitude into a range where FP16 can represent it.
Loss scaling directly solves this problem.
Loss Scaling
Loss scaling shifts gradient values into a range where FP16 can represent them accurately, then undoes the shift before applying them to the weights.
The Core Idea
Gradient values during training tend to be small, often clustering in the range . FP16 can represent values in this range, but values near the bottom of the range lose most of their significant bits. More critically, underflow to zero is common for the smallest gradients.
The trick is to multiply the loss by a large scalar (the scale factor) before the backward pass. By the chain rule, every gradient throughout the network is multiplied by . This shift moves all gradient values upward in magnitude, reducing underflow and making them comfortably representable in FP16. Think of it as temporarily stretching the gradient magnitude range so that the values that would have been too small for FP16 are now safely within its representable range.
The chain rule argument is worth spelling out explicitly. If the loss is and the scale factor is , the scaled loss is . The gradient of with respect to any parameter is:
So the backward pass on the scaled loss automatically produces gradients that are each multiplied by , with no additional code required. The chain rule propagates the scaling factor through the entire computation graph.
Concretely, if the unscaled gradient is and the scale factor is , then the scaled gradient stored in FP16 is . After the backward pass, before the optimizer update, we divide by :
where:
- : the true gradient that would have been computed in FP32
- : the scaled gradient stored in FP16
- : the loss scale factor, typically a power of 2 to avoid rounding
The division by is done in FP32 (on the master weight gradients), so it does not suffer from precision loss. The whole operation is mathematically a no-op: you multiply by before the backward pass and divide by after. The only effect is that the intermediate representation in FP16 uses a better portion of the available bit range.
Why use powers of 2 for ? Because powers of 2 can be divided exactly in floating point arithmetic. Dividing by is equivalent to subtracting 15 from the FP32 exponent, which introduces zero rounding error. Any other value of would introduce a small rounding error in the unscaling step, partially defeating the purpose.
A Numerical Walk-Through
Let us trace through the full loss scaling procedure with concrete numbers to see precisely what changes and what stays the same.
Suppose a gradient value computed during the backward pass, at some layer, is (approximately ). In FP32, this is representable with full precision. In FP16, the nearest representable normal value is approximately , a small rounding error of about 0.4%. So far, FP16 handles this reasonably well.
Now suppose the gradient update for a weight with value 1.0 is . The ULP for the value 1.0 in FP16 is approximately . A gradient update of is roughly 4 million times smaller than the ULP. In FP16, the weight simply does not change.
With loss scaling by :
The scaled gradient is . This is in the middle of FP16's normal range and carries full precision. The FP16 representation is , negligible rounding.
After the backward pass, the unscaling step divides in FP32: . This is the original gradient, recovered with essentially no error. The optimizer applies this gradient to the FP32 master weight, and the update proceeds correctly.
The effective change is that the gradient was temporarily multiplied by 32768 during the FP16 backward pass to move it from the subnormal/dangerous region of FP16's range into a well-represented middle region, then divided back to restore the original value in FP32 before the optimizer touches it.
Choosing the Scale Factor
The scale factor must be large enough to prevent underflow but small enough to prevent overflow. If is too large, the scaled gradients themselves overflow FP16 (exceeding 65504), producing Inf values. If is too small, gradients underflow to zero.
Static loss scaling uses a fixed , typically around or . This works reasonably well for many models but requires tuning. A scale that works for one architecture may fail for another.
Dynamic loss scaling addresses this by adjusting automatically during training. The algorithm is:
- Start with a large scale factor (e.g., ).
- After each backward pass, check whether any gradient is Inf or NaN.
- If overflow is detected: skip the optimizer step (the corrupted gradients would harm the model), reduce by a constant factor (e.g., halve it), and continue.
- If no overflow is detected: periodically increase (e.g., multiply by 2 every 2000 steps, or whenever overflow has not been seen for a stretch).
Dynamic loss scaling is the standard approach in modern frameworks. PyTorch's torch.amp.GradScaler implements exactly this algorithm. The scale starts high and self-adjusts so that it stays as large as possible without causing overflow, maximizing the precision of gradient representation throughout training.
The adaptive behavior of dynamic loss scaling is important in practice. Early in training, gradients can be large and scale should be lower. As training stabilizes, gradients become smaller and scale can safely increase. A static scale chosen for late-stage training would overflow early; a scale chosen for early training would underflow late.
The skipped steps when overflow is detected are not wasted. The model state is preserved; only the corrupted update is discarded. The scale factor is immediately reduced, which prevents the next step from overflowing as well. In practice, overflow events are rare after the initial warmup period, and the skipped steps represent a tiny fraction of total training compute.
The growth interval (the number of consecutive clean steps before the scale increases) is a subtle but important tuning parameter. The default of 2000 in PyTorch reflects a conservative balance: you want the scale to grow quickly enough to keep gradients well-represented, but not so aggressively that it oscillates between overflow and recovery. If your training shows frequent overflow events throughout, consider reducing the growth factor or increasing the growth interval. If overflow events disappear quickly but you suspect underflow later in training, a shorter growth interval allows the scale to rise faster and capture fine-grained gradient information.
Loss Scaling with Gradient Clipping
Gradient clipping (see the chapter on Gradient Clipping in Part X) caps gradient norms before they are applied to weights, preventing catastrophically large updates. When using loss scaling, gradient clipping must happen on the unscaled gradients. The GradScaler.unscale_() method in PyTorch handles this by dividing the scaled gradients by in-place before you call the gradient clip function.
The correct order of operations is:
- Scaled backward pass:
scaled_loss = scaler.scale(loss); scaled_loss.backward() - Unscale gradients:
scaler.unscale_(optimizer) - Clip gradients:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) - Optimizer step:
scaler.step(optimizer)(internally checks for Inf/NaN and skips if found) - Update scale:
scaler.update()
Applying gradient clipping before unscaling would clip the scaled gradients at the wrong threshold, effectively tightening the clip by a factor of . If your loss scale is 65536 and your intended clip norm is 1.0, clipping before unscaling would apply a clip of , cutting off nearly all gradient signal.
BF16: A Better Design for Deep Learning
BF16 (Brain Float 16) is a 16-bit floating point format developed at Google Brain specifically for deep learning. It uses 1 sign bit, 8 exponent bits, and 7 fraction bits. The key insight: BF16 has the same number of exponent bits as FP32.
Understanding why Google Brain chose to allocate bits this way requires stepping back and asking which property of floating point is most critical for stability during training. The answer, as the FP16 challenges above illustrated, is range, not precision. Overflow to Inf or NaN corrupts training irreversibly. Reduced precision in the gradient, by contrast, simply adds noise to each update, and neural network optimization is already highly tolerant of noisy gradients. The entire rationale for stochastic gradient descent is that you can make progress using noisy, batched estimates of the true gradient. The same tolerance extends to numerical noise from reduced floating point precision. By preserving the FP32 exponent, BF16 eliminates the most dangerous failure mode while accepting a modest precision reduction.
The origins of BF16 trace to Google's TPU development. When designing the second generation of TPUs around 2018, the Google Brain team needed a format that could run efficiently on custom hardware while supporting the large models they were training. FP16's overflow problems made it problematic at the scale of their workloads. Rather than building complex loss scaling infrastructure into the TPU stack, they designed a new format that simply did not have the overflow problem. BF16 was the result, and it was later adopted by NVIDIA for their Ampere GPUs (A100, released in 2020) and beyond.
Why the Exponent Allocation Matters
The 8-bit exponent gives BF16 the same numerical range as FP32: values from approximately to . Overflow is essentially impossible in BF16 during normal training, because any gradient that would overflow FP32 would overflow BF16 as well, and training in FP32 would also be broken. The fact that FP32 training works for a given model architecture and hyperparameters is a guarantee that BF16 will not encounter overflow problems.
This eliminates the need for loss scaling entirely. You can train in BF16 without tracking overflow, without skipping corrupted steps, and without the overhead of checking gradients for Inf/NaN values. The training loop becomes simpler, the failure modes are fewer, and the behavior is more predictable.
The cost is lower precision in the fraction: 7 bits versus 10 bits in FP16. BF16 provides only about 2-3 decimal digits of precision, compared to FP16's roughly 3. In practice, this reduced precision is rarely a problem for gradient values, because gradients carry directional information more than they carry precise magnitudes. What matters is that the optimizer step moves weights in roughly the right direction, not that it does so with 10 significant bits of accuracy. The noise introduced by 7-bit mantissa precision is similar in character to the noise introduced by using a random mini-batch instead of the full dataset: the optimizer is stochastic anyway, and another source of small noise does not materially change the optimization trajectory.
There is a deeper reason for this tolerance. The gradient at any given step is a stochastic estimate of the true gradient because it is computed on a finite batch of examples. This stochasticity is not a bug; it is essential for escaping sharp minima and finding flat, generalizable solutions. Adding a small amount of quantization noise on top of this existing stochasticity is unlikely to change training dynamics significantly. The quantization noise from BF16's 7-bit mantissa is typically smaller than the variance from batch sampling, so it is "absorbed" into the existing noise floor.
BF16 vs FP16: The Practical Tradeoff
The choice between BF16 and FP16 depends primarily on the hardware and the training stability requirements.
BF16 advantages:
- No loss scaling required, simplifying the training loop
- No overflow risk (same range as FP32)
- More numerically stable with large models and high learning rates
- First-class support on TPUs (BF16 was designed for Google's hardware)
- Simpler debugging: NaN values in BF16 training always indicate an underlying numerical problem rather than an FP16 range issue
FP16 advantages:
- Higher fraction precision (10 bits vs 7 bits)
- Broader hardware support: available on older GPUs (Pascal, Volta) that lack BF16
- Slightly better accuracy for models that are precision-sensitive
Hardware availability:
- BF16 requires Ampere GPUs (A100, A30, RTX 30 series) or newer on NVIDIA hardware
- TPUs support BF16 natively and have done so since the early TPU generations
- Intel's newer Xeon processors and Habana Gaudi accelerators also support BF16
- H100 and H200 GPUs support both BF16 and FP8 (an even more compressed format for inference)
In the transformer training community, BF16 has become the default for large model training on modern hardware. The simplicity of eliminating loss scaling, combined with consistent numerical stability, outweighs the small precision reduction. Major models including PaLM, Gemini, LLaMA, and most recent large-scale training runs use BF16 as their primary training precision.
Converting between BF16 and FP32 is essentially free. BF16 is the top 16 bits of an FP32 value (with the bottom 16 bits zero-padded). Truncating the mantissa from 23 to 7 bits is a bit shift operation. This makes BF16-FP32 conversion faster than FP16-FP32 conversion, which requires a full format reinterpretation.
The Evolution Toward FP8
It is worth briefly noting the direction the field is heading beyond BF16. NVIDIA's H100 GPU introduced support for FP8 (8-bit floating point), which comes in two variants: E4M3 (4 exponent bits, 3 mantissa bits) and E5M2 (5 exponent bits, 2 mantissa bits). FP8 offers another 2x reduction in memory and bandwidth compared to BF16, enabling even larger batch sizes and faster training. FP8 training requires even more careful management of numerical range than FP16, with per-tensor scaling factors rather than a single global loss scale. It is primarily used in production training pipelines at large organizations and is not yet as broadly accessible as BF16. The conceptual structure, however, is the same: trade precision for range where possible, and keep critical accumulators in higher precision.
Mixed Precision Architecture in Practice
The full mixed precision training setup combines the observations above into a coherent system. Let us walk through how it works end-to-end, then examine the memory implications carefully.
The Two-Copy Weight Scheme
The model maintains two copies of its parameters:
- FP16 (or BF16) working copies: Used for the forward pass and backward pass. The matrix multiplications and attention operations run on these copies using fast low-precision arithmetic.
- FP32 master copies: The authoritative version of the weights. Optimizer states (momentum, variance) are also kept in FP32 here.
At each training step, the working copies are created by casting the FP32 masters to the low-precision format. After gradients are computed and scaled back to FP32, they are applied to the FP32 masters. The working copies are discarded and recomputed at the next step.
This two-copy scheme has a memory cost: the FP32 masters exist alongside the FP16 working copies. For a model with parameters, you store bytes for the FP16 working copies and bytes for the FP32 masters, totaling bytes rather than the bytes of pure FP32 training. The gain comes from the activations during the forward and backward passes, which are the dominant memory consumer for large models. Activations in FP16 use half the memory of FP32 activations, and this saving more than compensates for the extra master copy.
The reason activations dominate memory rather than parameters is a consequence of the size of modern models and the batch sizes used during training. A transformer layer with and a sequence of length 2048 produces activation tensors with millions of elements per layer, multiplied by the batch size. The parameter count for that same layer is fixed at roughly million parameters, while the activations scale with both sequence length and batch size. At large batches and long sequences, activations dwarf parameters in memory, making FP16 activations the dominant source of memory savings.
The Flow of a Single Training Step
Tracing a single training step in mixed precision clarifies which operations happen in which precision and where the memory savings and costs accumulate.
First, at the start of the step, the FP32 master weights are cast to FP16 working copies. This casting is fast (a bit-level operation) and the FP16 copies live in GPU memory alongside the FP32 masters. The extra memory for the FP16 copies is bytes.
Second, the forward pass runs with the FP16 working copies. Activations, which include all the intermediate tensors produced by each layer's computation, are stored in FP16. For a batch size of and sequence length , the activations at each of the transformer layers occupy roughly memory in FP16. These activations must be retained in memory until the corresponding backward pass computes their gradients, which for a -layer model means the first layer's activations must survive until the last backward step. This is why activation memory scales with model depth and batch size, and why the FP16 savings here are so valuable.
Third, the backward pass computes gradients in FP16 (or with FP32 accumulation for tensor cores, as described earlier). If using loss scaling, the gradients are in scaled FP16. The GradScaler.unscale_() call converts these to FP32 by dividing by the scale factor, in-place, overwriting the FP16 gradient storage with FP32 values. This temporarily increases gradient memory by 2x for the brief duration of the unscaling, but the FP16 gradient memory is freed immediately after.
Fourth, the FP32 gradients are applied to the FP32 master weights by the optimizer. Adam updates its FP32 momentum and variance estimates and computes the updated master weights. The FP16 working copies are now stale and can be freed (they will be recast from the updated masters at the next step).
The net effect is that the peak memory consumption is lower than pure FP32 because the large activation tensors are stored in FP16, and this saving is much larger than the overhead of the FP32 master weight copy for typical transformer workloads.
Autocast and Automatic Mixed Precision
PyTorch's automatic mixed precision (AMP) framework handles the casting decisions automatically. The torch.autocast context manager wraps code regions where operations should run in lower precision. Within this context, PyTorch consults an internal dispatch table that maps each operation to its appropriate precision:
- Matrix multiplications (
torch.mm,torch.bmm,torch.linear): FP16/BF16 - Attention operations: FP16/BF16
- Convolutions: FP16/BF16
- Loss functions: FP32
- Reduction operations (sum, mean): FP32
- Batch norm, layer norm: FP32 (for statistics, though parameters can be in FP16)
The autocast mechanism inserts casts at operation boundaries automatically. You do not need to call .half() or .to(torch.float16) anywhere in your model code. Operations that benefit from low precision run in low precision; operations that need FP32 stay in FP32.
The dispatch table is based on empirical evidence from the PyTorch and NVIDIA teams about which operations are safe in lower precision. Operations that perform large reductions (softmax, layernorm) are kept in FP32 because their outputs are sensitive to accumulation error. Operations that perform independent, parallel computations (linear layers) are safe in FP16 because the per-element error does not compound within a single operation.
One practical consequence of the autocast dispatch table is that operations can change precision across calls even within the same model. A linear layer will run in FP16 under autocast, but a subsequent layer normalization will cast its inputs to FP32. This cast has a small overhead, typically negligible compared to the compute savings, but it means that the model's internal data flow involves more dtype transitions than a pure FP32 model. If you profile a mixed precision model with PyTorch's profiler, you will see autocast_cpu_cast or _cast operations interspersed between the main computations. These are expected and not a sign of a problem.
Autocast with BF16
Using BF16 instead of FP16 in PyTorch AMP is a single argument change: dtype=torch.bfloat16 instead of the default FP16. When using BF16, you should omit the GradScaler entirely, since loss scaling is unnecessary. The training loop becomes visually simpler and conceptually cleaner. Many practitioners who have switched from FP16 to BF16 report that the transition required only a handful of lines of code and improved training stability, particularly for large models.
The BF16 path through autocast is slightly different from FP16 at the hardware level: the CUDA kernels for BF16 tensor core operations use different instruction variants than FP16. The throughput is similar on Ampere and newer hardware, so the practical speedup is comparable. On some operations, BF16 may be slightly faster because the simpler exponent alignment (BF16 shares the FP32 exponent format) reduces conversion overhead.
Code Implementation
Let us implement mixed precision training in PyTorch, showing both FP16 with dynamic loss scaling and BF16 without scaling. We will train a small transformer-like model on a synthetic task to measure memory usage, throughput, and training stability.
Setup and Imports
import torch
# Check hardware support
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
has_fp16 = torch.cuda.is_available()
has_bf16 = torch.cuda.is_available() and torch.cuda.is_bf16_supported()Device: cpu FP16 support: False BF16 support: False
A Simple Transformer Block
We define a small transformer block that exercises the key components: multi-head attention and a feed-forward network. This model is deliberately small so it can run on a CPU if no GPU is available, but the mixed precision patterns are identical for large models.
import torch.nn as nn
class TransformerBlock(nn.Module):
def __init__(self, d_model=256, n_heads=8, d_ff=1024, dropout=0.1):
super().__init__()
self.attention = nn.MultiheadAttention(
d_model, n_heads, dropout=dropout, batch_first=True
)
self.ff = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.GELU(),
nn.Linear(d_ff, d_model),
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
attn_out, _ = self.attention(x, x, x)
x = self.norm1(x + self.dropout(attn_out))
ff_out = self.ff(x)
x = self.norm2(x + self.dropout(ff_out))
return x
class SimpleTransformer(nn.Module):
def __init__(
self,
vocab_size=1000,
d_model=256,
n_layers=4,
n_heads=8,
d_ff=1024,
seq_len=64,
):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_embedding = nn.Embedding(seq_len, d_model)
self.layers = nn.ModuleList(
[TransformerBlock(d_model, n_heads, d_ff) for _ in range(n_layers)]
)
self.norm = nn.LayerNorm(d_model)
self.head = nn.Linear(d_model, vocab_size)
self.seq_len = seq_len
def forward(self, x):
B, T = x.shape
positions = torch.arange(T, device=x.device).unsqueeze(0).expand(B, T)
h = self.embedding(x) + self.pos_embedding(positions)
for layer in self.layers:
h = layer(h)
h = self.norm(h)
return self.head(h)Model parameters: 921,064 FP32 model size: 3.7 MB FP16 model size: 1.8 MB
The model prints its parameter count and the theoretical size difference between FP32 and FP16 representations. In practice, the activation memory savings are even more significant at scale, but parameter size gives a sense of the baseline compression.
FP16 Training with Dynamic Loss Scaling
The standard PyTorch mixed precision training loop uses autocast for the forward pass and GradScaler for loss scaling. We demonstrate this with a complete training step.
from torch.amp import autocast
def train_step_fp16(model, optimizer, scaler, x, y):
"""Single training step with FP16 AMP and dynamic loss scaling."""
optimizer.zero_grad()
# Forward pass in FP16
with autocast(device_type=device.type, dtype=torch.float16):
logits = model(x)
# logits: (B, T, vocab_size); y: (B, T)
loss = nn.CrossEntropyLoss()(logits.view(-1, vocab_size), y.view(-1))
# Scale loss, backward, unscale, clip, step
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
return loss.item(), grad_norm.item(), scaler.get_scale()import torch.optim as optim
from torch.amp import GradScaler
# Run a short training loop and collect metrics
optimizer_fp16 = optim.AdamW(model.parameters(), lr=1e-4)
scaler = GradScaler(device=device.type)
losses_fp16 = []
scales = []
overflow_steps = []
n_steps = 30
rng = np.random.default_rng(42)
for step in range(n_steps):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
y = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
prev_scale = scaler.get_scale()
loss, grad_norm, current_scale = train_step_fp16(
model, optimizer_fp16, scaler, x, y
)
losses_fp16.append(loss)
scales.append(current_scale)
if current_scale < prev_scale:
overflow_steps.append(step)FP16 training completed: 30 steps Final loss: 7.0555 Final loss scale: 65536 Overflow events (scale reductions): 0 Mean loss (last 20 steps): 7.0747
The training loop runs cleanly. The GradScaler monitors gradient overflow at each step and reduces the scale only when it detects non-finite gradients. A short, well-behaved run may report no reductions at all, especially on the CPU execution path used for this notebook. That absence is a property of this smoke test, not evidence that dynamic scaling is unnecessary on production FP16 workloads.
BF16 Training without Loss Scaling
BF16 training simplifies the loop: no GradScaler, no overflow checking, no scale management.
def train_step_bf16(model, optimizer, x, y):
"""Single training step with BF16 AMP. No loss scaling needed."""
optimizer.zero_grad()
# Forward pass in BF16
with autocast(device_type=device.type, dtype=torch.bfloat16):
logits = model(x)
loss = nn.CrossEntropyLoss()(logits.view(-1, vocab_size), y.view(-1))
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
return loss.item()# Reset model and run BF16 training
model_bf16 = SimpleTransformer(
vocab_size=vocab_size, d_model=d_model, n_layers=n_layers, seq_len=seq_len
)
model_bf16 = model_bf16.to(device)
optimizer_bf16 = optim.AdamW(model_bf16.parameters(), lr=1e-4)
losses_bf16 = []
for step in range(n_steps):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
y = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
if has_bf16:
loss = train_step_bf16(model_bf16, optimizer_bf16, x, y)
else:
# Fallback to FP32 on hardware without BF16
optimizer_bf16.zero_grad()
logits = model_bf16(x)
loss = nn.CrossEntropyLoss()(logits.view(-1, vocab_size), y.view(-1))
loss.backward()
torch.nn.utils.clip_grad_norm_(model_bf16.parameters(), max_norm=1.0)
optimizer_bf16.step()
loss = loss.item()
losses_bf16.append(loss)BF16 training completed: 30 steps Final loss: 7.0786 Mean loss (last 20 steps): 7.0671 No loss scaling required
The BF16 loop is shorter and easier to reason about. There is no scale factor to monitor, no overflow detection, and no conditional step-skipping logic. The tradeoff is that you need compatible hardware (Ampere GPUs or newer, or TPUs).
Visualizing Training Loss
Let us visualize the loss values for both precision modes. The targets are freshly sampled random tokens at every step, so this 30-step run is a numerical-stability smoke test rather than a convergence experiment. On hardware without BF16 support, the second trace uses FP32.

Visualizing Dynamic Loss Scale Evolution
The recorded scale can remain flat when the short runtime produces no overflows, as it does on this notebook's CPU path. To make the controller mechanics visible on every platform, the next chart uses a deterministic illustrative schedule with four overflow decisions and a shortened five-step growth interval. PyTorch's production default growth interval is much longer at 2,000 successful steps.

Measuring the Throughput Benefit
The actual speedup from mixed precision depends on the GPU model, the operation sizes, and the model architecture. Let us measure time per step for FP32 vs FP16 on this model.
import time
def benchmark_step(
model, use_amp, amp_dtype, use_scaler, n_warmup=3, n_bench=5
):
"""Benchmark a training step and return mean time in milliseconds."""
optimizer = optim.AdamW(model.parameters(), lr=1e-4)
scaler = GradScaler(device=device.type) if use_scaler else None
# Warmup
for _ in range(n_warmup):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
y = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
optimizer.zero_grad()
if use_amp:
with autocast(device_type=device.type, dtype=amp_dtype):
logits = model(x)
loss = nn.CrossEntropyLoss()(
logits.view(-1, vocab_size), y.view(-1)
)
if scaler:
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
else:
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
else:
logits = model(x)
loss = nn.CrossEntropyLoss()(
logits.view(-1, vocab_size), y.view(-1)
)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
if device.type == "cuda":
torch.cuda.synchronize()
# Benchmark
times = []
for _ in range(n_bench):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
y = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
optimizer.zero_grad()
start = time.perf_counter()
if use_amp:
with autocast(device_type=device.type, dtype=amp_dtype):
logits = model(x)
loss = nn.CrossEntropyLoss()(
logits.view(-1, vocab_size), y.view(-1)
)
if scaler:
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
else:
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
else:
logits = model(x)
loss = nn.CrossEntropyLoss()(
logits.view(-1, vocab_size), y.view(-1)
)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
if device.type == "cuda":
torch.cuda.synchronize()
end = time.perf_counter()
times.append((end - start) * 1000)
return float(np.mean(times)), float(np.std(times))# Create fresh models for benchmarking
model_fp32_bench = SimpleTransformer(
vocab_size=vocab_size, d_model=d_model, n_layers=n_layers, seq_len=seq_len
).to(device)
model_fp16_bench = SimpleTransformer(
vocab_size=vocab_size, d_model=d_model, n_layers=n_layers, seq_len=seq_len
).to(device)
t_fp32_mean, t_fp32_std = benchmark_step(
model_fp32_bench, use_amp=False, amp_dtype=None, use_scaler=False
)
t_fp16_mean, t_fp16_std = benchmark_step(
model_fp16_bench, use_amp=True, amp_dtype=torch.float16, use_scaler=True
)
speedup = t_fp32_mean / t_fp16_meanFP32 step time: 21.30 +/- 0.27 ms FP16 step time: 346.84 +/- 1.37 ms Speedup (FP16 vs FP32): 0.06x Note: On CPU, speedup may be < 1x because FP16 tensor cores are GPU-only. On A100/H100 GPUs, typical speedups are 1.5-3x for transformer workloads.
The speedup on CPU will be small or negative because CPUs do not have dedicated FP16 hardware. On a modern GPU with tensor cores (A100, H100, RTX 30/40 series), the speedup for transformer matrix multiplications is typically 1.5x to 3x, depending on the batch and sequence dimensions. Memory savings are consistent regardless of hardware.
Checking Gradient Dtypes at Runtime
During debugging, it is useful to inspect the dtypes flowing through the model during a mixed precision forward pass.
def inspect_dtypes(model, x, amp_dtype):
"""Run one forward pass and print the dtype of each module's output."""
dtypes = {}
def hook(name):
def _hook(module, input, output):
if isinstance(output, torch.Tensor):
dtypes[name] = str(output.dtype)
return _hook
hooks = []
for name, module in model.named_modules():
if isinstance(module, (nn.Linear, nn.LayerNorm, nn.MultiheadAttention)):
hooks.append(module.register_forward_hook(hook(name)))
with autocast(device_type=device.type, dtype=amp_dtype):
_ = model(x)
for h in hooks:
h.remove()
return dtypesDtype inspection during FP16 autocast: -------------------------------------------------- 0.norm1 -> torch.float32 ff.0 -> torch.float16 ff.2 -> torch.float16 0.norm2 -> torch.float32 1.norm1 -> torch.float32 1.norm2 -> torch.float32 norm -> torch.float32 head -> torch.float16
The dtype inspection confirms that nn.Linear layers (which host matrix multiplications) run in FP16 within the autocast context, while nn.LayerNorm stays in FP32. This is the autocast dispatch table in action. The pattern aligns exactly with the theory: operations sensitive to precision stay in FP32, while compute-heavy projection operations use FP16.
Key Parameters
The key parameters for mixed precision training are:
dtypeinautocast:torch.float16for FP16 (requires loss scaling on older hardware),torch.bfloat16for BF16 (preferred on Ampere GPUs and newer).init_scaleinGradScaler: The initial loss scale value. Default is 65536. Lower values reduce early-training overflow at the cost of reduced gradient precision.growth_factorinGradScaler: How much to multiply the scale after a clean stretch. Default is 2.0.backoff_factorinGradScaler: How much to reduce the scale after overflow. Default is 0.5.growth_intervalinGradScaler: Number of consecutive steps without overflow before scaling up. Default is 2000. Lower values make the scale more aggressive.max_norminclip_grad_norm_: The gradient clipping threshold. Apply this after callingscaler.unscale_().
Debugging Mixed Precision Training
Mixed precision training introduces new failure modes that require specific diagnostic strategies. Understanding what can go wrong and how to observe it helps you resolve problems quickly.
Diagnosing NaN and Inf Gradients
If your training run produces NaN losses or diverges unexpectedly, the first step is to determine whether the problem originates from overflow in FP16 gradients or from an underlying numerical problem in the model.
With FP16 and GradScaler, monitor the scale value during training. If you see the scale rapidly decreasing from its initial value and continuing to decrease until it reaches very small values (around 1 or below), your gradients are consistently overflowing. This usually means the model architecture or hyperparameters produce very large gradient values. Common causes include a learning rate that is too high, missing or improperly scaled residual connections, or operations that produce large intermediate values (such as logit scales in attention that are not divided by ).
If the scale remains stable but you still see NaN values, the problem is not overflow in the loss-scaled gradients. Instead, look for custom operations that are not registered with autocast and may receive FP16 inputs they cannot handle, for operations that produce NaN in FP32 as well (like taking the log of zero, or dividing by zero), or for numerical instability in rarely-executed code paths (such as the mask filling in attention for padding tokens).
With BF16, there is no loss scaling, so any NaN values indicate an underlying numerical problem rather than a range issue. BF16's wider range means overflow is essentially impossible. If you see NaN values in BF16 training, you can rule out FP16-style overflow and focus on the model architecture.
Checking for Silently Stalled Weights
The vanishing update problem is harder to detect because it is silent. A useful diagnostic is to track the fraction of model parameters whose gradient norms fall below a threshold. You can log gradient norms per layer using PyTorch's model.named_parameters() loop after the backward pass and before the optimizer step.
If you observe that certain layers consistently show near-zero gradient norms while others show normal magnitudes, those layers may have stopped updating. Check whether the problematic layers are early in the network (where gradients can vanish due to multiplicative propagation through many layers) or whether they are layers that produce very small-magnitude features (such as layers operating on near-zero residuals).
Debugging Autocast Scope
A common mistake is placing the loss computation outside the autocast context. The loss function should generally be inside the autocast block for correct dtype handling, even though cross_entropy will internally use FP32 for its computation:
# Correct: loss inside autocast context
with autocast(device_type="cuda", dtype=torch.float16):
logits = model(x)
loss = criterion(logits, y)
# Incorrect: loss outside autocast may produce dtype mismatch warnings
with autocast(device_type="cuda", dtype=torch.float16):
logits = model(x)
loss = criterion(logits, y) # logits is FP16, criterion expects FP32The second pattern often works because PyTorch will implicitly upcast, but it produces efficiency warnings and can cause subtle issues with custom loss functions that do not handle mixed-dtype inputs.
Limitations and Practical Considerations
Mixed precision training is not without tradeoffs and failure modes. Understanding when it can go wrong helps you diagnose training instability and choose the right configuration.
Residual Precision Loss
Even with dynamic loss scaling, FP16 accumulates more rounding error than FP32. For most models, this error is negligible compared to the stochasticity of SGD. But for models that are numerically sensitive, such as very deep networks, models trained at very high precision for downstream tasks, or models with operations that involve large sums of small values, the reduced precision can measurably affect final performance.
The BF16 trade-off is slightly different: its wider range prevents overflow, but its 7-bit mantissa means each gradient computation is less precise than FP16. In practice, the stability advantage of BF16 dominates, and models trained in BF16 generally match FP32 quality. But if you observe that your BF16 model plateaus earlier or reaches slightly worse perplexity than your FP32 baseline, the precision reduction may be a factor worth investigating. Running a short FP32 ablation over the same hyperparameters for a few thousand steps is the most direct way to isolate precision as a variable.
The precision degradation from FP16 is most likely to surface in specific architectural patterns: softmax operations with very large or very small logits, layer normalizations over very wide hidden dimensions, and attention mechanisms with very long sequence lengths where many values are summed. These operations involve reducing a large number of values into a smaller result, and the accumulated error from each FP16 multiply-accumulate compounds over the length of the reduction. When training models with very long context windows (tens of thousands of tokens), the cumulative rounding error in attention can become noticeable. This is one reason why modern long-context models almost universally use BF16 rather than FP16: the wider range of BF16 tolerates the larger intermediate values that arise in attention over long sequences.
Hardware-Specific Behavior
The speedup from mixed precision depends heavily on the GPU. On older Volta GPUs (V100), FP16 tensor cores provide a 2x throughput advantage for matrix multiplications, but only when matrix dimensions are multiples of 8. On Ampere GPUs (A100), BF16 and FP16 tensor cores provide similar throughput, and dimensions should be multiples of 16 for peak efficiency. On consumer cards without tensor cores, mixed precision may be slower than FP32 due to the casting overhead.
This means that benchmark results transfer only between similar hardware. An FP16 speedup of 2.5x on an A100 does not guarantee any speedup on a GTX 1080. When planning training infrastructure, always benchmark the specific model architecture on the target hardware before committing to a precision strategy. The speedup also depends on the batch size and sequence length: operations that are already memory-bandwidth-limited see a larger benefit from the reduced memory footprint of FP16.
Achieving peak tensor core utilization requires attention to matrix dimensions. The tensor core units on NVIDIA GPUs operate most efficiently when the matrix dimensions are multiples of specific tile sizes: 8 for Volta FP16, 16 for Ampere BF16/FP16, and 32 for Hopper FP8. If your model's hidden dimension, vocabulary size, or feed-forward dimension is not aligned to these tile sizes, you lose a fraction of the potential speedup. This is why large language model architectures often choose hidden dimensions like 4096, 8192, or 16384 (all large powers of 2) rather than arbitrary round numbers like 4000 or 8000. The alignment is not accidental; it is a deliberate choice to maximize hardware utilization.
Incompatible Operations
Some operations cannot run in FP16 because their outputs would overflow or lose too much precision. PyTorch's autocast dispatch table handles the most common cases, but custom CUDA kernels or certain third-party operations may not be registered. If you see NaN values appearing in unusual places, inspect whether a custom operation is receiving FP16 inputs when it expects FP32.
You can force specific operations out of autocast with with torch.amp.autocast(enabled=False): inside the autocast context. This is useful for custom attention implementations, specialized normalization layers, or any operation where you have determined empirically that FP16 produces incorrect results.
Integration with Gradient Accumulation
Gradient accumulation (accumulating gradients over multiple small batches before an optimizer step) requires care with loss scaling. The scaled gradients should be divided by the number of accumulation steps when computing the loss, before scaling, to ensure the effective loss scale is consistent. Alternatively, use the scaler.step() and scaler.update() pattern only at the true optimizer step, not at each accumulation step. Calling scaler.update() at every accumulation step would incorrectly adjust the scale based on partially-accumulated gradients.
A correct pattern for gradient accumulation with loss scaling looks like this:
accumulation_steps = 4
model_ga = SimpleTransformer(
vocab_size=vocab_size, d_model=d_model, n_layers=n_layers, seq_len=seq_len
).to(device)
scaler_accum = GradScaler(device=device.type)
optimizer_accum = optim.AdamW(model_ga.parameters(), lr=1e-4)
optimizer_accum.zero_grad()
for step in range(20):
x_batch = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
y_batch = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
with autocast(device_type=device.type, dtype=torch.float16):
logits_accum = model_ga(x_batch)
# Divide the loss by accumulation_steps before backward
loss_accum = (
nn.CrossEntropyLoss()(
logits_accum.view(-1, vocab_size), y_batch.view(-1)
)
/ accumulation_steps
)
scaler_accum.scale(loss_accum).backward()
if (step + 1) % accumulation_steps == 0:
scaler_accum.unscale_(optimizer_accum)
torch.nn.utils.clip_grad_norm_(model_ga.parameters(), 1.0)
scaler_accum.step(optimizer_accum)
scaler_accum.update()
optimizer_accum.zero_grad()
print(
f"Gradient accumulation completed: {20} steps with accumulation_steps={accumulation_steps}"
)
print(f"Effective batch size: {batch_size * accumulation_steps}")Gradient accumulation completed: 20 steps with accumulation_steps=4 Effective batch size: 64
The division by accumulation_steps before the backward pass ensures that the accumulated gradient matches the gradient you would have gotten from a single step with the full effective batch. The scaler.update() call happens only once per true optimizer step, so the scale adapts at the correct granularity.
Memory Usage Breakdown
Mixed precision reduces memory for activations, which are the dominant consumer during training. For a model with parameters, the approximate memory breakdown is:
| Component | FP32 only | FP16/BF16 mixed |
|---|---|---|
| Parameters (working copy) | 4N bytes | 2N bytes |
| Parameters (master copy) | -- | 4N bytes |
| Optimizer state (AdamW) | 8N bytes | 8N bytes |
| Activations (depends on batch) | Varies | ~50% of FP32 |
| Gradients | 4N bytes | 2N (FP16) + 4N (FP32 master) |
The total parameter and optimizer overhead is 16N bytes for FP32-only (4N parameters + 8N Adam state + 4N gradients) and approximately 18N for mixed precision (2N working + 4N master + 8N Adam + 4N for FP16 gradients). The parameter memory difference is small. The activation memory, which scales with batch size and sequence length, is halved.
To put concrete numbers on this: a 7-billion parameter model like LLaMA-7B has parameters requiring 28 GB in FP32. With AdamW, the optimizer state adds another 56 GB, for a total of about 84 GB before counting activations. On a single A100 with 80 GB of memory, pure FP32 training is impossible for this model at any batch size. With BF16 mixed precision, parameters drop to 14 GB, and the optimizer state stays at 56 GB (maintained in FP32), totaling 70 GB before activations. You can then fit a small batch. For even larger models, combining mixed precision with ZeRO optimization (covered in the ZeRO Optimization chapter) partitions the optimizer state across GPUs, enabling training that would be otherwise infeasible.
Choosing Between FP16 and BF16 in Practice
For practitioners facing this decision today, the guidance is relatively clear. If you are training on NVIDIA Ampere or newer hardware (A100, A10, RTX 30/40 series), H100, or Google TPUs, use BF16. The simpler training loop, the elimination of loss scaling overhead, and the consistently stable behavior outweigh BF16's slightly lower fraction precision. The major open-source models released in the past several years (LLaMA, Mistral, Falcon, and their variants) all use BF16, and the training codebases for those models reflect this as the default.
If you are training on older hardware that lacks BF16 support (V100, older RTX 20 series, any Pascal-era GPU), FP16 with dynamic loss scaling is the correct choice. The GradScaler machinery handles the overflow detection reliably in practice, and the training stability of FP16 with loss scaling is well-established for transformer architectures.
If you are running on hardware without any dedicated low-precision units (which is increasingly rare for serious LLM training), mixed precision may offer little speedup and can introduce complexity without benefit. In this case, FP32 training is simpler and avoids the precision management overhead entirely.
Summary
Mixed precision training exploits the difference between numerical range and precision in floating point formats to achieve faster training with less memory, without materially harming model quality.
The key concepts are:
- Floating point formats: FP32 (32 bits, 8 exponent, 23 mantissa), FP16 (16 bits, 5 exponent, 10 mantissa), BF16 (16 bits, 8 exponent, 7 mantissa). Exponent bits control range; mantissa bits control precision.
- FP16 limitations: The narrow range (max ~65504) risks overflow; the limited precision risks underflow for small gradients. Both cause training instability.
- Loss scaling: Multiply the loss by a large constant before the backward pass to shift gradient magnitudes into FP16's representable range. Dynamic loss scaling adjusts automatically by detecting overflow and adapting.
- BF16 advantages: Same numerical range as FP32 (8 exponent bits), making overflow essentially impossible. No loss scaling required. First-class support on modern GPUs and TPUs. Has become the default for large-scale transformer training.
- Mixed precision implementation: Use
torch.autocastto run matrix multiplications and attention in low precision. Usetorch.amp.GradScalerwith FP16, omit it with BF16. Keep FP32 master weights and optimizer states. - The two-copy pattern: FP16 working copies for the forward and backward passes, FP32 master copies for the optimizer update. This preserves weight update precision while gaining activation memory savings.
- Practical considerations: Tensor core efficiency requires matrix dimension alignment; gradient accumulation must account for loss scaling; custom operations may need explicit FP32 coercion.
The next chapter on Communication Optimization extends these infrastructure ideas to distributed training, where the precision and format of gradients communicated between GPUs adds another layer of efficiency to consider.
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about mixed precision training.
Mixed Precision Training Quiz
Reference
Citation details
Cite or share this article.
Continue with the full handbook
This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.
Explore Language AI HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

Comments
No comments yet. Be the first to share your thoughts!