【问题标题】:How to understand model loss output and dice coef如何理解模型损失输出和骰子系数
【发布时间】:2020-10-27 21:40:09
【问题描述】:

我正在使用这个:

Python version: 3.7.7 (default, May  6 2020, 11:45:54) [MSC v.1916 64 bit (AMD64)]
TensorFlow version: 2.1.0
Eager execution: True

使用此 U-Net 模型:

inputs = Input(shape=img_shape)

    conv1 = Conv2D(64, (5, 5), activation='relu', padding='same', data_format="channels_last", name='conv1_1')(inputs)
    conv1 = Conv2D(64, (5, 5), activation='relu', padding='same', data_format="channels_last", name='conv1_2')(conv1)
    pool1 = MaxPooling2D(pool_size=(2, 2), data_format="channels_last", name='pool1')(conv1)
    conv2 = Conv2D(96, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv2_1')(pool1)
    conv2 = Conv2D(96, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv2_2')(conv2)
    pool2 = MaxPooling2D(pool_size=(2, 2), data_format="channels_last", name='pool2')(conv2)

    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv3_1')(pool2)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv3_2')(conv3)
    pool3 = MaxPooling2D(pool_size=(2, 2), data_format="channels_last", name='pool3')(conv3)

    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv4_1')(pool3)
    conv4 = Conv2D(256, (4, 4), activation='relu', padding='same', data_format="channels_last", name='conv4_2')(conv4)
    pool4 = MaxPooling2D(pool_size=(2, 2), data_format="channels_last", name='pool4')(conv4)

    conv5 = Conv2D(512, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv5_1')(pool4)
    conv5 = Conv2D(512, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv5_2')(conv5)

    up_conv5 = UpSampling2D(size=(2, 2), data_format="channels_last", name='up_conv5')(conv5)
    ch, cw = get_crop_shape(conv4, up_conv5)
    crop_conv4 = Cropping2D(cropping=(ch, cw), data_format="channels_last", name='crop_conv4')(conv4)
    up6 = concatenate([up_conv5, crop_conv4])
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv6_1')(up6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv6_2')(conv6)

    up_conv6 = UpSampling2D(size=(2, 2), data_format="channels_last", name='up_conv6')(conv6)
    ch, cw = get_crop_shape(conv3, up_conv6)
    crop_conv3 = Cropping2D(cropping=(ch, cw), data_format="channels_last", name='crop_conv3')(conv3)
    up7 = concatenate([up_conv6, crop_conv3])
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv7_1')(up7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv7_2')(conv7)

    up_conv7 = UpSampling2D(size=(2, 2), data_format="channels_last", name='up_conv7')(conv7)
    ch, cw = get_crop_shape(conv2, up_conv7)
    crop_conv2 = Cropping2D(cropping=(ch, cw), data_format="channels_last", name='crop_conv2')(conv2)
    up8 = concatenate([up_conv7, crop_conv2])
    conv8 = Conv2D(96, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv8_1')(up8)
    conv8 = Conv2D(96, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv8_2')(conv8)

    up_conv8 = UpSampling2D(size=(2, 2), data_format="channels_last", name='up_conv8')(conv8)
    ch, cw = get_crop_shape(conv1, up_conv8)
    crop_conv1 = Cropping2D(cropping=(ch, cw), data_format="channels_last", name='crop_conv1')(conv1)
    up9 = concatenate([up_conv8, crop_conv1])
    conv9 = Conv2D(64, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv9_1')(up9)
    conv9 = Conv2D(64, (3, 3), activation='relu', padding='same', data_format="channels_last", name='conv9_2')(conv9)

    ch, cw = get_crop_shape(inputs, conv9)
    conv9 = ZeroPadding2D(padding=(ch, cw), data_format="channels_last", name='conv9_3')(conv9)
    conv10 = Conv2D(1, (1, 1), activation='sigmoid', data_format="channels_last", name='conv10_1')(conv9)
    model = Model(inputs=inputs, outputs=conv10)

还有这个功能:

def dice_coef(y_true, y_pred):
    y_true_f = K.flatten(y_true)
    y_pred_f = K.flatten(y_pred)
    intersection = K.sum(y_true_f * y_pred_f)
    return (2.0 * intersection + 1.0) / (K.sum(y_true_f) + K.sum(y_pred_f) + 1.0)

def dice_coef_loss(y_true, y_pred):
    return 1-dice_coef(y_true, y_pred)

要编译我所做的模型:

model.compile(tf.keras.optimizers.Adam(lr=(1e-4) * 2), loss=dice_coef_loss, metrics=[dice_coef])

我在训练时得到这个输出:

Epoch 1/2
5/5 [==============================] - 8s 2s/sample - loss: 1.0000 - dice_coef: 4.5962e-05 - val_loss: 0.9929 - val_dice_coef: 0.0071
Epoch 2/2
5/5 [==============================] - 5s 977ms/sample - loss: 0.9703 - dice_coef: 0.0297 - val_loss: 0.9939 - val_dice_coef: 0.0061
Train on 5 samples, validate on 5 samples

我认为这个想法是让损失接近于零,但我不明白我得到的1.000(也许这是我能得到的最糟糕的损失值)。但我不明白 dice_coef 的值。

dice_coef 值是什么意思?

【问题讨论】:

    标签: python tensorflow keras


    【解决方案1】:

    Dice loss 是一种损失函数,可以防止普通交叉熵损失中存在的一些限制。

    交叉熵的限制:

    当使用交叉熵损失时,标签的统计分布对训练准确度起着重要作用。 标签分布越不平衡,训练就越困难。 虽然加权交叉熵损失可以缓解困难,但改善并不显着,也没有解决交叉熵损失的内在问题。在交叉熵损失中,损失被计算为每像素损失的平均值,每像素损失是离散计算的,不知道其相邻像素是否为边界。结果,交叉熵损失只考虑了微观意义上的损失,没有考虑全局,这对于图像级别的预测是不够的。

    骰子损失

    Dice Coef 函数可以描述为:

    这显然是您的函数dice_coef(y_true, y_pred) 正在计算的内容。更多关于Sørensen–Dice coefficient

    在上面的等式中,p_i 和 g_i 分别是预测和地面实况的对应像素值对。在边界检测场景中,它们的值为 0 或 1,表示像素是边界(值为 1)还是不边界(值为 0)。分母是预测和ground truth的总边界像素的总和,分子是正确预测的边界像素的总和,因为只有当pi和gi时总和才会增加匹配(均为值 1)。

    分母考虑全局范围内边界像素的总数,而分子考虑局部范围内两个集合之间的重叠。 因此,Dice loss 同时考虑了局部和全局的损失信息,这对于高精度至关重要。

    关于您的训练,由于您的损失值在整个训练过程中减少,您不必太担心,尝试增加 epoch 的数量并在网络通过模型时对其进行分析。

    骰子损失只是1 - dice coef。这是您的函数正在计算的内容。

    【讨论】:

    • 感谢您的回答,但是 loss 和 dice_coef 值是什么意思?损失值 1.0 比 0.0 最差?而且,如果 loss 为 1.0,为什么 dice_coef 非常小?谢谢。
    • 损失基本上是你离你的基本事实有多远。基本事实基本上是你的标签。随着模型的训练,它会学习特征并绘制输入数据和标签之间的关系。如此高的损失并不好,但是您不要期望任何模型一开始就具有低损失,因为网络仍然没有学到那么多。 dice_coef 是我在帖子中描述的,骰子损失定义为 1-dice_coef,这就是你的函数 dice_coef_loss 正在做的事情。如果你明白了,请接受答案
    猜你喜欢
    • 2022-07-25
    • 1970-01-01
    • 1970-01-01
    • 2021-03-08
    • 2019-01-29
    • 2021-05-07
    • 1970-01-01
    • 2021-03-15
    • 2022-09-24
    相关资源
    最近更新 更多