Build a Reasoning Model(From Scratch) 第四章 Improving reasoning with inference-time scaling
本章内容包含:
- 通过提示词引导大语言模型输出推理过程,以此提升答案准确率
- 修改文本生成函数,生成多样化的回复内容
- 对多条回复进行采样,提升模型推理结果的可靠性
无需重新训练或改动模型本身,就能提升模型的推理表现与答案准确率。如图4.1的整体框架所示,本章介绍两种推理时缩放方法。这类方法在模型生成文本的推理阶段生效。读者在本章后续内容中将会看到,两种方法都能将前几章所用基础模型的准确率提升一倍以上。
在讲解图4.1中提到的各类推理方法前,我们先介绍什么是推理时缩放。

【前情提要,qwen3用的结构全是作者本人写的代码,只是把权重移过来了】
4.1 推理时缩放概述
提升模型推理能力主要有两种方案:
- 增加训练阶段算力投入
- 增加推理阶段算力投入(也称作推理时缩放或测试时缩放)
注释:在机器学习与人工智能领域,“算力(compute)”指训练或运行模型所需的计算资源。
图4.2展示了上述两种方案。从图表直观来看,无论加大训练算力还是推理算力,都能提升模型推理效果。但实际工程中,大语言模型一般会结合两种方式共同优化推理能力:一方面投入大量算力完成模型训练(后续章节讲解),另一方面提升推理阶段算力开销(本章核心内容)。
https://openai.com/zh-Hans-CN/index/learning-to-reason-with-llms/ 2024年9月12号发布的
推理时算力缩放【inference-time scaling】(简称推理时缩放、测试时算力缩放)指:模型完成训练后,在生成回答的阶段额外增加计算量,以此提升输出质量。该方法无需修改模型权重,而是让已训练完成的模型针对单个问题做更多“思考工作”,例如生成更多文本token、延长推理过程、采样多组答案、反复迭代优化回答等。
本书将重点介绍三种实用且基础的推理时优化技术(见图4.3):
- 方法一:采用思维链提示,引导模型输出完整推理过程。该方法实现简单,却能大幅提升答题准确率。
- 方法二:基于自一致性的并行采样,让模型生成多份回答,选取出现频次最高的结果作为最终答案。
- 方法三:迭代式自精炼,让模型分多轮自查、优化自身的推理逻辑与答案。

