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

自定义Bart模型微调遇TypeError及后续问题求助(HuggingFace)

解决自定义Bart模型forward参数报错问题

问题1:TypeError: forward() got an unexpected keyword argument 'token_type_ids'

你的自定义forward函数未声明接收token_type_ids参数,但训练时数据加载器或Tokenizer会自动传入该参数(哪怕你没主动使用)。

解决方法

两种可行方案:

  1. 修改Tokenizer配置,禁止生成token_type_ids
    加载Tokenizer时直接设置不返回该参数:
    tokenizer = MyTokenizerFast.from_pretrained(tokenizer_repo, return_token_type_ids=False)
    
  2. 让自定义forward函数兼容额外参数
    在forward中用**kwargs兜底所有未声明的参数,再传给父类:
    class MyModel(BartForConditionalGeneration):
        def __init__(self, config):
            super().__init__(config)
    
        def forward(self, input_ids, attention_mask=None, decoder_input_ids=None, decoder_attention_mask=None, head_mask=None, **kwargs):
            return super().forward(
                input_ids=input_ids,
                attention_mask=attention_mask,
                decoder_input_ids=decoder_input_ids,
                decoder_attention_mask=decoder_attention_mask,
                head_mask=head_mask,
                **kwargs
            )
    

问题2:TypeError: forward() got an unexpected keyword argument 'encoder_outputs'

评估阶段Bart模型内部会调用forward并传入encoder_outputs参数(用于复用编码器输出提升效率),你的自定义forward未接收该参数导致报错。

解决方法

同样通过**kwargs兜底兼容:

class MyModel(BartForConditionalGeneration):
    def __init__(self, config):
        super().__init__(config)

    def forward(self, input_ids, attention_mask=None, decoder_input_ids=None, decoder_attention_mask=None, head_mask=None, **kwargs):
        return super().forward(
            input_ids=input_ids,
            attention_mask=attention_mask,
            decoder_input_ids=decoder_input_ids,
            decoder_attention_mask=decoder_attention_mask,
            head_mask=head_mask,
            **kwargs
        )

核心原因

HuggingFace Transformers的模型在训练、评估、生成等不同场景下,会根据逻辑传入不同的辅助参数。自定义子类的forward函数必须能接收这些参数——要么显式声明每个可能的参数,要么用**kwargs兜底,否则就会触发参数不匹配的报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 20:35:36