ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

深度学习 - 13 Attention 机制

深度学习 - 13 Attention 机制 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∑​Aij​vj​这个yiy_iyi​就是当前位置最终从整个序列“读取”到的信息,也通常称为 Attention Output 或 Context Vector。所以完整链路是:Q/K ↓ Score ↓ Softmax ↓ Attention Weight ↓ 加权 V ↓ Context2. 为什么需要 Attention2.1 RNN 的序列信息瓶颈传统 RNN 类模型通常按照:x1 → h1 → h2 → h3 → ... → hT传播信息。如果x1x_1x1​对xTx_TxT​有影响,中间需要经过很多次递归计算。这会带来两个问题:第一,远距离依赖传播路径长。第二,训练难以完全并行化。Attention 把这条路径缩短成:x1 ───────────────────────→ xT 的表示 x2 ───────────────────────→ x3 ───────────────────────→ ...理论上任意两个位置之间都可以直接建立联系。因此:Attention 的一个核心价值,就是把“长距离信息传播”从多步递归变成一次直接的全局交互。2.2 为什么 Attention 非常适合现代 ASRASR 输入天然是长序列。例如一段语音经过声学特征提取后: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 sizeTTT:序列长度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​=qi​kj⊤​如果两个向量方向比较接近,点积通常更大。例如: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}}dk​​QK⊤​这个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=
返回列表