如何识别BERT模型误分类测试数据的参与者ID与文本段
识别BERT模型误分类的转录文本段及对应参与者ID
问题背景
我正在使用BERT对ASR生成的转录文本中提取的文本段进行分类,每位参与者对应多个文本段。数据存储在包含Participant_ID、Segment_Text和Diagnosis(即标签)列的数据框中。模型已训练完成,现需识别出模型误分类的文本段及对应的Participant_ID,以下为原始自定义Dataset类与模型评估类代码:
原始自定义Dataset类
tokenizer = BertTokenizer.from_pretrained('bert-base-cased') labels = {'HC':0, 'PK':1} class Dataset(torch.utils.data.Dataset): def __init__(self, df): self.participantIDs = df['Participant_ID'] self.labels = [labels[label] for label in df['Diagnosis']] self.texts = [tokenizer(text, padding='max_length', max_length = 512, truncation=True, return_tensors="pt") for text in df['Transcript_Segment']] def classes(self): return self.labels def __len__(self): return len(self.labels) def get_batch_labels(self, idx): # Fetch a batch of labels return np.array(self.labels[idx]) def get_batch_texts(self, idx): # Fetch a batch of inputs return self.texts[idx] # def get_batch_participant_ids(self, idx): # # Fetch a batch of inputs # return self.participantIDs[idx] def __getitem__(self, idx): batch_texts = self.get_batch_texts(idx) batch_y = self.get_batch_labels(idx) # batch_participant_ids = self.get_batch_participant_ids(idx) return batch_texts, batch_y
原始BERT模型评估类
def evaluate(model, test_data): incorrect_samples = [] test = Dataset(test_data) test_dataloader = torch.utils.data.DataLoader(test, batch_size=8) use_cuda = torch.cuda.is_available() device = torch.device("cuda" if use_cuda else "cpu") if use_cuda: model = model.cuda() total_acc_test = 0 with torch.no_grad(): for test_input, test_label in test_dataloader: test_label = test_label.to(device) mask = test_input['attention_mask'].to(device) input_id = test_input['input_ids'].squeeze(1).to(device) output = model(input_id, mask) _, pred = torch.max(output,1) idxs_mask = ((pred == test_label) == False).nonzero() print(idxs_mask) incorrect_samples.append(input_id[idxs_mask].cpu().detach().numpy()) acc = (output.argmax(dim=1) == test_label).sum().item() total_acc_test += acc print(f'Test Accuracy: {total_acc_test / len(test_data): .3f}') print(incorrect_samples)
修改后的代码实现需求
原始代码仅保存了误分类样本的input_id,无法直接关联到参与者ID和原始文本。以下是修改后的代码,能够完整收集误分类样本的关键信息:
修改后的自定义Dataset类
tokenizer = BertTokenizer.from_pretrained('bert-base-cased') labels = {'HC':0, 'PK':1} class Dataset(torch.utils.data.Dataset): def __init__(self, df): self.participantIDs = df['Participant_ID'] self.labels = [labels[label] for label in df['Diagnosis']] self.texts_raw = df['Transcript_Segment'] # 保存原始文本段 self.texts = [tokenizer(text, padding='max_length', max_length = 512, truncation=True, return_tensors="pt") for text in self.texts_raw] def classes(self): return self.labels def __len__(self): return len(self.labels) def get_batch_labels(self, idx): return np.array(self.labels[idx]) def get_batch_texts(self, idx): return self.texts[idx] def get_batch_participant_ids(self, idx): return self.participantIDs.iloc[idx] # 通过iloc获取对应索引的参与者ID def get_batch_raw_text(self, idx): return self.texts_raw.iloc[idx] # 获取对应索引的原始文本 def __getitem__(self, idx): batch_texts = self.get_batch_texts(idx) batch_y = self.get_batch_labels(idx) batch_participant_ids = self.get_batch_participant_ids(idx) batch_raw_text = self.get_batch_raw_text(idx) return batch_texts, batch_y, batch_participant_ids, batch_raw_text
修改后的BERT模型评估函数
def evaluate(model, test_data): incorrect_samples = [] test = Dataset(test_data) test_dataloader = torch.utils.data.DataLoader(test, batch_size=8) use_cuda = torch.cuda.is_available() device = torch.device("cuda" if use_cuda else "cpu") if use_cuda: model = model.cuda() total_acc_test = 0 # 标签反向映射,将数字标签转回原始类别名称 label_mapping = {v: k for k, v in labels.items()} with torch.no_grad(): for test_input, test_label, participant_ids, raw_texts in test_dataloader: test_label = test_label.to(device) mask = test_input['attention_mask'].to(device) input_id = test_input['input_ids'].squeeze(1).to(device) output = model(input_id, mask) _, pred = torch.max(output, 1) # 筛选当前batch中误分类的样本索引 wrong_indices = (pred != test_label).nonzero(as_tuple=True)[0] # 收集每个误分类样本的详细信息 for idx in wrong_indices: true_label = label_mapping[test_label[idx].item()] pred_label = label_mapping[pred[idx].item()] incorrect_samples.append({ 'Participant_ID': participant_ids[idx], 'Segment_Text': raw_texts[idx], 'True_Diagnosis': true_label, 'Predicted_Diagnosis': pred_label }) acc = (output.argmax(dim=1) == test_label).sum().item() total_acc_test += acc print(f'Test Accuracy: {total_acc_test / len(test_data): .3f}') # 将误分类样本转为DataFrame,方便查看和保存 import pandas as pd incorrect_df = pd.DataFrame(incorrect_samples) print("\n误分类样本详情:") print(incorrect_df) # 可选:将结果保存到CSV文件 # incorrect_df.to_csv('incorrect_classifications.csv', index=False) return incorrect_df
关键修改说明
- Dataset类扩展:新增存储原始文本的
texts_raw属性,启用并修正get_batch_participant_ids方法,新增get_batch_raw_text方法,确保__getitem__返回参与者ID和原始文本。 - 评估函数优化:遍历每个batch时,筛选误分类样本的索引,收集对应的参与者ID、原始文本、真实标签和预测标签,最终整理成DataFrame格式,便于后续分析和查看。
内容的提问来源于stack exchange,提问作者csStudent2102
相关产品推荐
相关产品推荐

