一字一诗
LESSON 14 / 26

加工信息并让计算稳定

这一课完成什么

理解每个 GPT 层里除了注意力以外的三个部件:前馈网络、残差连接、归一化。它们分别负责进一步加工每个位置的信息、保留原来的信息通道,以及让数字尺度更容易控制。本项目把这些部件直接写在 Block 类里。

一条保留旧信息,一条加工新信息

残差连接把输入直接加回处理结果。归一化则调整同一位置内部数字的尺度,让后续计算更容易保持稳定。

一条保留旧信息,一条加工新信息
用文字逐步读这张图
  1. 输入 x:[1, 2, 3, 4];原来的表示
  2. 归一化:先减均值,再调尺度;再学习缩放与偏移
  3. 注意力处理:汇总不同位置;得到改动 f(x)
  4. 加回输入:x + f(x);保留直接通路
  5. 逐位置加工:放宽 → GELU → 缩回;不直接跨位置混合
  6. 再次相加:旧表示 + 新改动;交给下一层
前馈网络和注意力分别处理什么?

注意力混合不同位置的信息;当前前馈网络独立加工每个位置的特征。

逐位置加工

注意力让不同位置交换信息;前馈网络随后独立加工每个位置已经汇总的表示。所有位置共用同一组加工参数。先把宽度 C 扩大到 4C,经过非线性函数 GELU,再缩回 C。

python
nn.Sequential(
    nn.Linear(width, 4 * width),
    nn.GELU(),
    nn.Linear(4 * width, width),
    nn.Dropout(dropout),
)

如果两次线性变换之间没有非线性,它们可以合成一次线性变换,表达能力受到限制。GELU 对不同大小的输入以不同方式放行,使组合能够表达更复杂的关系。你不必先背它的公式,但应知道它会真正改变映射方式,而非只重排尺寸。

保留一条直接路径

残差连接的形式是 x + 新加工结果。例如原来是 [1,2,3,4],分支提供 [0.1,-0.1,0.2,-0.2],相加得到 [1.1,1.9,3.2,3.8]。分支可以学习增补和修正,原来的表示有一条直接通路继续向后传递。

它也为误差反传提供了直接路径,帮助较深的网络训练。残差连接要求相加双方尺寸相同,所以前馈网络最后必须缩回 C,注意力合并后也保持 C。

稳定数字尺度

LayerNorm 在每个位置内部,沿表示宽度求均值和方差,先减去均值,再除以标准差,之后还有可学习的缩放和偏移。分母加入很小的数以避免零方差问题。初始化的缩放为 1、偏移为 0,此时结果均值接近 0;训练以后不能要求每次结果均值都严格为 0。

bash
./poet lesson 14

把输入整体加上 100,再观察初始归一化输出,结果应接近原来。归一化并不意味着删除全部差异;各项之间的相对关系仍保留,并且模型还能学习缩放与偏移。

一个 Block 的实际顺序

python
x = x + self.attention(self.norm_attention(x))
x = x + self.feedforward(self.norm_feedforward(x))

这是先归一化再进入分支的设计。第一个分支做跨位置交流,第二个分支做逐位置加工;每个分支结束都与原输入相加。Dropout 训练时随机丢弃部分分支数值,减轻对特定通路的依赖;model.eval() 会关闭它。

验收与排错

小练习与答案

若前馈层把 C 扩成 4C 后没有缩回,为什么不能直接与 x 相加?

查看答案

x 最后一维为 C,而分支输出为 4C,两者含义与尺寸都不同,不能逐项相加。必须有一层把分支投影回 C,或者明确重新设计另一条路径。保持相同尺寸也是堆叠多个 Block 的便利所在。

动手补全一小段

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

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

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

展开本课补全练习的完整参考答案
python
"""第 14 课补全练习:保留残差通路。修改 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(x, branch):
    return x + branch(x)

x=torch.tensor([1.,2.])
torch.testing.assert_close(solve(x,lambda x:x*.1),torch.tensor([1.1,2.2]))
print("本课补全练习通过。")
展开本课完整、可独立运行的实验代码
python
"""本课独立实验;在仓库根目录执行 .venv/bin/python lessons/14/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)

x = torch.tensor([[1.0, 2.0, 3.0, 4.0]])
norm = nn.LayerNorm(4)
show('归一化后的均值', norm(x).mean().item())
show('整体加100后的结果', norm(x + 100).tolist())
branch = torch.tensor([[0.1, -0.1, 0.2, -0.2]])
show('残差相加', (x + branch).tolist())
下载本课实验