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

如何识别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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 15:53:14