如何在NLP填充掩码任务中计算指定候选词的概率
问题描述
我想用NLP完成文本掩码词填充,但不需要从所有词汇里筛选,仅需比较两个候选词的适配可能性。例如针对句子The [MASK] was stuck in the tree,我需要判断“kite”和“bike”哪个更适合填充掩码位置。
我已掌握使用Hugging Face的fill-mask pipeline查找全局概率最高词汇的方法,代码如下:
from transformers import pipeline, AutoTokenizer, AutoModelForMaskedLM # 定义带掩码的输入句子 input_text = "The [MASK] was stuck in the tree" # 加载预训练模型和分词器 model_name = "bert-base-cased" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForMaskedLM.from_pretrained(model_name) # 对输入句子分词 tokenized_text = tokenizer.tokenize(input_text) # 使用pipeline生成预测词和概率列表 mlm = pipeline("fill-mask", model=model, tokenizer=tokenizer) results = mlm(input_text) # 输出结果 for result in results: token = result["token_str"] print(f"{token:<15} {result['score']}")
但如果“bike”和“kite”不在pipeline返回的高概率结果中,这种方法就无法满足需求。请问如何计算指定掩码词的填充概率?
附:我不确定Overflow是否是发布此问题的最佳平台,似乎没有专门的NLP问题板块。
解决方案
要计算指定候选词的掩码填充概率,无需依赖pipeline返回的TopN结果,直接通过模型前向传播即可获取对应词的概率,具体步骤如下:
- 定位掩码位置:对输入文本分词后,找到
[MASK]对应的token索引。 - 获取模型输出logits:将分词结果转为模型可接受的张量输入,通过前向传播得到所有token的预测logits。
- 转换为概率并提取候选词值:对掩码位置的logits做softmax转换为概率分布,再提取候选词对应的概率值。
完整代码示例:
from transformers import AutoTokenizer, AutoModelForMaskedLM import torch # 输入句子与候选词 input_text = "The [MASK] was stuck in the tree" candidates = ["kite", "bike"] # 加载预训练模型与分词器 model_name = "bert-base-cased" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForMaskedLM.from_pretrained(model_name) model.eval() # 切换至评估模式 # 处理输入文本,获取张量与掩码位置 inputs = tokenizer(input_text, return_tensors="pt") mask_token_index = torch.where(inputs["input_ids"] == tokenizer.mask_token_id)[1][0] # 模型前向传播,获取logits with torch.no_grad(): # 关闭梯度计算,节省资源 outputs = model(**inputs) logits = outputs.logits # 提取掩码位置的logits并转换为概率分布 mask_logits = logits[0, mask_token_index, :] probs = torch.nn.functional.softmax(mask_logits, dim=-1) # 遍历候选词,输出对应概率 for candidate in candidates: token_id = tokenizer.convert_tokens_to_ids(candidate) # 若候选词不在词表中,概率设为0 if token_id == tokenizer.unk_token_id: print(f"{candidate:<15} 0.0 (不在模型词表中)") else: prob = probs[token_id].item() print(f"{candidate:<15} {prob:.6f}")
关键细节说明
- 评估模式:
model.eval()禁用模型的dropout、batch norm等训练层,保证预测结果稳定。 - 梯度关闭:
torch.no_grad()避免不必要的内存消耗,提升计算效率。 - 词表校验:若候选词不在模型词表中,会被转为未知词(UNK),此时概率直接设为0。
- 概率转换:softmax将模型输出的logits转换为0-1区间的概率值,直接对比候选词的概率即可判断适配性。
内容的提问来源于stack exchange,提问作者Shep Bryan
相关产品推荐
相关产品推荐

