使用归一化二元交叉熵损失的模型无法收敛问题排查
问题背景
我参考论文《Normalized Loss Functions for Deep Learning with Noisy Labels》,为CTR预测的二分类任务实现归一化二元交叉熵损失,目标公式为:
$$\mathcal{L}_{\text{NormBCE}} = -\frac{y\log p + (1-y)\log(1-p)}{H_0}$$
其中$H_0$是基准交叉熵(通常为随机猜测标签时的交叉熵,比如均匀分布下为$\log 2$)。
我基于TensorFlow实现了该损失函数,并将其用于训练一个多门混合专家(MMoE)架构的6任务二分类CTR预测模型,最终损失为6个任务的损失之和。但模型训练时损失始终不下降,ROC-AUC维持在0.49-0.5左右;而移除损失中的分母(即退化为普通二元交叉熵)后,模型训练完全正常。
实现代码
import tensorflow as tf from keras.utils import losses_utils class NormalizedBinaryCrossentropy(tf.keras.losses.Loss): def __init__( self, from_logits=False, label_smoothing=0.0, axis=-1, reduction=tf.keras.losses.Reduction.NONE, name="normalized_binary_crossentropy", **kwargs ): super().__init__( reduction=reduction, name=name ) self.from_logits = from_logits self._epsilon = tf.keras.backend.epsilon() def call(self, target, logits): if tf.is_tensor(logits) and tf.is_tensor(target): logits, target = losses_utils.squeeze_or_expand_dimensions( logits, target ) logits = tf.convert_to_tensor(logits) target = tf.cast(target, logits.dtype) if self.from_logits: logits = tf.math.sigmoid(logits) logits = tf.clip_by_value(logits, self._epsilon, 1.0 - self._epsilon) numer = target * tf.math.log(logits) + (1 - target) * tf.math.log(1 - logits) denom = - (tf.math.log(logits) + tf.math.log(1 - logits)) return - numer / denom def get_config(self): config = super().get_config() config.update({"from_logits": self._from_logits}) return config
示例验证代码
# Example Usage import numpy as np labels = np.array([[0], [1], [0], [0], [0]]).astype(np.int64) logits = np.array([[-1.024], [2.506], [1.43], [0.004], [-2.0]]).astype(np.float64) tf_nce = NormalizedBinaryCrossentropy( reduction=tf.keras.losses.Reduction.NONE, from_logits=True ) tf_nce(labels, logits) # <tf.Tensor: shape=(5, 1), dtype=float64, numpy= # array([[0.18737159], # [0.02945536], # [0.88459308], # [0.50144269], # [0.05631594]])>
我手动检查了极端情况,损失未出现NaN或0值,但模型始终无法收敛,请问问题出在哪里?
问题排查与修正
核心错误:对归一化分母的误解
你实现的分母- (tf.math.log(logits) + tf.math.log(1 - logits))是随模型输出动态变化的项,这完全不符合论文中归一化损失的定义:
- 论文中的分母$H_0$是固定的基准交叉熵,代表"随机猜测标签"时的交叉熵(比如二分类均匀分布下$H_0 = \log 2 ≈ 0.693$;如果数据集正样本比例为$\pi$,则$H_0 = -\pi\log\pi - (1-\pi)\log(1-\pi)$)。
- 你当前的分母实际是$-\log(p(1-p))$,当模型输出$p$接近0或1时,这个值会趋近于无穷大,导致损失值趋近于0,梯度也趋近于0,模型无法获得有效的更新信号;当$p$接近0.5时,分母值较小,但此时损失的梯度也远弱于普通BCE,不足以驱动模型收敛。
修正后的实现代码
以下是符合论文定义的归一化二元交叉熵实现,支持基于数据集先验或固定均匀分布的基准熵:
import tensorflow as tf from keras.utils import losses_utils class NormalizedBinaryCrossentropy(tf.keras.losses.Loss): def __init__( self, from_logits=False, label_smoothing=0.0, axis=-1, reduction=tf.keras.losses.Reduction.NONE, name="normalized_binary_crossentropy", baseline_entropy=None, pos_ratio=None, **kwargs ): super().__init__( reduction=reduction, name=name ) self.from_logits = from_logits self._epsilon = tf.keras.backend.epsilon() # 计算基准熵:优先使用传入的固定值,其次用数据集正样本比例计算,默认用均匀分布熵 if baseline_entropy is not None: self.baseline_entropy = tf.cast(baseline_entropy, tf.float32) elif pos_ratio is not None: pos_ratio = tf.clip_by_value(pos_ratio, self._epsilon, 1.0 - self._epsilon) self.baseline_entropy = -pos_ratio * tf.math.log(pos_ratio) - (1 - pos_ratio) * tf.math.log(1 - pos_ratio) else: self.baseline_entropy = tf.math.log(2.0) def call(self, target, logits): if tf.is_tensor(logits) and tf.is_tensor(target): logits, target = losses_utils.squeeze_or_expand_dimensions( logits, target ) logits = tf.convert_to_tensor(logits) target = tf.cast(target, logits.dtype) if self.from_logits: logits = tf.math.sigmoid(logits) logits = tf.clip_by_value(logits, self._epsilon, 1.0 - self._epsilon) # 计算普通二元交叉熵 bce = target * tf.math.log(logits) + (1 - target) * tf.math.log(1 - logits) # 除以固定基准熵,取负得到归一化损失 return -bce / self.baseline_entropy def get_config(self): config = super().get_config() config.update({ "from_logits": self.from_logits, "baseline_entropy": self.baseline_entropy.numpy() if hasattr(self.baseline_entropy, 'numpy') else self.baseline_entropy, "pos_ratio": None }) return config
使用建议
- 统计每个CTR任务的正样本比例(点击/曝光比),初始化损失时传入
pos_ratio参数,这样基准熵更贴合数据集分布。 - 如果没有数据集先验信息,可以直接使用默认的均匀分布基准熵
baseline_entropy=tf.math.log(2.0)。
内容的提问来源于stack exchange,提问作者Jatin Mandav
相关产品推荐
相关产品推荐

