从 generate 到 generate_with_logprobs:train-llm-from-scratch 如何用一个推理核心复用聊天、评测与RL训练
【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch
在从零训练大语言模型的开源项目train-llm-from-scratch中,"生成文本"看似是最简单的环节,却藏着整个项目的架构精华。从预训练后的一段generate文本续写,到强化学习阶段逐 token 记录 log 概率的generate_with_logprobs,项目用一条推理管线同时支撑了聊天 CLI、GSM8K 评测、PPO 与 GRPO 训练。本文将带你拆解这套 LLM 推理复用设计,看懂它是如何做到"写一次、用四处"的。
如果你手上正好有这个项目(git clone https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch),建议边读边对照源码,全文只涉及少量关键片段。
第一层:教学版generate——只够续写,不够生产
项目最初的推理能力来自 Transformer 模型类自带的生成方法,代码非常"教科书":
# src/models/transformer.py def generate(self, idx, max_new_tokens): for _ in range(max_new_tokens): idx_cond = idx[:, -self.context_length:] logits, _ = self(idx_cond) probs = F.softmax(logits[:, -1, :], dim=-1) idx = torch.cat((idx, idx_next), dim=1)完整实现见 src/models/transformer.py,配合 scripts/generate_text.py 就能对预训练基座模型做原始文本续写。但它有三个"生产级"短板:
- 没有采样参数:不支持 temperature / top_k / top_p,只能裸采样;
- 没有停止条件:不会在
</s>结束符处停下来,只能死等满max_new_tokens; - 没有批处理:一次只处理一条序列,跑几百条 GSM8K 评测题会慢到无法接受。
第二层:generate_with_logprobs——生成时顺手记录 log 概率
要支撑 PPO/GRPO 这类强化学习算法,光拿到生成文本远远不够——训练目标需要"每个生成 token 在策略下的 log 概率"。于是项目在 src/post_training/rollout.py 中重写了自回归循环,核心设计只有三句话:
- 同一份 logits,算两次:一次走完整分布算
log_softmax(用于记录 log 概率),一次经filter_logits做 temperature + top-k/top-p 截断后再采样(用于真正抽 token); - 按行独立停止:某一行一旦命中 stop token,后续位置填 pad 并在
response_mask中标 False,损失函数自动忽略; - 上下文自动截断:强制
prompt_len + max_new_tokens <= context_length,越界直接报错而不是静默出错。
生成结果被打包进一个结构清晰的数据类RolloutBatch(rollout.py):
| 字段 | 含义 |
|---|---|
sequences | prompt + 生成 token 的完整序列 |
response_mask | 只对"真实生成位置"为 True,pad 与 prompt 全部屏蔽 |
gen_logprobs | 每个生成 token 在采样温度下的全分布 log 概率 |
prompt_len | 共享的 prompt 长度 |
💡 这个循环没有用 KV cache——每一步都重跑整个前缀。作者明确说这是"为了教学清晰":短序列场景下,可读性比速度更值钱。
第三层:推理与评测如何"白嫖"RL 的推理核心
最有意思的复用发生在反向:聊天和评测功能并没有另写生成代码,而是直接调用 RL 的generate_with_logprobs,只是丢弃了gen_logprobs这个"副产品"。调用链自上而下分三层:
scripts/chat.py(命令行聊天) └─ generate_reply() 包装 chat template / raw 两种模式 └─ batched_generate() 按长度分桶 + 贪心/采样开关 └─ generate_with_logprobs() 真正的采样循环- src/post_training/inference.py 的
generate_reply负责"懂对话":SFT/DPO/PPO/GRPO 的指令模型走 chat 模板,基座模型走raw=True原始续写; - src/post_training/evaluation.py 的
batched_generate负责"懂批量":由于模型没有 padding 感知注意力掩码,它把 prompt按相同长度分桶后组微批解码,贪心模式(greedy=True→ top_k=1)保证了评测数字可复现。
这样带来的直接收益:Base → SFT → DPO → PPO → GRPO 五个阶段的 GSM8K 准确率,全部由同一条解码路径产生,指标天然可对比,不存在"训练用的解码器"和"评测用的解码器"两套行为。GRPO 训练同样走这个入口(rollout.py 的rollout_prompts只是再加一层长度分桶):
两个容易忽略的设计决策
🔍为什么是"自由函数"而不是模型方法?rollout.py 的模块注释写得直白:PPO/GRPO 期间,同一套 log 概率数学要对四组不同参数各跑一遍(可训练策略、冻结参考模型、旧策略快照、带 value head 的 actor-critic 包装器)。写成f(model, ...)的自由函数比绑定方法组合性好得多,也让教学用的模型文件保持干净。
🔢为什么 log 概率强制 fp32?PPO/DPO 的核心操作是对 log 概率做减法(算重要性采样比率),bf16 的舍入误差在这里会被指数放大。所以即使整个前向跑在 bf16 autocast 下,代码里也刻意写了logits.float()——这种细节正是"教育型"项目最值钱的部分。
快速索引:推理相关源码与文档
| 文件 | 职责 |
|---|---|
| src/models/transformer.py | 教学版generate,纯续写 |
| src/post_training/rollout.py | generate_with_logprobs、compute_logprobs、rollout_prompts |
| src/post_training/evaluation.py | batched_generate批量解码 + GSM8K 准确率 |
| src/post_training/inference.py | generate_reply,聊天模板 / raw 双模式 |
| scripts/chat.py | 一次性提问或交互式 REPL 聊天 |
| docs/09_inference.md | 推理与聊天完整文档 |
| docs/06_ppo.md / docs/07_grpo.md | PPO / GRPO 如何用 rollout 核心 |
小结
train-llm-from-scratch 的推理复用设计可以浓缩成一条主线:把"自回归采样"下沉为带 log 概率记录的generate_with_logprobs,再让聊天、评测、RL 各取所需——聊天层丢弃概率只留文本,评测层批量分桶只留准确率,RL 层则完整消费概率去做 PPO 的比率计算。对新手来说,这比任何 PPT 都直观地展示了:一个干净的生成内核,如何撑起一条完整的 LLM 训练流水线。
【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考