这一课完成什么
第 10 课直接给出了分数,这一课解释分数从哪里来。GPT 从当前每个位置的表示,经过三组可学习的变换,分别得到 Q、K、V。它们是 query、key、value 的简写,这里分别理解为“拿来比较的查询表示”“被比较的表示”“实际汇总的内容”。
Q、K、V 从同一组输入分别变换而来
Q、K、V 都是数字矩阵。三个可学习的变换分别回答:当前想找什么、各位置有什么线索、取回来什么信息。
用文字逐步读这张图
- 输入 X:每个位置一行;示例宽度为 2
- Q = X × Wq:查询用的表示;不是另一份文本
- K = X × Wk:匹配用的表示;与 Q 逐对比较
- V = X × Wv:被汇总的信息;与 K 用途不同
- Q × K 的转置:得到位置之间分数;除以 √每头宽度
- 权重 × V:变换为概率再汇总;输出仍按位置排列
为什么要把 K 转置?
这样每一行 Q 才能与每一行 K 做点积,得到当前位置对所有位置的分数。
三组变换来自同一份输入
Q = X × Wq
K = X × Wk
V = X × WvX 是输入,三个 W 是训练会更新的参数。Q 与 K 用来算每两个位置的匹配分数;V 才是第 10 课最后被加权相加的信息。三组变换可以学到不同用途,不需要人工规定哪个字应该关注谁。
用两项数字手算一次
X = [[1,0], [0,1], [1,1]]
Wq = [[1,1], [0,1]]
Wk = [[1,0], [1,1]]
Wv = [[2,0], [0,1]]相乘得到 Q=[[1,1],[0,1],[1,2]],K=[[1,0],[1,1],[2,1]],V=[[2,0],[0,1],[2,1]]。以 Q 第一行 [1,1] 分别与 K 三行点乘,得到 [1,2,3]。这就是第一个位置对三个位置的原始匹配分数。
此时还没有遮挡未来位置,所以先只看数学结构。第 12 课会把不允许看的格子去掉。
为什么除以一个数
每个注意力头的比较向量若有 d 项,匹配分数使用 Q @ K.T / sqrt(d)。向量越长,点乘累加的项越多,分数可能变得很大,让 softmax 过于集中。除以宽度的平方根用于控制这种尺度。它是模型结构中的固定缩放,不是“正确答案概率”。
本例 d=2,所以第一行分数是 [1,2,3] / 1.414。之后才进行遮挡与 softmax,再乘 V。完整顺序为:比较、缩放、遮挡、变成权重、汇总。
./poet lesson 11
.venv/bin/python lessons/11/experiment.py找到实际源码
正式实现用一个线性层一次输出三个宽度的结果,再拆成三份。这样与分别建立三个线性层具有相同的基本含义,但更便于执行。
self.qkv = nn.Linear(config.width, 3 * config.width)
q, k, v = self.qkv(x).chunk(3, dim=-1)这里包含偏置,因此更完整地写应是 XW+b。课程手算为了减少干扰没有加偏置,正式代码没有把它隐藏起来。chunk(3) 沿最后一维平均切成三块,位置数量和批数量不变。
验收与排错
- 能手算 X 第一行经过 Wq、Wk、Wv 后分别是什么。
- 知道分数表的行是查询位置,列是被看的位置。
.transpose(-2,-1)交换 K 最后两维,使位置与位置相互比较。- 若把 QK 相乘写成逐元素乘法,可能仍能得到数字,却不是本课要求的匹配表。
小练习与答案
把 Wq 改成单位矩阵 [[1,0],[0,1]],哪些输出会改变?
查看答案
Q 会变为 X 本身,K 和 V 不变。QK 转置得到的分数、随后权重以及汇总结果可能改变。第一行 Q 变为 [1,0],与三个 K 点乘得到 [1,1,2]。这说明改变“怎样比较”也会影响“最终汇总了什么”。