资讯详情

train-llm-from-scratch 评估指南:用 GSM8K 验证 Base → SFT → DPO → PPO → GRPO 全链路效果

📅 2026/9/15 15:02:53 | 华诺云谱 👁 阅读
train-llm-from-scratch 评估指南:用 GSM8K 验证 Base → SFT → DPO → PPO → GRPO 全链路效果
train-llm-from-scratch 评估指南用 GSM8K 验证 Base → SFT → DPO → PPO → GRPO 全链路效果【免费下载链接】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 仓库从零训练大模型的完整流水线中的统一评估体系如何用同一个 held-out GSM8K 测试集、同一种贪心解码策略度量 Base 预训练、SFT、DPO、PPO、GRPO 各阶段 checkpoint 的推理能力并把结果汇总成一张可比较的跨阶段精度表。读完本文你将掌握该仓库评估工具的调用方式、底层解码与答案解析实现原理以及如何复现仓库自带的打分正确性校验。评估的设计原则同一条尺子量到底一条流水线只有可测量才可信。本仓库的评估体系建立在一个核心假设上所有训练阶段都在同一份 held-out GSM8K 测试集上、以贪心解码greedy decoding评估从而让 Base → SFT → DPO → PPO → GRPO 的精度变化可以直接比较。整体评估流程见 docs/diagrams/08_evaluation.mmd 的可编辑 Mermaid 源码关键设计选择有三个固定的评测集与解码策略GSM8K 官方 test split 贪心解码top_k1保证数字确定、可复现、跨阶段可比可验证的奖励verifiable reward打分不依赖人工主观判断而是解析模型输出中的最终数字与 GSM8K 官方金标准#### N字段做精确比对一个脚本、一张表任意 checkpoint 都能通过同一个入口脚本打分结果以 JSONL 累积成跨阶段表格。这些目标对应仓库中的两个核心文件src/post_training/evaluation.py共享评估工具与 scripts/eval_post_training.py命令行入口。生成阶段length-bucketed 贪心解码评估的第一步是让模型对每个 GSM8K 问题生成答案。仓库的教学模型没有 padding-aware attention mask无法像主流框架那样直接对不等长序列做 mask 填充批处理。为此batched_generate采用按长度分桶length-bucketing策略按 prompt 长度对全部样本排序把长度完全相同的 prompt 聚成一个桶每个桶内至多micro_batch默认 32条样本一起送入模型解码解码预算为min(max_new_tokens, context_length - prompt_len)防止超出模型上下文默认 1024解码以|endoftext|EOT_ID 50256作为停止 token命中即截断输出见 src/post_training/chat_template.py。# src/post_training/evaluation.py 中的核心逻辑精简 if greedy: temperature, top_k, top_p 1.0, 1, None # 贪心 强制 argmax ... while i len(order): # 收集一段等长 prompt最多 micro_batch 条 ... budget min(max_new_tokens, cap - L) rb generate_with_logprobs(model, batch, budget, temperaturetemperature, top_ktop_k, top_ptop_p, stop_tokens(EOT_ID,), pad_idEOT_ID)底层的逐 token 解码由generate_with_logprobs实现它不做 KV cache为教学清晰性而刻意省略每步重跑前缀但对本仓库 1024 上下文的短序列足够快。greedyTrue时直接把temperature1.0, top_k1, top_pNone等价于逐位置取 argmax产出确定性的可比数字。打分阶段三级回退的答案解析 正确性主导的奖励精度计算gsm8k_accuracygsm8k_accuracy完成生成 → 解析 → 比对三步闭环先把每个问题包装成用户消息并 tokenize批量生成回答再用is_correct(resp, gsm8k_gold_answer(ans))逐条判定最终返回{accuracy, n, correct, samples}其中samples可保留少量(question, response, gold, correct)元组供人工抽查prompts [encode_prompt([{role: user, content: q}]) for q, _ in qa_pairs] responses batched_generate(model, prompts, max_new_tokens, devicedevice, greedygreedy) correct sum(is_correct(resp, gsm8k_gold_answer(ans)) for (q, ans), resp in zip(qa_pairs, responses))容错解析extract_answerSFT 阶段会让模型学习输出think…/thinkanswer…/answer的推理格式但小模型输出不稳定。因此 src/post_training/rewards/parsing.py 中的extract_answer采用三级优雅回退优先取answer…/answer标签内的数字否则取 GSM8K 风格的#### N之后的数字再退而求其次取文本中出现的最后一个数字正则-?\$?\d[\d,]*(?:\.\d)?兼容负号、美元符号与千分位逗号。金标准gsm8k_gold_answer则直接解析 GSM8K answer 字段结尾的#### N。奖励函数reward_gsm8k可验证奖励reward_gsm8k的设计原则是正确性主导 有界格式奖励以抑制 reward hacking小模型在格式奖励过高时会输出空answer/answer标签或重复 tokenr 0.0 if _answers_match(extract_answer(text), gold): r 1.0 # 真正重要的奖励 if has_well_formed_answer(text): r 0.2 # 小幅格式引导 return min(r, 1.2) # 截断其中_answers_match使用math.isclose以1e-4绝对容差做浮点比较CORRECT_BONUS1.0、FORMAT_BONUS0.2、REWARD_CLIP1.2。纯格式奖励reward_format仅当恰好一个格式良好的answer块时给 1.0供 RL 阶段单独使用算术热身任务复用同一逻辑reward_arithmetic reward_gsm8k。打分正确性校验与模型无关的独立验证值得强调的是打分逻辑本身经过了独立于任何模型的 sanity check给打分器喂正确答案得到 100/100 命中、喂错误答案得到 0/100 误报且金标准与在线 GSM8K 数据集交叉核验。这一校验以可运行测试的形式固化在 tests/verify_data_and_eval.py 的verify_eval_benchmark中构造think.../thinkanswer{gold}/answer格式的正确回答与gold7的错误回答分别断言 ≥98/100 判对、0 误报断言正确格式完整 仅正确 错误的奖励排序且奖励 ≤ 1.2 有界校验$、千分位逗号$1,234 → 1234.0与#### 18 → 18.0的容错解析。跨阶段评估表一个脚本吃遍所有 checkpointscripts/eval_post_training.py是评估流水线的统一入口。它的一个贴心设计是模型维度直接从 checkpoint 中存储的cfg读取无需在命令行重复指定即使是带 reward head 的奖励模型 checkpoint 也能加载因为只保留 backbone 的权重键filtered {k: v for k, v in state.items() if k in backbone_keys}strictFalse载入见model_from_ckpt。对每个阶段 checkpoint 各评一次再渲染成表格for s in base_pretrained sft dpo ppo grpo; do PYTHONPATH. python scripts/eval_post_training.py --ckpt /ephemeral/ckpts/$s.pt \ --label $s --limit 200 --append /ephemeral/logs/stage_table.jsonl done PYTHONPATH. python scripts/eval_post_training.py --table /ephemeral/logs/stage_table.jsonlstage GSM8K acc n ------------------------------------ base_pretrained ... 200 sft ... 200 dpo ... 200 ppo ... 200 grpo ... 200命令行参数一览对应argparse定义参数默认值说明--ckpt必填待评估的 checkpoint 路径--labelmodel写入表格时的阶段名--limit200评测样本数上限GSM8K test 抽子集--splittest数据集 split--max_new_tokens300每条回答的最大生成长度--devicecuda/cpu自动检测 CUDA 可用性--samples3打印多少条样本Q/gold/correct/A供人工抽查--append无把结果行追加到指定 JSONL--table无仅从 JSONL 渲染表格后退出单条执行时还会打印[{label}] GSM8K test accuracy: xx.x% (n/N)及若干样本详情。评测数据由load_gsm8k_eval从 HuggingFaceopenai/gsm8k加载HF_HOME默认/ephemeral/hf_cache返回(question, answer_field)列表。训练中的指标JSONL 常驻WB 可选跨阶段表格回答最终效果如何而训练过程中的指标由MetricsLogger负责每个训练器都会在/ephemeral/logs/下写入一个{stage}_{timestamp}.jsonl指标文件每步一行 JSON{step, wall, ...metrics}实时 flush不依赖任何外部服务即可离线绘图传入--use_wandb true时才会镜像到 Weights Biases且 WB 初始化失败不会中断训练打印wandb disabled ...; JSONL logging only。各阶段记录的指标见 docs/README.md 的阶段总览SFTtrain/dev lossReward Modelpreference accuracyDPOimplicit-reward accuracyPPO / GRPOreward、KL、clip fraction策略裁剪比例见 src/post_training/ppo.py 的 clipped surrogate loss以及 GSM8K accuracyGRPO 还会记录 group-relative 的kl与 reward 指标见 src/post_training/grpo.py。结果解读这个规模下好长什么样需要管理预期一个约 400M 参数、从零训练的模型不可能登顶 GSM8K 排行榜。评估的价值在于相对爬升——看 Base → SFT → DPO → PPO → GRPO 每步是否带来清晰、真实的增益以及 RL 阶段 KL 是否有界KL 失控通常意味着过度优化。绝对值适度偏低是正常的只要每个阶段呈现明确的 before/after 提升就说明流水线的每一环都在起作用。小结与下一步本文覆盖了该仓库从批量贪心解码到可验证奖励打分再到跨阶段表格的完整评估闭环并给出了打分正确性的独立测试依据。所有评估入口都集中在 src/post_training/evaluation.py 与 scripts/eval_post_training.py。评估完成后下一步自然是用 docs/09_inference.md 的推理/对话工具实际体验任意 checkpoint 的输出效果。【免费下载链接】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),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。