AI 工程基础体系 · 第 71/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。
Attention 数学细解:QKV、缩放点积、Mask、多头与复杂度
Attention(注意力)是一种根据查询内容,从一组候选信息中计算“应关注多少”的加权聚合机制。Transformer 将它作为核心计算单元:输入序列中的每个位置都可以根据与其他位置的相关性,动态读取其他位置的信息。
本文中的符号约定如下:
- :序列长度;
- :模型隐藏维度;
- :Query、Key 的维度;
- :Value 的维度;
- :Query、Key、Value 矩阵;
- :未归一化的注意力分数,也叫 logits;
- :经过 Mask 和 Softmax 后的注意力权重矩阵;
- :Attention 输出。
1. 从加权平均开始理解 Attention
假设有三个候选信息:
如果某个查询认为三个信息的重要性分别是:
那么输出就是加权平均:
计算得到:
Attention 的关键不在于“加权平均”本身,而在于权重 不是固定参数,而是由当前查询与候选信息之间的匹配关系动态计算出来的:
这也解释了为什么需要三组向量,而不是只使用一组向量:
- Query:当前请求“想找什么”;
- Key:每个候选项“具有什么索引特征”;
- Value:真正被读取和聚合的内容。
Key 用于匹配,Value 用于传递信息。Key 和 Value 可以来自同一个输入,也可以来自不同输入。
2. Q、K、V 的数学定义
设输入序列表示为矩阵:
其中第 行 是第 个位置的隐藏向量。
通过三个独立的线性投影得到:
如果是单头 Attention,通常有:
因此:
2.1 Self-Attention
Self-Attention 中, 都来自同一个输入 :
因此序列中的每个位置都可以读取同一序列中其他位置的信息。
例如句子:
小明把书放在桌子上,因为它很重。
“它”所在位置可以通过 Self-Attention 读取“小明”“书”“桌子”等位置的表示,模型据此判断指代关系。
2.2 Cross-Attention
Cross-Attention 中,Query 和 Key/Value 来自不同序列:
典型场景是编码器—解码器模型:
- 编码器输出源语言序列;
- 解码器当前状态生成 Query;
- 编码器输出生成 Key 和 Value;
- 解码器根据当前生成需求读取源序列。
如果目标序列长度为 ,源序列长度为 ,则:
最终注意力矩阵大小是:
它不必是方阵。
3. 点积如何产生匹配分数
对于第 个 Query 和第 个 Key,最基本的匹配分数是点积:
展开为:
如果 与 方向相近,点积较大;方向相反,点积较小;接近正交时,点积接近零。
把所有 Query 和 Key 的两两点积写成矩阵形式:
其中:
矩阵元素满足:
因此第 行表示第 个 Query 对所有 Key 的匹配分数。
3.1 点积不是概率
可以是任意实数,包括负数,也不满足行和为 1。因此它只是 logits,不是注意力权重。
需要对每一行应用 Softmax:
Softmax 后:
并且对每个未屏蔽的 Query 行:
于是每个输出位置都是 Value 的加权平均:
矩阵形式是:
4. 为什么需要缩放点积
Transformer 使用的核心公式是:
相比普通点积,多了一个缩放因子:
4.1 方差推导
假设 Query 和 Key 的每个分量独立,满足:
点积为:
每一项 的期望为 0,方差约为 1。独立相加后:
因此标准差约为:
当 增大时,未缩放点积的绝对值通常也会增大。
4.2 大 logits 会使 Softmax 饱和
考虑两个 logits:
其 Softmax 约为:
如果 logits 被放大为:
Softmax 约为:
这时几乎所有权重都集中到一个位置。输出虽然看似“明确”,但 Softmax 的梯度在极端区域会变得很小,导致训练信号变弱。
缩放后,点积的方差约为:
这不会保证 logits 永远处于理想范围,但能使不同维度规模下的初始化更稳定。
4.3 缩放不是归一化概率
只是对 logits 做尺度调整,之后仍然必须执行 Softmax。它不是除以行和,也不是 L2 归一化。
5. 一个完整的数值算例
设:
这里:
5.1 计算未缩放点积
5.2 缩放
因为:
所以:
5.3 对每一行应用 Softmax
第一行:
第二行:
所以:
5.4 加权读取 Value
第一行:
第二行:
最终:
注意力权重决定“读多少”,Value 决定“读到什么”。即使两个位置的 Key 相同,Value 不同,输出仍然可以不同。
6. Softmax 的数值稳定性
直接计算:
可能发生溢出。例如:
在浮点数中可能变成无穷大。
利用 Softmax 对整体平移不敏感的性质:
通常取:
于是:
所有指数的输入都不大于 0,因此最大指数为 1。
常见深度学习框架会在底层实现稳定版本,但自定义 Attention、量化内核或 NumPy 原型时不能忽略这一点。
7. Mask:在 Softmax 前限制可见范围
Mask 用于禁止某些 Query 读取某些 Key。最稳妥的数学形式是先对 logits 添加一个加性掩码:
其中:
然后:
禁止位置的指数项为:
因此它们的注意力权重为 0。
实际代码中不一定真的存储 ,也可能使用一个足够小的数,例如 。但具体数值需要结合数据类型:
- FP32 中 通常足够小;
- FP16、BF16 中应避免不必要的溢出或类型转换;
- 底层 fused kernel 可能使用专门的布尔 Mask 逻辑。
7.1 Padding Mask
不同样本长度不一致时,通常将 Batch 补齐到同一长度。例如:
样本 A:我 喜欢 猫
样本 B:我 喜欢 狗 很
补齐后:
样本 A:我 喜欢 猫 PAD
样本 B:我 喜欢 狗 很
Padding 位置不是有效输入,因此其他位置不应读取它对应的 Key/Value。Padding Mask 通常屏蔽 Key 维度上的 PAD 列。
如果某个 Query 本身是 PAD,是否还需要计算它的输出取决于后续实现:
- 后续损失是否忽略 PAD;
- 是否在残差连接或池化前清除 PAD 位置;
- 模型是否使用专门的 packed sequence。
只屏蔽 Key 不等于自动让 PAD Query 的输出无效。
7.2 Causal Mask
自回归生成要求位置 不能读取未来位置 。Causal Mask 定义为:
长度为 4 时,允许矩阵为:
第 1 个位置只能读取自己;第 2 个位置可以读取第 1、2 个位置;依此类推。
如果把方向写反,模型会读取未来 Token,训练时可能表现异常好,但生成时出现严重的信息泄漏。
7.3 将数值算例改为 Causal Attention
前面的未缩放分数为:
加入长度为 2 的 Causal Mask:
Softmax 后:
于是:
第一个位置不再读取第二个位置的信息。
7.4 为什么不能先 Softmax 再乘 Mask
错误做法是:
然后:
对于:
Softmax 是:
如果屏蔽第二列,直接乘 Mask 得到:
它的和是 0.269,不再是概率分布。正确做法是先屏蔽:
然后:
这表示“在可见候选集合中重新归一化”。
7.5 全部位置被屏蔽的边界
如果某个 Query 行全部是 ,那么:
会出现:
结果通常是 NaN。
因此实现中必须保证:
- 每个需要输出的 Query 至少有一个可见 Key;
- 或者对全屏蔽行做特殊处理;
- 不要仅通过“把所有无效位置都设为 ”就假设结果一定安全。
这类问题经常出现在 Padding Mask 与 Causal Mask 合并、变长批处理以及自定义 Attention 内核中。
8. Attention 的完整公式与张量形状
单头缩放点积 Attention 为:
逐步对应:
- :计算 Query 与 Key 的两两匹配;
- 除以 :控制 logits 的尺度;
- 加上 :删除不可见连接;
- Softmax:将每个 Query 的分数转成权重;
- 乘以 :聚合 Value。
设:
则:
矩阵乘法要求 Key 的数量和 Value 的行数相同,因为每个 Key 对应一个 Value。
9. Attention 的反向传播结构
训练时,Attention 不只是前向的加权平均,还需要对 传播梯度。
令:
如果上游梯度为:
那么对 Value 的梯度为:
这符合直觉:某个 Value 被关注得越多,输出对它的梯度通常越大。
对注意力权重的梯度为:
Softmax 的逐行 Jacobian 为:
其中 是 Kronecker delta:
用向量形式表示,若某一行 Softmax 输出为 ,上游梯度为 ,则:
这里 是逐元素乘法。
再考虑缩放点积:
令:
则:
Mask 位置不是可训练连接。实现中通常通过 或布尔逻辑使其梯度为零。
10. 多头 Attention
单头 Attention 只能在一个投影空间中建立匹配。Multi-Head Attention(多头注意力)让模型并行使用多个独立的投影空间。
设头数为 ,第 个头的投影为:
第 个头的输出为:
将所有头拼接:
最后经过输出投影:
完整公式为:
其中:
10.1 为什么需要多头
不同头可以学习不同类型的关系,例如:
- 局部相邻关系;
- 主谓关系;
- 指代关系;
- 句法结构;
- 长距离依赖;
- 特定语义特征。
这不是规范保证,而是模型训练后可能形成的行为。不能简单地把某个头固定解释为某种语法功能;头的语义通常需要通过可视化、消融实验或探针任务验证。
10.2 维度如何划分
Transformer 中常见设置是:
并且:
例如:
则:
每个头单独在 64 维空间中进行 Attention,最后拼回 512 维。
注意,除数应当是每个头的 Key 维度:
而不是:
这是实现多头 Attention 时常见的错误。
10.3 多头不是简单复制单头
如果所有头共享完全相同的 ,再复制结果,那么多头不会获得真正的表示多样性。多头的作用来自每个头拥有不同的投影参数和独立的注意力矩阵。
多头输出还需要 做跨头混合。没有输出投影时,各头只能在拼接后的固定分块中传递信息,表达能力和后续层的接口都会不同。
11. 从输入到多头输出的形状变化
以 Batch 输入为例:
经过线性层:
然后拆成头:
进行转置后,矩阵乘法得到:
Softmax 必须沿最后一个 Key 位置维度进行:
而不是沿头维度或 Query 维度。最终:
转回:
再合并头:
12. 可执行的 NumPy 实现
下面的实现包含:
- 缩放点积;
- 稳定 Softmax;
- 布尔 Mask;
- Causal Mask;
- 全屏蔽行检查。
import numpy as np
def stable_softmax(x, axis=-1):
"""数值稳定的 Softmax。"""
x_max = np.max(x, axis=axis, keepdims=True)
exp_x = np.exp(x - x_max)
return exp_x / np.sum(exp_x, axis=axis, keepdims=True)
def scaled_dot_product_attention(q, k, v, allow=None):
"""
q: [n_q, d_k]
k: [n_k, d_k]
v: [n_k, d_v]
allow: [n_q, n_k] 的 bool 矩阵。
True 表示允许关注,False 表示屏蔽。
"""
if q.ndim != 2 or k.ndim != 2 or v.ndim != 2:
raise ValueError("q、k、v 必须都是二维矩阵")
n_q, d_k = q.shape
n_k, d_k_k = k.shape
if d_k != d_k_k:
raise ValueError("q 和 k 的最后一维必须相同")
if v.shape[0] != n_k:
raise ValueError("k 的行数必须等于 v 的行数")
scores = q @ k.T / np.sqrt(d_k)
if allow is not None:
if allow.shape != (n_q, n_k):
raise ValueError("allow 的形状必须是 [n_q, n_k]")
# 每个 Query 至少需要一个可见 Key。
if np.any(np.sum(allow, axis=-1) == 0):
raise ValueError("存在完全被屏蔽的 Query 行")
scores = np.where(allow, scores, -np.inf)
weights = stable_softmax(scores, axis=-1)
output = weights @ v
return output, weights
q = np.array([
[1.0, 0.0],
[0.0, 1.0],
])
k = np.array([
[1.0, 0.0],
[0.0, 1.0],
])
v = np.array([
[10.0, 0.0],
[0.0, 20.0],
])
# 每个位置都可以看到全部位置
output_full, weights_full = scaled_dot_product_attention(q, k, v)
# Causal Mask:位置 i 只能看到位置 j <= i
causal = np.tril(np.ones((2, 2), dtype=bool))
output_causal, weights_causal = scaled_dot_product_attention(
q, k, v, allow=causal
)
np.set_printoptions(precision=4, suppress=True)
print("full weights:")
print(weights_full)
print("full output:")
print(output_full)
print("causal weights:")
print(weights_causal)
print("causal output:")
print(output_causal)
预期输出近似为:
full weights:
[[0.6698 0.3302]
[0.3302 0.6698]]
full output:
[[ 6.6976 6.6041]
[ 3.302 13.3959]]
causal weights:
[[1. 0. ]
[0.3302 0.6698]]
causal output:
[[10. 0. ]
[ 3.302 13.3959]]
这里的 allow 使用“允许矩阵”语义,和某些框架中使用的 1 表示保留、0 表示屏蔽的约定相似,但不同 API 对 Mask 的语义可能相反。调用框架接口时必须确认:
True是允许还是屏蔽;- 数值
1是保留还是屏蔽; - Mask 是加到 logits 上,还是传给底层布尔内核;
- Mask 的形状是否支持广播;
- 是否需要显式提供 Causal Mask。
不能仅凭参数名 attention_mask 推断语义。
13. 复杂度:时间、空间与真正的瓶颈
设:
- Batch 大小为 ;
- Query 长度为 ;
- Key 长度为 ;
- 模型维度为 ;
- 头数为 ;
- 每头维度为 。
13.1 投影复杂度
计算 的线性投影大致需要:
如果 都来自长度为 的序列,三组投影的常数大约是三次矩阵乘法;实际实现通常将它们合并为一个大的线性层以提高硬件利用率。
输出投影还需要:
13.2 分数矩阵复杂度
每个头需要计算:
其复杂度为:
由于:
可以写成:
Self-Attention 中 ,因此是:
这就是通常所说的 Attention 对序列长度呈二次复杂度。
13.3 加权聚合复杂度
计算:
复杂度同样为:
因此完整 Attention 的主要计算量可以概括为:
其中:
- 当 较小时,投影层的 项可能占主导;
- 当 很大时, 项通常成为瓶颈。
13.4 注意力矩阵的内存复杂度
分数矩阵和权重矩阵形状为:
因此显式保存它们的空间复杂度为:
训练时还可能需要保存:
- Softmax 前的 logits;
- Softmax 后的权重;
- Dropout 掩码;
- Q、K、V;
- 反向传播所需的中间结果。
所以实际显存压力不只来自一次矩阵乘法。
FlashAttention 等实现的核心思想之一,是通过分块和在线 Softmax 避免显式保存完整的 注意力矩阵,同时保持精确 Attention 的数学结果或在浮点误差范围内一致。它降低的是中间内存访问和显存占用,不代表理论上的全连接 Attention 已经从二次复杂度变成线性复杂度。
14. 自回归生成与 KV Cache
训练时,因果 Self-Attention 通常一次处理整个序列:
虽然未来位置被 Mask,但所有位置的计算仍然可以并行。
生成时,每一步只需要为新 Token 计算一个新的 Query。历史 Token 的 Key 和 Value 不再变化,因此可以缓存:
当前步:
只计算:
其复杂度从单步的全序列计算变为:
再加上当前 Token 的线性投影和输出层计算。
但是,如果生成 个 Token,缓存 Attention 的总读取量仍大致为:
因为第 步需要读取长度约为 的历史。KV Cache 主要避免了重复计算历史 Token 的 ,并使单步延迟可控;它没有消除长序列生成的全部成本。
KV Cache 的显存大致与以下因素成正比:
最后的 2 来自 Key 和 Value 两份缓存。实际大小还受数据类型、是否使用 MQA/GQA、张量并行布局等影响。
15. 多头 Attention 的复杂度是否增加
如果保持:
那么所有头的分数计算总复杂度仍然近似为:
不是简单地乘上一个额外的 ,因为每个头的维度同时缩小了。
但头数会影响:
- 中间张量的形状;
- Kernel 是否适合硬件;
- 每个头的并行粒度;
- 注意力权重的存储布局;
- KV Cache 的组织方式;
- 小头维度下的数值和表达能力。
因此理论 FLOPs 相近,不意味着不同头数的实际延迟、显存和吞吐完全相同。
16. Attention 的排列等变性与位置信息
如果只对输入 Token 做排列,而没有任何位置编码,Self-Attention 本身不会知道原始顺序。
设 是一个排列矩阵,输入变为:
则:
分数矩阵变为:
Softmax 按行应用后,注意力权重也相应排列,最终输出满足:
这叫排列等变性:输入顺序变化,输出按相同方式变化,但 Attention 没有独立能力判断“谁在前、谁在后”。
因此 Transformer 需要加入位置信息,例如:
- 绝对位置编码;
- 相对位置偏置;
- RoPE;
- ALiBi;
- 其他位置表示机制。
位置机制与 Causal Mask 解决的是不同问题:
- Causal Mask 限制“能否看见未来”;
- 位置机制告诉模型“当前 Token 位于什么位置、两个 Token 相距多远”。
有 Causal Mask 并不意味着模型自动获得了完整的位置信息。
17. 常见错误与失败表现
17.1 忘记除以
表现可能包括:
- Softmax 权重过早接近 one-hot;
- 注意力分布熵很低;
- 梯度变小;
- 训练初期不稳定;
- 增大隐藏维度后问题更明显。
但并非所有模型都显式使用同一种缩放方式;如果使用了余弦相似度、特殊归一化或专用内核,公式可能不同,不能机械叠加缩放。
17.2 在错误维度上做 Softmax
对于:
Softmax 应沿 Key 位置维度 进行,使每个 Query 对所有 Key 的权重之和为 1。
如果错误地沿 Query 维度做 Softmax,得到的含义变成“某个 Key 被不同 Query 关注的相对程度”,不再是标准 Attention。
诊断方法是检查:
weights.sum(axis=-1)
对于没有异常 Mask 的有效 Query 行,应接近 1。
17.3 Mask 方向写反
Causal Mask 的正确条件是:
若误写成:
模型就会看到未来信息。训练损失可能异常低,但真实逐 Token 生成时性能明显下降。
可以用一个长度为 3 的人工输入验证第一个 Query:
- 正确因果 Mask:第一行只有第一个 Key 可见;
- 若第一行能看到第二、第三个 Key,说明存在泄漏。
17.4 只屏蔽 Query,不屏蔽 Key
Padding Mask 主要阻止有效 Query 读取 PAD Key。若只把 PAD Query 的输出置零,却没有屏蔽 PAD Key,那么有效 Token 仍可能把 PAD 的 Value 聚合进来。
17.5 用 Mask 实现权限控制
Attention Mask 只能控制一次前向计算中的信息可见性,不能替代权限系统。
如果一个用户无权访问某条文档,不能仅依赖:
- Prompt 中的提示;
- 生成阶段的 Attention Mask;
- 模型“应该不会关注”的经验。
权限必须在数据检索、服务端授权和上下文组装阶段完成。模型内的 Mask 是计算约束,不是安全边界。
18. Dropout 在 Attention 中的位置
原始 Transformer 论文在注意力权重上使用 Dropout。若 是 Softmax 后的权重,训练时可以得到:
然后:
这与 Causal Mask 不同:
- Mask 是结构性禁止连接;
- Dropout 是训练时随机丢弃部分连接;
- Mask 在推理时仍然存在;
- Dropout 通常在推理时关闭。
Dropout 后的权重在实现中可能采用 inverted dropout 缩放,因此其行和不一定仍然严格等于 1。不能看到注意力权重经过 Dropout 后行和变化,就判断实现错误。
19. 长上下文的真实边界
理论复杂度为:
意味着序列长度翻倍时,两两交互相关的计算量和显式注意力矩阵大小约变为四倍。
这会影响:
- 训练显存;
- 推理延迟;
- KV Cache 体积;
- 带宽和显存访问;
- Batch size;
- 单请求成本;
- 并发能力。
工程上可以使用以下方向,但它们改变了计算或存储方式,不能自动视为与标准全连接 Attention 等价:
- FlashAttention:优化精确 Attention 的内存访问;
- 滑动窗口或局部 Attention:只连接邻近位置;
- 稀疏 Attention:只计算部分连接;
- MQA/GQA:减少 Key/Value 头数,降低 KV Cache;
- 线性 Attention:改变核函数或归一化形式,试图避免显式 ;
- 分块、检索或层次化上下文:减少一次性送入模型的 Token 数。
选择方案时应同时测量:
- 相同质量目标下的效果;
- 端到端延迟,而不是只看单个 Kernel;
- 峰值显存;
- 长度增长时的退化;
- 训练和推理是否使用不同实现;
- 异常长输入、空输入和全 Mask 输入的行为。
20. 评测时如何识别 Attention 泄漏
如果任务要求模型只能使用历史信息,评测不能只检查最终准确率。还应构造能暴露未来信息泄漏的测试。
一种简单方法是对未来 Token 做扰动:
- 固定当前位置及其历史;
- 修改未来位置的 Token;
- 重新计算当前位置输出;
- 比较两次输出。
对于严格的 Causal Attention,当前位置的输出不应依赖未来 Token。实际浮点计算可能存在极小误差,但不应出现有意义的系统变化。
训练数据也可能产生“伪泄漏”:
- 标签被拼入输入;
- 预处理使用了未来统计量;
- Padding 或特殊 Token 处理错误;
- 训练和评测的 Mask 方向不一致。
因此 Mask 正确只是必要条件,不是完整的数据泄漏检查。
21. 与生产系统的关系
Attention 的数学正确性最终会影响生产系统中的多个维度:
- 模型:缩放、Mask、头维度和位置机制决定模型实际计算;
- 数据:Padding、特殊 Token、序列截断会改变可见范围;
- 评测:训练时并行的 Causal Mask 必须与推理时逐 Token 生成一致;
- 权限:数据授权应在模型输入前完成,不能把 Attention Mask 当作访问控制;
- 成本:序列长度同时影响二次交互计算、显存和 KV Cache;
- 并发:长序列请求会占用更多显存,可能降低 Batch 合并和服务吞吐;
- 故障处理:全 Mask 行、错误广播、NaN logits 和 KV Cache 形状不匹配都应在运行时尽早检测。
一个可复现的 Attention 实现,至少应验证以下不变量:
这些条件分别对应形状、Mask、归一化、聚合和矩阵乘法的基本正确性。实现优化、混合精度或替换内核后,应重新检查这些数学性质,而不能只根据最终 loss 或吞吐判断正确。
Attention 的核心可以归结为一条链路:
Q 决定查询需求,K 决定匹配方式,V 决定被传递的内容;缩放控制 logits 的统计尺度,Mask 控制可见连接,多头提供多个投影空间,复杂度则决定了这种全连接交互在长序列和生产推理中的成本边界。
系列导航与关联阅读
- 系列入口:AI 工程完整学习路线:从机器学习与 Transformer 到 RAG、Agent 和生产治理
- 上一篇:ONNX 与模型互操作:导出、算子集、动态形状、验证和部署
- 下一篇:Transformer 位置编码:绝对位置、RoPE、ALiBi 与长度外推
官方资料
本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。

评论
0 条讨论