本文是「GPT 训练分步学习路径」的第三步,基于
GPT_teacher-3.37M-cn这个教学项目,讲清楚「模型是怎么从数据里一步步学会预测的」。文中所有代码均来自本项目真实的src/train.py和config/config.yml。分步学习路径共五步:① 看懂数据长什么样 → ② 模型核心(注意力+因果掩码)→ ③ 训练循环 → ④ 验收评估 → ⑤ 推理与部署。本文是第三步。
目录
- 这一步属于哪个环节
- 训练的本质:把学习变成一道优化题
- 训练四步:一个 batch 的完整生命周期
- 前向里到底发生了什么
- 反向传播:只有一行,却是引擎核心
- 调的到底是哪些参数
- 梯度累积:用小内存换大 batch
- 学习率调度:warmup + 余弦退火
- loss 是怎么来的
- 早停与最优模型保存
- 核心结论速览
1. 这一步属于哪个环节
整个项目的全流程是:
1 | 构建分词器 → 数据处理 → 训练 → 验收 → Web 演示 |
本文讲的是 训练 环节,对应代码 src/train.py。它的作用只有一个:
拿第二步搭好的模型(一堆随机初始化的参数),喂第一步加工好的数据(
x/y),通过反复调整参数,让模型学会「看到前文,预测下一个字」。
这一步是「数据处理」的下游(直接吃它产出的 x / y),也是「验收」的上游(产出训练好的 best.pt)。
2. 训练的本质:把学习变成一道优化题
在讲代码前,先讲清楚「为什么这么玩能学会」。整个深度学习的底层逻辑是三步转化:
- 把问题定义成函数:模型就是一个带 3.37M 个可调数字(参数) 的大函数——输入前文,输出「下个字是词表里每个字的概率」。
- 定义「错得多严重」= loss:用损失函数把「预测好不好」变成一个可量化的数字。
- 把学习变成搜索:在 3.37M 维空间里,找一组参数,让所有数据的 loss 最小。
于是「学习」就变成了一道纯数学题:怎么调这 3.37M 个数字,让 loss 最小? 答案就是下面的训练循环。
3. 训练四步:一个 batch 的完整生命周期
训练循环的核心在 src/train.py 的 while 循环里。每取一批数据,都会走这四步:
1 | logits, _ = model(xb) # ① 前向:算出预测 |
| 步骤 | 代码 | 干什么 | 类比 |
|---|---|---|---|
| ① 前向 | model(xb) |
数据从头流到尾,算出预测和 loss | 做一套卷子 |
| ② 算 loss | loss_fn(...) |
量化「预测离答案多远」 | 算错误率 |
| ③ 反向 | loss.backward() |
从 loss 倒推每个参数的梯度 | 复盘:谁该为扣分负责 |
| ④ 更新 | opt.step() |
按梯度把每个参数调一小步 | 针对性改进 |
顺序铁律:先 forward → 再 loss → 再 backward → 再 step,一步都不能颠倒。
本文后续章节,其实就是把这四步逐一放大来讲:下面先单独展开第 ① 步(第 4 节 前向内部链路),再讲第 ③ 步(第 5 节 反向传播)、第 ④ 步(第 6 节 调哪些参数、第 7 节 梯度累积、第 8 节 学习率调度)、第 ② 步(第 9 节 loss 怎么来)。看后面章节时,记得它们都挂在这四步之下,而不是新话题。
4. 前向里到底发生了什么:一次 model(xb) 的内部链路
上面表格里「数据从头流到尾」这一句,其实浓缩了第二步讲的整套注意力计算。展开看,model(xb) 内部对每一层 Block 都走一遍下面的流程(src/model.py):
1 | xb (一批字 id) |
关键点:Wq / Wk / Wv 这三个投影矩阵是固定的参数(训练学出来、存进 best.pt),但输入向量因上下文和位置而不同,所以算出的 Q/K/V 也随之不同——同一个字在不同句子里,QKV 是不一样的。前向的产出就是 logits:一批数据里每个位置对「下个字是词表里哪个」的打分。
注意区分「谁在变」:前向阶段 Wq/Wk/Wv 不变(它们是本轮要评估的参数),变的是输入数据算出的 QKV 和最终 logits。真正去调整 Wq/Wk/Wv 的,是后面第 ③④ 步(反向 + 更新)。
有了 logits,第 ② 步就能拿它和真实答案 yb 算 loss 了。
5. 反向传播:只有一行,却是引擎核心
反向传播是训练能「学」起来的关键,但在代码里它只体现为一行:
1 | loss.backward() |
「反向」这个名字,是相对「前向」而言的:
1 | 前向: 输入 ──→ 每一层 ──→ 预测 ──→ loss |
它背后靠两样东西:
- 计算图:前向时 PyTorch 偷偷记下了每一步运算(谁乘谁、谁加谁),连成一张图。
- 链式法则:
backward()沿这张图从 loss 往回走,算出「loss 对每个参数的偏导数」= 梯度。
算完的梯度存在每个参数的 .grad 里,等 opt.step() 来用。没有人手写求导公式——这整套计算由 PyTorch 的 autograd 引擎全自动完成,这也是为什么全世界训 GPT 都用 PyTorch。
类比:
loss.backward()就像按下「复盘」按钮。考完试从总扣分倒着追查,算出每道题(每个参数)该为这次扣分负多大责任。
6. 调的到底是哪些参数
opt.step() 更新的,是 model.parameters() 里的每一个数字。优化器一句就把全部参数接管了:
1 | opt = torch.optim.AdamW(model.parameters(), lr=..., weight_decay=...) |
对照 src/model.py,被训练的参数清单:
| 组件 | 参数 | 作用 |
|---|---|---|
词嵌入 tok_emb |
[vocab, 256] 向量表 |
字 id → 向量(与输出头权重共享) |
注意力 wq/wk/wv/proj |
4 个矩阵 × 4 层 | 投影出 Q/K/V,再把多头结果合回去 |
MLP w_gate/w_up/w_down |
3 个矩阵 × 4 层 | 前馈网络的非线性加工 |
| RMSNorm 的 γ | 每层 2 个 + 最终 1 个 | 归一化的可学习缩放 |
输出头 head |
与 tok_emb 共享 |
向量 → 词表打分 |
所以 tok_emb 不是被单独调的,它只是这张大网里的一块。反向传播会一路把「责任」分摊到上面每一个参数,opt.step() 让它们同时各挪一小步。全部加起来 sum(p.numel()) ≈ 3.37M——项目名里的数字就是这么来的。
7. 梯度累积:用小内存换大 batch
一个常见误解:以为「每取一批就更新一次参数」。本项目其实是攒够几批才更新一次。看 config.yml:batch_size=16、micro_batch=4,即攒 16÷4=4 个小批才调一次:
1 | loss.backward() # 每批都算,梯度「累加」到 .grad(不清空) |
关键机制:backward() 的梯度是逐元素累加的,不会自动清空。4 个小批各自算出的梯度相加成一个总和,再统一更新一次——效果等同于一次算 16 条,但内存只需装下 4 条。
| 策略 | 每轮更新次数 | 问题 |
|---|---|---|
| 每批就调 | 最多 | 方向太抖 |
| 攒 4 批调(本项目) | 适中 | 甜点区 |
| 全部累积完才调 | 太少 | 学不动 + 易卡局部最优 + 爆内存 |
类比:搬 16 箱货但推车一次只能装 4 箱。跑 4 趟全搬完才统一记一次账,结果和「一次搬 16 箱」一样,只是分了 4 趟省力气。
8. 学习率调度:warmup + 余弦退火
学习率(lr)决定「每一步迈多大」。本项目不用固定 lr,而是分两段变化(src/train.py 的 lr_lambda):
1 | def lr_lambda(step): |
- ① Warmup(前 100 步):lr 从 0 线性爬到最大值。刚初始化的模型参数是随机的,一上来迈大步容易「翻车」,所以先小步热身。
- ② 余弦退火:热身后沿 cos 曲线平滑降到 0。
| 训练进度 t | lr 系数 | 实际 lr |
|---|---|---|
| 0(刚热身完) | 1.0 | 满速 0.001 |
| 0.5(训练一半) | 0.5 | 半速 0.0005 |
| 1(训练结束) | 0.0 | 降到 0 |
为什么用余弦:它两头慢、中间快——开头保持大步长快速学,结尾用极小步长精细收敛。像开车进车位,越靠近越慢慢挪,最后稳稳停好。
9. loss 是怎么来的
loss 是每个 batch 算一个(loss_fn(logits, yb)),衡量「这一批数据上,预测离正确答案多远」。训练时刷屏的 step 100 loss 3.24,就是那一批的值——会上下抖动,但整体应震荡着往下走。
要区分三种 loss:
| 种类 | 何时算 | 含义 |
|---|---|---|
| 单批 train loss | 每个 batch | 刷屏看的就是它,会抖 |
| 验证 val loss | 每 50 步一次 | 跑整个验证集取平均,更可信,是早停依据 |
| 一次完整训练 | —— | 产出的不是单个值,而是一条 loss 曲线 |
用的损失函数是 CrossEntropyLoss(ignore_index=-100)——第一步讲过,-100 让模型只学答案部分、跳过问题部分。
10. 早停与最优模型保存
训练不是跑满 5000 步就好,而是盯着 val loss 见好就收:
1 | if eval_loss < best_val_loss: # 创新低 |
- best.pt:val loss 每创新低就存一次,保留「历史最好」的那份权重。
- 早停:连续 15 次评估都没改善,就提前停——继续练下去只会过拟合,浪费时间。
这也是为什么最终交付的是 best.pt 而不是 last.pt:我们要的是验证集上表现最好的那一刻,不是训到最后一步的那一份。
11. 核心结论速览
| 问题 | 一句话答案 |
|---|---|
| 训练在干什么 | 反复调 3.37M 个参数,让所有数据的 loss 最小 |
| 一个 batch 走几步 | 前向 → 算 loss → 反向 → 更新(四步) |
| 反向传播在哪 | 就是 loss.backward() 一行,autograd 自动算全部梯度 |
| 调哪些参数 | model.parameters() 全部:tok_emb、wq/wk/wv/proj、MLP、RMSNorm |
| 为什么攒 4 批才更新 | 梯度累积:用小内存达到大 batch 效果 |
| lr 怎么变 | 先 warmup 线性爬升,再余弦退火平滑降到 0 |
| loss 什么时候算 | 每个 batch 一个;val loss 每 50 步一次,是早停依据 |
| 什么时候停 | val loss 连续 15 次不改善就早停,交付 best.pt |
至此,模型已经从「一堆随机数」训练成「能预测下个字」的状态。下一步(第四步「验收评估」)将回答:这个训练好的模型,到底行不行?