Diagnosing and Fixing NaN Loss in Keras Mixed‑Precision Training
When Keras mixed‑precision training stalls with NaN loss, the culprit is often hidden in unsupported layers, missing loss scaling, or incompatible optimizers. This diagnostic guide walks through a concise checklist, quick checks, and concrete fixes to get your model back on track.
07 Feb 2026, 05:50 UTC

Problem Overview
When you enable tf.keras.mixed_precision and start training, you may see the loss drop to nan after a few steps. The training stalls, metrics diverge, and the model never converges. This guide walks through the most common culprits, a quick diagnostic checklist, and targeted fixes.
Root Causes in One Table
| Cause | Typical Symptoms | Quick Check |
|---|---|---|
Unsupported layers (e.g., Lambda, custom ops) | Loss becomes NaN immediately after the first batch. | Replace the layer with a Keras equivalent or cast to float32. |
| Missing or mis‑configured loss scaling | Gradients explode, loss spikes to inf then NaN. | Print tf.keras.mixed_precision.global_policy().loss_scale. |
| Optimizer incompatibility (e.g., AdamW) | Training stops after a few steps, loss stays at NaN. | Switch to Adam or add manual scaling. |
| Premature float16 data cast in the pipeline | Gradients are NaN from the start of training. | Ensure tf.data pipelines output float32 tensors. |
| Inconsistent dtype in loss/metrics | NaNs appear only when evaluating metrics. | Force float32 dtype in loss function. |
Step‑by‑Step Diagnostic Flow
- Confirm Mixed‑Precision Policy
import tensorflow as tf from tensorflow.keras import mixed_precision mixed_precision.set_global_policy('mixed_float16') print(mixed_precision.global_policy()) # Expected output: mixed_float16 with loss_scale=DynamicRunning this on a machine without a compatible GPU will silently fall back to float32, so verify you’re on a Volta or newer NVIDIA card.
- Check Loss Scaling State
policy = mixed_precision.global_policy() print('Loss scale:', policy.loss_scale) # Should be "Dynamic" or a numeric value, not NoneIf
loss_scaleisNone, loss scaling is disabled and NaNs are likely due to overflow. - Run a Minimal Model
Use a tiny model on a small dataset to isolate the problem.
import numpy as np from tensorflow.keras import layers, models x = np.random.rand(32, 28, 28, 1).astype('float32') y = np.random.randint(0, 10, 32) def build_model(): return models.Sequential([ layers.Conv2D(32, 3, activation='relu'), layers.Flatten(), layers.Dense(10) ]) model = build_model() model.compile(optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)) model.fit(x, y, epochs=1)If this runs without NaNs, the issue is in your custom layers or data pipeline.
- Inspect Custom Layers
Wrap any
Lambdaor user‑defined op with explicittf.casttofloat32before computation, then cast back.def my_layer(x): x32 = tf.cast(x, tf.float32) # compute in float32 y = tf.math.sin(x32) return tf.cast(y, tf.float16) - Validate Data Pipeline Dtypes
Check that
tf.datamapping functions do not cast tofloat16before the model consumes them.dataset = tf.data.Dataset.from_tensor_slices((x, y)) # Ensure dtype remains float32 for batch in dataset.take(1): print(batch[0].dtype) - Test with
mixed_float32Temporarily switch to
mixed_float32to see if NaNs disappear. If they do, the problem is definitely related to precision.mixed_precision.set_global_policy('mixed_float32') - Check Optimizer Compatibility
Some optimizers, like
AdamWfromtensorflow_addons, do not apply loss scaling automatically. Replace withAdamor wrap the optimizer withmixed_precision.LossScaleOptimizer:optimizer = tf.keras.optimizers.Adam() optimizer = mixed_precision.LossScaleOptimizer(optimizer) model.compile(optimizer=optimizer, loss=...) - Monitor Gradient Norms
Large batch sizes can still overflow even with scaling. Log the norm of gradients after each step:
@tf.function def train_step(x, y): with tf.GradientTape() as tape: logits = model(x, training=True) loss = loss_fn(y, logits) grads = tape.gradient(loss, model.trainable_variables) norm = tf.linalg.global_norm(grads) print('Grad norm:', norm) optimizer.apply_gradients(zip(grads, model.trainable_variables)) - Escalation Criteria
- If NaNs persist after all above checks, isolate by removing custom layers one by one.
- Run the same model on a CPU to confirm the issue is GPU‑specific.
- Submit a minimal reproducible example to the TensorFlow GitHub issue tracker with TensorFlow 2.16 and GPU details.
Practical Example: Fixing a NaN in a Custom Lambda Layer
Suppose you have the following layer in a model:
def custom_activation(x):
return tf.nn.relu(x) * tf.math.exp(-x)
model = models.Sequential([
layers.Dense(64, activation='relu'),
layers.Lambda(custom_activation),
layers.Dense(10)
])
When training with mixed_float16, the exponential can under‑flow, producing NaNs. Fix: cast to float32 inside the lambda and cast back.
def custom_activation(x):
x32 = tf.cast(x, tf.float32)
y = tf.nn.relu(x32) * tf.math.exp(-x32)
return tf.cast(y, tf.float16)
Re‑run training; the loss stays finite.
Key Takeaways
- Always verify the global mixed‑precision policy and loss scaling before training.
- Unsupported ops and premature dtype changes are the most common NaN triggers.
- Use
mixed_float32as a quick sanity check; if NaNs vanish, precision is the culprit. - Wrap incompatible optimizers with
LossScaleOptimizeror switch to a supported one. - Monitor gradient norms to catch hidden overflows early.
Final Checklist for Release
- [ ] Mixed‑precision policy set to
mixed_float16and loss scaling active. - [ ] All custom layers cast to
float32internally. - [ ] Data pipeline outputs
float32tensors. - [ ] Optimizer is compatible or wrapped with
LossScaleOptimizer. - [ ] Gradient norms stay below a reasonable threshold (e.g.,
1e4). - [ ] NaN loss is not present after the first epoch.
Conclusion
Mixed‑precision can boost performance dramatically, but it introduces subtle dtype pitfalls. By following the ordered checks above, you can quickly isolate NaN loss issues and apply the appropriate fix—be it a dtype cast, loss‑scaling adjustment, or optimizer change.
0 replies
A thoughtful contribution can make all the difference. Be the first to share one.