从 generate 到 generate_with_logprobs:train-llm-from-scratch 如何用一个推理核心复用聊天、评测与RL训练
2026/9/15 11:39:33 网站建设 项目流程

从 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 中重写了自回归循环,核心设计只有三句话:

  1. 同一份 logits,算两次:一次走完整分布算log_softmax(用于记录 log 概率),一次经filter_logits做 temperature + top-k/top-p 截断后再采样(用于真正抽 token);
  2. 按行独立停止:某一行一旦命中 stop token,后续位置填 pad 并在response_mask中标 False,损失函数自动忽略;
  3. 上下文自动截断:强制prompt_len + max_new_tokens <= context_length,越界直接报错而不是静默出错。

生成结果被打包进一个结构清晰的数据类RolloutBatch(rollout.py):

字段含义
sequencesprompt + 生成 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.pygenerate_with_logprobscompute_logprobsrollout_prompts
src/post_training/evaluation.pybatched_generate批量解码 + GSM8K 准确率
src/post_training/inference.pygenerate_reply,聊天模板 / raw 双模式
scripts/chat.py一次性提问或交互式 REPL 聊天
docs/09_inference.md推理与聊天完整文档
docs/06_ppo.md / docs/07_grpo.mdPPO / 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询