ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

从零训练 LLM:解码与采样机制全解析 —— 基于 train-llm-from-scratch 理解自回归生成、Temperature、Top-k/Top-p 与停止符

从零训练 LLM:解码与采样机制全解析 —— 基于 train-llm-from-scratch 理解自回归生成、Temperature、Top-k/Top-p 与停止符 从零训练 LLM解码与采样机制全解析 —— 基于 train-llm-from-scratch 理解自回归生成、Temperature、Top-k/Top-p 与停止符【免费下载链接】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 仓库 docs/foundations/generation.md 的深度展开。训练阶段模型为已知文本预测下一个 token而生成阶段模型会把自己采样出的 token 作为下一个输入——这个反馈闭环决定了 LLM 推理时的全部行为。读完本文你将掌握自回归生成循环的底层实现含仓库内Transformer.generate源码、贪心解码与采样的数学差异、温度 / Top-k / Top-p 三种采样控制的精确作用与源码实现、上下文裁剪与停止符EOT的处理方式并学会用仓库提供的chat.py/generate_text.py实战生成文本、用诊断表排查常见的生成质量问题。什么是生成从预测下一个 token到自回归闭环训练和生成共享同一个模型但输入来源截然不同训练teacher forcing输入是数据集里已知的文本损失只要求模型预测真正的下一个 token生成autoregressive decoding输入的前缀来自模型自己上一步采样的结果没有标准答案。训练中下一个 token永远是确定的而生成中下一个 token是一个概率分布上的随机变量。正因如此哪怕两个模型在某个位置的输出概率只差一点点随着 token 被不断回填进上下文最终生成的文本也可能完全不同——这就是反馈环路的放大效应。原文档用一张 Mermaid 流程图描述了完整的自回归循环循环中每个环节都对应仓库代码里的一个真实操作裁剪上下文 → 前向传播 → 取末位 logits → softmax/采样 → 追加 token → 再次裁剪。自回归循环的仓库实现Transformer.generate仓库中最简单的生成实现位于 src/models/transformer.pydef generate(self, idx: torch.Tensor, max_new_tokens: int) - torch.Tensor: for _ in range(max_new_tokens): idx_cond idx[:, -self.context_length:] logits, _ self(idx_cond) logits logits[:, -1, :] probs F.softmax(logits, dim-1) idx_next torch.multinomial(probs, num_samples1) idx torch.cat((idx, idx_next), dim1) return idx逐行拆解这段代码正好对应上面的流程图idx_cond idx[:, -self.context_length:]——只保留序列最后context_length个 token上下文裁剪见下文logits, _ self(idx_cond)——Transformer 前向传播返回形状为(B, T, vocab_size)的 logits这里的_是 loss生成时无目标、不计算logits logits[:, -1, :]——只取最后一个位置的 logits。这是整个生成循环最关键的假设最后一个位置通过因果自注意力已经看到了当前全部上下文因此它足以表示基于当前完整上下文预测下一个 tokenprobs F.softmax(logits, dim-1)——将 logits 转成词汇表上的概率分布idx_next torch.multinomial(probs, num_samples1)——从该分布中采样 1 个 token ididx torch.cat((idx, idx_next), dim1)——把新 token 追加到序列末尾作为下一轮循环的输入。generate的调用入口在 scripts/generate_text.py先把文本用 tiktokenr50k_base编码成 token id包成形状(1, T)的张量然后model.generate(context, max_new_tokensmax_new_tokens)最后enc.decode(generated_tokens)还原成文本。命令行用法为PYTHONPATH. python scripts/generate_text.py \ --model_path /path/to/checkpoint.pt \ --input_text Once upon a time \ --max_new_tokens 100注意generate_text.py直接以config/config.py中的默认配置重建模型n_head、n_embed、context_length、vocab_size、n_blocks并用checkpoint[model_state_dict]加载权重——它是面向基础模型base的原始续写入口没有聊天模板。贪心解码 vs 采样模型把 logits 变成 token 有两种基本方式贪心解码Greedy decoding——每一步都选概率最大的 token[ \arg\max_i p_i ]采样Sampling——按概率分布随机抽取[ x \sim \text{Categorical}(p) ]两者的行为差异非常直接贪心解码是确定性的同样的输入永远得到同样的输出但因为每次都取最可能 token文本往往重复、呆板陷入高频词的局部循环采样是随机的可能选错低概率 token却因此产生更多样、更有趣的文本。仓库里两种策略都有落点src/models/transformer.py 的基础generate方法直接对完整 softmax 分布做torch.multinomial采样属于纯采样后训练推理工具则显式提供贪心开关。在 src/post_training/evaluation.py 中greedyTrue会把采样参数强制设为temperature1.0, top_k1, top_pNone——top_k1即只保留概率最高的 1 个 token在实现上等价于 argmax。GSM8K 评估就是靠--greedy得到可复现、可比较的准确率数字。温度 Temperature在 softmax 之前重标定 logits温度(\tau)在 softmax之前对 logits 做缩放[ p_i \frac{\exp(z_i / \tau)} {\sum_j \exp(z_j / \tau)} ]不同 (\tau) 的效果(\tau 1)分布更尖锐高概率 token 更占优输出更安全但多样性下降(\tau 1)分布不变(\tau 1)分布更平坦低概率 token 被扶正输出更多样但也更易出错。一个容易被忽略的要点温度作用于 logits而不是概率。对概率直接开方或缩放会破坏归一化而先除温度再 softmax数学上正好等价于在原始分布上做幂次重标定Boltzmann 分布既保持归一化又可控分布锐度。仓库实现印证了这一点。src/post_training/rollout.py 的filter_logits第一步就是if temperature ! 1.0: logits logits / max(temperature, 1e-6)max(temperature, 1e-6)是为了防止temperature0造成除零贪心场景通常用top_k1而非temperature0表达。同文件还有一个数值细节值得注意generate_with_logprobs与compute_logprobs中log-prob 一律在fp32下计算logits.float() / max(temperature, 1e-6)后再log_softmax注释明确说明这是因为 PPO/GRPO/DPO 要对 log-prob 做减法bf16 的舍入误差在这里是有害的见 src/post_training/rollout.py。Top-k 与 Top-p采样前的候选集截断许多生成系统在采样之前会先限制候选 token 集合Top-k只保留概率最高的 (k) 个 token其余置为不可选Top-pNucleus Sampling核采样按概率从高到低累加保留累计概率刚好达到 (p) 的最小集合。需要特别强调这些控制不属于基础模型架构模型本身始终输出全词汇表 logitsTop-k / Top-p 只是叠加在 logits 之上的解码策略decoding policy。仓库在 src/post_training/rollout.py 的filter_logits中实现了完整的温度 Top-k Top-p 流水线核心逻辑if top_k is not None and top_k 0: k min(top_k, logits.size(-1)) kth torch.topk(logits, k, dim-1).values[..., -1, None] logits logits.masked_fill(logits kth, float(-inf)) if top_p is not None and 0.0 top_p 1.0: sorted_logits, sorted_idx torch.sort(logits, descendingTrue, dim-1) cumprobs sorted_logits.softmax(dim-1).cumsum(dim-1) remove cumprobs top_p remove[..., 1:] remove[..., :-1].clone() remove[..., 0] False remove remove.scatter(-1, sorted_idx, remove) logits logits.masked_fill(remove, float(-inf))实现细节值得推敲Top-k先取第 k 大的 logit 作为阈值kth把低于它的 logit 全部masked_fill为-infsoftmax 后这些 token 的概率为 0torch.multinomial永远不会选到它们Top-p先降序排序用 softmax 概率做累加cumsum凡是跨过阈值top_p的 token 都被标记删除remove[..., 1:] remove[..., :-1].clone()和remove[..., 0] False这两行保证至少保留概率最高的第一个 token避免极端分布下候选集为空过滤后的分布只用于实际采样而 RL 算法记录的重要性比值时用的是全分布的 log-prob见generate_with_logprobs中full_logprobs与probs的分工src/post_training/rollout.py。上下文裁剪模型只能记住最近context_length个 token模型有固定的最大上下文长度context_length即位置嵌入的维度上限。每次循环都执行idx_cond idx[:, -self.context_length:]如果对话/序列长度超过context_length最老的 token 会被直接丢弃模型无法 attend 到保留窗口之外的任何文本。这在 src/models/transformer.py 有结构性依据position_embed nn.Embedding(context_length, n_embed)只学了 0 到context_length-1的绝对位置嵌入超长序列根本没有位置向量可用。这一约束在后训练的 rollout 代码中被显式执行src/post_training/rollout.py 中max_new_tokens min(max_new_tokens, cap - P)若prompt_len max_new_tokens context_length会直接报错raise ValueError(...)src/post_training/evaluation.py 的budget min(max_new_tokens, cap - L)则把可生成 token 数压缩到上下文余量之内。这也是为什么文档强调上下文长度是产品约束而不只是训练超参数——部署时 prompt 多长、能续写多少 token都受它硬性限制。停止符模型必须学会何时停下来生成不是无限的模型需要某种信号来表示话已说完。仓库中的停止符定义在 src/post_training/chat_template.pyEOT_ID 50256这是 tiktokenr50k_base编码器唯一的特殊 token|endoftext|。仓库对停止符的用法贯穿训练与推理训练时EOT 出现在文档与文档之间预训练数据分隔也出现在每条 assistant 消息之后SFT 数据中chat 模板在每轮结尾追加|endoftext|见 src/post_training/chat_template.py推理时chat 循环可以在检测到 EOT 时停止生成或在格式化答案如answer.../answer完整时停止。仓库中的停止机制是逐行停止generate_with_logprobs以stop_tokens(EOT_ID,)为参数某一行一旦采到 EOT该行立即标记finished后续位置用pad_id也是 EOT填充并从response_mask中排除src/post_training/rollout.pybatched_generate在解码输出时则直接截断到第一个 EOT 之前src/post_training/evaluation.py。这里隐含一个重要结论如果模型从未被训练去输出清晰的停止符或答案分隔符解码就只能靠猜测什么时候该停——这正是一直重复停不下来类问题的根源之一。另外src/post_training/chat_template.py 的decode是防御性的它会丢弃 EOT 终止符以及所有 id ≥ 50256 的 token因为模型词汇表被 padding 到 50304而r50k_base只能解码 050255 的普通 token欠训练模型可能发出这些 padding id 导致解码崩溃。为什么生成文本会漂移train-test mismatch 的本质生成质量问题的深层原因在于训练与推理的输入分布不一致训练teacher forcing每个输入前缀都来自数据集是真实文本生成前缀来自模型自己。如果模型早期采样出一个较差的 token后续所有预测都建立在那个较差 token 之上错误被逐步放大——这就是经典的暴露偏差exposure bias/ 分布偏移。这个偏移是后训练post-training存在的核心理由之一SFT教会模型回答的格式通过 chat 模板与 loss mask只对 assistant 内容计算损失见 src/post_training/chat_template.pyreward / preference 方法奖励模型、DPO 等把模型推向更被偏好的补全RLVR / GRPO可以对满足外部验证器的最终答案直接给奖励仓库中对应 src/post_training/grpo.py、src/post_training/rewards/ 与 src/post_training/rewards/verifiers.py用 GSM8K 答案验证作为奖励信号。实用生成诊断表原文档给出一张高价值的故障排查表是排查生成质量问题时的第一手清单症状可能原因排查方向一直重复分布过尖或未学会停止行为调低 max tokens、检查 EOT 处理、调整采样参数无视指令基础模型 SFT 训练不足换用 SFT checkpoint 测试答案格式错误SFT 数据格式不匹配检查 chat 模板与 loss mask输出像乱码模型欠训练或温度过高对比 train/dev loss调低温度长 prompt 崩溃prompt 超出上下文或显存裁剪上下文检查context_length实战仓库中的两个生成入口基础模型与后训练模型分别对应两个生成工具1. 基础模型续写raw continuation如上文所述scripts/generate_text.py 直接对前缀做无模板续写。2. 任意阶段 checkpoint 的对话生成chat / raw 双模式scripts/chat.py 配合 src/post_training/inference.py 使用。generate_reply有两种模式src/post_training/inference.pychat默认把用户文本包进仓库的纯文本聊天模板|user|...|endoftext||assistant|见 src/post_training/chat_template.py可选 system 消息返回解码后的 assistant 回复——适用于 SFT / DPO / PPO / GRPO checkpointraw--raw把输入当作前缀直接续写不加模板——适用于base_pretrained.pt。load_model_from_ckptsrc/post_training/inference.py会从 checkpoint 里保存的cfg读取模型维度n_head、n_embed、context_length、vocab_size、n_blocks并容忍 DDP 的module.前缀与奖励头等额外 key因此任何阶段的 checkpoint 都能一键加载推理。实际命令行详见 POST_TRAINING.md 与 docs/09_inference.md# 指令微调模型自动套用聊天模板 PYTHONPATH. python scripts/chat.py --ckpt /ephemeral/ckpts/sft.pt --prompt What is 13 29? # GRPO 模型 贪心解码数学题推荐可复现 PYTHONPATH. python scripts/chat.py --ckpt /ephemeral/ckpts/grpo.pt --prompt ... --greedy # 基础模型续写 PYTHONPATH. python scripts/chat.py --ckpt /ephemeral/ckpts/base_pretrained.pt --raw --prompt Once upon a time # 交互式 REPL省略 --prompt PYTHONPATH. python scripts/chat.py --ckpt /ephemeral/ckpts/sft.pt采样控制参数--temperature、--top_p、--top_k或--greedy确定性 argmax最适合评估与数学题设备可用--device cuda或cpu仓库两种都验证过。推荐取值参考贪心用于可复现评估开放式聊天用温度约 0.71.0top_p/top_k用于截掉长尾低概率 token。下一步走向完整流水线生成与采样是理解整个 LLM 流水线的基石——从 logits 到文本的这段闭环会贯穿数据、预训练、SFT、奖励模型、DPO、PPO、GRPO 的每一个阶段后续所有阶段最终都服务于生成更好的文本。继续深入数据处理数据管线预训练SFT 指令微调若希望更系统地掌握本仓库的完整知识脉络可从 docs/foundations/README.md 的学习路径开始tokenization → transformer → attention → objectives → optimization → generation其中也给出了各概念对应的仓库源码位置表。【免费下载链接】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),仅供参考
返回列表