本文是「GPT 训练分步学习路径」的第四步,基于
GPT_teacher-3.37M-cn这个教学项目,讲清楚「训练好的模型怎么验收、能不能上线」。文中所有代码均来自本项目真实的src/evaluate.py和data/test.jsonl。分步学习路径共五步:① 看懂数据长什么样 → ② 模型核心(注意力+因果掩码)→ ③ 训练循环 → ④ 验收评估 → ⑤ 推理与部署。本文是第四步。
目录
- 这一步属于哪个环节
- 验收和训练里的 val loss 有什么区别
- 验收本质:批量跑推理 + 判分
- 第一步:加载模型
- 第二步:逐题生成答案
- 第三步:怎么判对错(关键词片段匹配)
- 通过率与验收结论
- 核心结论速览
1. 这一步属于哪个环节
整个项目的全流程是:
1 | 构建分词器 → 数据处理 → 训练 → 验收 → Web 演示 |
本文讲的是 验收 环节,对应代码 src/evaluate.py。它的作用只有一个:
拿一份训练时从没见过的测试集
data/test.jsonl,让模型逐题作答,自动对比标准答案,算出「通过率」,判断这个模型能不能上线。
这一步是「训练」的下游(吃它产出的 best.pt),也是「Web 演示」的门槛(不通过就别推给用户)。核心就一个函数 run_test()。
2. 验收和训练里的 val loss 有什么区别
初学者常有的疑惑:训练里不是已经有 val loss 了吗,为什么还要单独验收?
区别很关键:
| 训练里的 val loss | 第四步的验收 | |
|---|---|---|
| 用的数据 | val.jsonl(训练时反复用来早停) |
test.jsonl(全程没碰过) |
| 衡量的 | 预测每个字的平均「困惑度」 | 整句答对没有(贴近真实使用) |
| 输出 | 一个抽象的 loss 数字 | PASS/FAIL + 通过率百分比 |
打个比方:val loss 是「平时刷题的正确率」——题见过、用来调状态;验收是「从没见过的高考卷」。用见过的题评估会虚高,所以必须专门留一份 test 集做最终验收。
3. 验收本质:批量跑推理 + 判分
验收自己不生成文字,它 import 了第五步的推理函数 generate() 来干活(src/evaluate.py 开头):
1 | from .infer import generate |
所以验收 = 一套流程:对每道题调用一次推理,拿到答案后判分,最后汇总通过率。整个循环就三步动作:
1 | 加载模型 → 逐题 generate 生成答案 → 对比标准答案判 PASS/FAIL → 算通过率 |
下面逐步拆开。
4. 第一步:加载模型
load_model()(src/evaluate.py)从 best.pt 里读出权重和当时的配置 cfg,重建一个一模一样的 GPT 结构:
1 | checkpoint = torch.load(ckpt_path, map_location="cpu", weights_only=False) |
两个关键点:
dropout=0.0:评估要稳定、可复现,不能随机丢神经元。model.eval():切到评估模式,关掉 dropout、启用 KV Cache(见第五步),保证输出确定。
5. 第二步:逐题生成答案
run_test() 读入 test.jsonl 的每一题,逐条调用 generate():
1 | response = generate( |
两个为验收特意选的参数:
temperature=0.0——贪心解码:每次都选概率最高的字,保证同一道题每次答案完全一样(可复现)。判卷必须稳定,不能这次 PASS、下次 FAIL。这跟第五步聊天时「要花样」用高 temperature 正好相反。stop_strings=["。",";","用户:"]:遇到句号就停、冒出「用户:」就刹车,避免小模型啰嗦跑题、自问自答。
test.jsonl 里的题长这样:
1 | {"prompt": "RoPE 是什么?", "completion": "旋转位置编码。"} |
6. 第三步:怎么判对错(关键词片段匹配)
这是 evaluate.py 里最有意思的一段。为什么不逐字全等比对? 因为语言模型的答案「意思对就行」,措辞可能不同——标准答案「旋转位置编码」,模型答「是一种旋转位置编码方式」,逐字比对会判 FAIL,但明显答对了。
所以项目用的是关键词片段过半命中:
1 | # 去掉标点 |
逻辑三步:
- 去标点:把期望答案和模型输出都清掉句号、逗号等,只留文字。
- 切片段:把标准答案按每 2 个字切成一串片段(「旋转/位置/编码」)。
- 过半命中:数模型输出里出现了几个片段,命中数 ≥ 片段总数的一半就算 PASS。
这是一种务实的「抓大意」判分法——不苛求一字不差,只看关键信息点有没有答到。
| 标准答案 | 切成片段 | 模型输出 | 命中 | 结果 |
|---|---|---|---|---|
| 旋转位置编码 | 旋转/位置/编码 | 是一种旋转位置编码 | 3/3 | PASS |
| 旋转位置编码 | 旋转/位置/编码 | 一种编码方式 | 1/3 | FAIL |
注意:这种匹配法简单但会有误差——答案里恰好出现片段字、但整体意思不对,也可能误判 PASS。对教学项目够用;生产级验收会用更严格的语义评测或人工标注。
7. 通过率与验收结论
逐题判完,汇总打印:
1 | rate = passed / total * 100 |
三档结论:
| 通过率 | 结论 | 建议 |
|---|---|---|
| ≥ 80% | 验收通过 | 可以进入第五步部署 |
| 50%~80% | 部分通过 | 增加训练步数 / 扩充数据 |
| < 50% | 未通过 | 查数据质量、调配置重训 |
跑法:
1 | uv run python -m src.evaluate --ckpt checkpoints/best.pt --test data/test.jsonl |
8. 核心结论速览
| 问题 | 一句话答案 |
|---|---|
| 第四步在干嘛 | 拿没见过的 test 集给模型「考试判卷」,算通过率 |
| 和 val loss 啥区别 | val loss 是见过的题;验收是全新的题,且看「整句答对没」 |
| 验收自己会生成吗 | 不会,它 import 第五步的 generate() 来跑推理 |
| 为什么 temperature=0 | 贪心解码,保证答案可复现,判卷才稳定 |
| 怎么判对错 | 把标准答案切 2 字片段,模型输出命中过半算 PASS |
| 为什么不逐字比对 | 语言答案「意思对就行」,措辞可不同 |
| 多少算过关 | ≥80% 通过,50~80% 建议加训,<50% 重来 |
| 交付哪个模型 | 验收用的是训练产出的 best.pt |
至此,模型已经过验收,下一步就是第五步——把它包成能聊天的 Web 服务,真正用起来。