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

如何在TensorFlow/Keras中实现语义分割用归一化二元交叉熵损失?

TensorFlow/Keras实现语义分割用归一化交叉熵损失函数

根据你提供的论文定义和语义分割场景的PyTorch参考实现,以下是适配TensorFlow/Keras的归一化交叉熵(NCE)损失函数实现,支持忽略指定标签,适用于U-Net、FCN等语义分割任务:

import tensorflow as tf

class NormalizedCrossEntropy(tf.keras.losses.Loss):
    def __init__(self, ignore_label=-1, num_classes=None, name="normalized_cross_entropy"):
        super().__init__(name=name)
        self.ignore_label = ignore_label
        self.num_classes = num_classes

    def call(self, y_true, y_pred):
        # 语义分割任务中y_pred形状通常为(batch, height, width, num_classes)
        # 计算log softmax,axis=-1对应channels last维度
        logsoftmax = tf.nn.log_softmax(y_pred, axis=-1)
        
        # 生成one-hot编码,同时处理ignore标签
        y_true_int = tf.cast(y_true, tf.int32)
        # 创建mask:有效区域为1,ignore区域为0
        mask = tf.cast(tf.not_equal(y_true_int, self.ignore_label), tf.float32)
        # 将ignore标签位置的目标临时设为0,不影响后续计算(mask会过滤)
        y_true_clean = tf.where(tf.equal(y_true_int, self.ignore_label), 0, y_true_int)
        # 生成与y_pred维度匹配的one-hot编码
        ohot = tf.one_hot(y_true_clean, depth=self.num_classes, dtype=tf.float32)
        
        # 计算分子:-sum(one_hot * logsoftmax) over channels
        numerator = -1 * tf.reduce_sum(ohot * logsoftmax, axis=-1)
        # 计算分母:-sum(logsoftmax) over channels
        denominator = -1 * tf.reduce_sum(logsoftmax, axis=-1)
        
        # 计算NCE,添加epsilon避免除以0
        nce = numerator / (denominator + tf.keras.backend.epsilon())
        # 过滤ignore标签区域的损失
        nce = nce * mask
        
        # 计算有效区域的平均损失
        total_loss = tf.reduce_sum(nce)
        valid_pixels = tf.reduce_sum(mask)
        # 处理所有像素都是ignore标签的极端情况
        return tf.cond(valid_pixels > 0, lambda: total_loss / valid_pixels, lambda: 0.0)

使用示例

在U-Net模型编译时直接调用:

# 假设分割任务有10类,忽略标签为255
loss_fn = NormalizedCrossEntropy(ignore_label=255, num_classes=10)
model.compile(optimizer="adam", loss=loss_fn)

实现说明

  1. 适配TensorFlow的channels last数据格式(语义分割任务默认格式),对应PyTorch的channels first做了维度逻辑调整
  2. 完整处理ignore_label逻辑,通过mask过滤无效像素,避免其参与损失计算
  3. 添加epsilon防止分母为0的异常情况
  4. 极端场景下(所有像素都是ignore标签)返回0损失,避免训练报错

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 15:27:12