推理与部署——让模型开口说话(GPT 训练分步学习·第五步)

本文是「GPT 训练分步学习路径」的第五步,也是最后一步,基于 GPT_teacher-3.37M-cn 这个教学项目,讲清楚「训练好、验收过的模型,怎么一个字一个字地生成回答,又怎么包成一个能聊天的网页」。文中所有代码均来自本项目真实的 src/infer.pysrc/web_demo.py

分步学习路径共五步:① 看懂数据长什么样 → ② 模型核心(注意力+因果掩码)→ ③ 训练循环 → ④ 验收评估 → ⑤ 推理与部署。本文是第五步。


目录

  1. 这一步属于哪个环节
  2. 推理的本质:一个字一个字地续写
  3. 完整调用链路:从提示词到回答
  4. 采样参数:模型怎么挑下一个字
  5. KV Cache:为什么第二个字开始只喂一个字
  6. 什么时候停:四种停止条件
  7. 部署层做的三件事
  8. 和工业界生产做法的差距
  9. 核心结论速览

1. 这一步属于哪个环节

整个项目的全流程是:

1
构建分词器 → 数据处理 → 训练 → 验收 → Web 演示

本文讲的是 推理与部署 环节,对应代码 src/infer.py(推理内核)和 src/web_demo.py(部署外壳)。它的作用只有一个:

拿验收通过的 best.pt,接收用户的一句话,让模型自回归地一个字一个字生成回答,再包成一个网页让人能真正聊起来。

这一步是全流程的终点——前四步所有的努力(分词、数据、训练、验收),最终都是为了这一刻模型能开口说话。


2. 推理的本质:一个字一个字地续写

训练时模型学会的能力只有一个:看到前文,预测下一个字。推理就是把这个能力反复用起来:

1
2
3
4
5
输入「用户:你是谁\n助手:」
→ 模型预测下一个字「我」
→ 把「我」接到后面,再预测「是」
→ 接「是」,再预测「一」
→ ... 直到预测出「结束符」或撞到停止条件

这种「预测一个字→接上→再预测」的循环叫 自回归生成(autoregressive)。

打个比方:像玩「文字接龙」,每次只添一个字,添完再看整句重新想下一个字。模型本身没有记忆,它每一步看到的都是「到目前为止的完整文本」。

对应到代码,就是 generate() 里那个 for 循环 for step in range(max_new_tokens)


3. 完整调用链路:从提示词到回答

用户的一句话,在 generate() 里要走完这样一条链路:

阶段 代码位置 次数 做什么
① 拼提示词 infer.py:58 1 次 拼成 [BOS]用户:{问题}\n助手:
② 分词编码 infer.py:58 1 次 文本转成一串 token id
③ 推理循环 infer.py:66 N 次 每次生成 1 个字
④ 解码 infer.py:133 1 次 token id 转回文本
⑤ 返回 infer.py:137 1 次 攒完整句一次性 return

注意第 ③ 步是循环 N 次的核心,而 ①②④⑤ 都只做 1 次。

一个关键点:这是「非流式」返回。 虽然内部逐字循环了 N 次,但函数是攒完整段回答后 一次性 return text,不是生成一个字就吐一个字给你看。像 ChatGPT 那种「打字机效果」需要额外的流式改造,本项目没做。

提示词的拼法也有讲究——看 infer.py:58

1
prefix = [tok.bos_id] + tok.encode("用户:" + norm + "\n助手:", add_special_tokens=False)

结尾特意只写到「助手:」,逼模型去续写助手的话,这和训练数据的格式严格对齐。


4. 采样参数:模型怎么挑下一个字

模型每一步输出的不是一个确定的字,而是一张「概率表」(词表里每个字的概率)。怎么从这张表里挑一个字出来,就是采样参数在管的事。

infer.py:74 往后的逻辑:

参数 代码位置 作用 类比
temperature infer.py:74 温度,越低越保守越确定 调「胆子大小」
top_k infer.py:100 只从概率最高的 k 个字里挑 只在前 k 名里抽
top_p infer.py:105 只从累计概率达 p 的字里挑 只在「够格」的候选里抽
repetition_penalty infer.py:89 对最近 32 个字降权,防复读 「刚说过的别老说」

其中最关键的是温度这个分岔口 infer.py:96

1
2
3
4
5
if temperature <= 0:
next_id = torch.argmax(logits, ...) # 贪心:永远选概率最大的
else:
...
next_id = torch.multinomial(probs, 1) # 采样:按概率随机抽
  • temperature=0:贪心解码,每次都选概率最大的字,结果可复现——第四步验收就用这个(同一道题永远同一个答案,才能判分)。
  • temperature>0:按概率随机抽,有随机性——网页演示用这个(回答更有「人味」,不呆板)。

还有几个「安全阀」无论采不采样都生效:屏蔽 pad/bos/unk 特殊符号(infer.py:77)、强制最少生成几个字防秒断(infer.py:85)。


5. KV Cache:为什么第二个字开始只喂一个字

这是推理里最重要的一个优化。看 infer.py:68 这个 if/else:

1
2
3
4
if step == 0:
cur_input = x # 第一步:喂全量提示词
else:
cur_input = x[:, -1:] # 之后每步:只喂最新那一个字

为什么可以只喂一个字? 因为前面所有字算出来的注意力中间结果(Key 和 Value)被缓存下来了。看模型内部 model.py:94 附近:新字的 K/V 会和缓存里的历史 K/V 拼起来,历史部分不用重算。

