模块 09 · 第 2 课

注意力机制

一个字要理解自己的意思,得看看前面都有哪些字。从最简单的"平均前面所有的字"出发,一步步加上查询、键、值,因果掩码和缩放,写出一个完整的注意力头,并和 PyTorch 的实现核对。

  • 约 50 分钟
  • 难度:深入
  • 实测:2026-09-15 torch 2.14,CPU,固定随机种子

"明月"的"明"和"明天"的"明",意思不一样。一个字该怎么理解,要看它周围是什么字。

我们要训练的模型,每个位置的任务是预测下一个字。它在"床前明月"之后要写出"光",就必须知道前面出现过"明月"。所以每个位置都需要一个办法,从前面的字那里收集信息。

注意力(attention)就是这个办法,它是 Transformer 的核心。这一课从最简单的做法开始,一步一步把它搭出来。

python attention.py

最简单的办法:取平均

每个字先用一个向量表示(第 01 模块第 5 课讲过向量嵌入)。最简单的"收集前面的信息",就是把自己和前面所有字的向量取平均。

这件事可以用一次矩阵乘法完成。构造一个下三角的矩阵,每一行只有前面若干个位置有值,而且加起来是 1:

x = torch.randn(T, 4)  # 5 个词,每个词一个 4 维向量(先随便取)
weights = torch.ones(T, T).tril()  # 下三角:第 i 行只有前 i+1 个位置是 1
weights = weights / weights.sum(dim=1, keepdim=True)  # 每行除以个数,变成平均
out = weights @ x
== 1. 用一个矩阵乘法,让每个位置得到自己和前面所有位置的平均
tensor([[1.00, 0.00, 0.00, 0.00, 0.00],
        [0.50, 0.50, 0.00, 0.00, 0.00],
        [0.33, 0.33, 0.33, 0.00, 0.00],
        [0.25, 0.25, 0.25, 0.25, 0.00],
        [0.20, 0.20, 0.20, 0.20, 0.20]])
  第 3 个位置('明')的结果 = 前三个向量的平均?True

第 i 行表示第 i 个字从每个位置拿多少:第 1 个字只能看自己,第 3 个字平均地看前三个,第 5 个字平均地看全部五个。weights @ x 一次就算出了所有位置的结果。

这个矩阵乘法的写法请记住,后面的注意力也是这个形式:一个权重矩阵乘以一堆向量。不同的只是权重怎么来。

平均的问题很明显:所有字一视同仁。"光"想知道的是前面有没有"月","床"和"前"对它来说没那么重要。权重应该由内容决定。

查询和键:让每个字决定该看谁

办法是给每个字准备两个向量:

  • 查询(query,q):我想找什么样的信息。
  • (key,k):我这里有什么样的信息。

一个字对另一个字的"关注程度",就用自己的 q 和对方的 k 做点积来打分。两个向量方向越接近,点积越大。

先用手工造的向量看看效果。设向量有两维,第一维表示"和天空有关",第二维表示"和位置有关"。"光"的查询是 [1, 0],表示它想找和天空有关的东西:

k = torch.tensor([[0.1, 0.9],   # 床:一个地点
                  [0.0, 1.0],   # 前:一个方位
                  [0.9, 0.1],   # 明:和天空、光有关
                  [1.0, 0.0],   # 月:天空里的东西
                  [0.8, 0.2]])  # 光
q_guang = torch.tensor([1.0, 0.0])  # "光"想找的是:和天空有关的东西
scores = k @ q_guang
== 2. 点积打分:查询(q)和每个键(k)越像,分数越高
  '光' 对每个词的分数: {'床': 0.1, '前': 0.0, '明': 0.9, '月': 1.0, '光': 0.8}
  softmax 之后的权重:   {'床': 0.12, '前': 0.11, '明': 0.26, '月': 0.29, '光': 0.23}
  分数放大 5 倍再 softmax: {'床': 0.01, '前': 0.0, '明': 0.3, '月': 0.5, '光': 0.18}

分数经过 softmax 变成权重(第 08 模块第 5 课讲过 softmax),加起来是 1。"月"的分数最高,拿到的权重最大。

注意最后一行:分数整体放大 5 倍,softmax 之后的权重就变得集中得多,"月"一个拿走了一半。分数的大小决定了注意力是"雨露均沾"还是"只盯着一个"。这一点在后面讲缩放时很重要。

真实的模型里,q 和 k 不是手工造的,而是用两个矩阵从每个字的向量算出来的,矩阵里的参数通过训练学到。模型会自己学会,什么样的字该去找什么样的信息。

值:真正传递的内容

打完分,决定了从每个位置拿多少,那拿的是什么?是第三个向量:(value,v),表示"如果你关注我,我给你这些信息"。

为什么不直接拿原来的向量?因为"用来匹配的信息"和"用来传递的信息"不一定一样。把它们分开,模型更灵活。

三个向量合在一起,就是注意力的全部:

权重 = softmax(q 和每个 k 的点积)
输出 = 用这些权重,把每个位置的 v 加起来

因果掩码:不能偷看后面

训练时,模型要在每个位置预测下一个字。如果第 3 个位置能看到第 4 个字,它就可以直接"抄答案",什么都学不到。所以每个位置只能看自己和前面的位置。

做法是在 softmax 之前,把后面位置的分数设成负无穷。负无穷取指数是 0,softmax 之后权重就是 0:

