04-1|Attention 机制
Attention 是 Transformer、Conformer、Whisper 等现代 ASR 模型的核心计算机制之一。
对于已经有 ASR 模型训练经验的人,真正需要掌握的并不是“知道MultiheadAttention怎么调用”,而是能够从一个 Tensor 出发,把下面这条链路完整讲清楚:
输入序列 X ↓ Q / K / V 投影 ↓ Q 与 K 两两计算相似度 ↓ Scaled Dot-Product ↓ Softmax 得到 Attention Weight ↓ 按照 Weight 对 V 加权求和 ↓ 得到 Context / Attention Output ↓ Multi-Head 时重复多个 Head ↓ Concat ↓ Output Projection面试中经常沿着这条链路连续追问:
为什么需要 Q/K/V?
为什么要除以dk\sqrt{d_k}dk?
为什么要 Softmax?
为什么一个 Head 不够?
为什么 Attention 是O(T2)O(T^2)O(T2)?
为什么它适合 ASR,又为什么长语音会成为瓶颈?
Streaming ASR 又是怎么限制这个T2T^2T2的?
所以 Attention 的核心不是背公式,而是理解:
Attention 本质上是一种“根据当前 Query 动态决定应该从哪些位置读取多少信息”的可微分内容寻址机制。
1. 核心概念
1.1 Attention 到底在解决什么问题
假设输入是一段序列:
x1 x2 x3 x4 ... xT对于当前位置iii,模型在计算表示yiy_iyi时,不再只能依赖附近位置,也不需要像 RNN 那样把前面所有信息压缩到一个隐藏状态中,而是可以直接查看整个序列:
┌── x1 ├── x2 qi ──────────────┼── x3 ├── x4 └── ...但真正的问题是:
当前位置iii到底应该关注哪些位置?每个位置应该关注多少?
Attention 给出了一个数据依赖的答案。
对于第iii个 Query,先计算它和所有 Key 的相关程度:
qi 与 k1 → score(i,1) qi 与 k2 → score(i,2) qi 与 k3 → score(i,3) ... qi 与 kT → score(i,T)再把这些 score 转成归一化权重:
[0.05, 0.10, 0.70, 0.10, 0.05]最后用这些权重去读取 Value:
yi = 0.05 v1 + 0.10 v2 + 0.70 v3 + 0.10 v4 + 0.05 v5因此一个非常重要的理解是:
Attention 不是简单地“把所有输入加权平均”,而是先通过 Query-Key 计算“应该看谁”,再通过 Value 决定“从那里取什么信息”。
1.2 Query、Key、Value 分别是什么
这三个名字可以用一个直观类比理解:
Query:我现在想找什么? Key: 每个位置分别有什么特征,可以用来匹配? Value:真正被取出来、参与聚合的内容是什么?例如可以把 Attention 想象成一个数据库查询:
Query ↓ 去匹配所有 Key ↓ 得到匹配分数 ↓ 决定应该读取哪些 Value但需要注意:
Query / Key / Value 本身不是人为指定的“问题 / 索引 / 答案”。
在神经网络中,它们通常都是通过可学习的 Linear Projection 从输入特征中映射出来的。
所以更准确的说法是:
- Query:用于发起匹配的表示;
- Key:用于被匹配、决定相关性的表示;
- Value:被加权聚合的内容表示。
1.3 Attention Score、Attention Weight、Context Vector
这三个概念非常容易混淆。
Attention Score
表示 Query 和某个 Key 的匹配程度。
例如:
scores = [2.1, 0.3, 4.7, 1.2]此时它只是一个原始相关性分数。
Attention Weight
对 Score 做 Softmax 后得到:
weights = [0.06, 0.01, 0.87, 0.06]它具有概率分布形式:
∑jAij=1 \sum_j A_{ij}=1j∑Aij=1
这里AijA_{ij}Aij表示第iii个 Query 对第jjj个 Value 的注意力权重。
Context Vector
最终使用 Attention Weight 对 Value 进行加权求和:
yi=∑jAijvj y_i=\sum_j A_{ij}v_jyi=j∑Aijvj
这个yiy_iyi就是当前位置最终从整个序列“读取”到的信息,也通常称为 Attention Output 或 Context Vector。
所以完整链路是:
Q/K ↓ Score ↓ Softmax ↓ Attention Weight ↓ 加权 V ↓ Context2. 为什么需要 Attention
2.1 RNN 的序列信息瓶颈
传统 RNN 类模型通常按照:
x1 → h1 → h2 → h3 → ... → hT传播信息。
如果x1x_1x1对xTx_TxT有影响,中间需要经过很多次递归计算。
这会带来两个问题:
第一,远距离依赖传播路径长。
第二,训练难以完全并行化。
Attention 把这条路径缩短成:
x1 ───────────────────────→ xT 的表示 x2 ───────────────────────→ x3 ───────────────────────→ ...理论上任意两个位置之间都可以直接建立联系。
因此:
Attention 的一个核心价值,就是把“长距离信息传播”从多步递归变成一次直接的全局交互。
2.2 为什么 Attention 非常适合现代 ASR
ASR 输入天然是长序列。
例如一段语音经过声学特征提取后:
80-dim fbank ↓ T × 80如果是 16 kHz、10 ms frame shift:
10 秒语音 ≈ 1000 个 frame经过 4 倍 subsampling 后:
≈ 250 个 encoder time steps此时一个时间位置可以直接和其他所有时间位置建立关系。
这对于:
- 长距离音素关系;
- 上下文消歧;
- 跨词依赖;
- 语音中的远距离声学关联;
都很有价值。
但代价也非常明显:
序列越长,Attention 的计算量和尤其是 Attention Matrix 的内存开销会快速增加。
这也是后面 Conformer、Streaming ASR、Local Attention 等设计需要重点解决的问题。
3. 原理与底层机制
3.1 从输入 X 开始
假设输入:
X∈RB×T×D X\in\mathbb{R}^{B\times T\times D}X∈RB×T×D
其中:
- BBB:batch size
- TTT:序列长度
- DDD:模型特征维度,也可以理解为dmodeld_{\text{model}}dmodel
例如:
B = 4 T = 250 D = 512那么:
X.shape = [4, 250, 512]假设使用最简单的单头 Attention。
3.2 Q / K / V Projection
输入XXX不会直接拿来计算 Attention,而是分别经过三个可学习的线性投影:
Q=XWQ+bQ Q=XW_Q+b_QQ=XWQ+bQ
K=XWK+bK K=XW_K+b_KK=XWK+bK
V=XWV+bV V=XW_V+b_VV=XWV+bV
如果单头情况下:
WQ,WK∈RD×dk W_Q,W_K\in\mathbb{R}^{D\times d_k}WQ,WK∈RD×dk
WV∈RD×dv W_V\in\mathbb{R}^{D\times d_v}WV∈RD×dv
则:
Q∈RB×T×dk Q\in\mathbb{R}^{B\times T\times d_k}Q∈RB×T×dk
K∈RB×T×dk K\in\mathbb{R}^{B\times T\times d_k}K∈RB×T×dk
V∈RB×T×dv V\in\mathbb{R}^{B\times T\times d_v}V∈RB×T×dv
为什么不能直接拿XXX同时充当 Q/K/V?
因为三者承担的功能不同。
如果所有角色都使用完全相同的表示,那么模型必须在同一个表示空间中同时完成:
“我要查询什么” “别人拿什么来匹配我” “最后真正需要读取什么”独立 Projection 可以让模型学习三个不同的表示空间:
X ├── WQ → Query space ├── WK → Key space └── WV → Value space这是 Q/K/V 的一个重要设计意义。
3.3 Linear Projection 底层到底发生了什么
从 PyTorch 的角度:
q=self.q_proj(x)本质上就是线性层计算。
如果把 Batch 和 Time 展平:
[B, T, D] ↓ [B*T, D] ↓ 矩阵乘法 [B*T, D] × [D, d_k] ↓ [B*T, d_k] ↓ reshape [B, T, d_k]GPU 上的核心工作实际上就是高度优化的 Matrix Multiplication。
因此 Attention 并不是某种神秘的特殊计算,它主要由几个非常规则的 Tensor 运算组成:
Linear / GEMM + MatMul + Scale + Softmax + MatMul这也是为什么现代 GPU 很适合执行 Attention。
3.4 Dot Product 为什么可以表示相关性
对于 Queryqiq_iqi和 Keykjk_jkj:
sij=qikj⊤ s_{ij}=q_i k_j^\topsij=qikj⊤
如果两个向量方向比较接近,点积通常更大。
例如:
q = [1, 1] k1 = [1, 1] k2 = [1,-1]那么:
q · k1 = 2 q · k2 = 0说明qqq和k1k_1k1的匹配程度更高。
因此,对于所有位置,可以构建一个 Score Matrix:
S=QK⊤ S=QK^\topS=QK⊤
如果:
Q∈RB×Tq×dk Q\in\mathbb{R}^{B\times T_q\times d_k}Q∈RB×Tq×dk
K∈RB×Tk×dk K\in\mathbb{R}^{B\times T_k\times d_k}K∈RB×Tk×dk
那么:
S∈RB×Tq×Tk S\in\mathbb{R}^{B\times T_q\times T_k}S∈RB×Tq×Tk
对于 Self-Attention:
Tq=Tk=T T_q=T_k=TTq=Tk=T
因此:
S∈RB×T×T S\in\mathbb{R}^{B\times T\times T}S∈RB×T×T
这一步是 AttentionO(T2)O(T^2)O(T2)的根源。
因为每个 Query 都要和所有 Key 进行匹配:
T 个 Query × T 个 Key = T² 个 pair3.5 为什么一定要除以 sqrt(d_k)
这是 Attention 面试中最常被追问的问题之一。
原始 Dot Product 是:
QK⊤ QK^\topQK⊤
但实际使用的是:
QK⊤dk \frac{QK^\top}{\sqrt{d_k}}dkQK⊤
这个dk\sqrt{d_k}dk不是经验上随便加的,它和随机变量方差有关。
假设qqq和kkk的每个维度独立、均值为 0、方差为 1。
点积:
q⊤k=∑l=1dkqlkl q^\top k=\sum_{l=1}^{d_k}q_lk_lq⊤k=