BERT二元分类模型评估:预测数组含义解析求助
BERT讽刺检测分类器:预测结果解读与标签转换
我基于BERT构建了一个讽刺检测二元分类器,预期输出1代表讽刺文本、0代表非讽刺文本,但评估阶段输出的predictions数组不是这类标签,无法解读结果。以下是相关代码和输出:
模型定义
from transformers import BertForSequenceClassification, AdamW, BertConfig # Load BertForSequenceClassification, the pretrained BERT model with a single # linear classification layer on top. model = BertForSequenceClassification.from_pretrained( "bert-base-uncased", # Use the 12-layer BERT model, with an uncased vocab. num_labels = 2, # The number of output labels--2 for binary classification. # You can increase this for multi-class tasks. output_attentions = False, # Whether the model returns attentions weights. output_hidden_states = False, # Whether the model returns all hidden-states. attention_probs_dropout_prob=0.25, hidden_dropout_prob=0.25 ) # Tell pytorch to run this model on the GPU. model.cuda()
评估代码
from sklearn.metrics import confusion_matrix import seaborn as sn import pandas as pd print('Predicting labels for {:,} test sentences...'.format(len(eval_input_ids))) # Put model in evaluation mode model.eval() predictions , true_labels = [], [] # iterate over test data for batch in eval_dataloader: batch = tuple(t.to(device) for t in batch) # Unpack the inputs from our dataloader b_input_ids, b_input_mask, b_labels = batch # Telling the model not to compute or store gradients, saving memory and # speeding up prediction with torch.no_grad(): # Forward pass, calculate logit predictions. result = model(b_input_ids, token_type_ids=None, attention_mask=b_input_mask, return_dict=True) logits = result.logits # Move logits and labels to CPU logits = logits.detach().cpu().numpy() label_ids = b_labels.to('cpu').numpy() # Store predictions and true labels predictions.append(logits) true_labels.append(label_ids) true_labels[1] predictions[1]
输出结果
array([0, 0, 1, 1, 0, 1, 0, 0, 1, 0, 0, 1, 1, 1, 0, 1, 1, 0, 0, 1, 0, 1, 0, 1, 1, 0, 0, 0, 0, 1, 1, 1]) <-- true_labels[1] array([[ 2.9316974 , -2.855342 ], [ 3.4540875 , -3.3177233 ], [ 2.7424026 , -2.6472614 ], [-3.4326897 , 3.330751 ], [ 3.7238903 , -3.7757814 ], [-3.208891 , 3.175109 ], [ 3.0500402 , -2.8103237 ], [ 3.8333693 , -3.9073608 ], [-3.2779126 , 3.231213 ], [ 1.484127 , -1.2610332 ], [ 3.686339 , -3.7582958 ], [-2.1883147 , 2.205132 ], [-3.274582 , 3.2254982 ], [-1.606854 , 1.6213335 ], [ 3.7080388 , -3.6854186 ], [-2.351147 , 2.365543 ], [-3.7317555 , 3.4833894 ], [ 3.2413306 , -3.2116275 ], [ 3.7413723 , -3.7767386 ], [-3.6293464 , 3.4446163 ], [ 3.7779078 , -3.9025154 ], [-3.5576923 , 3.403335 ], [ 3.6226897 , -3.6370063 ], [-3.7081888 , 3.4720154 ], [ 1.1533121 , -0.8105195 ], [ 1.0573612 , -0.69238156], [ 3.4189024 , -3.4764926 ], [-0.13847755, 0.450572 ], [ 3.7248163 , -3.7781181 ], [-3.2015219 , 3.1719215 ], [-2.1409311 , 2.1202204 ], [-3.470693 , 3.358798 ]], dtype=float32) <-- predictions[1]
问题原因与解决方法
你看到的predictions数组是模型输出的logits(未归一化的原始得分),不是最终的0/1分类标签。BertForSequenceClassification在二元分类任务中,会为每个样本输出两个维度的logit,分别对应标签0和标签1的得分。
要得到预期的0/1标签,需要对logits做两步处理:
- 用Softmax函数将logits转换为概率(两个维度的概率和为1)
- 取概率最大的维度索引,即为对应的分类标签
修改后的评估代码
替换原评估代码中处理logits的部分,完整代码如下:
from sklearn.metrics import confusion_matrix import seaborn as sn import pandas as pd import torch print('Predicting labels for {:,} test sentences...'.format(len(eval_input_ids))) model.eval() predictions , true_labels = [], [] for batch in eval_dataloader: batch = tuple(t.to(device) for t in batch) b_input_ids, b_input_mask, b_labels = batch with torch.no_grad(): result = model(b_input_ids, token_type_ids=None, attention_mask=b_input_mask, return_dict=True) logits = result.logits # 对logits应用Softmax得到概率分布 probs = torch.nn.functional.softmax(logits, dim=1) # 取概率最大的索引作为预测标签 pred_labels = torch.argmax(probs, dim=1) # 转移到CPU并转换为numpy数组 pred_labels = pred_labels.detach().cpu().numpy() label_ids = b_labels.to('cpu').numpy() predictions.append(pred_labels) true_labels.append(label_ids) # 查看转换后的标签结果 print(true_labels[1]) print(predictions[1])
验证示例
以你输出的第一个样本为例:
- logit为
[2.9316974, -2.855342],经过Softmax后,第一个维度的概率远大于第二个,对应标签0,和true_labels[1]的第一个值一致,符合预期。
内容的提问来源于stack exchange,提问作者szhang04
相关产品推荐
相关产品推荐

