跳到主要内容

数据截至 (上游 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.generatenanochat/engine.py:176-275)四步:

  1. batch=1 prefillKVCache(batch_size=1, seq_len=len(tokens)),forward 整个 prompt(nanochat/engine.py:194-206)。
  2. 克隆:建 KVCache(batch_size=num_samples) + prefill(kv_cache_prefill),然后删掉小 cache(nanochat/engine.py:208-218)。
  3. 主循环:每步 sample_next_token 采一列 token,逐行走状态机(见 §4),yield (token_column, token_masks),再把这一列 forward 一步拿新 logits(nanochat/engine.py:223-275)。
  4. 终止:所有行 completed(采到 <|assistant_end|><|bos|>)或达到 max_tokens。

采样本身(sample_next_token nanochat/engine.py:141-156)有个细节:top-k 时只在 top-k 子集上做温度和 softmax,再 gather 回原 id——比全词表 softmax 省。

generate_batchnanochat/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),不采样

真实实现

  • 每行一个 RowStateforced_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_calculatornanochat/engine.py:46-79)是白名单求值器:

  • 纯数字算术(禁 ** 幂运算):eval,且 {"__builtins__": {}} 空内置 + 3 秒 SIGALRM 超时(eval_with_timeout nanochat/engine.py:35-44)。
  • 字符串 .count()(数 strawberry 里有几个 r 那类):字符白名单 + 危险模式黑名单(import / exec / eval / __ 等)。
  • 其余一律返回 None——静默放弃,不打断生成。

5. 两个外围件

代码沙箱 execution.py

HumanEval 评测要执行模型写的代码。execute_codenanochat/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-328 get_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.generateEngine.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_tokensnanochat/engine.py:209-216)——超长对话要自己注意 sequence_len 上限。

回到导读:index.md