How do you enable mixed_float16 in TensorFlow, and what must you fix afterwards?
answer
- two dtypes: compute and variable
- weights stay float32
- last layer back to float32
- float16 gradients underflow to zero
- bfloat16 keeps float32's exponent
basics
~10 sCall tf.keras.mixed_precision.set_global_policy("mixed_float16"): layers then compute in float16 while keeping float32 weights. Two fixes follow — force the output layer to float32, and apply loss scaling so small float16 gradients do not underflow to zero.
solid answer
~50 s`tf.keras.mixed_precision.set_global_policy("mixed_float16")` sets a dtype policy where every layer's compute dtype is float16 but its variable dtype stays float32, so matmuls and convolutions run on Tensor Cores while the master weights keep full precision. Two things then need attention. First, the model's final layer should emit float32 — set `dtype="float32"` on it — so the loss and any softmax are computed in a range where they are numerically well behaved. Second, float16's narrow exponent range flushes small gradients to zero, so you need loss scaling: multiply the loss by a large factor before the backward pass and divide the gradients back afterwards. In a hand-written training step you get that by wrapping the optimizer in `tf.keras.mixed_precision.LossScaleOptimizer`; with the high-level training API Keras handles it under this policy. Also check the hardware: without Tensor Cores the policy buys little and can be slower, and `mixed_bfloat16` needs no loss scaling at all because bfloat16 keeps float32's exponent range.
code
python · 10 linesimport tensorflow as tf
tf.keras.mixed_precision.set_global_policy("mixed_float16")
inputs = tf.keras.Input(shape=(784,))
h = tf.keras.layers.Dense(256, activation="relu")(inputs) # computes in float16
outputs = tf.keras.layers.Dense(10, dtype="float32")(h) # logits back in float32
model = tf.keras.Model(inputs, outputs)
print(model.layers[1].compute_dtype, model.layers[1].variable_dtype) # float16 float32go deeper
Know that the policy makes layers compute in float16 while keeping float32 weights, and that it is switched on globally before the model is built.
Explain the two dtypes in a policy, why the output layer is forced back to float32, and what loss scaling is protecting against — small gradients underflowing rather than overflowing.
Show the operational judgment: wrap the optimizer yourself in a custom loop, read the skipped-step and loss-scale signals correctly, force numerically delicate regions to float32, and verify the speedup by measurement rather than assumption.
Decide the default for the fleet. Weigh mixed_bfloat16's simplicity against mixed_float16's hardware reach, require an accuracy comparison against a float32 baseline before it becomes standard, and account for the debugging surface it adds for everyone after you.
## What the policy changes A Keras dtype policy has two dtypes: the **compute dtype**, used for a layer's arithmetic and activations, and the **variable dtype**, used to store its weights. `mixed_float16` sets compute to float16 and variables to float32. Set it once, before you build the model, with `tf.keras.mixed_precision.set_global_policy("mixed_float16")`, and every layer created afterwards picks it up. Any layer exposes `compute_dtype` and `variable_dtype` so you can verify what it inherited. The payoff is twofold: half-precision matmuls and convolutions run on dedicated hardware units several times faster than float32, and activations — usually the dominant term in training memory — halve in size, which lets you raise the batch size. The weights themselves are unchanged in size, so "mixed precision halves my model's memory" is wrong. Keeping variables in float32 is what makes it *mixed*. Weight updates are typically tiny relative to the weights; accumulating them in float16 would lose the update entirely to rounding, and the model would stop improving. ## Fix one: a float32 output Softmax over float16 logits, and the crossentropy that follows, sit exactly where float16's limited range hurts: large logits overflow to infinity, small probabilities flush to zero. The convention is to end the model with a linear layer built as `dtype="float32"`, so the head's output — and everything the loss does with it — is full precision. If your architecture ends in an activation, put it in its own layer with `dtype="float32"`. The cost is negligible, since only the last, narrow layer is affected. The same applies to any numerically delicate op you write by hand inside the model: reductions over very large tensors, normalization by a small denominator, or accumulating a running sum. Cast those regions to float32 explicitly. ## Fix two: loss scaling Float16 has a much narrower exponent range than float32. Gradients late in a backward pass are often small enough to fall below the smallest representable float16 value and become exactly zero — they *underflow*, and the corresponding weights stop learning. The loss looks fine; the model just converges worse. Loss scaling fixes it arithmetically. Multiply the loss by a large factor before the backward pass; by linearity every gradient is multiplied by the same factor, lifting the small ones back into representable range. Divide the gradients by that factor before applying them, and the update is mathematically identical to the unscaled one. Dynamic loss scaling automates the choice of factor: it starts high, and whenever any gradient comes back infinite or NaN it skips that update and halves the factor, raising it again after a stretch of clean steps. Skipped steps early in training are expected, not a bug. In TensorFlow this is `tf.keras.mixed_precision.LossScaleOptimizer`, which wraps a normal optimizer. In a hand-written step you scale the loss before `tape.gradient` and let the wrapper unscale and validate the gradients when you apply them. If you drive training through the built-in high-level loop, Keras arranges loss scaling for you under the mixed_float16 policy — which is precisely why teams that move from the high-level API to a custom loop hit a mysterious accuracy regression: the wrapper they never wrote is now missing. ## bfloat16 is the easier variant `mixed_bfloat16` uses bfloat16 for compute. bfloat16 has the same 8-bit exponent as float32 and buys the range back by giving up mantissa bits, so gradients do not underflow and **no loss scaling is needed**. Where the hardware supports it well, it is the lower-friction choice. Where it does not, mixed_float16 with loss scaling remains the option. ## Verify the win, do not assume it Mixed precision pays off on hardware with half-precision tensor units; on older GPUs the casts cost more than the arithmetic saves. Even on capable hardware, a small model or a starved input pipeline can be bound elsewhere, so the step time barely moves. Measure step time before and after on the same data, and confirm accuracy over a real run rather than a few hundred steps. Two more practical gotchas. Custom layers that mix a hard-coded float32 constant into a float16 activation raise a dtype mismatch; cast with the layer's `compute_dtype` instead of hardcoding. And because the policy is global and read at layer construction time, setting it *after* building the model does nothing — which produces the confusing report that "mixed precision made no difference at all".
- Why does mixed_bfloat16 not need loss scaling?bfloat16 keeps float32's 8-bit exponent and spends the saved bits from the mantissa instead. Since underflow is a range problem, not a precision problem, small gradients stay representable and no scaling is required. The tradeoff is coarser precision per value, which neural network training tolerates well. It is the simpler policy where the hardware supports it.
- After enabling the policy, a custom layer raises a dtype mismatch. Why?Its inputs now arrive as float16 while a constant or intermediate you created is still float32, and TensorFlow does not implicitly promote. Cast explicitly to the layer's `compute_dtype` — available as `self.compute_dtype` inside a Layer subclass — rather than hardcoding a dtype, so the same layer works under any policy including plain float32.
- Your loss scale keeps halving and many steps are skipped. What does that tell you?Dynamic scaling halves the factor whenever gradients come back inf or NaN. A few skips at the start are normal while the scale finds its level. Persistent skipping means real overflow — usually exploding gradients or a numerically unstable op left in float16. Investigate the model, not the scaler: add gradient clipping, or force the offending region to float32.
- You enabled the policy but step time is unchanged. What do you check?First, whether the policy was set before the model was constructed — layers capture it at build time. Second, whether the hardware has half-precision tensor units at all. Third, whether training is bound by the input pipeline or by many small ops rather than by large matmuls, in which case the arithmetic was never the bottleneck.
saying these in an interview costs you the question
- Claims mixed precision halves the weight memory
- Skips loss scaling in a hand-written training step
- Leaves the softmax or final layer in float16
- Expects a speedup on hardware without tensor cores
- Thinks bfloat16 also requires loss scaling
- Sets the global policy after building the model