以上三种技术都属于推理时缩放范畴——它们都会使模型生成更多token,进而提高推理阶段的算力消耗。推理时缩放技术的共性特点是:以更高的推理成本换取更高的答题准确率。
本章仅讲解前两种方法,第三种自精炼技术将在第5章展开介绍。
第一种方法通过提示模型输出推理过程来提升答案准确率,该方案简单却效果极佳。
第二种方法针对同一输入生成多条回复,再选取出现频次最高的结果。我们将基于上一章介绍的文本生成函数进行扩展,实现自一致性——一种基于投票机制的推理时缩放技术。
我累计开展数千组实验,评测了十余种不同的推理时缩放方案。最终选定图4.3中的三种方法,原因是它们同时具备三大优势:能有效提升本文所用模型的答题准确率、覆盖推理时缩放的主流实用范式,且无需复杂额外组件,从零搭建即可落地。这三种技术还囊括两类核心实现思路:生成更长推理文本、并行生成多份回答。
方法一、三会输出带有完整推导说明的长文本,这也是推理类大模型的典型输出特征;而第二种并行采样方案已广泛应用于线上生产系统,例如2025年5月Anthropic发布的Claude 4(官方介绍链接:https://www.anthropic.com/news/claude-4)。
4.2 加载预训练模型
在动手实现各类推理时缩放方法之前,我们需要先加载预训练基础模型与分词器。下方代码和上一章所使用的代码逻辑相近。
import torch
from reasoning_from_scratch.ch02 import get_device
from reasoning_from_scratch.ch03 import (
load_model_and_tokenizer
)
device = get_device()
# Use CPU for the first run of this chapter
device = torch.device("cpu")
model, tokenizer = load_model_and_tokenizer(
which_model="base",
device=device,
use_compile=False
)
本章代码的计算开销很低,因此我建议初次运行时使用CPU设备,这样能得到和本章展示完全一致的运行结果(如果在GPU上运行,结果会产生细微偏差)。
我们使用上一章用到的MATH-500数据集里的一条提示词,来测试该模型:
from reasoning_from_scratch.ch03 import render_prompt
raw_prompt = (
"Half the value of $3x-9$ is $x+37$. "
"What is the value of $x$?"
)
prompt = render_prompt(raw_prompt)
print(prompt)
You are a helpful math assistant.
Answer the question and write the final result on a new line as:
\boxed{ANSWER}
Question:
Half the value of $3x-9$ is $x+37$. What is the value of $x$?
Answer:
我们可以将上述提示词作为输入,传入上一章定义的generate_text_stream_concat文本流式拼接生成函数。
我们想要比较几种推理时扩展(inference-time-scaling)策略,因此会对文本生成的封装代码(wrapper)做一个小改动,让我们能够在不改变周围的提示处理和输出代码的前提下,替换不同的生成函数和设置。具体改动在下面清单的代码注释中已高亮标出。
from reasoning_from_scratch.ch02 import generate_text_basic_stream_cache
def generate_text_stream_concat_flex(
model, tokenizer, prompt, device, max_new_tokens,
verbose=False,
generate_func=None, # 新增:允许传入自定义的生成函数
**generate_kwargs # 新增:允许传入额外的生成参数(会转发给生成函数)
):
# 如果没有指定生成函数,则使用默认的带缓存的流式生成函数
if generate_func is None: # 新增
generate_func = generate_text_basic_stream_cache
# 将提示词编码为 token id,并转成张量后放到指定设备上
# unsqueeze(0) 用于增加 batch 维度,形状变为 (1, seq_len)
input_ids = torch.tensor(
tokenizer.encode(prompt), device=device
).unsqueeze(0)
# 用于收集生成出的所有 token id
generated_ids = []
# 逐个 token 地调用生成函数进行流式生成
for token in generate_func( # 新增:使用可替换的生成函数
model=model,
token_ids=input_ids,
max_new_tokens=max_new_tokens,
eos_token_id=tokenizer.eos_token_id,
**generate_kwargs, # 新增:转发额外的生成参数
):
# 去掉 batch 维度,得到单个 token id
next_token_id = token.squeeze(0)
# 将当前 token id 记录下来
generated_ids.append(next_token_id.item())
# 如果开启了 verbose,则实时地解码并打印当前 token
if verbose:
print(
tokenizer.decode(next_token_id.tolist()),
end="", # 不换行,让输出连续显示
flush=True # 立即刷新缓冲区,实现流式输出效果
)
# 将所有生成的 token id 一起解码为最终的文本字符串并返回
return tokenizer.decode(generated_ids)
简而言之,前面定义的generate_text_stream_concat_flex函数和上一章的generate_text_stream_concat函数大体相似,区别在于现在我们可以把文本生成函数(例如generate_text_basic_stream_cache)以入参形式传入,而非在代码里写死。
这样做的实际好处是:外层封装函数统一负责提示词编码、流式输出、解码全流程,内部的生成步骤则可以更换为不同的解码策略。我们仅改动生成逻辑、而非整套运行链路,能够更公平地对比各类优化方案。后续小节中,我们会把generate_text_basic_stream_cache替换为更高级的生成函数。
调用generate_text_basic_stream_cache的方式和以往基本一致,只是现在需要显式传入该函数作为参数。
response = generate_text_stream_concat_flex(
model, tokenizer, prompt, device,
max_new_tokens=2048, verbose=True,
generate_func=generate_text_basic_stream_cache # NEW
)
\boxed{20}
注意这个答案是错误的,正确结果是83。在本章剩余内容以及下一章中,我们将实现各类推理时缩放方法,让模型输出正确答案。
4.3 借助思维链提示生成更优质的回复
我们已经完成预训练基座模型的加载,并且配置好了文本生成函数;接下来我们重点介绍思维链提示,以此优化模型的输出效果。这是一种经典、简洁且高效的提示技巧:如图4.4所示,它修改输入提示词,引导大语言模型输出推导过程,也就是所谓的思维链(也叫推理链路)。

最简单的思维链使用方式,是在提示词末尾追加一句指令,要求模型分步推理。该指令有多种表述方式,相关原始论文使用的是“Let’s think step by step.”(论文链接:https://arxiv.org/abs/2205.11916)。不过不同模型、不同任务适配的话术效果各不相同。本书这里选用“Explain step by step.”,它在本章的实验中表现更佳,同时示例也更简洁。
prompt_cot = prompt + " \n\nExplain step by step."
response_cot = generate_text_stream_concat_flex(
model, tokenizer, prompt_cot, device,
max_new_tokens=2048, verbose=True,
)
To solve the problem, we need to find the value of \( x \) such that half the value of \( 3x - 9 \) is equal to \( x + 37 \).
### Step 1: Set up the equation
We are given that half the value of \( 3x - 9 \) is equal to \( x + 37 \). This can be written as:
\[
\frac{1}{2}(3x - 9) = x + 37
\]
### Step 2: Eliminate the fraction
To eliminate the fraction, multiply both sides of the equation by 2:
\[
2 \cdot \frac{1}{2}(3x - 9) = 2(x + 37)
\]
Simplifying both sides:
\[
3x - 9 = 2x + 74
\]
### Step 3: Solve for \( x \)
Subtract \( 2x \) from both sides to isolate \( x \):
\[
3x - 2x - 9 = 74
\]
Simplify:
\[
x - 9 = 74
\]
Add 9 to both sides to solve for \( x \):
\[
x = 74 + 9
\]
\[
x = 83
\]
### Final Answer:
\[
\boxed{83}
\]
可以看到,模型现在会输出一段详尽的分步推导过程,在本例中也得出了正确答案。
这个简单的思维链提示很好地体现了推理时缩放存在的取舍关系:虽然模型答对了,但生成的token数量相比之前大幅增加。正如第二章所讲,大语言模型逐一生成token,每多生成一个token都需要完整执行一次模型前向传播。因此这些中间推理步骤不只是拉长了回答篇幅,还会直接增加推理延迟、计算开销,往往也会拉高接口调用成本。
需要注意的是,并非所有问题都能靠思维链提示得到优化。面对简单问题时,该方法有时反而会降低模型效果——模型可能生成错误的推导逻辑,进而误导自身,这种现象被称作“过度思考”。
最后,并非所有模型都能从“分步说明”这类指令中获益。本文实验使用的是基础基座模型,它本身不会主动输出推导过程,因此思维链提示能起到明显作用;而像上一章用到的通义千问3推理专用版这类经过专项训练的推理模型,本身就会在回答中附带推导步骤,这类思维链提示对它们而言没有增益,甚至没有使用必要。
大语言模型按训练阶段和用途,大致分为以下几类:
base(基座模型):只经过预训练,学的是"预测下一个词"。它没有对话或指令的概念,你给一段文本它就顺着往下补全。要让它"答题",得靠少样本示例或提示词去引导。特点是能力最原始、最灵活,但不听话、不安全,需要使用者自己想办法激发。
instruct(指令微调模型):在 base 基础上,用大量"指令—回答"数据做微调(SFT),再常配合人类反馈强化学习(RLHF/DPO)对齐。它学会了"听懂人话、按要求做事",你直接下命令它就回答。日常聊天、写作、问答用的基本都是这类,比如 ChatGPT、通义千问的对话版。
reasoning(推理模型):在 instruct 之上进一步专项训练,强化"先思考、再作答"的能力。回答时会自带一段推导过程(常用
<think>之类的标记包起来),擅长数学、代码、逻辑等需要多步推理的任务。因为它天生就会推理,再给它"请一步步思考"这类思维链提示往往是多余的。代表如 OpenAI o1、DeepSeek-R1、通义千问3推理版。除此之外常见的还有:
chat:概念上和 instruct 高度重叠,很多厂商直接叫 chat 版,特指为多轮对话优化、带有对话模板(如
<|im_start|>)的模型。code:面向编程专门训练/微调的模型,代码补全和生成能力更强,如 Code Llama、Qwen-Coder。
multimodal(多模态):能同时处理文本、图像、音频等,如 GPT-4o、Qwen-VL。
distilled(蒸馏版):用大模型的输出去训练小模型,让小模型以更低成本获得接近大模型的能力,常见于 reasoning 模型的小尺寸版本。
一句话串起来:base 是原料,instruct/chat 是"会听话"的成品,reasoning 是"会思考"的增强版,其余(code、multimodal 等)则是针对特定场景的专门化分支。
为何思维链能够提升答题准确率
思维链提示会要求模型写出推导出最终答案的中间步骤,这能从两方面切实提升效果:
第一,分步推演的过程给了模型更多自我纠错的机会。
第二,分步推理的形式和大量训练样本的书写逻辑相契合。例如各类大型数学、逻辑数据集里都附带完整详细的解题过程,要求模型输出思维链,能让模型贴合它已经学到的数据模式。
但与此同时,思维链并不能保证结果一定正确。模型依然有可能产出错误推导;对于十分简单的题目,额外的推理步骤反而会画蛇添足,催生更多失误。也就是说,思维链可以提升多数推理类任务的正确率,却并非万能方案。
总而言之,思维链提示不会给模型注入新知识,只是改变模型调用已有知识的方式。这种模式转换往往能产出更可靠的答案,在数学、代码、多步骤逻辑类问题上效果尤为突出。
练习4.1:在MATH-500数据集上使用思维链提示
修改代码清单3.15中的evaluate_math500_stream函数,验证思维链提示能否提升基础模型在MATH-500数据集上的答题准确率。
4.4 通过温度缩放控制输出多样性
上一节我们借助思维链提示延长模型的回答篇幅,初步体验了推理时缩放技术。思维链属于串行优化手段,它会增加模型预测下一个词元的步骤数量。
本章后续内容将实现自一致性技术(又称自一致性采样),该方法能够生成多份答案,具体流程见图4.5。这项技术由谷歌研究院在论文《Self-Consistency Improves Chain of Thought Reasoning in Language Models》中正式提出,论文地址:https://arxiv.org/abs/2203.11171。

使用该方法生成的多份答案彼此相互独立,因此可以采用并行采样的方式实现(若具备多块GPU等配套算力资源),这种方式不会增加用户的等待耗时。(思维链提示可以和该方法搭配使用,后文会详细介绍。)
在实现图4.5中第5步的自一致性采样之前,我们需要先对文本生成函数进行扩展,使其能够针对同一条提示词输出不同答案。为此,我们将实现两项可让模型生成多样化回复的技术:温度缩放(第3步)与Top-p过滤(第4步)。
本节重点讲解温度缩放,它是本章后续内容中调控输出多样性的核心手段之一,也是下文Top-p过滤方法的实现基础。不过在此之前,我们先来搞懂大语言模型是如何采样生成下一个词元的。这类底层采样控制逻辑看似偏技术细节,但正是依靠它们,我们才能为自一致性采样生成多版候选答案。
4.4.1 理解下一词元的选取流程
我们来深入剖析目前已经实现的文本生成流程,搞清楚底层的下一词元选择机制。这能帮你理解引入温度缩放的设计初衷。
假设我们有一条简单提示词:
ex_prompt = "The capital of Germany is"
response = generate_text_stream_concat_flex(
model, tokenizer, ex_prompt, device,
max_new_tokens=1, verbose=True
)
模型输出结果为“ Berlin”。这个过程看似简单,但如图4.6所示,底层实则包含多个执行步骤。

文本生成时,输入文本首先会被转换为词元ID(第二章已讲解相关内容):
# 1) Convert input text into token IDs
input_token_ids = torch.tensor(
tokenizer.encode(ex_prompt), device=device
).unsqueeze(0)
print(input_token_ids)
本例对应的词元ID张量如下:tensor([[ 785, 6722, 315, 9856, 374]])
如图4.6的第二步所示,我们获取待生成输出词元的分值,这类输出分值也叫作logits(对数几率)。大语言模型会针对每一个输入词元都输出一个结果词元,但我们只关心最后一个位置的输出,通过张量索引[:, -1]取出该位置,这个位置就对应我们要预测生成的下一个词元:
# 2) Get scores for next token
with torch.inference_mode():
next_token_logits = model(input_token_ids)[:, -1]
print(next_token_logits.shape) # Shape: [1, vocab_size]
打印得到输出张量形状为[1, 151936],其中151936代表该分词器与模型的词表大小。词表包含分词器能够处理、同时大模型可以生成的全部独立词元。
为得到最终生成的下一个词元(本例中为“ Berlin”),我们需要选出分值最高的词表条目(图4.6第三步):
# 3) Find vocabulary index with highest score
max_token_id = torch.argmax(next_token_logits)
print(f"Token ID: {max_token_id}")
print(f"Decoded token: '{tokenizer.decode([max_token_id])}'")
Token ID: 19846
Decoded token: ' Berlin'
至此我们已经讲完生成下一个词元的三大主要步骤。进入下一小节之前,我们进一步观察传入torch.argmax函数、用于获取词元ID的next_token_logits张量的分数分布,并用 Matplotlib 将其绘图可视化。
import matplotlib.pyplot as plt
def plot_scores_bar(
next_token_logits, start=19_800, end=19_900,
arrow=True, ylabel="Logit 值"
):
# 选取词表的一个子区间
x = torch.arange(start, end)
# .cpu() 是 to(torch.device("cpu")) 的简写
logits_section = next_token_logits[0, start:end].float().cpu()
# 绘制 logits
plt.bar(x, logits_section)
plt.xlabel("词表索引")
plt.ylabel(ylabel)
# 标注最大的 logit
if arrow:
max_idx = torch.argmax(logits_section)
plt.annotate(
"Berlin",
xy=(x[max_idx], logits_section[max_idx]),
xytext=(x[max_idx] - 25, logits_section[max_idx] - 2),
arrowprops={
"facecolor": "black", "arrowstyle": "->", "lw": 1.5
},
fontsize=10,
)
plt.grid(alpha=0.3)
plt.tight_layout()
plt.show()
plot_scores_bar(next_token_logits)

4.4.2 通过温度参数对词元分数(logits)做重新缩放
既然你已经了解模型如何选择下一个token,我们就可以引入温度这一概念。温度,更准确地说是温度参数,会对logits做重新缩放处理,让对应的概率分布变得更尖锐或者更平缓,进而影响下一token的选取结果。
如图4.8所示,本节的核心是:在执行采样之前,使用温度参数对下一token的logits进行缩放。这里所说的缩放,指调整分数的幅值,使得采样步骤对分数之间的差异变得更加敏感,或是敏感度降低。
实现温度缩放步骤(图4.8中的3.2步骤)的代码简短易懂。本质上,下面这段代码会先对logit数值做缩放,之后再将其转换成概率(下一节会展开讲解)。
def scale_logits_by_temperature(logits, temperature):
# 温度必须为正数,否则会导致除以零或反转分布
if temperature <= 0:
raise ValueError("Temperature must be positive")
# 用温度对 logits 进行缩放:温度越高分布越平滑,越低越尖锐
return logits / temperature
实际使用中,温度参数应当取正数;因此如果温度等于0或者小于0,我们就抛出报错。后续把温度缩放集成到文本生成函数时,还会补充更多温度缩放的安全校验逻辑。
对于其余合法数值,直接用 logits 除以温度即可。温度取1.0代表不做任何变换,因为任意数字除以1结果保持不变。温度小于1.0会让概率分布变得更尖锐,模型在选择下一词元时会表现得更加“确信”。温度大于1.0会拉平 logit 分布,能够让采样(图4.8中3.4步骤)得到更多样的输出。换句话说,更高的温度会缩小最高分词元和其余排名靠后词元之间的概率差距,增大采样选中非最高分词元的概率。

我们来实际运行scale_logits_by_temperature函数,选用0.5和5.0这两个相对极端的温度值,以便观察更加明显的效果。
def plot_logits_with_temperature(
next_token_logits, start=19_800, end=19_900,
temps=(0.5, 5.0),
):
# 生成横坐标:从 start 到 end 的词表索引
x = torch.arange(start, end)
# 取出指定区间的原始 logits,转成 float 并移到 CPU 上
logits_orig = next_token_logits[0, start:end].float().cpu()
# 应用温度缩放:对每个温度值分别缩放 logits
logits_scaled = [
scale_logits_by_temperature(logits_orig, T) for T in temps
]
# 绘制原始 logits 曲线
plt.plot(x, logits_orig, label="Original logits", lw=2)
# 绘制低温缩放曲线(T<1,分布更尖锐)
plt.plot(
x, logits_scaled[0],
label=f"T={temps[0]} (sharper)", ls="--", lw=1
)
# 绘制高温缩放曲线(T>1,分布更平坦)
plt.plot(
x, logits_scaled[1],
label=f"T={temps[1]} (flatter)", ls=":", lw=3
)
# 标注最大 logit 对应的位置(此处为 "Berlin")
max_idx = torch.argmax(logits_orig)
plt.annotate(
"Berlin",
xy=(x[max_idx], logits_orig[max_idx]), # 箭头指向的点
xytext=(x[max_idx] - 25, logits_orig[max_idx] + 2), # 文字位置
arrowprops={"facecolor": "black", "arrowstyle": "->", "lw": 1.5},
fontsize=12,
)
plt.xlabel("Vocabulary index") # x 轴标签:词表索引
plt.ylabel("Logit value") # y 轴标签:logit 值
plt.legend() # 显示图例
plt.grid(alpha=0.3) # 添加半透明网格
plt.tight_layout() # 自动调整布局,防止元素重叠
plt.show() # 显示图像
# 调用函数:绘制不同温度下的 logits 对比图
plot_logits_with_temperature(
next_token_logits,
temps=(0.5, 5.0)
)
如图 4.9 所示,较大的温度值(5.0)会得到平缓得多的分数分布,而较小的温度值(0.5)则会得到尖锐得多的分数分布。

代码清单4.5中的绘图代码和代码清单4.3非常相似。除了执行温度缩放之外,我们这里改用折线图(plt.plot),不再使用柱状图(plt.bar)。虽然从技术角度来说,x轴为离散词表索引时柱状图会是更合适的选择,但在本场景下,折线图更便于可视化、对比经过缩放后的logits。
关键不在于单个logit的绝对数值大小,而在于各个logit之间的差值如何变化。如果像“Berlin”这类词元相比其他候选词元优势更加突出,那么它后续被选中的概率就会更高。如果分值之间的差距缩小,排名靠后的词元被选中的概率就会上升。下一节会将logits转换为概率,以此直观展示这一现象。
为什么取名为温度(temperature)?
该术语源自物理学,物理学中温度用于控制系统内部的随机程度与粒子运动剧烈程度。在大语言模型中沿用了这一思想,用来控制模型选择下一词元时的确信程度与创作发散程度,后续章节会进一步讲解。
4.4.3 从概率分布中采样下一词元
上一节中,我们使用不同温度值对 logit 数值进行了重新缩放,以此干预模型对下一词元的选择方式。在进入下一节的下一词元采样环节之前,还需要增加一个中间步骤:如图4.10所示,将经过缩放的logits转换为概率分值。

为演示缩放后的logits如何转换成概率分值(图4.10的3.3步骤),我们选用温度5.0,该设置可以更方便地在图表中观察输出的概率结果。举个例子,词元“ Berlin”原本的logit数值极高,会在数值尺度上占据绝对主导,导致很难看清周边其他词元的概率大小。
如下面的代码清单所示,只需调用torch.softmax这一个函数,即可完成从缩放后logits到概率分值的转换。
torch.softmax 函数会将 logit 数值归一化到 0‑1 的区间,并且所有数值相加总和为 1。我们可以通过下面这段代码验证该特性。
# Step 3.2: Rescale next-token scores
rescaled_logits = scale_logits_by_temperature(next_token_logits, 5.0)
# Step 3.3 Convert rescaled logits into probability scores
next_token_probas = torch.softmax(
rescaled_logits, dim=-1
)
print("Probability sum:", torch.sum(next_token_probas))
Probability sum: tensor(1., dtype=torch.bfloat16)
此外,我们复用代码清单4.3中的plot_scores_bar函数,对转换后的分值做可视化:
plot_scores_bar(
next_token_probas, arrow=False, ylabel="Probability value"
)

生成的图表y轴为概率值,如图4.11所示。可以看到,在选取的这段词表范围内,词元ID 19846(对应“ Berlin”)的数值最高。它的概率为0.0003,可以用下面代码验证:
print("Token ID 19,846 probability:", next_token_probas[:, 19846])
注意,尽管0.0003这个数值看上去很小,但在该图表未展示的其余词表区间里,不存在比它更大的概率值。例如执行如下代码,可以确认0.0003确实就是最大概率:
print("Highest probability:", max(next_token_probas.squeeze(0)))
我们可以把该分值理解为模型的确信度。这代表相比其他所有词元,模型更确信应当选择词元“ Berlin”(ID:19846)作为下一个词元。
该概率数值之所以很小,是因为我们使用了较高的温度参数,这样便于将该词元的分值和其他词元的分值放在同一张图中展示。如果把温度从5调整为0.5,该词元的概率分值会从0.0003上升至0.3398,而其余词元的概率分值则会进一步趋近于0。
softmax 函数的底层原理
softmax 函数将原始分值向量(logits)转换为概率分布:输出每个数值介于0到1之间,并且全部数值之和等于1。经过该转换后,数值更便于理解,也为后续采样提供条件。
如果你熟悉数学符号,softmax 的计算公式如下:
softmax(zi)=exp(zi)∑jexp(zj) \mathrm{softmax}(z_i) = \frac{\exp(z_i)}{\sum_j \exp(z_j)} softmax(zi)=∑jexp(zj)exp(zi)
其中 zzz 为实数输入向量:
z=[z1, z2, …, zn] z = [z_1,\ z_2,\ \dots,\ z_n] z=[z1, z2, …, zn]
- nnn:向量内元素的总个数
- iii:当前元素的索引(1≤i≤n1 \le i \le n1≤i≤n)
- jjj:用于遍历全部元素做求和运算的索引(1≤j≤n1 \le j \le n1≤j≤n)
该式为每一个元素 ziz_izi 输出归一化后的概率,满足:
∑isoftmax(zi)=1 \sum_i \mathrm{softmax}(z_i) = 1 i∑softmax(zi)=1
用代码实现 softmax 仅需三行:
def softmax_with_temperature(logits, temperature):
scaled_logits = logits / temperature
return torch.softmax(scaled_logits, dim=0)
实际开发中更推荐直接使用torch.softmax,该函数内置了数值稳定性优化,能够更加可靠地处理极大值与极小值。
将 logits 转换为这类概率分值,一方面是概率结果更便于解读,另一方面我们就可以调用torch.multinomial函数基于该分布进行采样。举例来说,温度设为5.0时,从该概率分布采样,抽到词元“ Berlin”的概率为0.03%;而温度设为0.5时,抽到“ Berlin”的概率为33.98%。
对应图4.10中3.4步骤的采样过程,实现代码如下:
# Step 3.4: Sample token according to probabilities
torch.manual_seed(123)
print(
"Sampled token:",
torch.multinomial(next_token_probas.cpu(), num_samples=1)
)
Sampled token: tensor([[65094]])
这段代码返回词元ID 65094,对应单词“ mistress”。可以看到,该词放在提示词“The capital of Germany is(德国的首都是)”这个语境下完全不通顺。出现该结果是因为使用了高温度配置,模型会更倾向于采样“ Berlin”之外的其他词元,本次只是随机抽到了这个词。
torch.multinomial函数会按照概率大小的比例来采样词表索引。换言之,概率越高的词表索引,被采样选中的可能性越大。在温度为5的条件下,如果我们重复非常多次采样,抽到词元“ Berlin”对应词表索引的概率就是0.03%。
torch.multinomial 做的是"按权重的有放回/无放回抽签"。它的核心思路是把一组概率值当成一段被切分成若干块的数轴,每块的长度等于对应词元的概率,然后随机丢一个点落在数轴上,看落在哪一块,就选哪个词元。
具体过程分三步:
累积概率(构建 CDF)
把概率分布累加起来,形成一个从 0 到 1 单调递增的区间划分。举例,假设只有 4 个词,概率是:词A: 0.1 -> 区间 [0.0, 0.1) 词B: 0.6 -> 区间 [0.1, 0.7) 词C: 0.05 -> 区间 [0.7, 0.75) 词D: 0.25 -> 区间 [0.75, 1.0)生成随机数
从 [0, 1) 上均匀采样一个随机数 r,比如 r = 0.42。查找落点
看 r 落在哪个区间里。0.42 落在 [0.1, 0.7),所以选中词B。
在看更多示例之前需要注意:我们设置了随机种子,保证本章代码可复现。但torch.multinomial()在CUDA、MPS设备上仍然可能输出不一样的结果;当采样数量较大时甚至会发生程序崩溃(我在PyTorch 2.9版本的CUDA与MPS设备上均观测到该问题),这也是我们通过.cpu()把采样运算放到CPU上执行的原因。
接下来我们再多采样几组下一词元候选,获取更有代表性的采样结果。
def count_samples(probas, num_samples=1000, threshold=1, tokenizer=None):
# 根据概率分布进行采样
# torch.multinomial 会按照 probas 中的概率抽取 num_samples 个索引
# replacement=True 表示有放回抽样,同一个索引可以被多次抽到
samples = torch.multinomial(
probas.cpu(), num_samples=num_samples, replacement=True
)
# 统计每个索引被抽中的次数
# squeeze(0) 去掉第 0 维(批次维),得到一维的采样结果
# bincount 会统计每个整数出现的次数,minlength=1 保证结果至少有 1 个元素
counts = torch.bincount(samples.squeeze(0), minlength=1)
# 打印统计结果
for i, c in enumerate(counts):
# 只打印出现次数超过阈值 threshold 的索引
if c > threshold:
if tokenizer is None:
# 没有提供分词器时,直接打印词表索引和出现次数
print(f"Vocab index {i}: {c.item()}x")
else:
# 提供了分词器时,把索引解码成对应的 token 文本再打印
print(f"'{tokenizer.decode([i])}': {c.item()}x")
代码清单4.7中的count_samples函数会从概率分布中采样词元索引,并统计每个词元被抽中的次数。该函数调用torch.multinomial,按照概率占比随机选出num_samples个索引。replacement=True参数开启有放回采样,允许同一个词元被多次抽取。
随后使用torch.bincount统计各个索引的出现频次。函数仅打印出现次数超过指定阈值的词元,避免大量低频采样结果把输出日志弄得杂乱。
samples和counts是两个 tensor,我用具体例子说明。假设
probas是一个词表大小为 5 的概率分布:probas = tensor([[0.1, 0.5, 0.2, 0.1, 0.1]]) # 形状 [1, 5]
samples:抽样得到的索引序列samples = torch.multinomial(probas.cpu(), num_samples=10, replacement=True) # 结果可能是: # tensor([[1, 1, 2, 1, 4, 0, 1, 2, 3, 1]]) 形状 [1, 10]这里每个数字都是一个"词表索引"。因为索引 1 的概率最高(0.5),所以它出现得最频繁。
samples记录的是"每一次抽样分别抽到了哪个索引",长度等于num_samples。> minlength 的作用是:如果按这个规则算出来的长度小于 minlength,就在后面补 0 补到 minlength 那么长。
counts:每个索引被抽中的次数counts = torch.bincount(samples.squeeze(0), minlength=1) # tensor([1, 5, 2, 1, 1]) # ↑ ↑ ↑ ↑ ↑ # 索引0 索引1 索引2 索引3 索引4 的出现次数
注意:
count_samples函数仅用于演示说明。在真实文本生成场景中,我们每次只采样一个token。该函数会从概率分布中抽取大量样本,用来展示各个词元被选中的频次,从而让底层的概率分布更便于可视化和理解。
首先,我们对温度等于5得到的概率分值运行count_samples函数:
torch.manual_seed(123)
count_samples(next_token_probas, tokenizer=tokenizer)
输出结果如下:
'}': 2x
' </': 2x
' represent': 2x
' Inf': 2x
'()*': 2x
' beside': 2x
' Kob': 2x
'�': 2x
可以看到,即便默认采样数量为1000次,所有被采样出来的词元最多只出现2次。而且对于提示词“The capital of Germany is(德国的首都是)”这个上下文来说,以上全部都是无意义的词元。产生这类不合理输出的原因是我们使用的温度值过高。
接下来尝试更低的温度值0.35,该参数会让分数分布变得更尖锐,提升采样得到有意义下一词元的概率:
torch.manual_seed(123)
probas_lowT = torch.softmax(
scale_logits_by_temperature(next_token_logits, 0.35), dim=-1
)
此时得到如下输出:
' __': 158x
' Berlin': 435x
' ____': 169x
' ______': 209x
' Munich': 3x
' Hamburg': 3x
' _____': 18x
针对提示词“The capital of Germany is”,这一组下一词元候选结果就合理很多。在1000次采样当中,词元“ Berlin”被抽中435次。执行print(probas_lowT[0, 19_846])可以验证:温度0.35条件下,该词元的理论概率大约为42%。如果想要进一步提高采样出“ Berlin”的概率,可以继续调低温度。
可以注意到,其余一些出现频次较低的候选词元同样具备合理性。例如“慕尼黑”与“汉堡”都是德国的大城市,因此它们和该问题并非完全无关。带下划线的输出(____),大概率是因为模型训练时见过测验填空题格式的文本,例如:“德国的首都是____”。
你可能会产生疑问:既然“柏林”才是正确答案,那为什么还要引入温度缩放和多项式采样,让模型偶尔输出错误结果? 针对不同类型的问题,在采样阶段引入随机性,可以让模型探索其他备选回复,而不是永远选择概率最高的那一个词元。这种输出的可变性对创意类、开放式任务十分有用,这类任务往往存在多种合理的续写结果。
具体到推理任务中,我们可以利用采样带来的多样性,使用自一致性等技术(4.6节会讲解):生成多份候选答案并做比对,以此提升答案的准确率。
4.4.4 将温度缩放加入文本生成函数
在为文本生成流程新增token概率选择过滤器这一改进之前,我们先对文本生成函数完成温度缩放的改造,这样就能更方便地调用模型生成新的词元。
# 练习 2.2 附录 B
from reasoning_from_scratch.qwen3 import KVCache
@torch.inference_mode() # 推理模式,禁用梯度计算,节省显存并加速
def generate_text_temp_stream_cache(
model,
token_ids,
max_new_tokens,
eos_token_id=None,
temperature=0.
):
model.eval() # 将模型设为评估模式(关闭 dropout 等)
# 创建 KV 缓存,用于存储各层的 Key/Value,避免重复计算
cache = KVCache(n_layers=model.cfg["n_layers"])
model.reset_kv_cache() # 重置模型内部的 KV 缓存状态
# 步骤 3.1:把整个输入序列送入模型,取最后一个位置的 logits
out = model(token_ids, cache=cache)[:, -1]
for _ in range(max_new_tokens): # 循环生成,最多生成 max_new_tokens 个新 token
########################################
# 新增部分:
orig_device = token_ids.device # 记录原始设备(CPU/GPU),后续把结果搬回来
if temperature is None or temperature == 0.0:
# 温度为 0 时使用贪心解码:直接选概率最大的 token
next_token = torch.argmax(out, dim=-1, keepdim=True)
else:
# 步骤 3.2:用温度对 logits 进行缩放(温度越高,分布越平滑/随机)
logits = scale_logits_by_temperature(out, temperature)
# 步骤 3.3:通过 softmax 把 logits 转换为概率分布
probas = torch.softmax(logits, dim=-1)
# 步骤 3.4:按概率分布进行随机采样,得到下一个 token
# multinomial 在 CPU 上执行,采样后再搬回原设备
next_token = torch.multinomial(probas.cpu(), num_samples=1)
next_token = next_token.to(orig_device)
#########################################
# 如果生成了结束符(EOS),提前停止生成
if (eos_token_id is not None
and torch.all(next_token == eos_token_id)):
break
yield next_token # 以流式方式逐个返回生成的 token
# 只把新生成的 token 送入模型,配合 KV 缓存高效计算下一步的 logits
out = model(next_token, cache=cache)[:, -1]
代码清单4.8中的generate_text_temp_stream_cache函数与第2章的generate_text_stream_cache函数结构相近。新增内容为温度缩放与采样逻辑(位于代码里# New注释标记的下方)。
下面的代码演示如何把温度缩放和采样集成到词元生成循环当中。该函数可以传入代码清单4.2介绍的灵活性更强的包装函数generate_text_stream_concat_flex。外层提示词处理逻辑保持不变,我们只可以替换接入不同的解码策略:
torch.manual_seed(123)
response = generate_text_stream_concat_flex(
model, tokenizer, prompt, device,
max_new_tokens=2048, verbose=True,
generate_func=generate_text_temp_stream_cache,
temperature=1.1
)
程序输出结果为 \boxed{$x = \frac{90}{7}$}。正确答案是83,因此模型本次输出依旧错误。但本段代码仅用作演示,目的是说明我们可以借助温度缩放和采样来调整模型输出结果。
温度参数的选择
实际使用时,温度参数的选取取决于业务目标。温度设为0.0对应贪心解码,即始终选择概率最高的词元。
如果希望输出具备少量多样性,同时又不至于生成太过混乱的内容,通常选用0.3‑0.8这类较小的非零温度值。设置很高的温度会让模型进行更广范围的探索,适合创意生成或者宽泛搜索类场景;但对于需要输出唯一最优答案的任务,高温往往会降低结果可靠性。
下一节,我们将对采样流程做进一步优化。
4.5 通过 top‑p 采样平衡多样性与文本连贯性
你已经了解温度缩放以及(借助torch.multinomial实现的)采样机制能够提升大语言模型回复的多样性,但这种效果有利也有弊。具体来说,该方式有可能采样出和用户问题毫无关联的词元。
接下来我们引入 top‑p 过滤器(见图4.12)对采样流程进行优化,避免意外采样到置信度极低的词元。本节介绍的 top‑p 采样也叫作核采样(nucleus sampling)。

图4.12概括了构成top‑p过滤器的四个步骤:对token概率排序、计算累积概率和、筛选满足top‑p截断条件的子集,以及对剩余概率重新归一化,使其再次构成合法的概率分布。
重新归一化这一步必不可少:一旦剔除全部低概率词元,剩下的概率值相加将不再等于1。为保证采样结果正确,我们需要对保留下来的概率数值做缩放,得到一个合规有效的概率分布。
接下来我们将详细讲解图4.12中的4.1至4.4每一个步骤。
4.5.1 选取 top‑p 词元子集
在实现 top‑p 过滤器之前,我们先构建一个简单的模拟数据集,便于演示温度缩放与采样流程。首先做一个假设:模型和分词器的词表仅有10个条目,而不是原本的151936个。
# 步骤 3.1: 获取 logits(这里用 10 个 token 的玩具 logits 作为示例)
toy_logits = torch.tensor(
[-0.7, -3.0, 0.1, -1.2, 2.0, -1.0, -0.5, -2.0, 0.3, 1.5]
)
# 步骤 3.2: 应用温度缩放(temperature=1.0 表示不改变原始分布)
toy_logits_scaled = scale_logits_by_temperature(toy_logits, 1.0)
# 步骤 3.3: 通过 softmax 将 logits 转换为概率分布(所有值之和为 1)
toy_probas = torch.softmax(toy_logits_scaled, dim=-1)
# 绘制柱状图:x 轴为词表索引,y 轴为对应 token 的概率
plt.bar(
torch.arange(len(toy_logits_scaled)), toy_probas,
alpha=0.5 # 设置柱子的透明度
)
plt.ylim([0, 1]) # 将 y 轴范围固定在 [0, 1],方便观察概率大小
plt.xlabel("Vocabulary index") # x 轴标签:词表索引
plt.ylabel("Probability") # y 轴标签:概率
# plt.savefig("12.pdf") # 如需保存为 PDF 文件,取消此行注释
plt.show() # 显示图像
注意:变量toy_logits存储的是模型输出的下一词元logit分值示例,对应代码清单4.8中model(token_ids, cache=cache)[:, -1]的调用结果,这里假设词表大小为10。
图4.13输出的柱状图展示了经过转换为概率之后的下一词元logit分值。

到这里为止,内容基本是对上一节的回顾。接下来我们执行图 4.12 所示的 top‑p 过滤的前两个步骤(步骤 4.1 与 4.2),也就是将概率分值按降序排序,并计算累积和。
# Step 4.1: 按概率降序排序
# torch.sort 返回排序后的值和对应的原始索引
sorted_probas, sorted_idx = torch.sort(toy_probas, descending=True)
# Step 4.2: 计算累积和(前缀和)
# cumsum[i] 表示概率最高的前 i+1 个 token 的概率之和
cumsum = torch.cumsum(sorted_probas, dim=-1)
# 用柱状图绘制每个 token 的概率(已按从高到低排序)
plt.bar(
torch.arange(len(sorted_probas)), sorted_probas,
alpha=0.5 # 设置透明度,方便与阶梯图叠加观察
)
# 用阶梯图绘制累积概率曲线
plt.step(
torch.arange(len(cumsum)), cumsum,
where="mid", color="C1", label="Cumulative sum" # where="mid" 让阶梯居中对齐
)
plt.ylim([0, 1]) # y 轴范围固定为 0~1(概率区间)
plt.xlabel("Token rank (sorted by probability)") # x 轴:token 排名(按概率排序)
plt.ylabel("Probability") # y 轴:概率
plt.show() # 显示图像
代码清单4.10中使用的torch.cumsum函数用于计算指定维度上元素的累积和。在本例中,该函数接收经过排序的词元概率,逐步累加;输出结果的每一个位置,代表累加至该词元为止的总概率。
举例来说,输出的第一个元素等于最大的概率值,第二个元素等于概率最高的前两项之和,以此类推,直到最后一个数值等于1。结合代码清单4.10生成的累积步长图(图4.14),可以把这个过程理解得更透彻。

现在我们已经得到概率累积和,就可以实现 top‑p 的核心过滤步骤。top‑p 里的 p 代表概率,top‑p 可以理解为:保留满足累积概率小于或等于 p 的最小词元集合。下面的代码清单给出了 top‑p 过滤的简易实现。
# Step 4.3.1: Apply top-p threshold (e.g., keep tokens until cumulative mass > 0.8)
top_p = 0.8
keep_mask = cumsum <= top_p
n_kept = torch.sum(keep_mask).item()
print("Cumulative sum:", cumsum)
print("Tokens kept:", n_kept)
在代码中,通过keep_mask = cumsum <= top_p实现该逻辑:该语句会标记所有累积概率质量尚未超过阈值p(本例中top_p=0.8)的词元。随后计算被保留的词元数量,并赋值给变量n_kept。输出如下:
Cumulative sum: tensor([0.4538, 0.7290, 0.8119, 0.8798,
0.9170, 0.9475, 0.9701, 0.9886, 0.9969, 1.0000])
Tokens kept: 2
观察返回的累积概率可以发现,第二个数值(0.7290)略低于0.8的top‑p阈值,而第三个数值(0.8119)已经超出阈值,因此仅保留前两个词元。
top‑p过滤还有一种更常用的变体:会把那个超出阈值的词元也一并保留,具体见下一份代码清单。
# A more common variant is to include the token that crosses the threshold
print(cumsum)
print(sorted_probas)
print(cumsum - sorted_probas)
keep_mask = (cumsum - sorted_probas) < top_p
n_kept = keep_mask.sum().item()
print("Tokens kept:", n_kept)
tensor([0.4538, 0.7290, 0.8119, 0.8798, 0.9170, 0.9475, 0.9701, 0.9886, 0.9969,
1.0000])
tensor([0.4538, 0.2752, 0.0829, 0.0679, 0.0372, 0.0305, 0.0226, 0.0185, 0.0083,
0.0031])
tensor([0.0000, 0.4538, 0.7290, 0.8119, 0.8798, 0.9170, 0.9475, 0.9701, 0.9886,
0.9969])
Tokens kept: 3
这段代码返回的保留token数量为3。这符合top‑p过滤的定义:选取最小的token集合,使集合的累积概率质量大于或等于p。
下面我们通过一张图表来演示该过滤过程。
plt.bar(
torch.arange(len(sorted_probas)), sorted_probas,
alpha=0.5, label="Sorted probabilities"
)
plt.step(
torch.arange(len(cumsum)), cumsum, where="mid",
color="darkorange", label="Cumulative sum"
)
# Highlight cutoff
plt.axhline(
top_p, color="red", linestyle="--",
label=f"top_p = {top_p}"
)
plt.axvline(
n_kept - 0.5, color="gray", linestyle=":",
label=f"Top-p cutoff at {n_kept} tokens"
)
plt.xlabel("Token rank (sorted by probability)")
plt.ylabel("Probability")
plt.legend()
plt.grid(alpha=0.3)
plt.ylim(0, 1.05)
# plt.savefig("14.pdf")
plt.show()

要实现该阈值截断逻辑,可以使用下面这段代码:首先将截断位置(图4.15中的竖直虚线)右侧的全部数值置零,之后恢复原本的排序顺序。
# Step 4.3.2: 将截断点之外的词元概率置零
# keep_mask 为 True 的位置保留原始概率,False 的位置替换为 0
kept_sorted = torch.where(
keep_mask, # 条件:是否保留该词元
sorted_probas, # 条件为 True 时取排序后的概率
torch.zeros_like(sorted_probas) # 条件为 False 时置零
)
# Step 4.3.3: 将结果映射回原始词元顺序
# 之前为了做 top-p 过滤,概率是按从大到小排序过的(对应 sorted_idx)
# 这里用 scatter 按 sorted_idx 把保留下来的概率放回它们在原始序列中的位置
filtered = torch.zeros_like(toy_probas).scatter(
0, # 沿第 0 维进行散射
sorted_idx, # 每个概率值应回填到的原始索引
kept_sorted # 要回填的(已过滤)概率值
)
print(filtered)
得到的张量如下:
tensor([0.0000, 0.0000, 0.0000, 0.0000, 0.4538, 0.0000, 0.0000, 0.0000,
0.0829, 0.2752])
可以看到,除索引位置4、8、9以外,其余全部数值都被置零。这意味着此时调用多项式采样函数,只会从这三个词元当中做选择。
最后对这些数值执行重新归一化,使概率总和再次等于1:
# Step 4.4: 重新归一化,使概率总和为 1
# 经过 top-p 过滤后,部分词元概率被置零,剩余概率之和不再等于 1,
# 因此需要重新归一化,才能作为合法的概率分布进行采样。
# 先求出保留下来的概率之和作为分母
# clamp_min(1e-12) 用一个极小值兜底,防止分母为 0 导致除零错误
denom = torch.sum(filtered).clamp_min(1e-12)
# 每个概率都除以总和,使过滤后的分布重新归一化为总和为 1
renormalized = filtered / denom
print(renormalized)
归一化之后的张量输出:
tensor([0.0000, 0.0000, 0.0000, 0.0000, 0.5589, 0.0000, 0.0000, 0.0000,
0.1021, 0.3390])
top‑p过滤的目的是剔除低概率词元,避免后续采样抽到这些词元。该机制有助于减少特定上下文下无意义的词元输出。
4.5.2 在文本生成函数中加入top‑p过滤器
我们已经通过模拟示例完整讲解了top‑p过滤的各个步骤。现在,我们要把top‑p过滤的四步逻辑(图4.16中的步骤4.1‑4.4)添加到现有的文本生成函数中,放在概率转换与采样操作之间。
首先,我们把上一节介绍的top‑p过滤逻辑封装成一个便于调用的独立函数。

def top_p_filter(probas, top_p):
# 如果未指定 top_p,或 top_p >= 1.0(表示不做过滤,保留全部词元),
# 则直接返回原始概率分布
if top_p is None or top_p >= 1.0:
return probas
# Step 4.1: 按概率从大到小排序
# sorted_probas 是排序后的概率,sorted_idx 记录每个值在原始序列中的索引
sorted_probas, sorted_idx = torch.sort(probas, dim=1, descending=True)
# Step 4.2: 计算累积概率和
cumprobas = torch.cumsum(sorted_probas, dim=1)
# Step 4.3.1: 保留“加入该词元之前”前缀累积质量 < top_p 的词元
# 例如:[0.5, 0.41, 0.09],top_p=0.9 时应保留前两个词元
prefix = cumprobas - sorted_probas # 每个词元之前(不含自身)的累积质量
keep = prefix < top_p
# 至少保留一个词元(当 top_p 极小或非正数时的兜底处理)
keep[:, 0] = True
# Step 4.3.2: 将截断点之外的词元概率置零
# keep 为 True 的位置保留原始概率,False 的位置置为 0
kept_sorted = torch.where(
keep, sorted_probas,
torch.zeros_like(sorted_probas)
)
# Step 4.3.3: 将结果映射回原始词元顺序
# 用 scatter 按 sorted_idx 把保留下来的概率回填到原始位置
filtered = torch.zeros_like(probas).scatter(1, sorted_idx, kept_sorted)
# Step 4.4: 重新归一化,使每一行概率总和为 1
# clamp_min(1e-12) 防止分母为 0 导致除零错误
denom = torch.sum(filtered, dim=1, keepdim=True).clamp_min(1e-12)
return filtered / denom
代码清单4.15中的top_p_filter函数首先对token概率排序,并计算其累积和。随后保留前缀累积概率(每个词元之前的累积概率) 低于top_p阈值的词元,同时把首个越过阈值的词元也一并保留。函数将其余词元概率置零,把保留下来的概率映射恢复为原始顺序,再对剩余概率重新归一化,使概率总和重新等于1。
在把top_p_filter函数接入文本生成函数之前,我们先做测试,观察它和前面的温度缩放配合时的运行效果。首先获取logits:
with torch.inference_mode():
next_token_logits = model(input_token_ids)[:, -1]
print(next_token_logits.shape)
上述代码输出张量next_token_logits的维度为[1, 151936]。我们现在使用真实数据与完整词表,该词表一共有151936个条目。
接下来,将logits转换为概率分值,和之前一样,设置温度为0.35做温度缩放:
torch.manual_seed(123)
probas_lowT = torch.softmax(
scale_logits_by_temperature(next_token_logits, 0.35), dim=-1
)
count_samples(probas_lowT, threshold=1, tokenizer=tokenizer)
这段代码和之前温度缩放示例所用代码相近,打印出的采样输出如下:
' __': 158x
' Berlin': 435x
' ____': 169x
' ______': 209x
' Munich': 3x
' Hamburg': 3x
' _____': 18x
现在加入top‑p过滤器,查看结果发生的变化:
torch.manual_seed(123)
probas_lowT = torch.softmax(
scale_logits_by_temperature(next_token_logits, 0.35), dim=-1
)
probas_lowT_filtered = top_p_filter(probas_lowT, top_p=0.8)
count_samples(probas_lowT_filtered, threshold=1, tokenizer=tokenizer)
top_p=0.8是一个常用典型阈值,该设置会剔除累积概率占剩余20%的那些低概率词元。此时采样输出变为:
' Berlin': 534x
' ____': 217x
' ______': 249x
可以看到,采样候选中已经去掉了“慕尼黑”和“汉堡”,模型只会输出正确城市,或是输出下划线词元;下划线是模型学到的格式,用于测验填空题里的占位符______。
回到数学问题的示例,我们给代码清单4.8实现的generate_text_temp_stream_cache函数增加top‑p过滤器。更新后的新函数命名为generate_text_top_p_stream_cache,见下文。
@torch.inference_mode() # 禁用梯度计算,推理专用,比 no_grad 更快、更省内存
def generate_text_top_p_stream_cache(
model,
token_ids, # 初始输入的 token 序列,形状通常为 (batch, seq_len)
max_new_tokens, # 最多生成的新 token 数量
eos_token_id=None, # 结束符 token id,遇到时提前停止生成
temperature=0., # 温度系数,控制随机性;0 表示贪心解码
top_p=None # top-p(核采样)阈值,只从累积概率达到 top_p 的候选中采样
):
model.eval() # 切换到评估模式(关闭 dropout 等训练专用层)
# 创建 KV 缓存,用于存储各层的 key/value,避免重复计算历史 token,加速生成
cache = KVCache(n_layers=model.cfg["n_layers"])
model.reset_kv_cache() # 重置缓存,确保从干净的状态开始
# 步骤 3.1:先把完整的提示(prompt)喂给模型,取最后一个位置的 logits
# [:, -1] 表示只保留序列最后一个 token 对应的输出,即预测下一个 token 的依据
out = model(token_ids, cache=cache)[:, -1]
for _ in range(max_new_tokens): # 循环生成,每次产生一个新 token
orig_device = token_ids.device # 记录原始设备(cpu/gpu),后面把采样结果搬回来
if temperature is None or temperature == 0.0:
# 温度为 0:贪心解码,直接取概率最大的 token
next_token = torch.argmax(out, dim=-1, keepdim=True)
else:
# 步骤 3.2:用温度系数对 logits 进行缩放(温度越高,分布越平滑越随机)
logits = scale_logits_by_temperature(out, temperature)
# 步骤 3.3:把 logits 转换为概率分布
probas = torch.softmax(logits, dim=-1)
# (新增)步骤 4:应用 top-p 过滤,只保留累积概率在 top_p 内的候选 token
probas = top_p_filter(probas, top_p)
# 步骤 3.4:根据概率分布进行多项式采样,抽取下一个 token
# 放到 cpu 上采样(multinomial 在 cpu 上兼容性/稳定性更好),再搬回原设备
next_token = torch.multinomial(probas.cpu(), num_samples=1)
next_token = next_token.to(orig_device)
# 如果生成了结束符,则提前终止生成
if (eos_token_id is not None
and torch.all(next_token == eos_token_id)):
break
yield next_token # 以流式方式逐个返回生成的 token
# 只把新生成的这一个 token 喂给模型(历史信息已存在 KV 缓存中),得到下一步的 logits
out = model(next_token, cache=cache)[:, -1]
我们把这个新的文本生成函数接入generate_text_stream_concat_flex,就像之前接入带温度缩放的文本生成函数一样:
torch.manual_seed(123)
response = generate_text_stream_concat_flex(
model, tokenizer, prompt, device,
max_new_tokens=2048, verbose=True,
generate_func=generate_text_top_p_stream_cache,
temperature=0.5,
top_p=0.8,
)
输出仍然是 "\boxed{18}",这依然是错误的。我们目前所做的一切都是为了让我们能够采样出不同的输出。在下一节中,我们将使用 generate_text_stream_concat_flex 配合我们增强后的文本生成函数(generate_text_top_p_stream_cache),在实现自一致性(self-consistency)推理时扩展技术时采样出不同的输出。
Top‑k过滤
Top‑k过滤是采样阶段限制候选下一词元集合的另一种方案。
和top‑p(保留累积概率不超过阈值的全部词元)不同,top‑k仅保留基于模型logits得到的概率最高的k个词元。将词表按照概率排序后,剔除前k项之后的所有内容,再对剩余k个词元做重新归一化,之后从中采样。
简言之:top‑k保留固定数量的最高概率词元;top‑p根据累积概率质量,保留数量不固定的词元。
top‑k实现更加简单,我在《Build a Large Language Model (From Scratch)》(《从零构建大语言模型》)一书中已经讲解过。而top‑p采样在当下更为流行。
练习4.2:在MATH‑500数据集上使用温度缩放与top‑p过滤
修改代码清单3.15中的evaluate_math500_stream函数,观察增加温度缩放与top‑p采样是否会改变基础模型在MATH‑500数据集上的准确率。温度与top‑p均可设置为0.9。

4.6 借助自一致性提升回复准确率
我们已经完成温度缩放与top‑p过滤的代码编写,现在可以实现本章第二种推理时缩放技术:自一致性采样。
自一致性采样的核心思想出自谷歌研究院论文《Self‑Consistency Improves Chain‑of‑Thought Reasoning in Language Models》(https://arxiv.org/abs/2203.1117)。虽然名字听起来很精巧,但自一致性采样本质上是一种多数投票机制:利用温度缩放和top‑p过滤生成多份答案,随后选出出现频次最高的答案,具体如图4.17所示。

图4.17中展示的自一致性采样(self-consistency sampling)技术属于一种推理时扩展(inference-time-scaling)技术,因为我们并不更新模型本身,而是通过消耗更多的计算资源来提升响应的准确性。
为了构建基于投票的推理时流水线,我们将复用第2章的生成代码、第3章的答案提取逻辑,以及本章的采样控制。得益于 generate_text_stream_concat_flex 函数,自一致性代码的实现相对简单。其主要流程可以概括为三个步骤,如图4.18所示。
为了实现图4.18中所示的三个步骤,我们只需重复调用 generate_text_stream_concat_flex 来生成多个答案,并复用第3章的 extract_final_candidate 函数。唯一需要新增的代码,就是实现基于答案频率的多数投票(majority voting)。

from reasoning_from_scratch.ch03 import extract_final_candidate
from collections import Counter
def self_consistency_vote(
model, tokenizer, prompt, device,
num_samples=10, temperature=0.8, top_p=0.9, max_new_tokens=2048,
show_progress=True, show_long_answer=False, seed=None,
):
# ============================================================
# 贯穿全程的示例:
# 假设 prompt 是「小明有3个苹果,又买了2个,一共几个?」
# 我们采样 num_samples=5 次
# ============================================================
# full_answers:保存每次采样的完整回答(含推理过程)
# short_answers:保存从每个完整回答中抽取的最终简短答案
full_answers, short_answers = [], []
# 1) 多次采样:对同一 prompt 生成 num_samples 个回答
# 由于 temperature/top_p 采样,每次结果略有不同
for i in range(num_samples):
# 指定种子时,为每次采样设置不同但可复现的种子
if seed is not None:
torch.manual_seed(seed + i + 1)
# 生成一条完整回答
# 示例产出 answer(第 1 次):
# "先有3个,再加2个,3+2=5,所以答案是5"
answer = generate_text_stream_concat_flex(
model=model, tokenizer=tokenizer, prompt=prompt, device=device,
max_new_tokens=max_new_tokens, verbose=show_long_answer,
generate_func=generate_text_top_p_stream_cache,
temperature=temperature, top_p=top_p,
)
# 2) 从完整回答中抽取最终简短答案
# fallback="number_then_full":优先取数字,取不到再回退到完整文本
# 示例产出 short(第 1 次):"5"
short = extract_final_candidate(
answer, fallback="number_then_full"
)
full_answers.append(answer)
short_answers.append(short)
if show_progress:
print(f"[Sample {i+1}/{num_samples}] → {short!r}")
# ------------------------------------------------------------
# 循环结束后,假设 5 次采样抽取到的简短答案为:
# short_answers = ["5", "5", "4", "5", "6"]
# ------------------------------------------------------------
# 3) 自洽性投票:选出出现次数最多的答案
# 统计每个答案出现的次数
# 示例产出 counts:Counter({"5": 3, "4": 1, "6": 1})
counts = Counter(short_answers)
# 按答案分组,记录每个答案对应的采样索引(便于溯源是哪几次得到的)
# 示例产出 groups:{"5": [0, 1, 3], "4": [2], "6": [4]}
groups = {s: [] for s in counts}
for idx, s in enumerate(short_answers):
groups[s].append(idx)
# most_common() 返回按次数从高到低排序的 (答案, 次数) 列表
# 示例产出 mc:[("5", 3), ("4", 1), ("6", 1)]
mc = counts.most_common()
if not mc:
# 采样为空时(理论极端情况)
majority_winners, final_answer = [], None
else:
# 最高出现频次
# 示例产出 top_freq:3
top_freq = mc[0][1]
# 所有达到最高频次的答案(可能平局)
# 示例产出 majority_winners:["5"]
majority_winners = [s for s, f in mc if f == top_freq]
# 只有唯一最高频答案时才确定最终结果;平局则返回 None
# 示例产出 final_answer:"5"
final_answer = mc[0][0] if len(majority_winners) == 1 else None
# 返回结果字典,示例整体产出:
# {
# "full_answers": [完整回答1, ..., 完整回答5],
# "short_answers": ["5", "5", "4", "5", "6"],
# "counts": {"5": 3, "4": 1, "6": 1},
# "groups": {"5": [0, 1, 3], "4": [2], "6": [4]},
# "majority_winners": ["5"],
# "final_answer": "5",
# }
return {
"full_answers": full_answers,
"short_answers": short_answers,
"counts": dict(counts),
"groups": groups,
"majority_winners": majority_winners,
"final_answer": final_answer,
}
简而言之,清单4.17中的自一致性方法执行如下操作:
- 在温度大于0并开启top‑p的条件下采样生成多个答案
- 从每一条回复中提取方框包裹的最终答案
- 选择出现频次最高的最终答案
注意在我们的实现中,使用for循环按顺序生成这些答案。实际工程中,也经常在不同设备上生成答案,从而实现采样过程并行化。
此外,在上述代码中,我们为每一轮分别设置随机种子:
if seed is not None:
torch.manual_seed(seed + i + 1)
严格来说,这一步并非必需;仅设置一次随机种子就足以生成多样化样本。如果需要单独复现其中某几轮结果,这种显式设置种子的方式会很有用。
下面调用该函数,查看输出结果:
results = self_consistency_vote(
model,
tokenizer,
prompt,
device=device,
num_samples=5,
temperature=0.8,
top_p=0.9,
max_new_tokens=2048,
seed=123,
show_progress=True,
)
[Sample 1/5] → '83'
[Sample 2/5] → '22'
[Sample 3/5] → '54'
[Sample 4/5] → '83'
[Sample 5/5] → '61'
由于83是出现次数最多的答案,因此最终结果为83(可通过代码results["final_answer"]获取该数值)。如果想要查看推理过程与完整回答,可以选取输出83的其中一条正确结果;例如执行print(results["full_answers"][0]),会打印出如下内容:
To find the value of \( x \), let's solve the equation step by step.
1. **Given Equation:**
\[
\frac{1}{2} \times (3x - 9) = x + 37
\]
2. **Multiply Both Sides by 2 to Eliminate the Fraction:**
\[
2 \times \frac{1}{2} \times (3x - 9) = 2 \times (x + 37)
\]
\[
3x - 9 = 2x + 74
\]
3. **Subtract \( 2x \) from Both Sides to Get:**
\[
3x - 2x - 9 = 74
\]
\[
x - 9 = 74
\]
4. **Add 9 to Both Sides to Solve for \( x \):**
\[
x = 74 + 9
\]
\[
x = 83
\]
**Final Answer:**
\[
\boxed{83}
\]
借助这套自一致性方案,我们终于得到了正确答案。
(注意:在 mps 或 cuda 设备上运行代码,得到的结果可能会存在差异。)
目前self_consistency_vote函数无法处理票数相同的平局情况:如果多个答案出现频次一致,函数会将最终答案返回为None。下一章我们将实现一套打分机制,用于计算每个答案的确信度,以此作为平局时的判定依据。
习题4.3:在MATH‑500数据集上使用自一致性采样
修改清单3.15中的evaluate_math500_stream函数,评估自一致性采样能否提升基础模型在MATH‑500数据集上的准确率。采样数量设为3,温度与top‑p均设置为0.9。
作为本习题要求,需要实现平局处理逻辑:出现票数相同时,选取采样列表中最先出现的答案。例如采样得到的答案为13、15、13、15、16,则应当选定答案13。
注意:实现该平局规则无需直接修改
self_consistency_vote函数本身,可以基于该函数返回的结果字典来完成处理。


习题4.4:自一致性采样中的提前终止
为提升计算效率,实现自一致性的提前终止版本:当超过半数采样输出同一个答案时,直接停止后续采样。

温度与top‑p参数的选择
前面已经介绍过温度参数,实际使用时需要确定自一致性采样应当引入多大的随机性。我们的目标不是让模型输出尽可能随机,而是生成具备适度多样性的答案,让多数投票可以发挥作用,同时保证绝大多数采样结果逻辑合理。
工程实践中,温度取0.5‑0.9、top‑p取0.7‑0.9是比较合适的初始参数。
经验参考:
如果全部完整输出几乎一模一样,可以适度调高温度或top‑p,以此增加结果多样性。
如果输出开始出现逻辑错乱、无意义内容,代表参数设置过于激进,需要调小,优先降低温度。
此处所说的“完整输出”,指提取方框最终答案之前的全部推理文本。可以查看self_consistency_vote返回结果字典里的完整输出;也可以调用函数时开启show_long_answer=True来查看。
大家或许还记得,介绍该技术的论文标题为《自一致性提升大语言模型的思维链推理能力》。那其中的思维链推理是如何发挥作用的呢?这里的思维链推理,就是本章前面介绍的思维链提示词技术:我们通过在提示词末尾添加"\n\n Explain step by step.",让大模型生成更长的回复。
下面我们将思维链提示词与自一致性采样结合起来:
results = self_consistency_vote(
model,
tokenizer,
prompt + "\n\nExplain step by step.",
device=device,
num_samples=5,
temperature=0.8,
top_p=0.9,
max_new_tokens=2048,
seed=123,
show_progress=True,
)
[Sample 1/5] → 'x = 83'
[Sample 2/5] → '83'
[Sample 3/5] → '83'
[Sample 4/5] → '83'
[Sample 5/5] → '83'
在本例中,五次采样得到的答案全部为83(正确),这很可能是因为借助思维链之后,该问题对大语言模型而言难度较低。如果观察模型在整个MATH‑500数据集上的整体表现,会更有参考意义。各类实验的结果如表4.1所示。
为简洁起见,表 4.1 中标注为"Top-p"的方法同时使用了温度缩放(temperature scaling)和 top-p 采样。表中所示的准确率值是在 MATH-500 测试集的全部 500 个样本上计算得出的,使用了一块 “cuda” GPU(DGX Spark)。
让我们逐一分析这些结果。n = 3 这一缩写表示我们在自洽性采样(self-consistency sampling)中使用了大小为 3 的样本量。
n = 3 就表示对每一道题让模型生成 3 次(采样 3 条不同的推理路径)
第 1 行和第 2 行展示了使用第 3 章代码的基础变体(base)和推理变体(reasoning)的结果;也就是说,我们使用了不带温度缩放或 top-p 过滤的文本生成函数。这个文本生成函数只是在每一步选择得分最高的 token(贪婪解码,greedy decoding)。我们可以看到,推理变体的准确率大约是基础变体的三倍,但运行时间也长得多,因为它生成了更多的 token。
在第 3 行中,我们看到了使用 “\n\nExplain step by step.” 提示修改的思维链(chain-of-thought)提示方法的结果。如你所见,这将基础模型的准确率从大约 15% 提升到了 40%。
第 4 行展示了当我们向基础模型添加温度缩放和 top-p 采样时会发生什么。(表 4.1 中所有涉及温度缩放和 top-p 采样的实验都将两者设置为 0.9。)相比第 1 行的基础模型,准确率仅有小幅提升,从 15.2% 提高到 17.8%。这是意料之中的,因为温度缩放和 top-p 采样只是帮助我们控制采样多样性,其本身并不是推理时扩展(inference-time-scaling)技术。我们还可以看到,运行时间从 10.1 分钟增加到了 30.7 分钟。这并不是由于采样代码的开销,而是因为模型现在在某些情况下会生成更长的响应。
贪婪解码总是选最高分的 token,模型想结束时会直接选中结束符(EOS),所以停得早、输出短。
而温度缩放和 top-p 采样引入了随机性,即使 EOS 得分最高,也可能随机抽中别的 token,导致模型错过结束时机、继续生成更长的内容。
token 生成得越多,计算步数越多,时间自然就越长——所以慢不是采样代码的开销,而是响应变长了。
第 5–7 行展示了自洽性扩展的结果。将样本数量从 3 增加到 10 进一步将准确率提升到了 31.6%,但也显著增加了运行时间。在这种情况下,使用 10 个样本相比 3 个样本几乎没有准确率上的优势。
第 8 行展示了温度缩放和 top-p 采样与思维链提示相结合的结果。在这种情况下,采样反而使思维链的结果变差了(33.4%,相比第 3 行的 40.6%)。将思维链提示与自洽性相结合,在样本量为 10 时将准确率提升到了 52%,但也大幅增加了运行时间,达到了惊人的 862.6 分钟。
在实践中,如果你能使用多块 GPU,自洽性采样中的不同样本可以并行计算,而不是顺序计算。这仍然会使用相同的计算量,但可以分散到多块 GPU 上,从而更快地生成结果。
最后,在第 12 行中,我们可以看到推理变体同样受益于自洽性采样,其准确率从第 2 行的 48.2% 提升到了 55.2%。同样,这也伴随着运行时间的增加。
要点在于:表 4.1 中的结果凸显了推理时扩展的权衡——用更多的计算换取更高的准确率。
本章中自洽性采样方法的一个主要缺点是,它依赖于一个可以提取出来用于多数投票的最终框选答案(boxed answer)。这种方法较难应用于那些没有数值或简短最终答案的问题。
在下一章中,如图 4.19 所示,我们将实现一种不同且更通用的推理时扩展方法,称为自我完善(self-refinement),在这种方法中,模型会迭代式地改进自己的答案。

总结
- 无需重新训练模型,通过在推理阶段增加计算量(推理时扩展,inference-time scaling),即可提升推理能力和答案准确率。
- 灵活的文本生成包装器(generate_text_stream_concat_flex)允许在不改动周边代码的情况下,插入不同的采样策略。
- 下一个 token 通过 softmax 从 logits 中生成。
- 温度缩放(Temperature scaling)通过改变 logits 来控制生成文本的多样性。
- Top-p(核采样,nucleus sampling)会过滤掉低概率的 token,从而降低生成无意义答案的几率。
- 思维链提示(Chain-of-thought prompting,例如“请逐步解释”)通常能得到更准确的答案,因为它鼓励模型写出中间推理过程;不过这会增加生成的 token 数量,进而提高运行时的开销。
- 自洽性采样(Self-consistency sampling)会生成多个答案,从每个答案中提取最终的方框(boxed)结果,并通过多数投票选出出现频率最高的答案,从而提升答案准确率。
- 在 MATH-500 数据集上的实验表明,将思维链提示与自洽性采样相结合,相比无采样的基线可以大幅提升准确率,但代价是运行时间大幅增加。
- 推理时扩展的核心权衡在于:以更多的计算量换取更高的准确率。
更多推荐




所有评论(0)