happy-llm 实战:基于 vLLM 实现 s1 论文的 Test-time Scaling 思考预算(Thinking Budget)
【免费下载链接】happy-llm📚 从零开始构建大模型项目地址: https://gitcode.com/GitHub_Trending/ha/happy-llm
本指南以 s1-vllm-thinking-budget 实验文档 为主体,完整复现并深度讲解如何用 vLLM 实现《s1: Simple test-time scaling》论文提出的"思考预算"(Thinking Budget)机制:通过迭代生成 + 强制插入
Wait!等特定词,把模型的思考 token 数量控制在预算之内。读者将掌握s1.py中每一行核心代码的作用、SamplingParams的配置含义、完整的运行方式,以及真实实验中出现的问题与结论。
一、背景:从 s1 论文看 Test-time Scaling
1.1 论文核心思想
《s1: Simple test-time scaling》 由李飞飞教授团队提出,核心是测试时间缩放(Test-time scaling):在推理阶段通过调整模型的"思考预算"(即允许模型思考/生成推理链的 token 数量),来提高复杂问题的推理效率与准确性。
简单来说:对于需要多步推理链才能解决的问题,模型的性能与思考预算存在正相关关系——当思考预算增加时,模型的推理准确率会有明显提升。
如上图所示,在 MATH500、AIME24、GPQA Diamond 等推理密集型任务上,随着平均思考时间的增加,模型准确率逐步提升。这也解释了为什么"让模型多想一会儿"能带来更好的推理效果。
1.2 论文中的两个关键技巧
论文在实现中还引入了两个重要的设计,本仓库实验同样沿用了这些思路:
(1)用模型回答对错来判断问题难易。论文通过让 Qwen2.5-32B-Instruct 回答问题来区分任务难度:模型答对的是简单问题,答错的是复杂问题。从仓库中的相关说明可以看出,数据筛选层面也使用了类似策略——利用 Qwen2.5 系列模型过滤掉简单样本,从而让思考预算更集中于复杂任务。
(2)插入特定词(如Wait!)强制模型继续思考。论文做了消融实验,探讨在未满足思考预算时插入不同的特定词对模型性能的影响。结果表明,插入特定词可以有效地引导模型进行更深入的思考,其中"Wait,Wait"的效果最好。
上表正是论文中关于"预算强制外推(Budget forcing extrapolation)"的消融实验(Table 4),对比了"不追加字符串(No extrapolation)""2 倍预算不加字符串""追加 Alternatively/Hmm/Wait 等特定词"等多种策略。可以看到,不同的字符串追加方式对准确率有明显影响,这也为后续 vLLM 实现中的Wait!注入提供了理论依据。
二、总体思路:vLLM 如何实现思考预算
vLLM 是一个高性能推理引擎,支持大规模语言模型的高效推理。本仓库使用 vLLM 来实现论文中的思考预算机制。整体流程如下:
对比图左侧是不使用思考预算的推理过程(Prompt → 构建输入 → 生成 → 响应);右侧是使用思考预算的推理过程(Prompt → 构建输入 → 生成 → 检查思考 token 数是否超过预算 → 未超过则在文本末尾追加Wait!重新生成 → 最终响应)。可以看到,使用思考预算后,模型会在推理过程中插入特定词来引导自己进行更深入的思考。
环境提示:考虑到部分同学配置环境可能会遇到问题,作者在 ucloud 平台准备了环境镜像,可直接创建 ucloud 实例使用(镜像链接见原文档)。
三、核心代码实现:逐步拆解 s1.py
完整的可运行代码在 s1.py,下面按模块逐步拆解。
3.1 依赖与工具函数
代码依赖vllm与transformers:
from vllm import LLM, SamplingParams from transformers import AutoTokenizer import time其中:
LLM:vLLM 的模型加载与推理入口;SamplingParams:采样参数配置对象,控制temperature、max_tokens、stop等生成行为;AutoTokenizer:用于 token 计数与 chat template 构建。
构建输入。模型使用 chat template 构造输入,build_input使用tokenizer.apply_chat_template处理 system/user 消息,并显式开启思考模式:
def build_input(prompt, tokenizer): messages = [ {"role": "system", "content": "Please reason step by step, and put your final answer within \\boxed{{}}."}, {"role": "user", "content": prompt} ] input_text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, enable_thinking=True ) return input_text其中enable_thinking=True会启用模型的思考标签(如<think>/</think>),这是后续统计"思考 token"的基础;system 提示要求模型逐步推理并用\boxed{}输出最终答案。
token 统计。思考 token 数的统计逻辑是:把 prompt 与生成结果拼接后,截取<think>\n之后的内容进行 token 计数:
def count_thinking_token(outputs, tokenizer): total_token = outputs[0].prompt + outputs[0].outputs[0].text thinking_token = total_token.split("<think>\n")[-1] thinking_token_id = tokenizer(thinking_token)["input_ids"] return total_token, len(thinking_token_id) def count_token(string, tokenizer): return len(tokenizer(string)["input_ids"])3.2 主函数:run_thinking_budget_sample
def run_thinking_budget_sample(llm_model, tokenizer, user_input, thinking_budget): input_text = build_input(user_input, tokenizer) input_token_count = count_token(input_text, tokenizer) iteration_count= 0 max_token = input_token_count + thinking_budget sampling_params = SamplingParams( temperature=0.7, max_tokens=4096, skip_special_tokens=False ) think_token_count = 0 while True: wait_sampling_params = SamplingParams( temperature=0.7, max_tokens=thinking_budget - think_token_count, stop='</think>', skip_special_tokens=False ) outputs = llm_model.generate( input_text, wait_sampling_params ) total_token, think_token_count = count_thinking_token(outputs, tokenizer) print(f'第{iteration_count}次迭代,思考token数:{think_token_count}') if think_token_count > thinking_budget: break input_text = total_token + "\nWait!\n" # \nWait a moment. Was there any loophole in my thought just now?!\n # \nWait!\n iteration_count += 1 final_outputs = llm_model.generate( outputs[0].prompt + outputs[0].outputs[0].text + "\n</think>\n", sampling_params ) total_content = final_outputs[0].prompt + final_outputs[0].outputs[0].text thinking_content = total_content.split("<think>")[-1].split("</think>")[0] print(total_content) print(f"迭代次数:{iteration_count}, 输入token数:{input_token_count}, 思考token数:{count_token(thinking_content, tokenizer)}, 总token数:{count_token(total_content, tokenizer)}")第一步:参数准备。函数接收模型、tokenizer、用户输入和思考预算(thinking_budget,即允许的思考 token 上限)四个参数。先构建输入文本并计算输入的 token 数量;同时为最终答案生成准备好独立的sampling_params。
因为max_tokens参数表示"生成的最大 token 数量",所以需要把思考预算转换成每次迭代的生成上限:每次迭代允许模型新生成的 token 数为thinking_budget - think_token_count(总预算减去已思考的 token 数)。同时,需要在SamplingParams中设置stop='</think>',这样模型在生成到</think>时会自动停止,便于分轮统计思考内容。
wait_sampling_params = SamplingParams( temperature=0.7, max_tokens=thinking_budget - think_token_count, stop='</think>', skip_special_tokens=False )第二步:循环生成与 Wait! 注入。核心循环逻辑如下:
while True: wait_sampling_params = SamplingParams( temperature=0.7, max_tokens=thinking_budget - think_token_count, stop='</think>', skip_special_tokens=False ) outputs = llm_model.generate( input_text, wait_sampling_params ) total_token, think_token_count = count_thinking_token(outputs, tokenizer) print(f'第{iteration_count}次迭代,思考token数:{think_token_count}') if think_token_count > thinking_budget: break input_text = total_token + "\nWait!\n" # \nWait a moment. Was there any loophole in my thought just now?!\n # \nWait!\n iteration_count += 1每次迭代中:
- 用当前已思考 token 数更新
max_tokens上限; - 调用
llm_model.generate生成一段思考内容; - 统计累计思考 token 数并打印日志;
- 若累计思考 token 数超过思考预算,则跳出循环;否则把已生成的完整文本作为新的输入,并在末尾追加
\nWait!\n,引导模型继续深入思考。
注释中还保留了两个可替换的提示词变体(\nWait a moment. Was there any loophole in my thought just now?!\n与\nWait!\n),方便读者对照论文消融实验做不同字符串的 A/B 测试。从源码结构看,仓库最终选用了更简洁的\nWait!\n。
第三步:生成最终答案。当思考达到预算后,模型还需要把思考过程总结成最终答案,这一步同样需要单独的采样参数:max_tokens设置为 4096。如原文档所述,模型根据思考过程进行总结得出答案也需要很多 token,这个值设置为多少都可以,通常设置为一个较大的值即可:
sampling_params = SamplingParams( temperature=0.7, max_tokens=4096, skip_special_tokens=False )随后在已生成的思考内容末尾补上\n</think>\n闭合思考标签,再次调用llm_model.generate生成最终答案:
final_outputs = llm_model.generate( outputs[0].prompt + outputs[0].outputs[0].text + "\n</think>\n", sampling_params ) total_content = final_outputs[0].prompt + final_outputs[0].outputs[0].text thinking_content = total_content.split("<think>")[-1].split("</think>")[0] print(total_content) print(f"迭代次数:{iteration_count}, 输入token数:{input_token_count}, 思考token数:{count_token(thinking_content, tokenizer)}, 总token数:{count_token(total_content, tokenizer)}")最后打印完整输出,并输出关键指标:迭代次数、输入 token 数、思考 token 数、总 token 数。s1.py中还额外实现了把结果写入output_{int(time.time())}.txt文件的逻辑,便于留存实验记录。
3.3 对照实验:run_sample(无思考预算)
为了对比思考预算的效果,s1.py还提供了一个不使用思考预算的基准函数run_sample:直接以max_tokens=32768生成,不做Wait!注入,也不做思考 token 控制:
def run_sample(llm_model, tokenizer, user_input): input_text = build_input(user_input, tokenizer) input_token_count = count_token(input_text, tokenizer) sampling_params = SamplingParams( temperature=0.7, max_tokens=32768, skip_special_tokens=False ) final_outputs = llm_model.generate( input_text, sampling_params ) total_content = final_outputs[0].prompt + final_outputs[0].outputs[0].text thinking_content = total_content.split("<think>")[-1].split("</think>")[0] print(total_content) print(f"输入token数:{input_token_count}, 思考token数:{count_token(thinking_content, tokenizer)}, 总token数:{count_token(total_content, tokenizer)}")从源码结构看,该函数与run_thinking_budget_sample的差异正是思考预算机制的全部差异:前者"一次生成到底",后者"分轮迭代 + 强制注入Wait!"。
3.4 主程序入口与模型配置
if __name__ == "__main__": model_path = "/model/ModelScope/Qwen/Qwen3-14B" tokenizer = AutoTokenizer.from_pretrained(model_path) llm = LLM( model=model_path, gpu_memory_utilization=0.9, trust_remote_code=True ) print("=================================== 思考预算采样 ===================================") run_thinking_budget_sample( llm_model=llm, tokenizer=tokenizer, user_input="There are exactly three positive real numbers $ k $ such that the function\n$ f(x) = \\frac{(x - 18)(x - 72)(x - 98)(x - k)}{x} $\n defined over the positive real numbers achieves its minimum value at exactly two positive real numbers $ x $. Find the sum of these three values of $ k $.", thinking_budget=32768 )主程序中的关键配置:
| 配置项 | 值 | 说明 |
|---|---|---|
model_path | /model/ModelScope/Qwen/Qwen3-14B | 实验所用模型为 Qwen3-14B(约 14B 参数),通过 ModelScope 路径加载 |
gpu_memory_utilization | 0.9 | 允许 vLLM 使用 90% 的 GPU 显存 |
trust_remote_code | True | 信任远端代码,允许加载模型自定义实现 |
thinking_budget | 32768 | 思考预算为 32768 个 token |
| 测试题目 | 一道求三正实数 k 之和的数学题 | 属于需要长推理链的复杂问题,适合验证思考预算效果 |
注意:
model_path为作者实验环境的本地路径,实际使用时需替换为自己环境中模型的实际存放路径,并确保 GPU 显存足够(gpu_memory_utilization=0.9意味着显存占用较高)。
四、结果分析:思考预算的实际效果与局限性
使用思考预算后,模型在推理过程中能够更深入地思考问题,从而提高推理效率和准确性。从 output 目录 中的实验记录可以看到真实运行结果:
| 输出文件 | 迭代次数 | 输入 token 数 | 思考 token 数 | 总 token 数 | 模型给出的答案 |
|---|---|---|---|---|---|
| output_1754208752.txt | 8 | 109 | 32785 | 33755 | 45 + 58 + 85 = 188 |
| output_1754209653.txt | 17 | 109 | 32772 | 33697 | 44 + 126 + 152 = 322 |
从运行记录看,两次实验在 32768 的思考预算下分别迭代了 8 次和 17 次,思考 token 数最终均超过了预算(分别为 32785 和 32772),总 token 数维持在 3.3 万左右。这也印证了文档中的观察:模型在思考过程中可能出现重复生成相同内容,导致思考 token 数量超过思考预算的情况。
此外,实验还发现了一些有趣的现象:
(1)Wait!未必触发真正的"反思"。在某些情况下,就算插入了Wait!,模型并不会按照论文中所示进行多种不同方式的解答尝试,或是反思之前的思考过程是否正确。
(2)重复思考导致预算超支。模型会在思考过程中重复生成相同的内容,导致思考 token 数量超过思考预算。
如上图所示,模型的思考过程中会出现Wait a moment. Was there any loophole in my thought just now?!这类循环式自我检查语句——虽然表面上在"反思",但实际可能只是在重复确认已有的思路,而非真正换一种思路求解。
(3)强插特定词可能"一条道走到黑"。经过测试,强行使用特定词(如Wait!)来引导模型进行更深入的思考,可能会促使模型产生"一条道走到黑"的想法——即沿原有错误思路继续深入,而不是跳出来换一个方向。
当然,也有一个客观原因值得考虑:本实验使用的模型只有 14B 参数(Qwen3-14B),思考过程中的推理能力可能受到模型规模限制。
五、可复现实验建议
基于以上分析,读者在实际使用中可以参考以下建议:
- 调整思考预算:
thinking_budget可以根据任务复杂度调整,简单问题可以设小(如 2048~8192),复杂数学推理题可以设大(如本实验的 32768)。论文实验中的思考预算一般取 512、1024、2048、4096、8192 等档位。 - 替换注入词:代码注释中提供了
\nWait a moment. Was there any loophole in my thought just now?!\n等变体,可以对照论文消融实验对比不同特定词的效果(论文结论是"Wait,Wait"效果最好)。 - 更换模型与显存:主程序中
gpu_memory_utilization=0.9适用于显存充足的环境;模型路径需替换为本地实际路径。从output记录的运行结果看,14B 模型在长思考预算下仍可能出现重复思考,若追求更高质量的推理,可尝试更大规模的模型。 - 保存实验记录:
s1.py已内置将输出与统计信息写入output_{时间戳}.txt的逻辑,建议每次实验后保留输出文件,便于对比不同配置的效果。
六、小结
本文以 s1-vllm-thinking-budget 文档 为主线,结合 s1.py 完整源码,系统讲解了:
- Test-time scaling 与思考预算的理论来源:源于 s1 论文,通过增加思考 token 提升复杂推理任务的准确率;
- vLLM 的完整实现方案:迭代生成 +
stop='</think>'+Wait!注入 + 预算检查的闭环流程,以及两个SamplingParams各自的作用; - 真实实验的量化结果与观察:8~17 次迭代、约 3.3 万思考 token 的完整记录,以及"重复思考""一条道走到黑"等局限性。
思考预算机制的关键价值在于:它把"让模型多想一会儿"从不可控的直觉,变成可配置、可度量的工程参数。读者可以基于本仓库代码继续调整预算、注入词和模型规模,进一步探索 Test-time scaling 的上限。
【免费下载链接】happy-llm📚 从零开始构建大模型项目地址: https://gitcode.com/GitHub_Trending/ha/happy-llm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考