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

关于tf.nn.softmax_cross_entropy_with_logits_v2运算原理的技术问询

我最近仔细拆解了TensorFlow里的tf.nn.softmax_cross_entropy_with_logits_v2(labels, logits)函数,发现它本质上是按三个核心步骤执行的,给大家详细拆解下:

核心操作步骤

  • 第一步:对输入的logits(也就是模型的原始输出y_hat)应用softmax归一化,将其转换为概率分布:
    y_hat_softmax = tf.nn.softmax(y_hat)
    
  • 第二步:结合真实标签y_true计算交叉熵损失:
    y_cross = y_true * tf.math.log(y_hat_softmax)
    
  • 第三步:对单个样本的所有类别维度求和,并取负值得到最终的样本损失:
    loss_per_sample = -tf.reduce_sum(y_cross, axis=1)
    

完整验证代码

下面的代码可以完美验证这个逻辑,手动计算的结果和直接调用tf.nn.softmax_cross_entropy_with_logits_v2的结果完全一致:

import tensorflow as tf
import numpy as np

# 定义真实标签(one-hot格式)
y_true = tf.convert_to_tensor(np.array([[0.0, 1.0, 0.0], [0.0, 0.0, 1.0]]), dtype=tf.float32)
# 定义模型输出的logits(未经过softmax的原始值)
logits = tf.convert_to_tensor(np.array([[1.0, 2.0, 0.5], [0.3, 0.1, 3.0]]), dtype=tf.float32)

# 手动模拟函数的三步操作
y_hat_softmax = tf.nn.softmax(logits)
y_cross = y_true * tf.math.log(y_hat_softmax)
manual_loss = -tf.reduce_sum(y_cross, axis=1)

# 直接调用官方函数计算损失
official_loss = tf.nn.softmax_cross_entropy_with_logits_v2(labels=y_true, logits=logits)

# 打印结果对比
print("手动计算的样本损失:", manual_loss.numpy())
print("官方函数计算的样本损失:", official_loss.numpy())

运行这段代码你会发现两个结果完全相同,这就证明了我们对函数内部逻辑的拆解是正确的。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 11:52:02