打个比方:像做数学题打草稿。没有缓存时,每算一步都要把前面所有步骤从头抄一遍再往下算;有了草稿纸(KV Cache),前面的中间结果记在纸上,每步只算新增的那一点。

收益:把每步的计算量从「重算整句」降到「只算一个字」,复杂度从 O(n²) 降到 O(n),句子越长省得越多。

一个容易搞混的点——这个缓存的作用域

  • 它只在单次 generate() 调用内有效。infer.py:62 每次进函数都 kv_caches = None 重置,函数返回后这个局部变量就被丢弃。
  • 所以跨轮对话不复用。多轮对话靠 build_multi_turn_prompt() 把全部历史重新拼一遍、从头再算,对话越长越慢。
  • 这和商业 API 的「前缀缓存 / Prompt Caching」(跨请求持久化、相同前缀命中打折)不是一回事——同源不同域,本项目只做了最基础的单次生成内缓存。

6. 什么时候停:四种停止条件

模型不会自己知道该收尾,得靠停止条件掐断。generate() 里有四道闸:

停止条件 代码位置 说明
生成够了 max_new_tokens infer.py:66 for 循环上限,硬顶
遇到结束符 EOS infer.py:125 模型自己觉得说完了
撞到停止串 infer.py:127 结尾出现指定字串就停
最少长度保护 infer.py:85 反向条件:太短时禁止 EOS

其中「停止串」在部署时特别有用。看 web_demo.py:91

1
stop_strings=["用户:", "\n用户", "。", ";"]

"用户:" 是防自问自答——模型续写时可能自己又编出「用户:…」继续对话,撞到这个串就立刻停,只留助手这一句。句号分号则让回答干净利落收在一句话。


7. 部署层做的三件事

部署层 src/web_demo.py 没有发明任何新的推理逻辑——它开头就 from src.infer import generate,只做三件事:

① 拼历史——build_multi_turn_prompt()

把多轮对话拼成模型认识的格式,因为模型无记忆,每轮都得把全部历史重新喂一遍:

1
用户:Q1\n助手:A1\n用户:Q2\n助手:

② 调 generate 并做兜底——do_generate()

薄封装 generate() + 计时 + 空回复兜底。当模型吐空串或只吐一个标点时,web_demo.py:97 随机挑一句 FALLBACK_ANSWERS 卖萌话术顶上(纯 UI 体验设计,不是模型能力)。

③ 包界面——gr.Blocks + demo.queue().launch

用 Gradio 搭一个聊天网页,跑在 7860 端口。这里还有两个教学彩蛋:

  • 置信度web_demo.py:123):把每个字的 softmax 概率取平均,>0.8 标「高」、>0.5 标「中」、其余「低」,让你直观看到模型「有多确定」。
  • 自洽性检测web_demo.py:146):同一道题用 temperature=0.5 连问 5 次,看几次答案一致。这和第四步验收的 temperature=0 恰好相反——故意用随机性去暴露模型的不确定:稳定的问题 5 次都一样,含糊的问题每次都变。

8. 和工业界生产做法的差距

本项目是教学项目,很多地方为了讲清原理而刻意手写。和工业界的生产系统对照,边界很清晰:

维度 本项目 工业界生产
推理内核 手写 for 循环 推理引擎(vLLM / TGI / SGLang / llama.cpp)
返回方式 非流式,攒完再返回 流式(打字机效果)
缓存 单次生成内 KV Cache + 跨请求前缀缓存持久化
API 无 OpenAI 兼容接口 OpenAI 兼容 REST API
部署位置 Python 后端(Gradio) 也可纯前端(transformers.js + ONNX + WebGPU)
吞吐优化 连续批处理 / PagedAttention / 量化 / 张量并行

类比:本项目的推理像自行车——结构全裸露、每个零件都看得见,适合学原理;生产引擎像高铁——把所有优化封装进黑盒,追求极致吞吐。学会了骑自行车,才更容易理解高铁到底在优化什么。

顺带一提:本项目这个 3.37M 的小模型,参数量极小,转成 ONNX 后其实天然适合 transformers.js,完全能在浏览器里纯前端跑起来,做一个无需服务器的 demo。


9. 核心结论速览

问题 一句话答案
推理的本质是什么 自回归——一个字一个字地续写,模型本身无记忆
输出是流式的吗 不是,内部循环 N 次但攒完整句一次性返回
采样参数管什么 从概率表里怎么挑字:温度/top_k/top_p/重复惩罚
验收和演示的温度为何不同 验收 temperature=0 求可复现;演示 >0 求人味
KV Cache 省了什么 第二步起只算新字,O(n²)→O(n)
KV Cache 跨轮复用吗 不,只在单次 generate 内有效,多轮靠重拼历史
部署层做了什么 只做三件事:拼历史 / 调 generate 兜底 / 包 Gradio 界面
停止条件有哪些 达上限 / 遇 EOS / 撞停止串 / 最短长度保护
和生产系统差在哪 无推理引擎、无流式、无前缀缓存、无 OpenAI API

至此,「GPT 训练分步学习路径」五步全部走完:从数据长什么样,到模型核心机制,到训练循环,到验收判卷,再到今天的推理与部署,已经完整走过了一个微型 GPT 从数据到开口说话的全流程。