ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

从 checkpoint 到对话:train-llm-from-scratch 推理与 Chat 指南

从 checkpoint 到对话:train-llm-from-scratch 推理与 Chat 指南 从 checkpoint 到对话train-llm-from-scratch 推理与 Chat 指南【免费下载链接】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训练只有在你真正能和结果「对话」时才令人满足。本指南以 docs/09_inference.md 为主体讲解本仓库在预训练基座之上新增的统一推理层它如何读取任意阶段base / SFT / DPO / PPO / GRPO的 checkpoint 并正确对话、chat 与 raw 两种模式的区别、scripts/chat.py的命令行用法以及温度、top-p、top-k、贪婪解码等采样旋钮的底层原理。读完你既能用一条命令和训练产物对话也能理解其背后的解码循环、上下文裁剪与停止 token 机制。图中流程对应 src/post_training/inference.py 的实现任意 checkpoint →load_model_from_ckpt维度取自 checkpoint 中保存的 cfg→ 判断 chat 或 raw 模式 →batched_generate采样 →decode输出回复。Mermaid 源文件见 docs/diagrams/src/09_inference.mmd。为什么需要一层新的推理封装仓库最早期的generate_text.py只能对基座模型做「原始续写」raw continuation它硬编码绑定在旧的配置上并且没有任何聊天模板。这带来两个问题训练链路预训练 → SFT → DPO/PPO/GRPO产出的 checkpoint 结构各不相同——DPO/PPO/GRPO 阶段保存的权重带有module./transformer.前缀或额外奖励头参数经过 SFT 指令微调的模型期望输入带有角色标记|user|/|assistant|直接喂裸文本会导致模型忽略指令。因此本仓库在 src/post_training/inference.py 新增了一个统一推理层由load_model_from_ckpt与generate_reply两个函数构成被聊天 CLI 与评测脚本共用。其核心目标是加载任何阶段的 checkpoint 并用正确的方式与之对话。按 checkpoint 中保存的配置加载任意阶段模型基座模型超参n_embed、n_blocks、context_length、vocab_size在预训练时确定而load_model_from_ckpt直接从 checkpoint 里保存的cfg字段读出这些维度因此调用方永远不需要重复指定n_embed/n_blocksck torch.load(ckpt_path, map_locationcpu, weights_onlyFalse) cfg {**(ck.get(cfg) or {}), **(overrides or {})} model Transformer( n_headcfg.get(n_head, 16), n_embedcfg.get(n_embed, 1024), context_lengthcfg.get(context_length, 1024), vocab_sizecfg.get(vocab_size, 50304), N_BLOCKScfg.get(n_blocks, 24), ) state ck[model_state_dict] if model_state_dict in ck else ck state {k.removeprefix(module.).removeprefix(transformer.): v for k, v in state.items()} model.load_state_dict({k: v for k, v in state.items() if k in keys}, strictFalse) return model.to(device).eval()这一小段代码体现了三个关键设计配置自描述模型维度来源于 checkpoint 内的cfg而非外部配置文件推理时只需给路径--ckpt天然规避了配置漂移前缀容错DDP 训练会把权重键包装成module.前缀RL 阶段的包装器可能引入transformer.前缀这里统一removeprefix剥离后再加载奖励头容忍load_state_dict(..., strictFalse)只载入与Transformer骨干匹配的键reward / value head 等额外参数被安全忽略——这是「wrap, dont rewrite」设计规则的推论所有后训练头都环绕在 Transformer.forward_hidden 之上骨干结构不变。模型随后被置为eval()模式并搬到目标设备这也是 src/post_training/inference.py 返回前的最后一步。Chat 与 Raw两种对话模式generate_reply复用与训练/评测同一套经过测试的生成核心batched_generate对外暴露两种模式chat默认把文本包装进聊天模板可选带system消息返回解码后的 assistant 回复。适用于 SFT / DPO / PPO / GRPO 产出的指令微调 checkpoint。raw--raw把你的文本当作前缀直接返回基座模型的续写不套模板。适用于base_pretrained.pt。两种模式的差异在源码中一目了然if raw: ids get_tokenizer().encode_ordinary(user_text) else: messages [] if system: messages.append({role: system, content: system}) messages.append({role: user, content: user_text}) prompt_ids encode_prompt(messages) # 结尾是 |assistant|提示模型开始生成注意encode_ordinary与encode_prompt的区别raw 模式用普通编码直接处理输入chat 模式则经过 chat_template.encode_prompt把消息渲染为|user|\n{内容}|endoftext||assistant|\n的 token 序列——prompt 形式以 assistant 头结尾模型被提示接下来应当生成 assistant 内容。聊天模板的底层原理纯文本角色标记为什么不用真正的特殊 token原因在 chat_template.py 的模块注释中写得很清楚本仓库的 tokenizer 是 tiktokenr50k_base其唯一特殊 token 是|endoftext|id 50256无法注册新的聊天角色特殊 token。因此实现采用纯文本角色标记|user|、|assistant|、|system|它们只是普通的多 token 字符串SFT 时模型会像学习其他文本一样学会它们。|endoftext|EOT被复用为回合/序列终止符与唯一的生成停止 token|user| {user content}|endoftext||assistant| {assistant content}|endoftext|对推理而言encode_chat在add_generation_promptTrue时生成全零 loss mask 的 prompt 序列因为没有任何 assistant 内容需要训练这正是 rollout 与推理使用的形式。测试 tests/test_post_training_smoke.py 与 tests/verify_data_and_eval.py 分别验证了模板的 ids/mask 对齐与往返解码正确性。防御性解码EOT 截断与填充 vocab 清洗模型词表被填充到 50304但r50k_base只能解码普通 token 0–50255未训练充分的模型可能吐出填充区的 id 导致解码崩溃。因此 decode 做了防御处理clean [t for t in ids if 0 t EOT_ID] return get_tokenizer().decode(clean)即丢弃 EOT50256终止符以及任何 50256的 id。而 batched_generate 在生成侧同样会在遇到 EOT 时提前截断 token 序列两侧配合保证输出干净可读。命令行一键对话与交互式 REPLscripts/chat.py提供一次性问答与交互式 REPL 两种形态# 指令微调模型自动套用聊天模板 PYTHONPATH. python scripts/chat.py --ckpt /ephemeral/ckpts/sft.pt --prompt What is 13 29? 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输入 exit 或 Ctrl-D 退出 PYTHONPATH. python scripts/chat.py --ckpt /ephemeral/ckpts/sft.pt启动时脚本会打印模型概况loaded {ckpt} ({参数量}M params) on {device} | moderaw/chat greedy或 T.. top_p..见 scripts/chat.py便于快速确认加载路径与采样配置是否正确。参数一览参数默认值说明--ckpt必填checkpoint 路径模型维度从其中保存的cfg读取--promptNone一次性提问省略则进入交互式 REPL--systemNone可选 system 消息仅 chat 模式生效--rawFalse基座模型续写模式不套聊天模板--max_new_tokens256最多新生成的 token 数--temperature0.8采样温度--top_p0.95nucleus 采样阈值chat 模式默认启用--top_kNonetop-k 截断默认关闭--greedyFalse确定性 argmax 解码--device自动有 CUDA 用cuda否则cpu两种设备均已验证注意--greedy的一个细节REPL 内调用generate_reply时top_pargs.top_p if args.top_p 1 else Nonescripts/chat.py即 top_p 为 1 时视为关闭避免语义上「全概率」与「不截断」的歧义。采样旋钮从公式到实现对话质量高度依赖采样策略本仓库的生成核心 rollout.generate_with_logprobs 完整实现了这些旋钮filter_logits 负责把它们应用到 logits 上。原理详见 docs/foundations/generation.md。greedy贪婪解码取argmax确定、可复现最适合评测与数学题--greedy。在实现上等价于top_k1——batched_generate收到greedyTrue时直接改写为temperature1.0, top_k1, top_pNoneevaluation.py。这也是整个流水线用固定 GSM8K 集做横向对比时统一使用的模式见 evaluation.py。temperature温度生成前先把 logits 除以温度再 softmaxτ1分布更尖锐更保守、τ1不变、τ1更平坦更多样但更易错。开式对话常用0.7–1.0。实现上logits / max(temperature, 1e-6)rollout.py并对temperature0做了除法保护。top_p核采样保留累积概率达到p的最小 token 集合砍掉长尾的低概率 token。实现中先按概率降序排序、计算累积概率、把超过p的项置-inf并保证至少保留概率最高的那个 tokenrollout.py。top_k只保留概率最高的 k 个 token其余置-infrollout.py。一个值得注意的工程细节filter_logits只用于构造采样分布而记录到gen_logprobs的用于 RL 比率的对数概率始终取自全分布logits.float()下用 fp32 计算见 rollout.py 与模块注释——bf16 下做对数概率减法会有害这也是所有 RL 算法共享同一套生成内核的原因。解码循环、上下文裁剪与停止底层自回归循环在 Transformer.generate 中最为直白每步取idx[:, -context_length:]裁剪上下文取最后位置的 logitssoftmax 后torch.multinomial采样一个新 token 追加到序列。generate_with_logprobs在其上增加了逐行独立停止某行一旦吐出 EOT 即标记 finished后续位置填充 pad id 并从response_mask中排除rollout.py。由于模型使用学习到的绝对位置编码序列长度被context_length硬性封顶超长 prompt 必须裁剪又因模型没有 padding-aware 注意力掩码batched_generate采用「按相同长度分桶 micro-batch」策略让无掩码模型也能成批解码evaluation.py。budget min(max_new_tokens, cap - L)保证了 prompt 与生成长度之和不超过上下文上限。常见问题速查结合 docs/foundations/generation.md 的诊断表对话阶段常见异常与排查方向如下症状可能原因排查无限重复分布过尖或未学会停止调低 max tokens检查 EOT 处理调整采样参数无视指令基座模型未经充分 SFT换用 SFT checkpoint 测试回答格式错误SFT 数据格式不匹配检查聊天模板与 loss mask输出像随机文本模型欠训练或温度过高对比 train/dev loss降低温度长 prompt 崩溃超出上下文或显存裁剪上下文核对context_length基座模型套了模板答非所问用错模式基座模型应加--raw至此完整的闭环已经打通预训练 → SFT → DPO/PPO/GRPO 对齐 → GSM8K 评测 → 对话。想回顾整条流水线回到 docs/README.md 总览所有命令行操作的速查表见 docs/howto/commands.md。【免费下载链接】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),仅供参考
返回列表