Multi-head Attention
设 Transformer 某一层的输入为
这里:
:batch size,一次处理多少条序列; :sequence length,每条序列多少个 token; :hidden size,每个 token 的隐藏向量维度。
Attention 首先对
如果暂时不拆 head,它们通常都是
然后进入 Multi-Head Attention。
设 head 数量为
原始:
因为
然后为了方便矩阵乘法,把 head 维提前:
于是:
例如:
可以理解为每个 batch 有
对于一个 head:
而
为了做矩阵乘法,对
所以:
得到:
的矩阵。
因此完整 batch + multi-head 情况下:
为了避免和 sequence length
假设一句话有 4 个 token:
1 | 我 喜欢 深度 学习 |
那么
所以每一个 Query 都要和所有 Key 做点积。得到:
其中:
表示第
因此:
- 行
:当前 Query token; - 列
:它正在观察的 Key token。
例如
这里出现了:
因此标准 Attention 的 score matrix 元素数量:
当 sequence length 增大时,增长速度是平方级。这是长上下文昂贵的原因之一。
不过 Attention 不是直接计算
原因和数值尺度有关。假设 Query 和 Key 每一个分量大致满足均值 0、方差 1。点积
的方差会随着
分布会非常尖锐,容易导致梯度不好训练。除以
这就是 Scaled Dot-Product Attention 中 “Scaled” 的来源。
自回归语言模型有一个非常重要的限制:第
例如:
1 | 我 喜欢 深度 学习 |
预测“喜欢”时,不能提前看到后面的“深度 学习”。因此需要 causal mask。
对于
然后计算:
第一行:
第二行:
因为后面马上做 Softmax:
而$e^{-}= 0,所以一个被 mask 的位置最终概率就是 0。
例如:
第三项概率:
因此 causal mask 最终实现:
现在:
对于每一个:
- batch
- head
- query token
都有一整行:
Softmax 对最后一个维度,也就是 Key/token 维度 进行:
所以:
每一个 Query token 都得到一组对所有 Key token 的概率权重。例如:
表示 token
得到 Attention probability:
而:
计算:
观察最后两个矩阵维度:
得到:
所以:
跟原来的
对于 token
例如 Attention 权重:
那么:
所以 Attention 的真正含义是 根据 Query-Key 相似度,对不同 Value 做加权求和。
因此:
:我在寻找什么? :我有什么特征可供匹配? :如果你关注我,你真正获得什么信息?
现在:
先 transpose:
然后把最后两个维度拼起来:
所以:
最终:
和原始输入
shape 完全一致。
Flash Attention
Flash Attention 的核心不是改变 Attention 数学公式,而是不把完整
中间矩阵写回 HBM。
普通实现逻辑上可能类似:
把巨大的
再写回。然后:
大量时间浪费在:
而不仅仅是 FLOPs。Flash Attention 不一次生成整个:
而是把
1 | Q 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 对。所以主要计算复杂度仍然近似:
多头总体近似:
讨论
评论