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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:39:31