使用logits时Negative log likelihood损失的正确实现选择问题
负对数似然TensorFlow实现正确性判断
基础定义说明
多分类/序列标注任务中,负对数似然(NLL)的核心计算逻辑是:取真实标签对应模型输出的概率,求对数后取负,再按需求做序列维度求和、批次维度求平均。TensorFlow提供的tf.keras.losses.sparse_categorical_crossentropy接口在设置from_logits=True时,输出结果本身就是单个位置真实标签对应的负对数似然值。
两种实现对比
- 第一种实现为错误实现:
代码里多余引入了tf.reduce_logsumexp和两次负号变换,完全不符合NLL的计算逻辑,输出结果和真实NLL值偏差极大,不建议使用。
对应代码:negative_log_likelihood = tf.reduce_mean( -tf.reduce_logsumexp(-tf.keras.losses.sparse_categorical_crossentropy( targets, logits, from_logits=True), axis=1) - 第二种实现为正确实现(适配序列标注/序列级多分类场景):
当输入logits维度为(batch_size, seq_len, class_num)、标签维度为(batch_size, seq_len)时,sparse_categorical_crossentropy输出形状为(batch_size, seq_len),对axis=1求和可以得到每个样本整条序列的总NLL,再对批次维度求平均就是全局平均NLL,完全符合计算要求。
对应代码:negative_log_likelihood = tf.reduce_mean( tf.reduce_sum(sparse_categorical_crossentropy( targets, logits, from_logits=True), axis=1))
特殊场景适配
如果是普通单样本多分类任务(非序列场景),logits维度为(batch_size, class_num),标签维度为(batch_size,),此时不需要做序列维度的求和操作,直接使用以下代码即可得到正确的平均NLL:
negative_log_likelihood = tf.reduce_mean( tf.keras.losses.sparse_categorical_crossentropy(targets, logits, from_logits=True) )
内容的提问来源于stack exchange,提问作者wuannnn
相关产品推荐
相关产品推荐

