注意力机制深入:从加性注意力到多头与 FlashAttention

注意力机制是当代大语言模型的算力核心。它让模型在处理一个 token 时,不再被动接受一段被压扁的固定向量,而是按需要从整段输入中动态地挑出与当前任务最相关的部分。理解注意力,等于理解了 Transformer 为何能扩展到数千亿参数、又为何在长上下文下举步维艰。

注意力要解决什么问题

在注意力出现之前,神经机器翻译普遍采用编码器-解码器结构。编码器把整句源语言压缩成一个固定长度的向量,解码器再从这个向量生成译文。Bahdanau 等人在 2014 年的论文中指出,用一个定长向量承载整句话的全部语义,是一个明显的瓶颈:句子越长,信息损失越严重,长程依赖也越难保留。

所谓长程依赖,是指句子中相距很远的成分之间存在的语法或语义关联。例如主谓一致、指代消解。循环神经网络理论上能传递这类信息,但梯度在长序列上容易衰减或爆炸,实际很难学到。注意力机制给出的解法是:放弃对全部信息做一次性压缩,改为在处理每一个输出位置时,重新计算输入各位置的加权组合。权重由当前查询与每个键的相关程度决定,因此模型可以「回看」任意远处的内容。

加性注意力与乘性注意力

最早的注意力来自 Bahdanau 2014,称为加性注意力(additive attention)。它的打分函数把一个查询向量与一个键向量拼接后送入一个小的前馈网络:

score(s, t) = v^T * tanh(W_a * [h_s; h_t])

其中 h_s 是源端隐藏状态,h_t 是目标端隐藏状态,v 与 W_a 是可学习参数。之所以叫加性,是因为它通过一个带 tanh 的加性变换来融合两个向量。加性注意力在维度较大时数值稳定,但计算无法被矩阵乘法高效加速,因为要对每一对查询-键单独过网络。

乘性注意力则直接用点积衡量相似度:score = q^T k。点积可以整体写成矩阵乘法 QK^T,对硬件非常友好。Vaswani 等人在 2017 年的 Transformer 中采用的正是乘性路线,并通过除以维度的平方根来解决点积方差过大的问题,这就是缩放点积注意力。

缩放点积注意力

缩放点积注意力的完整公式为:

Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V

其中 Q、K、V 分别是查询、键、值矩阵,d_k 是键向量的维度。为什么要除以 sqrt(d_k)?假设 Q 与 K 的各分量独立且均值为 0、方差为 1,那么点积 q^T k 的均值为 0、方差为 d_k。维度越大,点积的绝对值越大,softmax 的输入进入饱和区后梯度会趋近于 0,训练难以推进。除以 sqrt(d_k) 把方差拉回 1,使 softmax 处于梯度较灵敏的区间。

下面给出一段可直接运行的 NumPy 风格实现,完整呈现 softmax、缩放与点积三步:

import numpy as np

def softmax(x, axis=-1):
    # 减去最大值以保证数值稳定,避免 exp 溢出
    x = x - np.max(x, axis=axis, keepdims=True)
    e = np.exp(x)
    return e / np.sum(e, axis=axis, keepdims=True)

def scaled_dot_product_attention(Q, K, V):
    d_k = Q.shape[-1]
    # score = QK^T / sqrt(d_k)
    scores = (Q @ K.transpose(-1, -2)) / np.sqrt(d_k)
    weights = softmax(scores, axis=-1)
    return weights @ V, weights

# 演示:序列长度 4,每个向量维度 8
np.random.seed(0)
Q = np.random.randn(4, 8)
K = np.random.randn(4, 8)
V = np.random.randn(4, 8)
out, w = scaled_dot_product_attention(Q, K, V)
print("输出形状:", out.shape)
print("注意力权重每行之和(应为 1):", w.sum(axis=-1))

若使用 PyTorch,只需把 np 换成 torch,并调用 torch.softmax(scores, dim=-1),逻辑完全一致。

多头注意力

单一的缩放点积注意力只能学到一种固定的相似度模式。多头注意力(multi-head attention)把 Q、K、V 分别用不同的可学习投影矩阵映射到 h 个子空间,在每个子空间独立计算注意力,再把 h 个结果拼接并线性变换:

