如何为PyTorch中BERT分类任务准备适配Trainer的数据集?
解决Hugging Face Trainer与自定义PyTorch Dataset的兼容问题
这个错误的根源很明确:Hugging Face Trainer默认使用的数据收集器(DataCollator)会尝试调用vars()函数来处理每个样本,将其转换为字典格式。但你的BloodDataset的__getitem__方法返回的是普通Python字典(或错误时返回字符串),而字典没有__dict__属性,所以触发了TypeError。
下面是两种可行的解决方案,按推荐优先级排序:
方案一:用SimpleNamespace包装样本(最简便)
我们可以用Python内置的types.SimpleNamespace将样本字典转换为具有__dict__属性的对象,这样vars()就能正常处理它。
修改后的BloodDataset代码
首先导入需要的模块:
from types import SimpleNamespace from torch.utils.data import Dataset import torch
然后更新数据集类的__init__和__getitem__方法(顺便优化索引查找效率):
class BloodDataset(Dataset): """MIMIC Blood dataset.""" def __init__(self, arff_file): """ Args: arff_file (string): Path to the arff file with annotations. """ self.indices, self.contents, self.labels = read_arff(arff_file) self.labels = torch.as_tensor(self.labels) self.inputs = encode(self.contents) self.input_ids = self.inputs['input_ids'] self.attention_mask = self.inputs['attention_mask'] # 创建索引映射字典,将样本index快速映射到列表下标(O(1)查找) self.idx_to_sample = {idx: i for i, idx in enumerate(self.indices)} def __len__(self): return len(self.indices) def __getitem__(self, idx): if torch.is_tensor(idx): idx = idx.tolist() # 用映射字典快速查找样本下标 sample_index = self.idx_to_sample.get(idx) if sample_index is not None: sample_dict = { 'index': idx, 'content': self.contents[sample_index], 'label': self.labels[sample_index], 'input_ids': self.input_ids[sample_index], 'attention_mask': self.attention_mask[sample_index] } # 将字典转换为SimpleNamespace对象,赋予__dict__属性 return SimpleNamespace(**sample_dict) else: # 改用抛出异常,符合PyTorch Dataset的规范 raise ValueError(f"Sample with index {idx} not found!")
这样修改后,Trainer的默认数据收集器就能正常处理你的数据集了。
方案二:自定义数据收集器(更灵活)
如果你希望保持返回字典格式的样本,可以自定义一个数据收集器,跳过vars()调用,直接处理字典。
自定义DataCollator代码
from transformers import DataCollatorForSequenceClassification import torch class CustomDataCollator(DataCollatorForSequenceClassification): def __call__(self, features): # 直接从字典列表中提取张量并堆叠成批量 batch = {} batch['input_ids'] = torch.stack([f['input_ids'] for f in features]) batch['attention_mask'] = torch.stack([f['attention_mask'] for f in features]) # 注意这里用'labels'(Trainer期望的字段名),而不是原字典中的'label' batch['labels'] = torch.stack([f['label'] for f in features]) return batch
初始化Trainer时指定自定义DataCollator
model = BertForSequenceClassification.from_pretrained(model_type) training_args = TrainingArguments( output_dir='./results', # output directory logging_dir='./logs', # directory for storing logs ) # 初始化时传入自定义数据收集器 trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=test_dataset, data_collator=CustomDataCollator(tokenizer=tokenizer) # 若需动态padding,传入你的tokenizer )
额外注意事项
- 避免在
__getitem__中返回字符串错误信息,改用抛出异常,这是PyTorch Dataset的标准做法。 - 原代码中
self.indices.index(idx)的时间复杂度是O(n),大数据集下会很慢,用映射字典优化为O(1)查找能显著提升性能。
内容的提问来源于stack exchange,提问作者Zingg
相关产品推荐
相关产品推荐

