PyTorch 2.x 深度学习专题【左扬精讲】—— 注意力机制到底在做什么,Q/K/V 怎么来的?一文读懂 Attention 注意力机制

PyTorch 2.x 深度学习专题【左扬精讲】—— 注意力机制到底在做什么,Q/K/V 怎么来的?一文读懂 Attention 注意力机制

本文是 PyTorch 2.x 深度学习专题【左扬精讲】 系列:把 Transformer 里的 Self-Attention 一次性讲透——Scaled Dot-Product Attention 的核心公式到底在算什么,Q/K/V 到底是怎么来的,为什么要做 Scale,为什么最后要乘 V,以及多头(Multi-Head)到底把 Q/K/V 拆成了几份。

本文的思路借鉴 PP 鲁 知乎专栏 "注意力机制到底在做什么,Q/K/V 怎么来的?一文读懂 Attention 注意力机制"(原文发布于 https://zhuanlan.zhihu.com/p/414084879)一文从向量点乘出发、逐步拆解 Attention 的方式,并把所有可计算的步骤落到 PyTorch 2.x 的最小可运行示例上。

torch.matmul             ← 向量点乘 / 矩阵乘法的最基础 API
torch.softmax            ← 把相似度归一化为概率分布
torch.nn.functional.softmax  ← 与 torch.softmax 等价,函数式写法
torch.nn.Linear          ← 用作 Q / K / V 的线性投影 W^Q / W^K / W^V
torch.bmm / torch.matmul ← 批次矩阵乘:Attention 的矩阵化主战场
torch.nn.MultiheadAttention  ← PyTorch 官方提供的多头注意力实现(参考)

AttentionQ / K / VScaled Dot-ProductMulti-HeadTransformerPyTorch 2.x

学习重点(必读)

    • Attention 的本质:把"序列里哪些 token 之间相关、相关度多少"算成一个加权求和。
    • Q / K / V 的来源:同一个输入矩阵 X 分别乘上三个可训练参数矩阵 W^QW^KW^V 得到。
    • Scaled Dot-Product 的四步Q K^T/ sqrt(d_k)softmax乘 V
    • Multi-Head:把 d_model 维度切成 h 个头并行做 Attention,再拼回去。

先修知识(了解即可)

    • 向量点乘(Dot Product)的几何意义。
    • 矩阵乘法 A @ B 的基本运算规则。
    • Softmax 函数把一个向量归一化为概率分布。

一、从向量点乘说起

如果我们不直接面对 Q/K/V 这三个看上去很神秘的字母,而是先看一个更朴素的表达式 Softmax(X X^T) X,会发现 Attention 其实就是把这件事做到了极致。这个公式看起来吓人,但拆开看,每一步都只是 "点乘 + 归一化 + 加权" 三件套。

What — 向量点乘是什么?

向量点乘(Dot Product)是两个长度相同的向量之间最基本的运算。给定行向量 image 和 image ,点乘结果是一个标量:

向量点乘公式 image 

几何上,x · y 等于 xy 方向上的投影长度,乘以 y 自己的长度。它衡量的是两个向量方向上的对齐程度:点乘结果越大,两个向量越相似(方向越接近);点乘结果接近 0,说明两者几乎正交;点乘为负,说明方向大致相反。

1.1 矩阵 X 与自身的转置相乘:得到"相似度矩阵"

n 个行向量堆成一个 n × n 的方阵 X,再用 X 与自身的转置 X^T 做矩阵乘法,得到的结果 X X^T 是一个 n × n 的矩阵,元素 (i, j) 就是向量 x_ix_j 的点乘,也就是它们的相似度:

相似度矩阵 

image 

图 1-1  词向量矩阵 X 与自身转置相乘,得到"谁跟谁相似"的矩阵
┌─────────────────────────────┐
│   x_1   x_2   x_3   x_4     │  ← X (n 行,每行是一个词的向量)
│  ┌──┐  ┌──┐  ┌──┐  ┌──┐    │
│  │  │  │  │  │  │  │  │    │
│  │  │  │  │  │  │  │  │    │
│  └──┘  └──┘  └──┘  └──┘    │
└─────────────────────────────┘
       │    │
       │    │   X @ X.T   (矩阵乘法)
       ▼    ▼
┌─────────────────────────────────────┐
│ x1·x1  x1·x2  x1·x3  x1·x4         │  ← 相似度矩阵 S
│ x2·x1  x2·x2  x2·x3  x2·x4         │     每行 = 某个词与所有词的相似度
│ x3·x1  x3·x2  x3·x3  x3·x4         │
│ x4·x1  x4·x2  x4·x3  x4·x4         │
└─────────────────────────────────────┘

1.2 套上 Softmax:相似度变权重

对相似度矩阵的每一行做 Softmax,得到的就是一个 行和为 1 的权重矩阵。第 i 行第 j 列的值,可以直接解释成"在考虑第 i 个 token 时,应该分配给第 j 个 token 的注意力比例"。

图 1-2  Softmax 行归一化:每一行变成一个"在词 i 上的注意力分布"
       相似度矩阵 S                        权重矩阵 A (按行 Softmax)
┌────────────────────┐                ┌──────────────────────────┐
│  3.2   0.5  -1.0   │   ────►       │ 0.83  0.14  0.03          │  ← "看到词 1 时
│  0.8   2.1   0.2   │   按行         │ 0.10  0.78  0.12          │     关注 1 的概率最大"
│ -0.5   0.1   1.4   │   Softmax     │ 0.11  0.17  0.72          │
└────────────────────┘                └──────────────────────────┘
                                       每行和 = 1,元素 ∈ (0, 1)

1.3 再乘回 X:一次完整的"加权求和"

把权重矩阵 A = Softmax(X X^T)X 本身相乘,得到 Softmax(X X^T) X。它的第 i 行就是"所有 token 的向量"按第 i 行的权重做的加权平均——也就是把上下文信息聚合进了第 i 个 token 的新表示里。

图 1-3  Softmax(X X^T) X 就是"按权重把 X 的行加在一起"
   权重矩阵 A              X                  输出 Z = A @ X
┌───────────────┐    ┌───────────────┐    ┌─────────────────────────────┐
│0.83 0.14 0.03 │    │ x_1 x_2 x_3   │    │ 0.83·x_1+0.14·x_2+0.03·x_3 │
│0.10 0.78 0.12 │ @  │ y_1 y_2 y_3   │ =  │ 0.10·y_1+0.78·y_2+0.12·y_3 │
│0.11 0.17 0.72 │    │ z_1 z_2 z_3   │    │ 0.11·z_1+0.17·z_2+0.72·z_3 │
└───────────────┘    └───────────────┘    └─────────────────────────────┘
                                            每一行都是"按权重混合后的新表示"
How — 用 PyTorch 把"向量点乘 + Softmax + 加权求和"三件套写成可运行代码

下面这段代码就是上面三张图对应的最小实现:构造 X,先算相似度 X @ X.T,再按行 Softmax 得到权重,最后用权重乘回 X。运行后能看到每一行的权重加起来等于 1。

import torch                            # 导入 PyTorch 主包
import torch.nn.functional as F         # 导入函数式 API(下面用 F.softmax)

# 构造一个 3×3 的"词向量"矩阵 X,3 行代表 3 个 token
x = torch.tensor(
    [[1.0, 3.0, 2.0],                  # 词 1 的向量
     [1.0, 1.0, 3.0],                  # 词 2 的向量
     [1.0, 2.0, 1.0]],                 # 词 3 的向量
    dtype=torch.float64,               # 用 float64 让数值打印更稳定
)

# 第 1 步:X @ X.T 得到相似度矩阵(行 = 查询 token,列 = 被比 token)
similarity = torch.matmul(x, x.transpose(-1, -2))   # 转置最后两个维度再做矩阵乘

# 第 2 步:按行 Softmax,把相似度变成"加起来等于 1"的权重分布
attn_weights = F.softmax(similarity, dim=-1)        # dim=-1 表示对每一行做归一化

# 第 3 步:用权重矩阵乘回 X,得到加权求和后的新表示
output = torch.matmul(attn_weights, x)              # (3,3) @ (3,3) → (3,3)

print("similarity =", similarity)         # 打印相似度矩阵
print("attn_weights =", attn_weights)     # 打印权重矩阵(每行和 = 1)
print("output =", output)                 # 打印加权求和后的结果

运行后 attn_weights 的每一行相加都恰好等于 1.0,且 output 的每一行就是 x 三行的加权和。整个流程就是 Attention 最朴素的雏形:只不过 Attention 把"用 X 自己做自己的查询"换成"用 Q 去查 K",并且引入了可学习的参数矩阵。

本章小结

  • 点乘是相似度X X^T 的每个元素,就是两个向量的相似度。
  • Softmax 是归一化:把相似度按行归一化成"加起来等于 1"的权重分布。
  • 乘回 X 是加权求和:用权重把所有 token 的向量聚合,得到带上下文的新表示。
  • Attention 的雏形Softmax(X X^T) X 已经是一个"按注意力加权"的运算,只是它没有可学习参数,也没有缩放。

二、Q / K / V 从哪儿来

上一节的 Softmax(X X^T) X 用的是同一份 X 同时承担 "查询" "被查" 两个角色。Transformer 把它分成了三份:对同一个输入 X,分别用三组可学习的参数矩阵投影到三个不同的子空间,得到 QKV

2.1 讲讲公式

Q / K / V 的来源其中,Q 为 Query、K 为 Key、V 为 Value。Q、K、V是从哪儿来的呢?Q、K、V 其实都是从同样的输入矩阵X线性变换而来的。
我们可以简单理解成:  image 
其中,W^Q, W^K, W^V 是三个可训练的参数矩阵(Linear 层),W^Q、W^K、W^V 都是模型参数,训练时通过反向传播学习。
注意这里一定要区分清楚两个 W:下面公式里的 W^Q/W^K/W^V 是网络参数,下面的 W(出现在"加权求和"语境里)指的是样本权重——两者完全不同。
图 2-1  同一份输入 X 并行乘三个参数矩阵,得到 Q / K / V
                        ┌───────────────┐
                  ┌────▶│   W^Q         │──▶ Q   (n × d_q)
                  │     └───────────────┘
                  │
┌─────────────┐   │     ┌───────────────┐
│   X         │───┼────▶│   W^K         │──▶ K   (n × d_k)
│ (n × d_model)│   │     └───────────────┘
└─────────────┘   │
                  │     ┌───────────────┐
                  └────▶│   W^V         │──▶ V   (n × d_v)
                        └───────────────┘
              ↑ X 是输入矩阵,每行一个 token 的表示向量(Embedding + Positional Encoding)
              ↑ W^Q / W^K / W^V 是三个 Linear 层的权重,是模型可学习参数
              ↑ Q / K / V 与 X 同形状(行数 = 序列长度 n),列数可以相同也可以不同

2.2 Q、K、V 三个角色到底怎么分工

QKV 是同一份输入经过三种不同线性投影后的产物,分工明确:

  • Q(Query,查询):代表"我当前这个位置想去查什么样的信息"。在 Self-Attention 里,q_i 就是第 i 个 token 提出的查询向量。
  • K(Key,键):代表"我能被什么样的查询匹配到"。k_j 是第 j 个 token 的标签,用来跟所有 q_i 做相似度。
  • V(Value,值):代表"如果你关注到我,我会提供什么样的内容"。v_j 是第 j 个 token 真正要被聚合进输出的内容。

注意:Q / K 必须同维度。

为了能计算 Q K^TQK 的列数(即向量维度)必须相等,记作 d_kV 的列数 d_v 可以不同——实践里通常取 d_k = d_v。Self-Attention 中 Q / K / V 都来自同一个 X,自然同维度;Cross-Attention 中 Q 来自一种输入、K/V 来自另一种输入,只要 QK 同维度即可。

How — 用 PyTorch 实现 Q / K / V 的线性投影

下面这段代码演示了从输入矩阵 X 投影得到 Q/K/V 的全过程。PyTorch 没有专门叫 W^Q 的类,常规做法就是用三个 torch.nn.Linear,把它们设成不带 bias、矩阵乘形式的线性层。

import torch                                  # 导入 PyTorch 主包
import torch.nn as nn                         # 导入 nn 模块,里面有 Linear 层

torch.manual_seed(0)                          # 固定随机种子,让参数矩阵可复现
n, d_model, d_k = 4, 8, 6                     # 序列长度 4、模型维度 8、Q/K/V 维度 6

# 构造输入矩阵 X:4 行 token,每行 d_model=8 维
x = torch.randn(n, d_model)                   # 4×8 随机矩阵,模拟 embedding 输出

# 定义三个 Linear 层,分别作为 W^Q / W^K / W^V
w_q = nn.Linear(d_model, d_k, bias=False)     # 形状 (d_k, d_model),不带 bias
w_k = nn.Linear(d_model, d_k, bias=False)     # 形状 (d_k, d_model),不带 bias
w_v = nn.Linear(d_model, d_k, bias=False)     # 形状 (d_k, d_model),不带 bias

# 用三个 Linear 把 X 投影到 Q / K / V
Q = w_q(x)                                    # (n, d_k) = (4, 6)
K = w_k(x)                                    # (n, d_k) = (4, 6)
V = w_v(x)                                    # (n, d_k) = (4, 6)

print("Q.shape =", Q.shape)                   # 打印 Q 的形状,确认列数 = d_k
print("K.shape =", K.shape)                   # 打印 K 的形状
print("V.shape =", V.shape)                   # 打印 V 的形状

代码里 nn.Linear(d_model, d_k, bias=False) 就是一个标准的 y = x · W^T 形式的线性层,其中 W(d_k, d_model) 的可学习参数。它实现的就是 X · W^Q 这一步,只不过 PyTorch 把矩阵放在转置位置上——数学上等价。

本章小结

  • 来源相同QKV 都是同一个输入 X 经过三种不同线性投影得到的。
  • 三个参数矩阵W^QW^KW^V 是网络参数,训练时学得。
  • 三种角色Q 查,K 被查,V 提供内容。
  • 维度约束QK 必须同维度 d_k,才能算 Q K^T

三、Scaled Dot-Product Attention 四步走

有了 QKV,Transformer 论文 Attention Is All You Need 给出 Attention 的最终公式:

Scaled Dot-Product Attention(Transformer 原论文公式) Attention(Q, K, V) = Softmax( Q · K^T / sqrt(d_k) ) · V

把这个公式按运算顺序拆开,就是四步。

Step 1:Q · K^T 得到相似度矩阵

QK 的转置做矩阵乘,结果是一个 n × n 的相似度矩阵,元素 (i, j) 表示"q_ik_j 的相似度",也就是"在第 i 个位置看来,第 j 个位置有多相关"。

图 3-1  Q K^T 计算相似度矩阵
       Q (n × d_k)              K^T (d_k × n)             Q K^T (n × n)
   ┌───────────────┐         ┌────────────────┐        ┌──────────────────┐
   │ q_1           │         │ k_1^T  k_2^T  …│        │ s_11 s_12 ... s_1n│
   │ q_2           │    ×    │                │   =    │ s_21 s_22 ... s_2n│
   │ ...           │         │                │        │ ...               │
   │ q_n           │         │                │        │ s_n1 s_n2 ... s_nn│
   └───────────────┘         └────────────────┘        └──────────────────┘
                                                            ↑ 每个 s_ij = q_i · k_j
   ↑ 行为查询(query),列数 = d_k
   ↑ 行为键(key),做转置后变成列向量
How — Step 1 在 PyTorch 里的最小代码
# 沿用上一节已经算好的 Q、K
similarity = torch.matmul(Q, K.transpose(-1, -2))   # (n, d_k) @ (d_k, n) → (n, n)
print("similarity.shape =", similarity.shape)       # 应该是 (4, 4)

注意 K.transpose(-1, -2) 转置的是最后两个维度。在 Self-Attention 中 K 是 2D((n, d_k)),等价于 K.T;但写成 -1, -2 在更高维度张量里也能通用。

Step 2:除以 sqrt(d_k) 的"Scale"

d_kQ / K 的维度大小。当 d_k 比较大(比如 64 或 512),Q K^T 元素的方差会随 d_k 线性增长,数值容易跑到极端值,导致 Softmax 出来的分布近似 one-hot,梯度反向传播到 W^Q / W^K 时变得很小,训练不稳定。

解决方案很简单:除以 sqrt(d_k)。直觉上,可以认为点乘结果的方差大约是 d_k,除以 sqrt(d_k) 后方差就回到 1 附近,Softmax 的输入回到一个温和的量级。这就是 "Scaled"(缩放)的来源。

Scale 的作用 scores = (Q · K^T) / sqrt(d_k) d_k 越大,Q·K^T 的数值方差越大,除以 sqrt(d_k) 让方差回到 ~1, Softmax 不会饱和,梯度不会消失。
How — Step 2 的最小代码
d_k = K.shape[-1]                                 # 取 K 的最后一维作为 d_k
scaled = similarity / (d_k ** 0.5)                 # 等价于 similarity / sqrt(d_k)
print("scaled.shape =", scaled.shape)              # 形状不变,仍是 (n, n)

使用 d_k ** 0.5math.sqrt(d_k) 更通用:d_k 是 0-D 张量时也能正常工作。

Step 3:Softmax 把相似度变成权重

scaled 每一行做 Softmax,得到行和为 1 的权重矩阵。第 i 行的第 j 列元素 a_ij,就是"在第 i 个位置上分配给第 j 个位置的注意力比例"。

图 3-2  Scale + Softmax:从相似度到注意力权重
        scaled 矩阵                              attention weights
   ┌────────────────────────┐                ┌─────────────────────────┐
   │  2.4   0.3  -0.5  -1.1 │                │  0.71  0.11  0.05  0.13 │  ← 行和 = 1
   │  0.8   1.6   0.2   0.5 │  ──Softmax──▶  │  0.20  0.42  0.10  0.28 │  ← 行和 = 1
   │ -0.2   0.4   1.9   0.7 │   按行         │  0.07  0.13  0.61  0.19 │  ← 行和 = 1
   │  0.1  -0.4   0.3   2.2 │                │  0.08  0.05  0.10  0.77 │  ← 行和 = 1
   └────────────────────────┘                └─────────────────────────┘
       数值大小任意,                            范围 (0, 1),每行加起来 = 1
       正负皆有
How — Step 3 的最小代码
import torch.nn.functional as F                     # 导入函数式 API

attn = F.softmax(scaled, dim=-1)                   # 按最后一行归一化
print("attn row sums =", attn.sum(dim=-1))         # 验证每一行加起来等于 1

dim=-1 在 2D 张量上等价于 dim=1,但在更高维张量(带 batch、head 维)上写法更通用:永远对最后一个维度(也就是"被查询的 token 序列"这一维)做 Softmax。

Step 4:用权重乘 V 完成加权求和

把上一步得到的权重矩阵 attnV 做矩阵乘,就完成了"按注意力比例聚合所有 token 的内容"。结果矩阵的形状与 V 一致,行数还是 n,每行就是该位置对应的、融合了所有上下文信息的新表示。

输出矩阵第 i 行 output_i = Σ_j attn[i, j] · V[j, :] 即"用第 i 行的权重,对 V 的所有行做加权求和"。
图 3-3  attn @ V:把"权重"作用到"内容"上
       attention weights                  V                     output
   ┌────────────────────┐            ┌──────────────┐      ┌────────────────────┐
   │ 0.71 0.11 0.05 0.13│            │ v_1 v_2 v_3  │      │ Σ a_1j·v_j          │
   │ 0.20 0.42 0.10 0.28│     @      │ ...          │  =   │ Σ a_2j·v_j          │
   │ 0.07 0.13 0.61 0.19│            │ ...          │      │ Σ a_3j·v_j          │
   │ 0.08 0.05 0.10 0.77│            │ v_n ...      │      │ Σ a_4j·v_j          │
   └────────────────────┘            └──────────────┘      └────────────────────┘
   (n × n)                            (n × d_v)              (n × d_v)
How — Step 4 与完整四步的最小代码
output = torch.matmul(attn, V)                      # (n, n) @ (n, d_v) → (n, d_v)
print("output.shape =", output.shape)              # 与 V 同形状 (4, 6)

# === 完整四步连起来:一次 Scaled Dot-Product Attention ===
def scaled_dot_product_attention(Q, K, V):          # 定义为可复用函数
    d_k = K.shape[-1]                              # 取 K 的最后一维作为 d_k
    scores = torch.matmul(Q, K.transpose(-1, -2))  # Step 1:Q K^T 得到相似度
    scaled = scores / (d_k ** 0.5)                  # Step 2:除以 sqrt(d_k)
    attn = F.softmax(scaled, dim=-1)                # Step 3:按行 Softmax 得到权重
    out = torch.matmul(attn, V)                     # Step 4:权重乘 V 加权求和
    return out, attn                                # 同时返回权重,方便后续可视化

final, weights = scaled_dot_product_attention(Q, K, V)
print("final.shape =", final.shape)                # 形状与 V 一致
print("weights[0] =", weights[0])                  # 第 0 行:第 0 个位置对所有位置的关注度

把四步封装成函数后,调用一次就能拿到聚合后的输出和中间权重。注意函数返回 attn 是为了方便后续画"注意力热力图"或调试,不是必须返回的。

Why — 为什么一定要 Softmax?

Softmax 做三件事:

  • 归一化:把任意范围的分数压到 (0, 1),让"权重"含义成立。
  • 行和 = 1:保证加权求和是一个"加权平均",不会因为序列长度变化而改变数值量级。
  • 可微分:相比 hardmax(只保留最大),Softmax 平滑可导,梯度能顺利回传到 W^QW^K

没有 Softmax 会发生什么?

  • 没有 Softmax,"加权求和"可能变成"任意线性组合",数值范围不可控,后续层极难训练。
  • 没有 Softmax,无法解释为"概率/权重",可视化也无法直接看出"关注谁"。
  • 没有 Softmax,梯度无法稳定反向传播到 Q / K 的来源参数矩阵。

本章小结

  • 四步固定Q K^T/ sqrt(d_k)softmax乘 V
  • Scale 的存在是为了稳定训练,不是可有可无的工程细节。
  • Softmax 给结果赋予了"概率/权重"语义,也让反向传播平稳。
  • 最终输出形状与 V 相同:仍然是 (n, d_v),每行是一个融合上下文后的 token 表示。

四、Multi-Head Attention:同一份 X,多组 W

单头 Self-Attention 把所有相关性塞进一个 d_k 维空间里,模型只能学到一种"关系模式"。Transformer 的做法是把 d_model 切成 h 份,每份独立做一次 Scaled Dot-Product Attention,最后拼回去。这就是 Multi-Head Attention。

4.1 为什么需要多头

Why — 单头不够用,为什么?

问题一:单头只能学一种"关注模式"。

单头把 QKV 都限制在同一个 d_k 维子空间里,意味着所有"谁该关注谁"的判断共用同一组参数。语法关系、指代关系、语义相似关系挤在一起,互相干扰。

问题二:单头的表达能力受限。

对一个 d_model 维的输入,单头 Attention 输出仍是 d_model 维,但内部只有一个相似度矩阵,等于用一张 n × n 的图同时建模所有关系。

没有多头会发生什么?

  • 模型对长距离依赖和局部依赖一视同仁地挤压到同一个注意力分布,训练困难。
  • 同一个 token 对其他 token 的关注被强行"打包"成一种平均行为,无法在不同语义层做精细区分。
  • 实验表明,去掉多头后 BLEU / 困惑度都会明显变差。

多头的解法:用 h 个独立的小 Attention 并行算,每头只负责 d_model / h 维子空间内的相似度,最后把 h 个结果拼起来再线性投影一次。

4.2 多头的计算流程

Multi-Head Attention(Transformer 论文公式) head_i = Attention( X · W_i^Q, X · W_i^K, X · W_i^V ) i = 0..h-1 MultiHead(Q, K, V) = Concat(head_0, ..., head_{h-1}) · W^O

参数矩阵说明:

  • W_i^Q ∈ R^{d_model × d_k}W_i^K ∈ R^{d_model × d_k}W_i^V ∈ R^{d_model × d_v}:第 i 个头专属的投影矩阵。
  • W^O ∈ R^{h·d_v × d_model}:把拼接后的输出再投影回 d_model 维。
  • 实践里常取 d_k = d_v = d_model / h,这样总参数量与单头版本相当。
图 4-1  Multi-Head Attention 的完整数据流
                            X (n × d_model)
                                │
        ┌──────────┬───────────┼───────────┬──────────┐
        ▼          ▼           ▼           ▼          ▼
       W^Q_0     W^Q_1       W^Q_2       ...        W^Q_{h-1}
        │          │           │           │          │
        ▼          ▼           ▼           ▼          ▼
       Q_0        Q_1         Q_2         ...        Q_{h-1}
        │          │           │           │          │
   ┌────┴────┬─────┴────┬──────┴──────┬────┴────┬────┴────┐
   ▼         ▼         ▼            ▼        ▼         ▼
  K_0       K_1       K_2          ...      K_{h-1}
   │         │         │            │        │
   ▼         ▼         ▼            ▼        ▼
  V_0       V_1       V_2          ...      V_{h-1}
   │         │         │            │        │
   ▼         ▼         ▼            ▼        ▼
 Scaled   Scaled    Scaled        ...    Scaled
 Dot-Pro  Dot-Pro   Dot-Pro              Dot-Pro
   │         │         │            │        │
   ▼         ▼         ▼            ▼        ▼
 head_0    head_1    head_2       ...    head_{h-1}     每个 head_i: (n × d_v)
   │         │         │            │        │
   └────┬────┴────┬────┴──────┬────┴────┬───┘
        ▼         ▼           ▼         ▼
            Concat (按最后一维拼接)
                     │
                     ▼
                (n × h·d_v)
                     │
                     ▼
                   W^O
                     │
                     ▼
              Output (n × d_model)
How — Multi-Head Attention 的最小 PyTorch 实现

下面的代码从零实现了一个多头注意力,思路是:先把 X 一次性投影成 (h × d_k) 的组合形式,再 reshape / permute 把"头"这个维度提到 batch 维,最后并行算 h 个 Attention。输出再 reshape(n, d_model),过一层 W^O

import torch                                       # 导入 PyTorch 主包
import torch.nn as nn                              # 导入 nn 模块
import torch.nn.functional as F                    # 导入函数式 API

class MultiHeadAttention(nn.Module):                # 定义一个多头注意力模块
    def __init__(self, d_model, num_heads):         # 初始化:模型维度 + 头数
        super().__init__()                          # 调用父类构造函数
        assert d_model % num_heads == 0             # 断言 d_model 必须能被头数整除
        self.h = num_heads                          # 头数
        self.d_k = d_model // num_heads             # 每个头的维度 = d_model / h
        self.w_q = nn.Linear(d_model, d_model, bias=False)  # 合并的 W^Q
        self.w_k = nn.Linear(d_model, d_model, bias=False)  # 合并的 W^K
        self.w_v = nn.Linear(d_model, d_model, bias=False)  # 合并的 W^V
        self.w_o = nn.Linear(d_model, d_model, bias=False)  # 输出投影 W^O

    def forward(self, x):                           # 前向计算:输入 (n, d_model)
        n = x.shape[0]                              # 取出序列长度 n
        Q = self.w_q(x)                             # (n, d_model)
        K = self.w_k(x)                             # (n, d_model)
        V = self.w_v(x)                             # (n, d_model)

        # 把最后一维拆成 (h, d_k),再挪到第 1 维
        def split_heads(t):                         # 辅助函数:拆多头
            return t.view(n, self.h, self.d_k).transpose(0, 1)  # → (h, n, d_k)

        Qh, Kh, Vh = split_heads(Q), split_heads(K), split_heads(V)  # 各自拆成多头
        # 在每个头上独立做 Scaled Dot-Product Attention
        scores = torch.matmul(Qh, Kh.transpose(-1, -2))   # (h, n, d_k) @ (h, d_k, n) → (h, n, n)
        scores = scores / (self.d_k ** 0.5)              # Scale
        attn = F.softmax(scores, dim=-1)                 # (h, n, n),行和 = 1
        out = torch.matmul(attn, Vh)                      # (h, n, n) @ (h, n, d_k) → (h, n, d_k)

        # 把头维度拼回去:先 (h, n, d_k) → (n, h, d_k) → (n, h·d_k) = (n, d_model)
        out = out.transpose(0, 1).contiguous().view(n, self.h * self.d_k)  # 合并多头
        return self.w_o(out)                             # 过 W^O,输出 (n, d_model)

mha = MultiHeadAttention(d_model=8, num_heads=2)  # d_model=8, h=2
x = torch.randn(4, 8)                            # 4 个 token,每个 8 维
y = mha(x)                                       # 前向计算
print("y.shape =", y.shape)                      # 应该是 (4, 8)

这段代码的核心是"先合并投影、再拆头":用一个 d_model × d_model 的大矩阵替代 hd_model × d_k 的小矩阵,数学上完全等价,计算上更高效(一次大矩阵乘 vs h 次小矩阵乘)。最后用 W^O 把拼接结果压回 d_model 维,让整个模块可以像一层普通神经网络那样插入 Transformer。

本章小结

  • 多头 = 多个独立的小 Attention 并行,每头只看 d_model / h 维子空间。
  • 每头参数独立,能学到不同的"关注模式"(语法、指代、长距等)。
  • 拼接 + W^O 投影:让 h 个子空间的结论重新组合成 d_model 维。
  • 总参数量基本不变h 个小矩阵 vs 1 个大矩阵,乘起来都是 d_model × d_model

五、形状速查:一次 Self-Attention 的张量维度

把上一节的四步和 Multi-Head 都放到一起,给一张"形状速查表"。其中 B 表示 batch size,L 表示序列长度(也叫 n),E 表示每个 token 的 embedding 维度(d_model),H 表示头数。

符号含义未拆头形状(Self-Attn)拆 h 头后形状(单头)最终输出形状
X 输入 token 表示 (B, L, E) (B, H, L, E/H)
Q 查询矩阵 (B, L, E) (B, H, L, E/H)
K 键矩阵 (B, L, E) (B, H, L, E/H)
V 值矩阵 (B, L, E) (B, H, L, E/H)
Q K^T 相似度矩阵 (B, L, L) (B, H, L, L)
attn 注意力权重 (B, L, L) (B, H, L, L)
output Attention 输出 (B, L, E)

小贴士:理解"维度"比记公式更重要。

实际工程里调 Attention 最常见的 bug 就是维度不匹配:

  • Q K^T 之前忘了转置 K 的最后两维。
  • 拆多头时 view / reshape 没考虑 contiguous,permute 后忘了 .contiguous()
  • Softmax 的 dim 取错——应该是最后两维里的"被查 token"那一维。

任何时候报形状错误,都回到这张表查一遍对应阶段的正确维度。

本章小结

  • Self-Attention 的核心公式Softmax(Q K^T / sqrt(d_k)) V
  • 四步固定流程:相似度 → Scale → Softmax → 加权求和。
  • Q / K / V 同源:都从 X 经三个可训练 Linear 投影得到。
  • Multi-Headh 个独立小 Attention 并行,Concat 后过 W^O
  • 输出形状:与输入 X 完全一致((B, L, E)),可作为 Transformer block 的子层继续堆叠。

六、FAQ 20 问

FAQ 分组说明

  • Q1~Q5 概念类:澄清 Self-Attention 的本质和"在做一件什么事"。
  • Q6~Q10 矩阵 / 公式类:拆解 Q K^T、Scale、Softmax 的数学细节。
  • Q11~Q15 多头与参数类:Multi-Head 的设计动机和实现细节。
  • Q16~Q20 工程与对比类:PyTorch 2.x 的实现、性能与变体对比。

Q1. Attention 到底在做什么?一句话能讲清楚吗?

一句话结论:把"序列里哪些位置彼此相关、相关度多少"算成一个加权聚合,每个位置的输出 = 所有位置内容的加权和,权重由内容之间的相似度决定。展开讲,Self-Attention 的输入是一组 token 的向量表示 X,输出是与 X 同形状的新表示;新表示里每一行,都融入了其他所有 token 的信息,融合比例由"谁更相关"决定。

Q2. 为什么要叫"Self"-Attention?和普通 Attention 有区别吗?

"Self"指的是 Q / K / V 三个矩阵都来自同一个输入 X。Encoder 的 Self-Attention 是这样;Decoder 的第二个 Attention(Cross-Attention)则是 Q 来自 Decoder、X 来自 Encoder 的输出,那就不叫 Self 而是 Cross。两者数学形式完全一样,差别只在"Q/K/V 同不同源"。

Q3. Attention 和全连接(MLP)到底有什么本质区别?

全连接的权重是固定的、对所有输入一视同仁;Attention 的权重是动态的、随输入计算。MLP 的权重矩阵 W 在训练后是常数,对任意 x 永远是 x · W;Attention 的"权重"是 Softmax(Q K^T / sqrt(d_k)),每次前向计算都要重新根据当前 batch 算一遍,因此能根据上下文动态调整"谁该关注谁"。

Q4. Attention 是"注意力"吗?和人类视觉/认知的注意力是一回事吗?

只是借用了"注意力"这个比喻,机制完全不同。人类注意力是认知资源分配;Attention 是一种加权聚合运算。两者都属于"在不同位置上分配不同的重要性",但没有直接的神经科学对应关系。它不模拟人脑的工作机制。

Q5. Attention 的输出和输入形状一样,那它"改造"了什么?

形状不变,但每个位置的"含义"变了:每一行向量从"只看自己"变成"融合了所有位置的上下文"。比如输入是词向量时,输出就变成了"这个词考虑到其他词之后的新表示",可以更好地被下游层使用。

Q6. Q · K^T 里"转置 K"到底在转什么?

转置的是最后两维,让 K 的每一行(每个 token 的 key 向量)变成 K^T 的列,才能和 Q 的行做点乘。数学上 Q K^T(L, d_k) @ (d_k, L) = (L, L)。如果写成 Q @ K,维度对不上、矩阵乘法根本做不出来。

Q7. 为什么 Q 和 K 必须是同一个维度 d_k?

因为点积需要两个向量长度相等。q_i · k_jd_k 维点 d_k 维,结果是标量;如果维度不同,点积没有定义。这是 Transformer 设计 d_k = d_model / h 的硬性原因之一。

Q8. 为什么要除以 sqrt(d_k),而不是直接除 d_k?

因为 d_k 维独立随机向量的点积方差约为 d_k,开根号后能归一化到 ~1。假设 q / k 各分量是零均值、方差为 1 的独立随机变量,那么 q · k 的方差 ≈ d_k,标准差 ≈ sqrt(d_k)。除以 sqrt(d_k) 后方差稳定在 1 附近,Softmax 输入温和、梯度稳定。如果直接除以 d_k,数值过小、Softmax 输出趋向均匀分布,反而丢失"尖锐关注"的能力。

Q9. 能不能不除以 sqrt(d_k)?

能算,但训练会不稳定。原论文 §3.2.1 的脚注从数学上严格证明了点积方差随 d_k 线性增长("assume the components of q and k are independent random variables with mean 0 and variance 1, then their dot product has mean 0 and variance d_k"),所以 d_k 越大点积数值越大,Softmax 越容易进入饱和区、梯度越小。原论文 §3.2.1 还引用 Britz et al. (2017) 的实验结论:当 d_k 较大时,不带缩放的点积注意力表现差于加性注意力(additive attention)。实践中若被迫去掉 Scale,可以用更小的学习率或更小的初始化 std 缓解,但最干净的做法仍是 / sqrt(d_k)

Q10. Softmax 之前是不是可以加 Mask?怎么加?

可以,在 Softmax 之前把要屏蔽的位置加一个很大的负数(如 -1e9)即可。Decoder 中的 Masked Self-Attention 就是这样:翻译第 i 个词时,把 i 之后的位置在 Q K^T 中加 -inf,Softmax 后这些位置的权重自然为 0。代码写法是 scores = scores.masked_fill(mask == 0, -1e9)

Q11. 多头注意力到底比单头强在哪?

多头把"关系建模"分到 h 个独立的子空间,每个头可以学到不同语义的关注模式(语法、指代、上下位等)。单头必须用一张 n × n 的相似度矩阵同时解释所有关系,互相挤占;多头相当于让 h 个子网络并行建模,最后 Concat 融合。

Q12. 头数 h 越多越好吗?

不是。h 太大时每头的维度 d_model/h 太小,表达力反而下降。原论文 Transformer-base 用 h=8、d_model=512、d_k=64;Transformer-big 用 h=16、d_model=1024、d_k=64。经验上 d_k ≈ 64 是个甜点,再小就学不到东西,再大就开始过参数化。

Q13. Multi-Head 里的 W^O 是必须的吗?

W^O 不是数学必需,但对工程表现几乎是必需的。它的作用是把 Concat 出来的 h · d_v 维重新线性混合到 d_model 维,让不同头的结论做一次跨头融合。不带 W^O 模型也能跑,但收敛慢、效果差。

Q14. Self-Attention 的参数量是多少?

4 个 d_model × d_model 的矩阵:W^Q、W^K、W^V、W^O。d_model = 512,参数量约为 4 × 512 × 512 = 1.05M(不计 bias)。去掉多头拆分、合并成一个大矩阵也是 4 个,参数量与单头等价。

Q15. W^Q、W^K、W^V 是不是绑定的?

Multi-Query Attention / Grouped-Query Attention 才会共享 K/V 的投影矩阵,标准 Transformer 里三个是独立的。前者用更少的 KV 头(如 GQA 把 K/V 头数减半)来省显存,常见于大模型推理;标准 Self-Attention 里 W^Q ≠ W^K ≠ W^V,三者各自学习。

Q16. PyTorch 里有官方 Attention 实现吗?

有,torch.nn.MultiheadAttention 就是 PyTorch 2.x 提供的官方实现。它接受 querykeyvalue 三个张量,num_heads 指定头数,batch_first=True 时张量布局为 (B, L, E)。本节末尾给了一段最简调用示例。

Q17. torch.nn.MultiheadAttention 和 nn.functional.scaled_dot_product_attention 有什么区别?

前者是面向层的封装(带参数 W^Q/W^K/W^V/W^O),后者是无参数纯函数。F.scaled_dot_product_attention(Q, K, V) 只做"四步",需要自己提供 Q/K/Vnn.MultiheadAttention 自己管理参数,从输入 x 一路投影到多头输出,且内部已经默认走 F.scaled_dot_product_attention 的 fast path(PyTorch 2.0+)。日常写模型优先用 F.scaled_dot_product_attention,写可复用模块用 nn.MultiheadAttention

Q18. self-attention 的时间复杂度是多少?

O(L² · d),核心是 Q K^T 这个 (L, L) 的矩阵乘。L 是序列长度、d 是头维度。当 L 变大(如 8k、32k)时平方项占主导,这正是 Longformer、Linformer 等"线性 Attention"变体想解决的问题。

Q19. Self-Attention 能处理任意长度的序列吗?

数学上可以,物理上受显存和算力约束。Self-Attention 本身不假设 L 是定值——只要能构造 (L, d_model) 的矩阵,就能跑。但相似度矩阵是 L × L,L=4096 时已经是 16M 个浮点数,L=32768 时到 1G 个,工程上必须借助 Flash Attention、稀疏 Attention 等技术才能跑。

Q20. Self-Attention 能完全替代 RNN 吗?

在大多数主流 NLP 任务上已经替代了 RNN,但不能完全替代所有场景。Self-Attention 没有递归结构,并行度高、长距离依赖建模好;但它天然不携带"位置信息",必须显式加 Positional Encoding;而 RNN 天然有顺序偏置。在流式/超长上下文场景,State Space Model(如 Mamba)反而是 Self-Attention 的替代者。

FAQ 总纲

  • Attention = 加权聚合,权重由输入动态计算(Q1~Q5)。
  • 四步公式:Q K^T → Scale → Softmax → × V(Q6~Q10)。
  • 多头 = 多视角并行,参数与单头相当但表达力强很多(Q11~Q15)。
  • 工程实现F.scaled_dot_product_attentionnn.MultiheadAttention 各有定位,复杂度 O(L²d)(Q16~Q20)。

七、下一篇预告

Roadmap:ML 系列下一站

讲完了 Attention,下一篇会基于同样的"图 + 代码 + 形状表"风格,拆解 Self-Attention 的两个"补丁":

  • Positional Encoding:Self-Attention 本身位置无关,必须显式注入位置信息——正弦位置编码、RoPE、ALiBi 各是什么思路。
  • Masked Self-Attention:Decoder 训练时如何用三角 Mask 实现"看不到未来"。

同一份 Q K^T、同一份四步公式,加上 Mask、改一下输入维度,就能从 Encoder 跑到 Decoder,敬请期待。

posted @ 2026-08-01 15:59  左扬  阅读(19)  评论(0)    收藏  举报