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

如何用Llama2生成式模型获取二分类任务的0-1正类置信分?

正确实现Llama2二分类并获取0-1区间正类置信度的方案

问题分析

你当前的核心问题在于:

  1. compute_transition_scores返回的是对数概率,直接取exp得到的是单个生成token的绝对概率,但没有对比正负类的概率分布,无法直接映射到0-1的正类置信度。
  2. 对负类样本直接取负概率的逻辑错误,这样得到的分数会超出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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 10:25:30