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

请求将论文中的Normalized Cross Entropy Loss转为Tensorflow/Keras实现

TensorFlow版Normalized Cross Entropy Loss实现

下面是对应PyTorch版本转换过来的TensorFlow实现,完全对齐原论文的公式逻辑:

import tensorflow as tf

class NormalizedCrossEntropy(tf.keras.losses.Loss):
    def __init__(self, num_classes, scale=1.0, name="normalized_cross_entropy"):
        super().__init__(name=name)
        self.num_classes = num_classes
        self.scale = scale

    def call(self, y_true, y_pred):
        # 对模型输出计算log softmax(类别维度为axis=1)
        log_softmax_pred = tf.nn.log_softmax(y_pred, axis=1)
        # 将类别索引标签转换为one-hot编码
        label_one_hot = tf.cast(tf.one_hot(tf.cast(y_true, tf.int32), depth=self.num_classes), tf.float32)
        # 计算分子:对应类别的负log概率
        numerator = -tf.reduce_sum(label_one_hot * log_softmax_pred, axis=1)
        # 计算分母:所有类别log概率的负和
        denominator = -tf.reduce_sum(log_softmax_pred, axis=1)
        # 计算每个样本的NCE,添加epsilon避免除以0
        nce = numerator / (denominator + tf.keras.backend.epsilon())
        # 返回缩放后的平均损失
        return self.scale * tf.reduce_mean(nce)

关键说明:

  • 完全对齐原PyTorch代码的逻辑,严格遵循论文中的数学公式
  • 继承tf.keras.losses.Loss类,可直接在Keras模型的compile方法中使用
  • 加入tf.keras.backend.epsilon()避免分母为0的异常情况,增强鲁棒性
  • 如果你的标签已经是one-hot格式,可移除tf.one_hot转换步骤,直接使用y_true

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 18:46:05