请求将论文中的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
相关产品推荐
相关产品推荐

