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

如何使用Hugging Face Trainer结合自定义collate函数训练及报错修复

自定义Collate函数在Hugging Face Trainer中触发IndexError的原因与修复方案

报错原因

核心问题在于数据集格式与Trainer的自定义collate函数不兼容:

  • 代码中对加载的数据集调用了with_format(type="torch"),将样本转换为PyTorch张量格式存储;
  • 但自定义的collate_tokenize函数仍然期望接收原始字典格式的样本,尝试通过键名(如example["generated informal statement"])访问字段时,张量格式的样本无法支持这种索引方式,导致内部触发索引越界错误;
  • 额外隐患:当前collate函数返回的结果缺少labels字段,GPT2LMHeadModel训练时必须依赖该字段(与input_ids一致),即使修复索引问题,后续也会触发训练报错。

修复方法

方案1:移除数据集的张量格式转换(推荐)

删除加载数据集时的with_format(type="torch"),让数据集保持原始字典格式,确保collate函数能正确访问字段:

# 修改数据集加载代码,移除with_format
train_dataset = load_dataset(path, name, streaming=False, split="train", token=token)
eval_dataset = load_dataset(path, name, streaming=False, split="test", token=token)

方案2:完善collate函数,补充labels字段

GPT2语言模型训练需要labels字段(与input_ids内容一致),在collate函数中添加该字段:

def collate_tokenize(data):
    text_batch = [f'informal statement {example["generated informal statement"]} formal statement {example["formal statement"]}' for example in data]
    tokenized = tokenizer(text_batch, padding='longest', max_length=128, truncation=True, return_tensors='pt')
    # 为GPT2添加labels字段,与input_ids一致
    tokenized["labels"] = tokenized["input_ids"].clone()
    return tokenized

可选:调整训练参数兼容小batch

由于eval_dataset的样本数(13)无法被batch_size(8)整除,可在TrainingArguments中添加参数确保小batch正常处理:

training_args = TrainingArguments(
    # 保留原有参数
    eval_accumulation_steps=1,
)

内容的提问来源于stack exchange,提问作者Charlie Parker

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 23:10:38