从 generate 到 generate_with_logprobs:train-llm-from-scratch 如何用一个推理核心复用聊天、评测与RL训练
发布时间:2026/9/15 11:39:40 作者:尧图编辑部 阅读量:1,286

从 generate 到 generate_with_logprobstrain-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), dim1)完整实现见 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越界直接报错而不是静默出错。生成结果被打包进一个结构清晰的数据类RolloutBatchrollout.py字段含义sequencesprompt 生成 token 的完整序列response_mask只对真实生成位置为 Truepad 与 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 模板基座模型走rawTrue原始续写src/post_training/evaluation.py 的batched_generate负责懂批量由于模型没有 padding 感知注意力掩码它把 prompt按相同长度分桶后组微批解码贪心模式greedyTrue→ top_k1保证了评测数字可复现。这样带来的直接收益Base → SFT → DPO → PPO → GRPO 五个阶段的 GSM8K 准确率全部由同一条解码路径产生指标天然可对比不存在训练用的解码器和评测用的解码器两套行为。GRPO 训练同样走这个入口rollout.py 的rollout_prompts只是再加一层长度分桶两个容易忽略的设计决策为什么是自由函数而不是模型方法rollout.py 的模块注释写得直白PPO/GRPO 期间同一套 log 概率数学要对四组不同参数各跑一遍可训练策略、冻结参考模型、旧策略快照、带 value head 的 actor-critic 包装器。写成f(model, ...)的自由函数比绑定方法组合性好得多也让教学用的模型文件保持干净。为什么 log 概率强制 fp32PPO/DPO 的核心操作是对 log 概率做减法算重要性采样比率bf16 的舍入误差在这里会被指数放大。所以即使整个前向跑在 bf16 autocast 下代码里也刻意写了logits.float()——这种细节正是教育型项目最值钱的部分。快速索引推理相关源码与文档文件职责src/models/transformer.py教学版generate纯续写src/post_training/rollout.pygenerate_with_logprobs、compute_logprobs、rollout_promptssrc/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),仅供参考