You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Flax训练神经网络时如何忽略输出NaN值进行训练?

问题解决:Flax训练中忽略NaN值训练的正确实现

这个需求完全可以实现,你遇到的损失变NaN问题,主要是因为jnp.nanmean在所有样本均为NaN时会返回NaN,进而导致梯度和参数更新出现NaN;另外也可能是模型输出本身产生了NaN,放大了这个问题。以下是具体修复方案:

核心修复思路

  1. 显式过滤NaN样本,计算有效样本的损失均值
  2. 处理全NaN的极端情况,避免除以0或返回NaN
  3. 排查模型输出的数值稳定性问题(根源性解决)

修复后的代码示例

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 15:02:45