使用Flax训练神经网络时如何忽略输出NaN值进行训练?
问题解决:Flax训练中忽略NaN值训练的正确实现
这个需求完全可以实现,你遇到的损失变NaN问题,主要是因为jnp.nanmean在所有样本均为NaN时会返回NaN,进而导致梯度和参数更新出现NaN;另外也可能是模型输出本身产生了NaN,放大了这个问题。以下是具体修复方案:
核心修复思路
- 显式过滤NaN样本,计算有效样本的损失均值
- 处理全NaN的极端情况,避免除以0或返回NaN
- 排查模型输出的数值稳定性问题(根源性解决)
修复后的代码示例
import jax.numpy as jnp def nanloss(params, inputs, targets): pred = model.apply(params, inputs) # 计算每个样本的均方误差 squared_error = (pred - targets) ** 2 # 生成有效样本的mask(非NaN标记为True) valid_mask = ~jnp.isnan(squared_error) # 计算有效误差总和与有效样本数量 total_valid_error = jnp.sum(squared_error * valid_mask) valid_sample_count = jnp.sum(valid_mask) # 避免除以0:无有效样本时返回0.0,否则计算均值 loss = jnp.where(valid_sample_count > 0, total_valid_error / valid_sample_count, 0.0) # 可选:临时处理模型输出的NaN(建议找到根源后移除) # pred = jnp.where(jnp.isnan(pred), 0.0, pred) return loss def train_step(state, inputs, targets): loss, grads = jax.value_and_grad(nanloss)(state.params, inputs, targets) # 可选:检查并修复梯度中的NaN(仅临时应急,优先解决根源) # grads = jax.tree_map(lambda g: jnp.where(jnp.isnan(g), 0.0, g), grads) state = state.apply_gradients(grads=grads) return state, loss
为什么原代码失效?
jnp.nanmean在所有输入均为NaN时会返回NaN,此时梯度计算结果也是NaN,参数更新后模型参数变为NaN,后续所有输出和损失都会变成NaN。- 如果模型本身存在数值不稳定(比如激活函数溢出、log(0)、除以0等操作),会导致输出pred出现NaN,即使targets有有效值,
(pred - targets)**2也会变成NaN,逐渐消耗有效样本,最终触发全NaN的情况。
额外排查建议
- 检查模型的层结构:比如是否使用了可能产生NaN的操作(如未加epsilon的log、softmax输入过大等),给这类操作添加数值稳定项(比如
jnp.log(x + 1e-8))。 - 检查初始化参数:极端的初始化值可能导致模型一开始就输出NaN。
- 在训练日志中加入有效样本数量的监控,确认是否存在大量样本变为NaN的趋势。
内容的提问来源于stack exchange,提问作者rhombidodecahedron
相关产品推荐
相关产品推荐

