Hugging Face Trainer训练时数据集丢失word_ids键问题求助
问题根因
Hugging Face Trainer默认开启remove_unused_columns配置(默认值为True),该机制会自动剔除不在模型前向传播参数列表中的字段,避免给模型传入不识别的输入引发报错。word_ids字段不属于预训练掩码语言模型默认接收的输入参数,因此在样本送入数据整理器之前就被Trainer自动过滤删除,这就是单独测试数据整理器正常、传入Trainer就报KeyError的核心原因——单独测试时直接取的是数据集完整样本,没有经过Trainer的字段过滤环节。
解决方案
方案1(操作最简单,优先推荐)
在定义TrainingArguments时显式关闭未使用列删除开关即可,其余代码无需改动:
from transformers import TrainingArguments training_args = TrainingArguments( # 原有其他配置(输出目录、batch size、学习率等)保持不变 remove_unused_columns=False )
方案2(兼容性更强,适合复杂项目)
如果担心全局关闭字段过滤会引入其他冗余字段影响模型运行,可以自定义适配的模型类,在forward方法的参数列表中声明word_ids即可(不需要实际使用该参数,Trainer只要识别到参数列表里有这个字段,就不会删除对应的列):
# 若使用的不是Bert类模型,替换为实际使用的MLM模型类即可 from transformers import BertForMaskedLM class WWMMaskedLM(BertForMaskedLM): def forward( self, input_ids=None, attention_mask=None, token_type_ids=None, labels=None, word_ids=None, **kwargs ): # 直接调用父类的forward逻辑,忽略传入的word_ids即可 return super().forward( input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids, labels=labels, **kwargs ) # 初始化模型时使用自定义的类即可 masked_model = WWMMaskedLM.from_pretrained("预训练模型路径或名称")
内容的提问来源于stack exchange,提问作者Rick Vink
相关产品推荐
相关产品推荐

