TensorFlow旋转目标检测自定义损失函数训练出现NaN值问题求助
问题原因与修复方案
核心原因:NaN梯度反向传播
你当前写法的问题出在tf.where的运算逻辑上:TensorFlow执行tf.where时会提前计算两个分支的所有运算,再根据条件选择输出值。当负样本的标签为NaN时,你依然会执行true_labels[:, 1:3] - pred_labels[:, 1:3]这类运算,结果本身就是NaN,虽然tf.where最终会把负样本的损失置为0,但反向传播时NaN梯度会流入预测值的更新逻辑,导致pred_labels被更新为NaN,后续所有batch的损失都会变成NaN。
其他潜在问题
- 最后一层输出无激活适配:你直接用无激活的
Dense(6)输出,第一个维度的目标存在概率直接喂给默认的二元交叉熵(默认from_logits=False,要求输入是0-1的概率),可能出现数值不稳定;角度、坐标的输出没有范围限制,也可能导致梯度过大。 - 损失平均逻辑错误:你当前对整个batch的损失做全局平均,当batch内负样本占比高时,正样本的损失量级会被过度压缩,影响训练效果。
修复方案
1. 损失函数重写(核心修复)
先通过掩码过滤负样本,提前替换标签中的NaN,避免NaN参与运算:
import tensorflow as tf binary_loss_func = tf.keras.losses.BinaryCrossentropy(from_logits=True) # 加from_logits适配无激活输出 def loss_func(true_labels, pred_labels): # 1. 计算二元交叉熵 binary_loss = binary_loss_func(true_labels[:, 0], pred_labels[:, 0]) # 2. 生成正样本掩码:1代表正样本,0代表负样本 pos_mask = tf.cast(~tf.math.is_nan(true_labels[:, 1]), tf.float32) num_pos = tf.maximum(tf.reduce_sum(pos_mask), 1.0) # 避免batch无正样本时除以0 # 3. 填充标签中的NaN为0,避免NaN参与运算 true_filled = tf.where(tf.math.is_nan(true_labels), 0.0, true_labels) # 4. 计算坐标MSE损失,仅正样本生效 xy_loss = tf.reduce_sum(tf.square(true_filled[:, 1:3] - pred_labels[:, 1:3]) * pos_mask[:, None]) / num_pos # 5. 计算宽高MSE损失,仅正样本生效 wh_loss = tf.reduce_sum(tf.square(true_filled[:, 4:] - pred_labels[:, 4:]) * pos_mask[:, None]) / num_pos # 6. 计算角度损失,仅正样本生效 angle_diff = true_filled[:, 3] - pred_labels[:, 3] angle_loss = tf.reduce_sum(tf.abs(tf.atan2(tf.sin(angle_diff), tf.cos(angle_diff))) * pos_mask) / num_pos return binary_loss + xy_loss + wh_loss + angle_loss
2. 模型输出适配(可选但推荐)
如果要进一步提升稳定性,可以把最后一层的输出拆分,适配不同任务的数值范围:
import numpy as np from tensorflow.keras import Sequential from tensorflow.keras.layers import Dense # 把原来的Dense(6)改成如下结构 model.add(Flatten()) x = Dense(6)(model.output) # 按任务适配输出范围 output = tf.concat([ x[:, 0:1], # 目标存在概率的logits,不用激活 tf.sigmoid(x[:, 1:3]) * IMAGE_SIZE, # x/y为像素值时,用sigmoid限制在图像尺寸范围内 tf.tanh(x[:, 3:4]) * np.pi, # 角度限制在[-pi, pi]区间 tf.sigmoid(x[:, 4:]) * IMAGE_SIZE # 宽高限制在合理尺寸范围内 ], axis=1) model = tf.keras.Model(inputs=model.input, outputs=output)
3. 额外稳定训练的配置
- 编译模型时加入梯度裁剪,比如设置
clipvalue=1.0或者clipnorm=1.0,避免梯度爆炸。 - 训练前检查标签是否存在异常值,比如角度是否超出你定义的范围,坐标是否超出图像尺寸。
内容的提问来源于stack exchange,提问作者Varun
相关产品推荐
相关产品推荐

