Multi-head Attention

设 Transformer 某一层的输入为

这里:

  • :batch size,一次处理多少条序列;
  • :sequence length,每条序列多少个 token;
  • :hidden size,每个 token 的隐藏向量维度。

Attention 首先对 做三个线性变换:

如果暂时不拆 head,它们通常都是

然后进入 Multi-Head Attention。

设 head 数量为 ,每个 head 的维度为 。通常有

原始:

因为 ,先 reshape:

然后为了方便矩阵乘法,把 head 维提前:

于是:

例如:

可以理解为每个 batch 有 个独立的 Attention head,每个 head 都看到 S 个 token,每个 token 在这个 head 中是 维。

对于一个 head:

为了做矩阵乘法,对 最后两个维度转置:

所以:

得到:

的矩阵。

因此完整 batch + multi-head 情况下:

为了避免和 sequence length 混淆,可以把 score matrix 记为

假设一句话有 4 个 token:

1
我  喜欢  深度  学习

那么 。每个 token 都要询问:“我应该关注哪些 token?”

所以每一个 Query 都要和所有 Key 做点积。得到:

其中:

表示第 个 token 的 Query,认为第 个 token 的 Key 有多相关。

因此:

  • :当前 Query token;
  • :它正在观察的 Key token。

例如 就是:第 3 个 token 对第 4 个 token 的注意力匹配分数。

这里出现了:

因此标准 Attention 的 score matrix 元素数量:

当 sequence length 增大时,增长速度是平方级。这是长上下文昂贵的原因之一。

不过 Attention 不是直接计算 ,而是:

原因和数值尺度有关。假设 Query 和 Key 每一个分量大致满足均值 0、方差 1。点积

的方差会随着 增长。如果直接进入 Softmax,例如:

分布会非常尖锐,容易导致梯度不好训练。除以 可以把点积的典型尺度稳定下来:

这就是 Scaled Dot-Product Attention 中 “Scaled” 的来源。

自回归语言模型有一个非常重要的限制:第 个 token 不能偷看未来 token。

例如:

1
我 喜欢 深度 学习

预测“喜欢”时,不能提前看到后面的“深度 学习”。因此需要 causal mask。

对于 ,mask 可以理解成:

然后计算:

第一行:

第二行:

因为后面马上做 Softmax:

而$e^{-}= 0,所以一个被 mask 的位置最终概率就是 0。

例如:

第三项概率:

因此 causal mask 最终实现:


现在:

对于每一个:

  • batch
  • head
  • query token

都有一整行:

Softmax 对最后一个维度,也就是 Key/token 维度 进行:

所以:

每一个 Query token 都得到一组对所有 Key token 的概率权重。例如:

表示 token 把 10% 注意力放第 1 个 token,20% 放第 2 个,60% 放第 3 个,10% 放第 4 个。

得到 Attention probability:

而:

计算:

观察最后两个矩阵维度:

得到:

所以:

跟原来的 shape 一样。

对于 token

例如 Attention 权重:

那么:

所以 Attention 的真正含义是 根据 Query-Key 相似度,对不同 Value 做加权求和

因此:

  • :我在寻找什么?
  • :我有什么特征可供匹配?
  • :如果你关注我,你真正获得什么信息?

现在:

先 transpose:

然后把最后两个维度拼起来:

所以:

最终:

和原始输入

shape 完全一致。

Flash Attention

Flash Attention 的核心不是改变 Attention 数学公式,而是不把完整 中间矩阵写回 HBM。

普通实现逻辑上可能类似:

把巨大的 矩阵写入 GPU HBM。然后再:

再写回。然后:

大量时间浪费在:

而不仅仅是 FLOPs。Flash Attention 不一次生成整个:

而是把 切成小块。例如概念上:

1
2
3
4
5
6
7
Q block

K block → 算局部 QKᵀ

局部 Softmax

和 V block 累积

只在 GPU 更快但容量更小的片上存储中保留当前 block 所需数据,例如:

  • registers;
  • shared memory / SRAM。

算完这个块后立即把贡献累积到输出。因此不需要:

普通 Softmax:

通常需要知道整行:

而数值稳定版本还需要:

然后:

看起来必须一次拿到所有 个元素。

Flash Attention 使用一种 online softmax 思想分块处理时持续维护当前行的最大值和归一化因子。

假设已经处理了一部分,保存:

和:

新 block 出现新的最大值之后,可以重新缩放旧累积值,从而得到正确的新 normalization。因此无需保存整行 score,也能精确完成稳定 Softmax。事实上,Flash Attention 没有消除标准 Attention 的平方计算复杂度。Flash Attention 主要是 IO-aware exact attention,显著降低了 HBM IO 带宽从而降低显存压力。

标准 self-attention 仍然需要比较:

个 Query-Key 对。所以主要计算复杂度仍然近似:

多头总体近似: