Hugging Face Trainer使用自定义CSV数据集报Invalid key: 0错误如何解决
报错原因
调用datasets.load_dataset()加载单个CSV文件时,默认返回的是DatasetDict(数据集字典)类型对象,所有数据会被归入默认的train切分下。从你打印的train_data输出也可以验证这一点:外层是DatasetDict结构,实际的训练数据存储在train键对应的取值中。
Hugging Face Trainer要求传入的train_dataset、eval_dataset参数必须是Dataset类型,你直接传入DatasetDict对象后,Trainer尝试按索引读取样本时无法匹配到对应键,就会抛出该KeyError。
解决方案
可以任选以下任意一种方法修改代码:
- 方法1:加载数据集时直接指定切分,直接获取Dataset类型对象
# 加载时指定split参数,返回结果直接为Dataset类型 train_data = datasets.load_dataset('csv', data_files="/gdrive/MyDrive/project/train.csv", split="train") test_data = datasets.load_dataset('csv', data_files="/gdrive/MyDrive/project/test.csv", split="train") # 后续初始化Trainer的代码无需修改 trainer = Trainer( model=model, args=training_args, train_dataset=train_data, eval_dataset=test_data ) trainer.train()
- 方法2:加载后手动提取对应切分的数据集
如果不修改数据集加载逻辑,仅需要在传入Trainer时取出DatasetDict中的对应切分即可:
trainer = Trainer( model=model, args=training_args, # 从DatasetDict中取train切分的实际数据集 train_dataset=train_data["train"], eval_dataset=test_data["train"] ) trainer.train()
内容的提问来源于stack exchange,提问作者J Luo
相关产品推荐
相关产品推荐

