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区间内。
根本原因
出现该现象是三个认知误区叠加导致的:
from_logits是损失函数的内部参数,不会修改模型结构:这个参数的作用仅仅是告诉交叉熵损失函数,传入的预测值是不是已经经过softmax归一化的概率值——如果设为True,损失内部会先对输入做softmax再算交叉熵;如果设为False,损失会直接把输入当概率算交叉熵。它全程不会给模型的输出层自动加softmax激活,两次实验的模型输出层都是无激活的Dense(units=10),天生就输出原始logits,不可能自动产出概率。- 第二次实验没有完成有效训练:代码中只是基于已有的输入输出张量重新创建了模型、修改了compile参数,没有调用
fit()执行训练流程,模型权重本身就是第一次实验里适配from_logits=True训练出来的原始logits权重,没有任何适配新损失的调整。 - 输入和损失预期不匹配导致数值异常:就算给第二个模型跑了训练,只要输出层不加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
相关产品推荐
相关产品推荐

