TF 2.6环境下WGAN处理大数据集时损失输出NaN问题求助
问题原因
从tf.debugging.enable_check_numerics()的输出可以定位到,NaN出现在训练迭代器返回的输入张量阶段,不是模型前向/反向传播计算过程中产生的,核心原因有两个:
- 全量数据集中存在缺失值/溢出值,读入后被转为NaN/Inf,但原有过滤无效行的逻辑存在漏洞:用
np.std(myline) !=0判断时,如果行内已经存在NaN,std返回值也为NaN,NaN与0的布尔比较结果为False,导致携带无效值的行没有被过滤,进入最终训练数据集。 - 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
相关产品推荐
相关产品推荐

