如何用Llama2生成式模型获取二分类任务的0-1正类置信分?
正确实现Llama2二分类并获取0-1区间正类置信度的方案
问题分析
你当前的核心问题在于:
compute_transition_scores返回的是对数概率,直接取exp得到的是单个生成token的绝对概率,但没有对比正负类的概率分布,无法直接映射到0-1的正类置信度。- 对负类样本直接取负概率的逻辑错误,这样得到的分数会超出0-1范围,且不符合置信度的定义(置信度应表示模型对正类的信任程度,而非符号翻转的概率)。
正确实现思路
二分类任务中,我们需要对比模型对正类token和负类token的预测概率,将两者的相对差异转化为0-1区间的置信度。常用的两种可靠方式:
- 方式一:计算正类token概率与正负类概率之和的比值(即softmax结果)
- 方式二:计算正类与负类对数概率的差值,再通过sigmoid函数映射到0-1区间(更稳定,避免极端概率下的数值问题)
修正后的代码实现
from transformers import AutoModelForCausalLM, AutoTokenizer import torch # 加载模型和tokenizer(假设peft_config已定义) model = AutoModelForCausalLM.from_pretrained( peft_config.base_model_name_or_path, torch_dtype='auto', device_map='auto', offload_folder="offload", offload_state_dict=True ) tokenizer = AutoTokenizer.from_pretrained(peft_config.base_model_name_or_path) # 确保pad_token和eos_token一致(Llama默认无pad_token) tokenizer.pad_token = tokenizer.eos_token pos_scores = [] # 定义二分类对应的token(根据你的prompt设计调整,比如"1"代表正类,"0"代表负类) pos_token = "1" neg_token = "0" pos_token_id = tokenizer.encode(pos_token, add_special_tokens=False)[0] neg_token_id = tokenizer.encode(neg_token, add_special_tokens=False)[0] # 处理单条测试样本(批量处理可增加循环) input_ids = tokenizer(test_sample, return_tensors="pt").input_ids.to(model.device) # 仅生成1个判定token outputs = model.generate( inputs=input_ids, do_sample=False, max_length=input_ids.shape[1] + 1, pad_token_id=tokenizer.eos_token_id, output_scores=True, return_dict_in_generate=True ) # 获取生成步骤的原始logits(scores列表的第一个元素对应唯一的生成token) gen_logits = outputs.scores[0][0] # 提取正类和负类token的对数概率 pos_logit = gen_logits[pos_token_id].item() neg_logit = gen_logits[neg_token_id].item() # 方式一:用softmax计算正类置信度(直接得到0-1区间的概率值) pos_confidence = torch.softmax(torch.tensor([pos_logit, neg_logit]), dim=0)[0].item() # 方式二:用sigmoid计算置信度(通过logit差值映射,数值稳定性更强) # pos_confidence = torch.sigmoid(torch.tensor(pos_logit - neg_logit)).item() # 直接记录正类置信度,无需符号翻转 pos_scores.append(pos_confidence)
关键细节说明
- token匹配:必须确保
pos_token和neg_token与你的prompt引导输出完全一致(比如prompt结尾是"请输出1或0:",模型就会输出"1"或"0")。 - logit提取:
outputs.scores存储了每个生成步骤的原始logits,直接提取正负类token对应的logit值即可,无需通过compute_transition_scores绕路(当然用compute_transition_scores也能得到相同结果,但提取特定token的logit更直接)。 - 置信度逻辑:无论模型最终生成的是正类还是负类token,正类置信度都会自然落在0-1区间——预测为正类时置信度趋近于1,预测为负类时趋近于0,无需手动翻转符号。
内容的提问来源于stack exchange,提问作者Ofir
相关产品推荐
相关产品推荐

