使用Huggingface Transformers做掩码语言任务预测结果不符合预期
BERT掩码语言模型任务非掩码内容丢失问题解决
问题原因
- BERT的掩码语言模型在训练阶段仅对标记为
[MASK]的位置计算损失、优化预测效果,非掩码位置的输出概率没有被强制要求和原输入token对齐,因此直接对模型输出的全局logits取argmax,会导致非掩码位置的内容被修改。 - 你当前的处理逻辑是直接用所有位置的预测结果解码,没有保留原输入的非掩码位置内容,才会出现原输入末尾的
detail丢失的情况。
正确实现方式
如果需要保留原输入的非掩码内容,仅替换掩码位置的预测结果,可以参考以下代码:
from transformers import BertTokenizer, BertForMaskedLM import torch model = BertForMaskedLM.from_pretrained("bert-base-uncased") tokenizer = BertTokenizer.from_pretrained("bert-base-uncased") text = "Read the rest of this [MASK] to understand things in more detail" inputs = tokenizer(text, return_tensors="pt") input_ids = inputs["input_ids"] # 前向传播获取预测结果,关闭梯度计算加速 with torch.no_grad(): outputs = model(**inputs) predictions = outputs.logits.argmax(dim=-1) # 定位掩码位置,仅替换掩码位置的id,保留原输入其他内容 mask_position = (input_ids == tokenizer.mask_token_id)[0] final_input_ids = input_ids.clone()[0] final_input_ids[mask_position] = predictions[0][mask_position] # 解码输出 print(tokenizer.decode(final_input_ids))
运行后输出结果为:<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> read the rest of this book to understand things in more detail [SEP]
符合预期,仅[MASK]位置被填充为预测值,其余原输入内容全部保留。
内容的提问来源于stack exchange,提问作者Roman Kazmin
相关产品推荐
相关产品推荐

