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

使用T5ForConditionalGeneration的generate方法推理时如何获取生成logits?

获取T5生成过程的Logits并计算置信度

实现步骤

要获取T5ForConditionalGeneration生成过程中的logits,只需在调用generate时添加两个关键参数,之后就能提取并处理每个生成token的置信度:

  1. 修改generate调用参数
    在model.generate()中加入return_dict_in_generate=True和output_scores=True,这样生成结果会返回包含logits(scores)的结构化对象,而不只是token IDs。

  2. 提取并处理logits
    生成结果的scores属性是一个列表,每个元素对应一个生成步骤的logits张量(形状为[batch_size, vocab_size])。我们可以对每个步骤的logits做softmax得到概率分布,再提取对应生成token的概率值。

  3. 计算整体置信度
    收集每个生成token的预测概率后,取平均值作为序列的整体置信度得分(相比直接平均logits,概率的平均值更直观,因为已经归一化到0-1区间)。

代码示例

from transformers import T5Tokenizer, T5ForConditionalGeneration
import torch

# 加载模型和tokenizer
tokenizer = T5Tokenizer.from_pretrained("t5-small")
model = T5ForConditionalGeneration.from_pretrained("t5-small")

# 示例输入(生成式分类任务的输入格式)
input_text = "classify: The movie was absolutely fantastic!"
tokenizer_outputs = tokenizer(input_text, return_tensors="pt")

# 生成并获取logits
generation_output = model.generate(
    **tokenizer_outputs,
    max_new_tokens=3,  # 根据分类任务的标签长度调整
    return_dict_in_generate=True,
    output_scores=True
)

# 提取生成的token IDs和各步骤的logits
generated_ids = generation_output.sequences
step_logits = generation_output.scores

# 计算每个生成token的置信度(softmax后的预测概率)
token_confidences = []
for step_idx, logits in enumerate(step_logits):
    # 将logits转换为概率分布
    token_probs = torch.softmax(logits, dim=-1)
    # 对应生成的token(generated_ids的第一个token是bos,所以取step_idx+1位置)
    current_token_id = generated_ids[0, step_idx + 1]
    # 获取该token的预测概率
    token_conf = token_probs[0, current_token_id].item()
    token_confidences.append(token_conf)

# 计算序列的平均置信度
avg_confidence = sum(token_confidences) / len(token_confidences)
print(f"生成分类标签的平均置信度: {avg_confidence:.4f}")

关于置信度计算的说明

直接平均logits的方式并不推荐,因为logits是未归一化的数值,不同生成步骤的logit范围差异较大,无法直接反映模型的置信程度。而通过softmax转换为概率后,每个token的概率值在0-1之间,代表模型对该token的预测置信度,取平均值能更合理地衡量整个生成序列的置信水平。

内容的提问来源于stack exchange,提问作者Sean

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 22:45:34