一字一诗
LESSON 07 / 26

模型如何根据错误学习

这一课完成什么

理解“学习”实际怎样修改数字。先用只有一个参数的例子把纠错过程算清楚,再对应到 GPT。PyTorch 能自动求出误差对参数的变化方向,我们仍然需要知道这个方向代表什么。

根据误差,给一个数字指方向

先只训练一个数字 w:希望它接近 3。误差取 (w−3)²,梯度 2(w−3) 告诉我们当前位置往哪边变化会增加误差。

根据误差,给一个数字指方向
用文字逐步读这张图
  1. 当前参数:w = 0;从一个数字开始
  2. 计算误差:(0 − 3)² = 9;越偏离越大
  3. 求梯度:2 × (0 − 3) = −6;当前变化方向
  4. 选调整幅度:学习率 = 0.2;步子大小
  5. 修改参数:0 − 0.2 × (−6);w 变为 1.2
  6. 再看误差:(1.2 − 3)² = 3.24;重复这个过程
学习率越大就一定越快学好吗?

不是。步子太大会越过最低点,甚至越走越远。

一个可以手算的目标

假设参数叫 w,希望它接近 3。误差用 (w-3)²。从 w=0 开始,误差为 9。误差对 w 的导数是 2(w-3),在 0 处为 -6。这个导数告诉我们:在当前位置略微增大 w,会让误差下降。

每次更新使用 新 w = 旧 w - 调整幅度 × 导数。如果调整幅度为 0.2,第一次变成 0 - 0.2×(-6) = 1.2;下一次导数为 -3.6,变成 1.92。数字逐渐靠近 3。

次数更新前 w导数更新后 w
10-61.2
21.2-3.61.92
31.92-2.162.352

导数推广到很多参数,得到每个参数对应的变化方向,通常叫“梯度”。这里下降的是误差;更新方向取梯度的反方向。

把四步写进代码

python
optimizer.zero_grad()
loss = (parameter - 3) ** 2
loss.backward()
optimizer.step()

zero_grad 清除上一次积累的梯度;前向计算得到本次误差;backward 沿着计算关系求各个参数的梯度;step 根据梯度更新参数。backward 本身没有更新参数,真正的更新发生在 step

bash
./poet lesson 7
.venv/bin/python lessons/07/experiment.py

检查第一行是否显示参数 0、误差 9、梯度 -6。实验打印的是更新前状态,所以第二行才会看到 w=1.2。最后应逐渐接近 3。

从一个参数到 GPT

GPT 的误差由许多乘法、加法、softmax 等运算共同产生。链式法则把“最后误差如何变化”逐层传回每个参数。例如 z=w×x,误差依赖 z,则误差对 w 的变化等于误差对 z 的变化乘以 x。PyTorch 根据实际计算保存这条关系链,因此无需手工为每一层写导数。

正式程序使用 AdamW 优化器,它会参考过去梯度的情况调整各参数的更新,并带有权重衰减。你仍可以用本课四步理解训练循环,只是 step 内的规则比固定幅度的普通梯度下降更细致。课程不要求你先手写 AdamW 才开始训练。

常见错误

小练习与答案

把本课的 0.2 改成 0.05,再改成 5,比较十次更新。先预测后运行。

查看答案

0.05 通常更慢地靠近 3;5 会大幅越过目标并使误差增大。具体例子可直接算:w=0 时若幅度为 5,更新后为 30,已经越过 3。模型训练中的学习率就是这种调整幅度的控制项,但复杂误差曲面不像这个平方函数一样简单。

动手补全一小段

先运行上面的完整实验,再复制本课起始文件为自己的练习。starter 有意留空 solve 函数;补全后运行文件,底部检查会告诉你是否符合本课要求。卡住时打开参考答案,比较每一步。

bash
cp lessons/07/starter.py lessons/07/my_exercise.py
# 编辑 my_exercise.py 中的 solve 函数,然后运行:
.venv/bin/python lessons/07/my_exercise.py
# 对照完整答案:
.venv/bin/python lessons/07/solution.py

下载起始代码 · 下载参考答案

展开本课补全练习的完整参考答案
python
"""第 07 课补全练习:根据梯度更新一次。修改 solve,保持下方检查不变。"""
import math, io, argparse
import torch
from torch import nn
from torch.nn import functional as F
torch.set_num_threads(2)
torch.manual_seed(26)

def solve(weight, gradient, rate):
    return weight - rate * gradient

assert abs(solve(0.,-6.,.2)-1.2)<1e-9
print("本课补全练习通过。")
展开本课完整、可独立运行的实验代码
python
"""本课独立实验;在仓库根目录执行 .venv/bin/python lessons/07/experiment.py。"""
import sys, json, math
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
import torch
from torch import nn
from torch.nn import functional as F
from poetry_gpt.common import ROOT, DEFAULT_DATA, read_json, read_jsonl
from poetry_gpt.model import ModelConfig, PoetryGPT, CausalAttention, Block
from poetry_gpt.data import Tokenizer, SPECIAL, clean_record, keywords_for
from poetry_gpt.labs import show
torch.set_num_threads(2)
torch.manual_seed(26)

parameter = nn.Parameter(torch.tensor(0.0))
opt = torch.optim.SGD([parameter], lr=0.2)
for step in range(10):
    opt.zero_grad()
    loss = (parameter - 3) ** 2
    loss.backward()
    print(f'第{step}次 数字={parameter.item():.4f} 误差={loss.item():.4f} 梯度={parameter.grad.item():.4f}')
    opt.step()
下载本课实验