关于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
相关产品推荐
相关产品推荐

