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

如何将PyTorch DataLoader传入Hugging Face Trainer?是否可行?

问题:能否直接将PyTorch DataLoader传入Hugging Face Trainer?

常规使用步骤

使用Hugging Face Trainer的标准流程为:

  • 加载数据
  • 对数据进行Tokenize处理
  • 将Tokenize后的数据集传入Trainer

最小可复现示例(MWE)

data = generate_random_data(10000)  # 生成10000个样本
df = pd.DataFrame(data)
dataset = Dataset.from_pandas(df)
tokenized_datasets = dataset.map(preprocess_function, batched=True)
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized_datasets,
    eval_dataset=tokenized_datasets,
)

错误操作及问题

若尝试直接传入PyTorch DataLoader:

train_dataset = convert_to_tensors(tokenized_datasets)
train_dataloader = DataLoader(train_dataset, batch_size=32, shuffle=True)
trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=train_dataloader,
        eval_dataset=train_dataloader,
)

会触发key 0错误,本质原因是Trainer期望train_dataset是支持索引、长度查询的数据集对象(如Hugging Face Dataset),而DataLoader是批量迭代器,无法通过索引(如dataset[0])获取单样本,且返回的批量数据格式不符合Trainer的预期。


解决方案

Hugging Face Trainer不支持直接传入DataLoader作为train_dataset/eval_dataset参数,但可以通过以下两种方式实现自定义数据加载逻辑:

1. 推荐:使用Hugging Face Dataset配合训练参数

继续沿用标准流程,通过training_args中的per_device_train_batch_size等参数控制批量大小,Trainer会自动处理数据加载,无需手动创建DataLoader。

2. 自定义Trainer类,重写数据加载方法

如果必须使用自定义DataLoader,可以继承Trainer并重写get_train_dataloader和get_eval_dataloader方法,返回你自己定义的DataLoader:

from transformers import Trainer

class CustomTrainer(Trainer):
    def get_train_dataloader(self):
        # 返回你的自定义训练DataLoader
        return train_dataloader

    def get_eval_dataloader(self, eval_dataset=None):
        # 返回你的自定义验证DataLoader
        return eval_dataloader

# 使用自定义Trainer初始化
trainer = CustomTrainer(
    model=model,
    args=training_args,
    # 无需传入train_dataset/eval_dataset(或传入任意占位对象,方法会覆盖逻辑)
)

内容的提问来源于stack exchange,提问作者Kirk Walla

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 18:05:01