TensorFlow模型logit置信度远低于PyTorch同架构模型的原因排查
TensorFlow与PyTorch版BERT-tiny模型置信度差异解析
问题背景
使用相同数据集与prajjwal1/bert-tiny架构,分别在TensorFlow和PyTorch中训练文本分类模型,二者验证准确率均为90%。但TensorFlow模型的输出置信度极低:针对类别0的输入文本,其softmax后输出为[0.3928, 0.2365, 0.1854, 0.1854],最高值不超过0.5;而PyTorch模型对应softmax后输出为[0.8778, 0.0532, 0.0056, 0.0635],置信度符合预期。已排除softmax本身问题,以下是差异原因分析:
核心差异原因
1. 输出层与损失函数的匹配逻辑
- TensorFlow版本在最后一层
Dense直接添加了softmax激活,输出的是归一化后的概率值,且损失函数SparseCategoricalCrossentropy设置为from_logits=False,与输出形式匹配。 - PyTorch版本最后一层
Dense无激活函数,输出的是原始logit值,损失函数CrossEntropyLoss内部会自动对logit执行softmax计算损失。
你提到的TensorFlow"logit"实际是softmax后的概率,而PyTorch的输出是原始logit——这是概念混淆点,但即使排除这点,两者置信度差异的核心在于:
2. 全连接层初始化策略不同
- TensorFlow的
Dense层默认使用glorot_uniform初始化权重,该方式更适配sigmoid/tanh类激活函数的场景。 - PyTorch的
nn.Linear层默认使用kaiming_uniform初始化权重,更适配ReLU类激活函数场景(BERT池化输出经tanh,但后续全连接层无激活,kaiming初始化会让权重分布更易产生大的logit值)。
不同初始化策略导致训练过程中权重收敛的分布不同,最终使得PyTorch模型的logit幅度更大,经softmax后置信度更高。
3. 预训练权重转换的细微差异
TensorFlow版本通过from_pt=True从PyTorch权重转换加载预训练模型,核心权重虽一致,但自定义全连接层是各自初始化的,加上初始化策略差异,进一步放大了输出的置信度差距。
验证与修正建议
若要让TensorFlow模型输出与PyTorch对齐的高置信度结果,可修改输出层与损失函数:
# 修改TensorFlow模型输出层,去掉softmax激活 outputs = tf.keras.layers.Dense(4)(x) # 损失函数设置为from_logits=True,内部自动处理softmax loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
此时模型输出原始logit,经softmax后即可得到与PyTorch量级接近的置信度概率。
附用户提供的模型代码
TensorFlow模型代码
tokenizer = AutoTokenizer.from_pretrained('prajjwal1/bert-tiny', from_pt = True) encoder = TFAutoModel.from_pretrained('prajjwal1/bert-tiny', from_pt = True) # Define input layer with token and attention mask input_ids = tf.keras.layers.Input(shape=(None,), dtype=tf.int32, name="input_ids") attention_mask = tf.keras.layers.Input(shape=(None,), dtype=tf.int32, name="attention_mask") # Call the ALBERT model with the inputs pooler_output = encoder(input_ids, attention_mask=attention_mask)[1] # 1 is pooler output # Define a dense layer on top of the pooled output x = tf.keras.layers.Dense(units=params['fc_layer_size'])(pooler_output) x = tf.keras.layers.Dropout(params['dropout'])(x) outputs = tf.keras.layers.Dense(4, activation='softmax')(x) # Define a model with the inputs and dense layer model = tf.keras.Model(inputs=[input_ids, attention_mask], outputs=outputs) loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False) optimizer = tf.keras.optimizers.SGD(learning_rate=0.0008) metrics = [tf.metrics.SparseCategoricalAccuracy()] # Compile the model model.compile(optimizer='sgd', loss=loss, metrics=metrics)
PyTorch模型代码
tokenizer = AutoTokenizer.from_pretrained('prajjwal1/bert-tiny') encoder = AutoModel.from_pretrained('prajjwal1/bert-tiny') loss_fn = nn.CrossEntropyLoss() criterion = nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.0008) class EscalationClassifier(nn.Module): def __init__(self, encoder): super(EscalationClassifier, self).__init__() self.encoder = encoder self.fc1 = nn.Linear(128, 312) self.fc2 = nn.Linear(312, 4) self.dropout = nn.Dropout(0.2) def forward(self, input_ids, attention_mask): pooled_output = self.encoder(input_ids, attention_mask=attention_mask)[1]# [0] is last hidden state, 1 for pooler output # pdb.set_trace() x = self.fc1(pooled_output) x = self.dropout(x) x = self.fc2(x) return x model = EscalationClassifier(encoder)
内容的提问来源于stack exchange,提问作者ben.jamin2hard
相关产品推荐
相关产品推荐

