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

基于加权损失的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 04:06:19