Your TensorFlow training loss turns NaN after a few hundred steps — how do you diagnose it?
answer
- find the first non-finite op
- check the inputs first
- rising norm versus one-shot jump
- log of zero, divide by zero, sqrt
- clipping treats explosion, not log(0)
basics
~20 sFind the first op that produces a non-finite value rather than guessing. tf.debugging.enable_check_numerics() raises at that op with a stack trace; then check inputs for NaN, watch the gradient global norm for a spike, and only then reach for clipping or a lower rate.
solid answer
~50 sWork backwards from the first non-finite value, not from a hypothesis. Call `tf.debugging.enable_check_numerics()` on a reproduction run: it makes the op that first emits NaN or Inf raise immediately, with a stack trace pointing at the line, which usually ends the investigation. In parallel, rule out the data — assert with `tf.debugging.assert_all_finite` on incoming batches, since a division or a missing value in the input pipeline reaches the loss unchanged. Then log `tf.linalg.global_norm(grads)` every step: a norm that climbs over several steps and spikes just before the NaN is exploding gradients, treated with `clipnorm` or `global_clipnorm` on the optimizer, a lower learning rate, or warmup. A norm that is flat until it jumps to NaN in one step points instead at an arithmetic hazard inside the loss — a `log` of zero, a division by an empty-class count, a `sqrt` of a negative — or at float16 overflow if mixed precision is on. Reproduce on a single fixed batch to iterate fast.
code
python · 17 linesimport tensorflow as tf
tf.debugging.enable_check_numerics() # raises at the first op emitting NaN or Inf
model = tf.keras.Sequential([tf.keras.layers.Dense(10)])
optimizer = tf.keras.optimizers.SGD(learning_rate=0.1, global_clipnorm=1.0)
loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
x = tf.random.normal((32, 4))
y = tf.random.uniform((32,), maxval=10, dtype=tf.int32)
tf.debugging.assert_all_finite(x, "non-finite input batch")
with tf.GradientTape() as tape:
loss = loss_fn(y, model(x, training=True))
grads = tape.gradient(loss, model.trainable_variables)
tf.print("loss:", loss, "grad global norm:", tf.linalg.global_norm(grads))
optimizer.apply_gradients(zip(grads, model.trainable_variables))go deeper
Recall the usual suspects — a learning rate that is too high, non-finite values in the data, and a log or division by zero in the loss — and know that TensorFlow has debugging helpers rather than guessing.
Explain how to localize the first non-finite op instead of tuning blindly, and describe the concrete fixes for each mechanism: clipping for explosion, safe arithmetic for hazards, finite assertions for data.
Demonstrate the full routine on a real system: a deterministic single-batch repro, per-step gradient-norm logging to classify the failure, check-numerics on the repro run, a float32 rerun to rule out half-precision range, and a permanent guard left behind.
Make it a property of the platform rather than a personal skill — non-finite detection and gradient-norm telemetry on by default in the standard training harness, so any team hitting this gets the diagnosis from the logs instead of restarting a multi-day run blind.
## Localize before theorizing NaN debugging goes wrong when people start changing hyperparameters. The productive first move is to find the exact op that produced the first non-finite value. `tf.debugging.enable_check_numerics()` instruments execution so the first op emitting NaN or Inf raises an error naming the op and giving a stack trace. It slows training substantially, so turn it on for a reproduction run and call `tf.debugging.disable_check_numerics()` afterwards. In most real cases this single call identifies the culprit line and you are done diagnosing. A useful companion is `tf.config.run_functions_eagerly(True)`, which forces compiled step functions to run op by op in Python so stack traces point at your source rather than at graph internals, and Python breakpoints work inside the step. ## Three families of cause **Bad input.** A NaN that arrives in the batch propagates straight through, and the model is blameless. It happens more than people expect: a normalization dividing by a zero standard deviation for a constant feature, a missing value that survived as a NaN, a corrupted record, a label outside the valid class range. Assert it away — `tf.debugging.assert_all_finite(batch, "non-finite input")` at the top of the step, or a one-off scan of the dataset with `tf.math.is_nan`. **Exploding gradients.** The loss is fine, then the weights blow up. The signature is a rising gradient norm: log `tf.linalg.global_norm(grads)` on every step and you will see it grow over tens of steps and then spike. Common triggers are a learning rate too high for the batch size, no warmup on a deep or attention-heavy model, recurrent architectures where the recurrence amplifies, or a bad initialization. Direct fixes are `tf.clip_by_global_norm(grads, max_norm)` before applying, or configuring `clipnorm`, `global_clipnorm` or `clipvalue` on the optimizer so it clips for you. Clipping is a real fix here, not a mask, because it targets the actual mechanism. **An arithmetic hazard in the loss or the model.** These produce a NaN in one step from a healthy state. `log(0)` from a hand-written crossentropy where a predicted probability rounded to zero — the reason the built-in crossentropy losses prefer to consume logits and fuse the softmax internally. `0/0` from averaging over a class or a mask that happens to be empty in that batch. `sqrt` of a value that went slightly negative through rounding, whose derivative is infinite at zero even when the forward value looks fine. Division by a variance before an epsilon is added. The fix here is arithmetic, not optimization: consume logits rather than probabilities, add a small epsilon *inside* the risky op, guard empty-denominator cases with `tf.math.divide_no_nan`, and clamp before `sqrt` and `log`. Note the subtlety that an op can be finite forward and non-finite in its gradient, so a forward-only inspection can come up empty while the backward pass is where it dies. ## The mixed-precision special case If a half-precision policy is on, a value too large for float16 becomes Inf where float32 would have coped. The quick discriminator: rerun the same steps in plain float32. If the NaN disappears, it is a range problem — force the offending region to float32, or check that loss scaling is actually configured, since dynamic scaling is supposed to skip overflowing steps rather than let them poison the weights. ## A repeatable routine 1. Capture a deterministic repro: fix the seed, and if possible loop on one saved batch so the failure arrives in seconds instead of minutes. 2. Assert inputs finite. Eliminate the data. 3. Log the gradient global norm and the loss every step. Rising versus one-shot separates exploding gradients from an arithmetic hazard. 4. Run once with `enable_check_numerics` (and eager execution if the trace is opaque) to name the op. 5. Apply the fix that matches the mechanism you found — clipping and warmup for explosion, epsilons and safe division for hazards, dtype changes for range. 6. Add a permanent, cheap guard: an assertion on the loss being finite, or a check that skips and logs a non-finite step rather than corrupting the weights, so the next occurrence is reported instead of discovered days later. ## What not to do Do not sprinkle large epsilons everywhere until the NaN stops — that changes the objective and hides the mechanism. Do not treat clipping as a universal cure; it does nothing for a `log(0)`. And do not lower the learning rate and declare victory without having identified why it exploded, because the same run will fail on a longer schedule or a bigger batch.
- How do you tell exploding gradients from a bad op in the loss?Log the gradient global norm every step. Exploding gradients show a norm climbing over many steps before the failure, so clipping or a lower rate addresses the mechanism. An arithmetic hazard shows a flat, healthy norm and then a single step where the loss goes non-finite from nowhere — clipping will not help, and you need to find the log, division or sqrt that broke.
- Why can an op be finite in the forward pass but produce NaN in the backward pass?Derivatives can be unbounded where the value is not. `sqrt(x)` at x=0 is 0 but its gradient is infinite; `log(x)` near zero is large-negative with a gradient of 1/x; a division whose denominator is tiny has a huge derivative. So inspecting only activations can come up clean while the tape's backward pass is where the non-finite value is born.
- What permanent guard would you leave in the training loop?A cheap finiteness check on the loss each step — assert or, better, detect and skip the update while logging the batch identifier and the gradient norm. That converts a silent weight corruption into an actionable log line, costs almost nothing, and preserves the run instead of losing hours of training to weights already filled with NaN.
saying these in an interview costs you the question
- Lowers the learning rate without finding the failing op
- Adds large epsilons everywhere until the NaN disappears
- Assumes NaN always means the learning rate is too high
- Never checks whether the input batches are finite
- Claims gradient clipping fixes a log of zero in the loss