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

TF 2.6环境下WGAN处理大数据集时损失输出NaN问题求助

问题原因

从tf.debugging.enable_check_numerics()的输出可以定位到,NaN出现在训练迭代器返回的输入张量阶段,不是模型前向/反向传播计算过程中产生的,核心原因有两个:

  1. 全量数据集中存在缺失值/溢出值,读入后被转为NaN/Inf,但原有过滤无效行的逻辑存在漏洞:用np.std(myline) !=0判断时,如果行内已经存在NaN,std返回值也为NaN,NaN与0的布尔比较结果为False,导致携带无效值的行没有被过滤,进入最终训练数据集。
  2. TF 2.4及以上版本对输入张量的数值校验严格度远高于TF 2.3,旧版本会自动忽略输入中的轻微数值异常,新版本会直接抛出NaN错误;同时全量数据集下10万维的大输入+1000的batch size触发了TF输入缓存的内存对齐问题,也可能导致数值异常。

解决方案

  • 修复数据读入与过滤逻辑
    在读入每个值、过滤无效行时增加数值有效性校验,避免无效值进入数据集:

    # 过滤无效行的逻辑修改
    myline_arr = np.asarray(myline).astype(np.float32)
    # 先检查所有值合法,再判断标准差不为0,加小阈值避免浮点精度问题
    mybool = np.all(np.isfinite(myline_arr)) and (np.std(myline_arr) > 1e-6)
    meaningful.append(mybool)
    
    # 读入每个cell时校验
    for mycell in myline : 
        val = float(mycell)
        # 处理非法值,可根据业务需求选择填充或者标记行无效
        if not np.isfinite(val):
            mybool = False
            val = 0.0
        sampleData[sampleIDs[i]].append(val)
        i+=1
    
  • 数据集全局校验
    所有数据读入转换为numpy数组后,统一处理无效值:

    real = np.array(real)
    # 方案1:替换所有NaN/Inf为固定值/特征均值
    finite_vals = real[np.isfinite(real)]
    real = np.nan_to_num(real, nan=0.0, posinf=np.nanmax(finite_vals), neginf=np.nanmin(finite_vals))
    # 方案2:直接删除携带无效值的样本
    # real = real[np.all(np.isfinite(real), axis=1)]
    
  • 适配TF 2.6版本特性
    代码开头加入配置,对齐旧版本TF的优化行为,同时调小batch size避免内存对齐问题:

    import tensorflow as tf
    # 关闭可能导致数值异常的布局优化
    tf.config.optimizer.set_experimental_options({"layout_optimizer": False})
    # 固定浮点精度
    tf.keras.backend.set_floatx('float32')
    
    # 训练时将batch size从1000调整为256或更低
    x_real, y_real = generate_real_samples(real, 256)
    x_fake, y_fake = generate_fake_samples(generator, latent_dim_size, 256)
    

内容的提问来源于stack exchange,提问作者Whitehot

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 22:27:04