使用预训练MLM训练AllenNLP AdversarialBiasMitigator时报参数错误
问题原因与修复方案
报错根因
你遇到的TypeError: forward() got an unexpected keyword argument 'tokens'错误,核心是token索引器配置与预训练BERT嵌入层的输入参数不匹配:
- 你当前配置里用了
single_id类型的token索引器,该索引器输出的张量键名为tokens - 但BERT对应的
PretrainedTransformerEmbedder的forward方法仅接受token_ids、attention_mask、token_type_ids这类参数,不识别tokens参数,因此触发参数不匹配报错
修复步骤
1. 修正token_indexers配置
删掉你之前加的single_id类型的索引器配置,替换为适配预训练Transformer的索引器,注意模型名要和你用的分词器、预训练MLM权重对应:
"token_indexers": { "tokens": { "type": "pretrained_transformer", "model_name": "bert-base-uncased" } }
同时确认你的text_field_embedder配置中,对应tokens键的嵌入器类型为pretrained_transformer,且模型名和上述配置一致。
2. 保留sorting_keys配置
你添加的"sorting_keys":["tokens"]没有问题,该配置作用是按句子长度排序优化batch生成效率,不需要修改。
3. 修正数据集读取逻辑的列映射错误
你的_read()方法存在列取反的问题:你提到CSV第一列是带特殊标记的句子,第二列是预测目标,但代码里你把第一列赋值给了targets,第二列赋值给了sentences,需要调整:
@overrides def _read(self, file_path: str): import pandas as pd data= pd.read_csv(file_path) # 调整列索引,第一列是句子,第二列是目标 sentences = data.iloc[:,0].tolist() targets = data.iloc[:,1].tolist() zipped = zip(sentences, targets) for s, t in zipped: tokens = self._tokenizer.tokenize(s) target = str(t) yield self.text_to_instance(s, tokens, [target])
内容的提问来源于stack exchange,提问作者Alice Saunders
相关产品推荐
相关产品推荐