scores = q @ k.T / math.sqrt(8)
mask = torch.ones(T, T, dtype=torch.bool).tril()
scores = scores.masked_fill(~mask, float("-inf"))
== 3. 因果掩码:把后面位置的分数设成负无穷,softmax 之后权重就是 0
tensor([[1.00, 0.00, 0.00, 0.00, 0.00],
        [0.96, 0.04, 0.00, 0.00, 0.00],
        [0.24, 0.26, 0.50, 0.00, 0.00],
        [0.09, 0.39, 0.43, 0.09, 0.00],
        [0.02, 0.14, 0.02, 0.09, 0.73]])

和第一节的平均矩阵对比:形状一样,都是下三角,每行加起来都是 1。区别是现在每行的权重不再平均,而是由 q 和 k 决定的。这里的 q 和 k 是随机的,所以权重看起来没有规律;训练之后,它们会变得有意义。

因为有这个掩码,这种注意力叫因果自注意力:自注意力是指 q、k、v 都来自同一段文字,因果是指只能看过去、不能看未来。GPT 这类"一个字一个字往后写"的模型用的都是它。

为什么要除以根号维度

上面的代码里有一个 / math.sqrt(8),8 是 q 和 k 的维度。这一步叫缩放,看看不做会怎样:

== 4. 为什么要除以 √d:维度越大,点积的数值越大,softmax 会变得非常极端
  d=   16  点积的标准差   3.89(√d =  4.00)  每行最大权重的平均:不缩放 0.79,除以 √d 后 0.53
  d=   64  点积的标准差   8.17(√d =  8.00)  每行最大权重的平均:不缩放 0.86,除以 √d 后 0.47
  d=  256  点积的标准差  16.08(√d = 16.00)  每行最大权重的平均:不缩放 0.95,除以 √d 后 0.47
  d= 1024  点积的标准差  32.53(√d = 32.00)  每行最大权重的平均:不缩放 0.92,除以 √d 后 0.39

两个随机向量的点积,是 d 个乘积加起来。加的项越多,结果的波动越大,标准差大约是根号 d:d=256 时点积的标准差是 16。

分数的波动大了,就像第二节"放大 5 倍"那样,softmax 会变得极端:不缩放时,每一行最大的那个权重平均在 0.8 到 0.95,几乎所有注意力都给了一个位置。这在训练刚开始时很糟糕:权重接近 0 或 1 的地方,softmax 的梯度非常小,模型很难学习。

除以根号 d,分数的标准差回到 1 左右,每行最大的权重在 0.4 到 0.5 之间,注意力是分散的,梯度正常,模型可以慢慢学会该关注谁。

一个完整的注意力头

把前面的东西合在一起:

d_model, d_head = 16, 8
x = torch.randn(T, d_model)
W_q, W_k, W_v = (torch.randn(d_model, d_head) / math.sqrt(d_model) for _ in range(3))
q, k, v = x @ W_q, x @ W_k, x @ W_v
scores = (q @ k.T / math.sqrt(d_head)).masked_fill(~mask, float("-inf"))
out = F.softmax(scores, dim=-1) @ v
== 5. 一个完整的注意力头:q、k、v 都由同一个输入经过不同的矩阵得到
  输入 (5, 16) → q、k、v 各 (5, 8) → 输出 (5, 8)
  和 PyTorch 自带的 scaled_dot_product_attention 比,最大差别 2.4e-07

五行代码:三个矩阵把输入变成 q、k、v,打分、缩放、掩码,softmax,加权求和。结果和 PyTorch 自带的 scaled_dot_product_attention 一致(差别在小数点后 7 位,是浮点数计算顺序不同带来的误差)。

W_qW_kW_v 就是这个注意力头要学的参数。训练时,它们和其他参数一样,通过反向传播和梯度下降更新。

这个公式最早出自 2017 年的论文《Attention Is All You Need》,写成数学形式是 softmax(QKᵀ/√d)V。现在你知道其中每一个符号是做什么的了。

注意力的代价

每个位置都要和前面所有的位置打分。序列长度是 T,分数表就是 T×T 的大小。长度翻一倍,计算量和这张表占的内存就变成四倍。

这是长上下文很贵的根本原因之一(第 01 模块第 4 课讲过上下文和成本)。各家模型为了支持几十万、上百万词元的上下文,在注意力的计算方式上做了大量优化,但基本的思路还是这一课的这几行。

练习

  1. 在第二节里,给"光"换一个查询向量 [0, 1](想找和位置有关的东西),权重会怎么变?
  2. 在第五节里,去掉因果掩码,再和 scaled_dot_product_attention(..., is_causal=False) 对比。
  3. 把第五节的序列长度从 5 改成 1000、2000、4000,用 time.time() 测一下每种长度要多少时间,是不是大约按平方增长?

自测

1. 查询、键、值分别起什么作用?

查询表示"我想找什么",键表示"我这里有什么"。用一个位置的查询和所有位置的键做点积,得到它对每个位置的关注分数,经过 softmax 变成权重。值表示"关注我时能拿到的信息",输出就是按这些权重把所有位置的值加起来。

2. 因果掩码是怎么实现的?为什么需要它?

在 softmax 之前把每个位置后面的分数设成负无穷,softmax 之后这些位置的权重就是 0。需要它是因为模型要学习预测下一个字,如果能看到后面的字,就能直接抄答案,什么也学不到。

3. 注意力分数为什么要除以根号维度?

d 维的两个向量的点积,标准差大约是根号 d,维度越大分数的波动越大,softmax 之后的权重会极端地集中在一个位置,梯度变得很小,模型难以训练。除以根号 d 让分数的标准差回到 1 左右。

提问与讨论

这一课没看懂的地方,在这里问。看到别人的问题,也欢迎你来回答。

提问 +3 积分,回答别人 +6 积分。内容经审核后公开。

正在加载讨论…