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

TensorFlow中Unet模型验证集Dice系数内外计算结果不一致问题

Unet变体训练异常排查:验证Dice指标异常与外部验证结果不符

我在TensorFlow中训练Unet变体网络时遇到异常:训练阶段TensorBoard显示训练损失高于验证损失,但验证集Dice指标仅0.25-0.30,可模型可视化输出表现并不差;重新加载模型做外部验证时,Dice系数却能超过0.9。怀疑问题出在损失函数和自定义Dice指标的实现上,相关代码如下:

损失类

class sce_dsc(losses.Loss):
    def __init__(self, scale_sce=1.0, scale_dsc=1.0, sample_weight = None, epsilon=0.01, name=None):
        super(sce_dsc, self).__init__()
        self.sce = losses.SparseCategoricalCrossentropy(from_logits=False) #while the last layer activation is sigmoid, logits needs to be false
        self.epsilon = epsilon
        self.scale_a = scale_sce
        self.scale_b = scale_dsc
        self.cls = 1
        self.weights = sample_weight

    def dsc(self, y_true, y_pred, sample_weight = None):
        
        true = tf.cast(y_true[..., 0] == self.cls, tf.int64)
        pred = tf.nn.softmax(y_pred, axis=-1)[..., self.cls]
        if self.weights is not None:
            #true = true * (sample_weight[...])
            true = true & (sample_weight[...] !=0)
            #pred = pred * (sample_weight[...])
            pred = pred & (sample_weight[...] !=0)
        A = tf.math.reduce_sum(tf.cast(true, tf.float32) * tf.cast(pred,tf.float32)) * 2
        B = tf.cast(tf.math.reduce_sum(true), tf.float32) + tf.cast(tf.math.reduce_sum(pred),tf.float32) + self.epsilon
        
        return (1.0 - A/B) 

    def call(self, y_true, y_pred):
        sce_loss = self.sce(y_true=y_true, y_pred=y_pred, sample_weight=self.weights) * self.scale_a
        dsc_loss = self.dsc(y_true=y_true, y_pred=y_pred, sample_weight=self.weights) * self.scale_b
        loss = tf.cast(sce_loss, tf.float32) + tf.cast(dsc_loss,tf.float32)     
        #self.add_loss(loss)
        return loss

自定义Dice指标类

class custom_dice(keras.metrics.Metric):
    
       def __init__(self, name = "dsc", **kwargs):
           super(custom_dice,self).__init__(**kwargs)
           self.dice = self.add_weight(name = 'dice_coef', initializer = 'zeros')
        
       def update_state(self, y_true,y_pred, sample_weight = None):
           true = tf.cast(y_true[...,0] == 1, tf.int64)
           pred = tf.math.argmax(y_pred == 1 , axis=-1) 
           if sample_weight is not None:
            true = true * (sample_weight[...])
            pred = pred * (sample_weight[...])
   
           A = tf.math.count_nonzero(true & pred) * 2
           B = tf.math.count_nonzero(true) + tf.math.count_nonzero(pred)
           value = tf.math.divide_no_nan(tf.cast(A, tf.float32),tf.cast(B, tf.float32))
           self.dice.assign(value)
        
       def result(self):
           return self.dice
    
       def reset_state(self):
           self.dice.assign(0.0)

外部验证Dice函数

def dsc(y_true, y_pred, sample_weight=None, c = 1):
       print(y_true.shape, y_pred.shape)
       true = tf.cast(y_true[...,0] == 1, tf.int64)
       pred = tf.math.argmax(y_pred== c , axis=-1) 
       print(true.shape,pred.shape)
       if sample_weight is not None:
           true = true * (sample_weight[...])
           pred = pred * (sample_weight[...])
   
       A = tf.math.count_nonzero(true & pred) * 2
       B = tf.math.count_nonzero(true) + tf.math.count_nonzero(pred)
       return A / B 

问题根源分析

  1. 自定义Dice指标的预测值计算逻辑错误
    pred = tf.math.argmax(y_pred == 1 , axis=-1) 完全不符合预期:y_pred == 1会生成布尔张量,对布尔张量执行argmax毫无意义,无法得到正确的类别预测结果。外部验证函数也存在同样的错误,但可能实际运行时被手动修正,才得到了正确的Dice结果。

  2. 指标更新方式错误
    self.dice.assign(value) 会直接覆盖当前指标值,导致最终结果仅保留最后一个batch的计算值,无法反映整个验证集的平均Dice水平。

  3. 损失与指标的计算逻辑不一致
    损失函数中用softmax后的连续概率值计算Dice损失,而自定义指标试图用离散类别计算,但逻辑错误,导致指标与模型实际性能完全脱节。


修正后的代码

修正后的自定义Dice指标类

class custom_dice(keras.metrics.Metric):
    def __init__(self, name="dsc", cls=1, **kwargs):
        super(custom_dice, self).__init__(name=name, **kwargs)
        self.cls = cls
        # 累加交叉区域的2倍值
        self.total_intersection = self.add_weight(name='intersection', initializer='zeros')
        # 累加真实目标区域和预测区域的总像素数
        self.total_true = self.add_weight(name='total_true', initializer='zeros')
        self.total_pred = self.add_weight(name='total_pred', initializer='zeros')

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 转换真实标签为目标类别的布尔张量
        true = tf.cast(y_true[..., 0] == self.cls, tf.float32)
        # 对预测值做softmax后取目标类别概率,用0.5阈值生成预测掩码
        pred_probs = tf.nn.softmax(y_pred, axis=-1)[..., self.cls]
        pred = tf.cast(pred_probs > 0.5, tf.float32)

        if sample_weight is not None:
            sample_weight = tf.cast(sample_weight, tf.float32)
            true = true * sample_weight
            pred = pred * sample_weight

        # 累加当前batch的计算值
        intersection = tf.reduce_sum(true * pred)
        self.total_intersection.assign_add(intersection * 2)
        self.total_true.assign_add(tf.reduce_sum(true))
        self.total_pred.assign_add(tf.reduce_sum(pred))

    def result(self):
        # 计算整体Dice,避免除以0
        return tf.math.divide_no_nan(self.total_intersection, self.total_true + self.total_pred)

    def reset_state(self):
        self.total_intersection.assign(0.0)
        self.total_true.assign(0.0)
        self.total_pred.assign(0.0)

修正后的外部验证Dice函数

def dsc(y_true, y_pred, sample_weight=None, c=1):
    print(y_true.shape, y_pred.shape)
    true = tf.cast(y_true[..., 0] == c, tf.float32)
    # 正确生成预测掩码:softmax取概率后阈值化
    pred_probs = tf.nn.softmax(y_pred, axis=-1)[..., c]
    pred = tf.cast(pred_probs > 0.5, tf.float32)
    print(true.shape, pred.shape)

    if sample_weight is not None:
        sample_weight = tf.cast(sample_weight, tf.float32)
        true = true * sample_weight
        pred = pred * sample_weight

    intersection = tf.reduce_sum(true * pred) * 2
    union = tf.reduce_sum(true) + tf.reduce_sum(pred)
    return tf.math.divide_no_nan(intersection, union)

补充说明

损失函数中存在一个小矛盾:注释提到最后一层用sigmoid激活,但使用了SparseCategoricalCrossentropy(适合多分类+softmax激活)。如果是二分类任务,建议将最后一层改为softmax(输出2个通道),或改用BinaryCrossentropy损失,避免逻辑不匹配。

内容的提问来源于stack exchange,提问作者zhilothebest

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 17:05:10