基于加权损失的Keras U-net虹膜分割训练异常求助
问题描述
尝试使用带加权损失函数的Keras U-Net实现人眼虹膜图像分割,训练过程中损失值变为NaN,精度始终保持不变;即使将输出层激活函数改为softmax,问题依然存在。
模型与训练代码
损失函数与模型定义
def my_loss(target, output): return - tf.reduce_sum(target * output, len(output.get_shape()) - 1) # Standard Unet model from blog post _epsilon = tf.convert_to_tensor(K.epsilon(), np.float32) def make_weighted_loss_unet(input_shape, n_classes): ip = L.Input(shape=input_shape) weight_ip = L.Input(shape=input_shape[:2] + (n_classes,)) conv1 = L.Conv2D(64, 3, activation='relu', padding='same', kernel_initializer='he_normal')(ip) conv1 = L.Conv2D(64, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv1) conv1 = L.Dropout(0.1)(conv1) mpool1 = L.MaxPool2D()(conv1) conv2 = L.Conv2D(128, 3, activation='relu', padding='same', kernel_initializer='he_normal')(mpool1) conv2 = L.Conv2D(128, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv2) conv2 = L.Dropout(0.2)(conv2) mpool2 = L.MaxPool2D()(conv2) conv3 = L.Conv2D(256, 3, activation='relu', padding='same', kernel_initializer='he_normal')(mpool2) conv3 = L.Conv2D(256, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv3) conv3 = L.Dropout(0.3)(conv3) mpool3 = L.MaxPool2D()(conv3) conv4 = L.Conv2D(512, 3, activation='relu', padding='same', kernel_initializer='he_normal')(mpool3) conv4 = L.Conv2D(512, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv4) conv4 = L.Dropout(0.4)(conv4) mpool4 = L.MaxPool2D()(conv4) conv5 = L.Conv2D(1024, 3, activation='relu', padding='same', kernel_initializer='he_normal')(mpool4) conv5 = L.Conv2D(1024, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv5) conv5 = L.Dropout(0.5)(conv5) up6 = L.Conv2DTranspose(512, 2, strides=2, kernel_initializer='he_normal', padding='same')(conv5) conv6 = L.Concatenate()([up6, conv4]) conv6 = L.Conv2D(512, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv6) conv6 = L.Conv2D(512, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv6) conv6 = L.Dropout(0.4)(conv6) up7 = L.Conv2DTranspose(256, 2, strides=2, kernel_initializer='he_normal', padding='same')(conv6) conv7 = L.Concatenate()([up7, conv3]) conv7 = L.Conv2D(256, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv7) conv7 = L.Conv2D(256, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv7) conv7 = L.Dropout(0.3)(conv7) up8 = L.Conv2DTranspose(128, 2, strides=2, kernel_initializer='he_normal', padding='same')(conv7) conv8 = L.Concatenate()([up8, conv2]) conv8 = L.Conv2D(128, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv8) conv8 = L.Conv2D(128, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv8) conv8 = L.Dropout(0.2)(conv8) up9 = L.Conv2DTranspose(64, 2, strides=2, kernel_initializer='he_normal', padding='same')(conv8) conv9 = L.Concatenate()([up9, conv1]) conv9 = L.Conv2D(64, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv9) conv9 = L.Conv2D(64, 3, activation='relu', padding='same', kernel_initializer='he_normal')(conv9) conv9 = L.Dropout(0.1)(conv9) c10 = L.Conv2D(n_classes, 1, activation='sigmoid', kernel_initializer='he_normal')(conv9) #Mimic crossentropy loss c11 = L.Lambda(lambda x: x / tf.reduce_sum(x, len(x.get_shape()) - 1, True))(c10) c11 = L.Lambda(lambda x: tf.clip_by_value(x, _epsilon, 1. - _epsilon))(c11) c11 = L.Lambda(lambda x: K.log(x))(c11) weighted_sm = L.multiply([c11, weight_ip]) model = Model(inputs=[ip, weight_ip], outputs=[weighted_sm]) return model
训练代码
model = make_weighted_loss_unet((256, 256, 3), 1) # shape of input, number of classes model.compile(optimizer='adam',loss=my_loss, metrics=['acc']) model.fit([X_train, wmap], y_train, validation_split=0.1, epochs=100)
数据说明
X_train: 输入图像列表,形状为(imgs, 256, 256, 3)wmap: 权重图列表,形状为(imgs, 256,256,1)y_train: 掩码列表,形状为(imgs,256,256,1)
解决思路与修正方案
1. 重构模型输出与损失计算逻辑
当前模型将损失的部分计算嵌入到网络层中,这种设计会干扰梯度的正常传播,且单分类场景下的归一化逻辑错误:
- 当
n_classes=1时,sigmoid输出已经是单类概率,无需执行x / tf.reduce_sum(x, ...)的归一化操作,该操作在单通道下会导致输出恒为1,进而引发log(1-epsilon)的异常计算。 - 正确的做法是让模型直接输出概率图,将加权损失的计算放在损失函数中,而非网络结构内。
修正后的模型定义(仅保留关键修改部分):
def make_weighted_loss_unet(input_shape, n_classes): ip = L.Input(shape=input_shape) weight_ip = L.Input(shape=input_shape[:2] + (n_classes,)) # 原有U-Net卷积、池化、上采样结构保持不变 # 直接输出sigmoid激活后的概率图 c10 = L.Conv2D(n_classes, 1, activation='sigmoid', kernel_initializer='he_normal')(conv9) model = Model(inputs=[ip, weight_ip], outputs=[c10]) return model
2. 实现稳定的加权二元交叉熵损失
针对单分类的虹膜分割场景,使用加权二元交叉熵,通过clip操作避免log输入趋近于0或1,同时融入像素级权重:
import tensorflow.keras.backend as K def weighted_binary_crossentropy(y_true, y_pred): # 获取权重输入 weight = K.in_train_phase(model.input[1], K.ones_like(y_pred)) epsilon = K.epsilon() # 裁剪输出,防止log(0)或log(1) y_pred = K.clip(y_pred, epsilon, 1. - epsilon) # 计算二元交叉熵 bce = -y_true * K.log(y_pred) - (1 - y_true) * K.log(1 - y_pred) # 应用权重并取均值 weighted_bce = bce * weight return K.mean(weighted_bce) # 模型编译时使用该损失 model = make_weighted_loss_unet((256, 256, 3), 1) model.compile(optimizer='adam', loss=weighted_binary_crossentropy, metrics=['acc'])
3. 验证数据与权重的有效性
- 检查
wmap:确保无NaN、无穷大值,权重非负,建议将权重归一化到[0, 10]区间,避免权重过大引发梯度爆炸。 - 检查
y_train:确认掩码取值为0/1的二元格式,非该范围会导致损失计算异常。 - 检查
X_train:必须做归一化处理(如除以255将像素值缩至[0,1]),未归一化的大数值会导致梯度异常。
4. 调整训练参数避免梯度爆炸
- 降低学习率:将Adam默认的0.001改为0.0001,减少梯度更新幅度。
- 添加梯度裁剪:限制梯度的范数,防止梯度值过大:
from tensorflow.keras.optimizers import Adam optimizer = Adam(learning_rate=0.0001, clipnorm=1.0) model.compile(optimizer=optimizer, loss=weighted_binary_crossentropy, metrics=['acc'])
内容的提问来源于stack exchange,提问作者Naven
相关产品推荐
相关产品推荐

