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

TensorFlow中SparseCategoricalCrossEntropy的from_logits参数未按预期生效

问题描述

经过调研,对logit相关参数的初始认知如下:

  • 当设置from_logits=True时,模型输出未做归一化处理(不属于概率分布)
  • 当设置from_logits=False时,预期输出会经softmax函数完成归一化,得到类别的概率分布
    但实际运行结果和上述认知不符,需要定位根本原因。
复现过程

实验配置1:from_logits=True

实现代码

(img_train, label_train), (img_test, label_test) = tf.keras.datasets.fashion_mnist.load_data()
train_ds = tf.data.Dataset.from_tensor_slices((img_train/255.0, label_train)).batch(32)
test_ds = tf.data.Dataset.from_tensor_slices((img_test/255.0, label_test)).batch(32)

inputs = tf.keras.Input(shape=(28, 28), batch_size=32)
flatten_layer = tf.keras.layers.Flatten()(inputs)
dense = tf.keras.layers.Dense(units=512, activation='relu')(flatten_layer)
outputs = tf.keras.layers.Dense(units=10)(dense)

model = tf.keras.Model(inputs=inputs, outputs=outputs)
model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.01),
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=[tf.keras.metrics.SparseCategoricalAccuracy()])
history = model.fit(train_ds, validation_data=test_ds, epochs=10)

metric = tf.keras.metrics.SparseCategoricalAccuracy()
for x, y in test_ds:
    logits = model(x)
    metric.update_state(y, logits)
metric.result()

输出结果

打印测试集最后批次的首个样本输出:

<tf.Tensor: shape=(10,), dtype=float32, numpy=
array([  1.3062842 ,   2.253938  ,  -5.295599  ,   7.0740013 ,
       -17.184162  , -19.801863  ,   0.29550672, -81.3132    ,
        -9.149338  , -46.353527  ], dtype=float32)>

该结果符合原始logits的特征:值分布在任意实数区间,和不为1,不属于概率分布。


实验配置2:from_logits=False

实现代码

logit_false_model = tf.keras.Model(inputs=inputs, outputs=outputs)
logit_false_model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.01),
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False),
              metrics=[tf.keras.metrics.SparseCategoricalAccuracy()])
metric2 = tf.keras.metrics.SparseCategoricalAccuracy()
for x, y in test_ds:
    false_logits = logit_false_model(x)
    metric2.update_state(y, false_logits)
metric2.result()

输出结果

打印测试集最后批次的首个样本输出:

<tf.Tensor: shape=(10,), dtype=float32, numpy=
array([-125.248  , -211.52843, -243.62004, -230.45828, -336.3651 ,
       -369.41864, -177.14006, -871.6252 , -401.76608, -529.4581 ],
      dtype=float32)>

该结果完全不符合概率分布的特征:所有值为负数,不在0-1区间内。

根本原因

出现该现象是三个认知误区叠加导致的:

  1. from_logits是损失函数的内部参数,不会修改模型结构:这个参数的作用仅仅是告诉交叉熵损失函数,传入的预测值是不是已经经过softmax归一化的概率值——如果设为True,损失内部会先对输入做softmax再算交叉熵;如果设为False,损失会直接把输入当概率算交叉熵。它全程不会给模型的输出层自动加softmax激活,两次实验的模型输出层都是无激活的Dense(units=10),天生就输出原始logits,不可能自动产出概率。
  2. 第二次实验没有完成有效训练:代码中只是基于已有的输入输出张量重新创建了模型、修改了compile参数,没有调用fit()执行训练流程,模型权重本身就是第一次实验里适配from_logits=True训练出来的原始logits权重,没有任何适配新损失的调整。
  3. 输入和损失预期不匹配导致数值异常:就算给第二个模型跑了训练,只要输出层不加softmax,传给from_logits=False损失的就是值域任意的原始logits,交叉熵计算时会对输入值做对数运算,负数输入会直接触发数值不稳定,导致梯度爆炸,权重最终会崩成看到的全是极大负数的异常状态,根本不可能收敛出合法概率。

正确获取概率输出有两种标准实现:

  • 方案1:输出层加softmax激活,即outputs = tf.keras.layers.Dense(units=10, activation='softmax')(dense),损失固定用from_logits=False,模型推理直接输出概率。
  • 方案2:保持输出层无激活输出原始logits,损失用from_logits=True(该方案数值稳定性更好,为官方推荐写法),推理时手动调用tf.nn.softmax(model(x))得到概率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 11:09:23