TensorFlow训练出现NaN损失:尝试调整logits与labels维度匹配
解决TensorFlow二分类训练中NaN损失的问题
刚上手TensorFlow遇到NaN损失确实挺头疼的,我帮你梳理下代码里的几个关键问题,一步步来解决:
1. 数据未归一化是NaN损失的头号元凶
你的图像像素值是0-255的范围,直接输入网络会导致激活值过大,进而引发梯度爆炸,最后损失直接变成NaN。这是新手最容易踩的坑之一。
解决方法很简单,把像素值归一化到0-1的范围,同时把数据类型从float16改成float32——float16精度太低,很容易出现数值溢出:
def load_images(path): # ... 其他代码保持不变 for img in os.listdir(directorypath): img = os.path.join(directorypath, img) if not img.endswith(".bmp"): continue a = ndimage.imread(img) if a is None: print ("Unable to read image: ", img) continue a = np.resize(a, [512, 512]) # 新增:归一化到0-1区间 a = a / 255.0 list_of_imgs.append(a.flatten()) # ... 其他代码保持不变 # 修改数据类型为float32 images = np.array(list_of_imgs, dtype="float32") labels = np.array(list_of_classes, dtype="int32") return images,labels
2. 网络结构和二分类任务不匹配
你做的是二分类任务,但最后一层logits居然输出了10个单元,这完全不对!二分类只需要输出2个单元(对应两类)就够了,否则损失函数计算时会出现逻辑错误,也可能导致数值异常。
修改logits层:
# Logits Layer - 二分类任务改成输出2个单元 logits = tf.layers.dense(inputs=dropout, units=2)
你的标签是0和1的整数,用sparse_softmax_cross_entropy损失函数是合适的,这部分不用改。
3. 过大的卷积核容易引发梯度不稳定
你用了16x16的大卷积核,对于512x512的输入来说,这个尺寸太大了——大卷积核会导致卷积层参数过多,计算出来的梯度很容易波动甚至爆炸。CNN里常规操作都是用3x3或者5x5的小卷积核,既能提取有效特征,又能保持梯度稳定:
# 修改第一层卷积核尺寸 conv1 = tf.layers.conv2d( inputs=input_layer, filters=32, kernel_size=[3, 3], # 从16x16改成3x3 padding="same", activation=tf.nn.relu) # 修改第二层卷积核尺寸 conv2 = tf.layers.conv2d( inputs=pool1, filters=64, kernel_size=[3, 3], # 从16x16改成3x3 padding="same", activation=tf.nn.relu)
4. 可选:添加梯度裁剪防止极端情况
如果上面的修改后还是出现NaN,可以试试添加梯度裁剪,限制梯度的最大范数,彻底杜绝梯度爆炸的可能:
if mode == tf.estimator.ModeKeys.TRAIN: optimizer = tf.train.GradientDescentOptimizer(learning_rate=0.001) # 计算梯度并裁剪 gradients, variables = zip(*optimizer.compute_gradients(loss)) gradients, _ = tf.clip_by_global_norm(gradients, 5.0) # 把梯度最大范数设为5 train_op = optimizer.apply_gradients( zip(gradients, variables), global_step=tf.train.get_global_step()) return tf.estimator.EstimatorSpec(mode=mode, loss=loss, train_op=train_op)
5. 小细节:图像resize更合理的方式
你用np.resize直接拉伸图像,可能会导致图像严重变形,影响模型学习。建议用PIL的resize方法,用高质量插值来调整尺寸:
from PIL import Image # 替换原来的np.resize部分 img_pil = Image.fromarray(a) # 用LANCZOS插值保持图像质量 img_pil = img_pil.resize((512, 512), Image.LANCZOS) a = np.array(img_pil)
按照上面的步骤修改后,应该就能解决NaN损失的问题了。先从数据归一化和网络结构匹配这两点入手,这是最关键的!
内容的提问来源于stack exchange,提问作者Adrián Arroyo Perez
相关产品推荐
相关产品推荐