MultiHead(Q, K, V) = Concat(head_1, ..., head_h) W^O
head_i = Attention(Q W_i^Q, K W_i^K, V W_i^V)

不同头可以专注于不同关系:有的头捕捉局部相邻 token 的关联,有的头负责跨句长距离指代,有的头关注语法结构。这种并行多子空间的设计,是 Transformer 表达能力强于单层注意力的关键。

自注意力与交叉注意力

注意力的查询来源不同,便区分为自注意力与交叉注意力。

自注意力中,Q、K、V 全部来自同一序列。编码器里的自注意力让每个 token 同时看见句中所有其他 token,从而建立上下文相关的表示。

交叉注意力出现在编码器-解码器之间,其键与值来自编码器输出,查询来自解码器当前状态。换言之,解码器用自身的查询去检索编码器提供的源端表示,实现源语言与目标语言的对齐。

因果掩码

在解码器做自回归生成时,第 t 个位置只能依赖第 1 到第 t 个位置,不能「偷看」未来。因果掩码(causal mask)通过把上三角的注意力分数设为负无穷,使 softmax 之后这些位置的权重为 0,从而强制信息单向流动。没有因果掩码,模型在训练时就会泄露未来信息,推理时也无法自洽。

KV 缓存加速自回归生成

逐 token 生成时,若每次都重新计算全部历史 token 的键与值,开销随序列长度线性累积。KV 缓存的作法是:把已经算过的 K 与 V 暂存下来,每来一个新 token,只计算该 token 对应的查询,并与缓存中的历史 K、V 做注意力。这样每个新 token 的计算量从 O(n) 降为常数级,显著加速生成。代价是显存需常驻整段历史,这也是长上下文推理显存吃紧的根源之一。

FlashAttention 与 IO 感知

标准注意力的计算会先把巨大的注意力分数矩阵 QK^T 完整写入 GPU 的高带宽显存(HBM),再从 HBM 读回做 softmax,最后再写回。序列长度 n 时,这个中间矩阵大小为 n 乘 n,显存占用与读写量都是平方级,成为长序列下的瓶颈。

Dao 等人在 2022 年提出 FlashAttention,核心动机是让算法感知 GPU 的内存层级。它用分块(tiling)把 Q、K、V 切成能放进片上 SRAM 的小块,在 SRAM 内就地完成 softmax 与加权求和,避免把完整 n 乘 n 分数矩阵落盘到 HBM。由于 SRAM 的带宽远高于 HBM,减少 HBM 读写既提速又省显存。FlashAttention 是精确注意力,不损失精度,却把显存占用从平方级近似降为常数级,使模型得以处理更长的上下文。

计算复杂度与长上下文挑战

自注意力的计算与显存复杂度相对序列长度均为 O(n^2)。这意味着序列从 2K 翻倍到 4K,开销约变为四倍。随着长文档、长对话需求增长,平方复杂度成为主要约束:训练更慢、推理显存更紧张、成本更高。这正是稀疏注意力、线性注意力、以及 FlashAttention 类 IO 感知算法持续演进的动力。

小结

注意力机制用动态的加权聚合取代了定长向量瓶颈,使模型能直接建模任意位置间的依赖。加性注意力开启了这一思路,缩放点积注意力以可并行点积成为主流,多头注意力通过多子空间提升表达力,因果掩码保证生成的自洽,KV 缓存加速自回归解码,而 FlashAttention 从内存 IO 视角把平方级开销压了下来。理解这些层累的设计,是读懂现代大模型工程与瓶颈的前提。

参考与延伸阅读

  • arXiv:1409.0473 — Neural Machine Translation by Jointly Learning to Align and Translate(Bahdanau, Cho, Bengio, 2014),加性注意力的提出。已核验
  • arXiv:1706.03762 — Attention Is All You Need(Vaswani et al., 2017),缩放点积注意力与多头注意力、Transformer 的提出。已核验
  • arXiv:2205.14135 — FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., 2022),IO 感知的精确注意力。已核验
本文累计阅读