数据截至 (上游 commit 92d63d4e8bb4)
04 · Engine 推理与工具调用
这一章讲什么:
nanochat/engine.py怎么把「训练好的 GPT」变成「能并发采样、会用计算器、能多轮对话」的推理引擎;外加execution.py的代码沙箱和scripts/infer_bench.py的 roofline 基准。
1. 它要解决的小问题
训练好的模型只会「给一段 token 序列,算下一个 token 的 logits」。对话 / 评测 / RL 需要的是:
- 快:不能每个 token 都重算整段前缀——要 KV cache。
- 并发采样:同一 prompt 要采 N 个答案(RL 一题 16 个),不能重复 prefill N 次。
- 工具:模型写到一半可能要算
237*41,引擎得接住、算完、把结果「强制注入」回生成流。
2. KVCache:FA3 原生布局
直觉
KV cache 就是「每层每个历史 token 的 k/v 向量」的缓存。nanochat 的 KVCache 有两个不寻常之处:布局跟着 FA3 走 (B, T, H, D)(而非常见的 (B, H, T, D)),且 cache 的写入由 FA3 kernel 在注意力计算时原地完成——上层不用管插入。
真实实现
- 预分配两块
(n_layers, B, T, H, D)全零张量 + 一个 int32 的cache_seqlens记录每行当前长度(KVCache.__init__nanochat/engine.py:92-104)。 - 模型 forward 里,注意力层
get_layer_cache(layer_idx)拿到本层视图,直接调flash_attn_with_kvcache(..., k=k, v=v, cache_seqlens=...);最后一层处理完才advance(T)统一推进位置(nanochat/gpt.py:113-124)。 - SDPA 兜底路径在 Python 侧手工完成同样的插入与窗口裁剪(
nanochat/flash_attention.py:163-185)——上层零感知。 prefill(other):把另一个 cache 的内容整体拷入(nanochat/engine.py:123-137)——这是「一次 prefill、N 路克隆」的关键;连 smear 的prev_embedding也一起 expand 复制。
3. 生成循环:prefill 一次,克隆 N 份
图示
prompt ──► batch=1 prefill(建小 cache)──► logits expand 成 N 份
│
克隆 ──► batch=N decode cache ────┤
▼
循环:采样一列 token ──► 逐行状态机(工具?)──► forward 一步
真实实现
Engine.generate(nanochat/engine.py:176-275)四步:
- batch=1 prefill:
KVCache(batch_size=1, seq_len=len(tokens)),forward 整个 prompt(nanochat/engine.py:194-206)。 - 克隆:建
KVCache(batch_size=num_samples)+prefill(kv_cache_prefill),然后删掉小 cache(nanochat/engine.py:208-218)。 - 主循环:每步
sample_next_token采一列 token,逐行走状态机(见 §4),yield(token_column, token_masks),再把这一列 forward 一步拿新 logits(nanochat/engine.py:223-275)。 - 终止:所有行
completed(采到<|assistant_end|>或<|bos|>)或达到 max_tokens。
采样本身(sample_next_token nanochat/engine.py:141-156)有个细节:top-k 时只在 top-k 子集上做温度和 softmax,再 gather 回原 id——比全词表 softmax 省。
generate_batch(nanochat/engine.py:277-299)是非流式封装:收集每行序列与 mask、剥掉终止 token——第三章 RL 的「一题 16 答」和 chat_eval 都用它。
4. 工具调用状态机
直觉
模型生成到 <|python_start|> 后,接下来的 token 是「写给计算器的表达式」;到 <|python_end|> 时引擎要:解码表达式 → 本地求值 → 把 <|output_start|>结果<|output_end|> 强制注入该行后续 token(forced tokens,mask=0)。从模型视角看,就像它「看到了」计算结果。
状态转移
采样到 <|python_start|> ──► 进入块,清空表达式缓冲
块内 token ──► 全部记入 python_expr_tokens
采样到 <|python_end|> ──► use_calculator(表达式)
├─ 成功 ──► output_start+结果+output_end 压入 forced 队列
└─ 失败 ──► 静默跳过
forced 队列非空 ──► 该行下一步强制弹出(mask=0),不采样
真实实现
- 每行一个
RowState:forced_tokens队列、in_python_block标志、python_expr_tokens缓冲、completed标志(nanochat/engine.py:160-167)。 - 每行先查 forced 队列:有则弹出当本行 token(mask=0),没有用采样值(mask=1)(
nanochat/engine.py:240-245)。 - 状态转移在
nanochat/engine.py:251-267;求值成功才注入 output 三件套。 - mask 的意义在第三章见过:RL 训练时 forced token 不进 loss——工具输出是「环境给的」,不该模型背。
计算器本身
use_calculator(nanochat/engine.py:46-79)是白名单求值器:
- 纯数字算术(禁
**幂运算):eval,且{"__builtins__": {}}空内置 + 3 秒 SIGALRM 超时(eval_with_timeoutnanochat/engine.py:35-44)。 - 字符串
.count()(数 strawberry 里有几个 r 那类):字符白名单 + 危险模式黑名单(import/exec/eval/__等)。 - 其余一律返回 None——静默放弃,不打断生成。
5. 两个外围件
代码沙箱 execution.py
HumanEval 评测要执行模型写的代码。execute_code(nanochat/execution.py:74-134)的做法:
- 子进程跑
python -c,先执行GUARD脚本(nanochat/execution.py:47-71):rlimit 限 256MB 内存、把os.system/shutil.rmtree/subprocess.Popen等危险函数置 None、环境只留 PATH。 - 然后
exec(compile(code));cwd 是临时目录,跑完即删;stdin 关闭,stdout/stderr 捕获;超时硬杀。
文件头明确列出「不防」清单(nanochat/execution.py:14-21):不挡网络、不防 ctypes、无内核级隔离——防意外不防恶意,不是安全沙箱。
推理基准 infer_bench.py
scripts/infer_bench.py 把「这个模型在这张卡上理论能跑多快」算给你看:
- decode 是带宽受限:理论 tok/s ≈ 峰值带宽 ÷ 每步读取字节(权重 + KV cache 读取)。模型侧数字由
GPT.kv_bytes_per_token/kv_read_bytes给出(nanochat/gpt.py:374-388),滑窗层只读 min(context, window)。 - 指标:TTFT / TPOT / tok/s;MBU(带宽利用率,绑小 batch 的 decode)与 MFU(算力利用率,绑大 batch / prefill)双 roofline(
scripts/infer_bench.py:97-232)。 - 硬件峰值 FLOPS / 带宽硬编码成两张表(
nanochat/common.py:228-328get_peak_flops/get_peak_bandwidth)。
6. 关键细节 / 坑
- rotary 的 cache 偏移:KV cache 存在时,cos/sin 要按
kv_cache.get_pos()切片(nanochat/gpt.py:466-469)——位置编码必须和历史对齐。 - smear 的推理分支:decode 单 token 时从
kv_cache.prev_embedding取上一个嵌入;prefill 多 token 时走训练同款切片(nanochat/gpt.py:480-492)——这就是 KVCache 上要存prev_embedding的原因。 get_pos假设全 batch 同步推进(nanochat/engine.py:111-113):完成的行只被标记、不摘出 batch——简单,但白算一点。- 一致性有内建对拍:
engine.py的__main__(nanochat/engine.py:302-352)会拿 naive 的model.generate和Engine.generate逐 token 比对。 - chat_cli 没有对话对象:多轮对话就是「自己维护一个 token 列表,每轮追加 user/assistant 边界 token 再交给 Engine」(
scripts/chat_cli.py:71-96);若生成被 max_tokens 截断,手动补一个<|assistant_end|>保持状态机闭合(scripts/chat_cli.py:92-95)。token 就是状态。 - KV cache 长度是预分配的:
kv_length_hint = len(prompt) + max_tokens(nanochat/engine.py:209-216)——超长对话要自己注意 sequence_len 上限。
回到导读:index.md。