RoPE
RoPE(Rotary Position Embedding)的作用,是把位置信息加入到 Attention 中。
它和直接给 hidden states 加位置向量的方式不同。RoPE 的核心思路是:
根据 token 所在的位置,对 Q 和 K 做旋转,让 Q 和 K 的相似度中自然包含相对位置信息。
RoPE 主要作用于 Q 和 K,不作用于 V,因为 Attention Score 是通过 Q @ K^T 计算出来的。
为什么可以用复数表示旋转
RoPE 会把最后一个维度中的两个数分成一组:
1 | [x0, x1, x2, x3] |
然后通过:
1 | torch.view_as_complex(...) |
把:
1 | [a, b] |
看成复数:
1 | a + bi |
再乘上:
1 | cosθ + i sinθ |
就相当于在复平面中旋转一个角度。
因此 RoPE 本质上就是:
把 hidden dimension 两两配对,然后根据 token 的位置进行旋转。
reshape 和 view_as_complex
假设 Q 的 shape 是:
1 | [B, S, H, D] |
执行:
1 | x.reshape(*x.shape[:-1], -1, 2) |
会得到:
1 | [B, S, H, D/2, 2] |
这里最后的 2 表示两个实数组成一组,而 -1 会由 PyTorch 自动推导成 D/2。
接下来:
1 | torch.view_as_complex(...) |
会把最后的两个实数解释成一个复数,所以 shape 会变成:
1 | [B, S, H, D/2] |
这也是为什么 head_dim 必须是偶数,否则无法两两配对。
完成旋转以后,再通过:
1 | torch.view_as_real(...) |
和:
1 | flatten(-2) |
恢复成:
1 | [B, S, H, D] |
所以 RoPE 最终不会改变 Q、K 的 shape,只会改变其中的数值。
Broadcasting
RoPE 中预计算的位置频率通常类似:
1 | [S, D/2] |
而 Q、K 的复数形式可能是:
1 | [B, S, H, D/2] |
因此会把频率 reshape 成:
1 | [1, S, 1, D/2] |
这样就可以利用 broadcasting:
- batch 之间共用频率;
- head 之间共用频率;
- 不同 position 使用不同角度;
- 不同 hidden dimension 使用不同旋转频率。
容易踩的坑
第一个坑是 head_dim 必须是偶数。
因为 RoPE 会把最后一个维度两两组成复数。
第二个坑是把 reshape 和 view_as_complex 混在一起理解。
reshape(..., -1, 2) 只是把数据两两分组,而:
1 | view_as_complex |
才是真正把:
1 | [实部, 虚部] |
解释成一个复数。
第三个坑是忽略 dtype。
RoPE 中通常会先:
1 | x.float() |
使用 FP32 完成复数和旋转计算,最后再恢复到原来的 dtype,以减少低精度计算带来的数值问题。
一句话理解 RoPE
RoPE 可以理解成:
把 Q、K 的 hidden dimension 两两组成二维向量,再根据 token 位置旋转这些向量,从而把位置信息加入 Attention Score。
Attention
Attention 的核心作用,是让一个 token 根据当前的 Query,从其他 token 中选择需要的信息。
核心过程可以理解成:
1 | Q 和 K |
Attention 的核心公式是:
$$
\operatorname{Attention}(Q,K,V)
\operatorname{softmax}
\left(
\frac{QK^T}{\sqrt{D}}
\right)V
$$
Multi-Head Attention 的 Shape
假设:
1 | B = batch size |
原始输入通常是:
1 | [B, S, hidden_size] |
其中:
1 | hidden_size = H × D |
经过 Q、K、V 投影以后,会把最后一个 hidden dimension 拆成多个 head:
1 | [B, S, H, D] |
再通过:
1 | transpose(1, 2) |
变成:
1 | [B, H, S, D] |
这样每个 Attention Head 就可以单独进行矩阵乘法。
Attention Score 的 Shape
这里最容易犯的错误,是直接死记:
1 | scores = [B, H, S, S] |
更准确的写法应该是:
1 | Q: [B, H, Sq, D] |
K 转置最后两个维度:
1 | K^T: [B, H, D, Skv] |
所以:
1 | Q @ K^T |
得到:
1 | scores: [B, H, Sq, Skv] |
例如:
1 | Q: [2, 4, 3, 32] |
那么:
1 | scores: [2, 4, 3, 11] |
这里:
3表示当前有 3 个 Query;11表示每个 Query 可以关注 11 个 Key。
只有当:
1 | Sq == Skv |
时,才可以简单写成:
1 | [B, H, S, S] |
Softmax 为什么不改变 Shape
Attention Score 经过:
1 | probs = torch.softmax(scores, dim=-1) |
只是对最后一个维度进行归一化。
因此:
1 | scores: [B, H, Sq, Skv] |
shape 不会发生变化。
Softmax 做的只是让每一个 Query 对所有 Key 的注意力权重之和变成 1。
probs @ V 的 Shape
假设:
1 | probs: [B, H, Sq, Skv] |
矩阵乘法关注最后两个维度:
1 | [Sq, Skv] @ [Skv, D] |
中间的 Skv 会被消掉,所以结果是:
1 | [Sq, D] |
最终:
1 | output: [B, H, Sq, D] |
例如:
1 | probs: [2, 4, 3, 11] |
得到:
1 | output: [2, 4, 3, 32] |
可以理解成:
每一个 Query 使用 Attention 权重,对 11 个 Value 做加权求和,最终得到一个 32 维向量。
恢复 Attention 输出 Shape
Attention 计算结束以后:
1 | [B, H, S, D] |
但 Transformer 主干需要的是:
1 | [B, S, hidden_size] |
所以首先:
1 | output.transpose(1, 2) |
得到:
1 | [B, S, H, D] |
再把最后两个维度合并:
1 | [B, S, H × D] |
通常代码会写成:
1 | output.transpose(1, 2).contiguous().view(B, S, -1) |
这里 contiguous() 很重要,因为 transpose 以后 tensor 的内存布局通常不再连续,而 view 要求连续的内存布局。
GQA
MHA 中:
1 | num_q_heads == num_kv_heads |
每个 Query Head 都有对应的一组 K、V。
而 GQA(Grouped Query Attention)中:
1 | num_q_heads > num_kv_heads |
多个 Query Head 会共享同一个 KV Head。
例如:
1 | 4 个 Q Head |
那么:
1 | Q0、Q1 → KV0 |
真正进行 Attention 计算时,可以通过:
1 | repeat_kv(...) |
把 KV Head 临时扩展到和 Query Head 相同的数量。
GQA 的主要作用不是改变 Attention 的基本公式,而是:
减少 K、V 的 Head 数量,从而降低 KV Cache 的显存和带宽开销,提高大模型推理效率。
KV Cache
在自回归生成过程中,之前 token 的 K、V 不会发生变化。
例如已经处理了 10 个 token,现在生成第 11 个 token,如果重新计算前面所有 token 的 K、V,就会产生很多重复计算。
因此可以把之前计算好的:
1 | K |
保存起来,也就是 KV Cache。
新的 K、V 只需要沿 sequence dimension 拼接:
1 | 旧 KV: [B, H, 10, D] |
在 GQA 中需要注意:
应该缓存 repeat 之前的少量 KV Head,再在真正做 Attention 时执行
repeat_kv。
否则如果缓存已经 repeat 后的 KV,就失去了 GQA 节省 KV Cache 的意义。
总结
RoPE 和 Attention 是 Transformer 中紧密联系的两个部分。
RoPE 负责 位置信息。它通过把 Q、K 的 hidden dimension 两两组成复数,再根据 token position 进行旋转,使 Attention Score 能够感知相对位置。
Attention 负责 token 之间的信息选择。学习时最重要的是理解:
1 | Q @ K^T |
的 shape:
1 | [B, H, Sq, D] |
以及:
1 | probs @ V |
的 shape:
1 | [B, H, Sq, Skv] |
GQA 则是在 MHA 的基础上减少 KV Head 数量,主要用于降低 KV Cache 开销。
这一部分代码真正容易出问题的地方并不是 Attention 公式,而是 reshape、transpose、broadcasting、矩阵乘法以及 KV Cache 带来的 Shape 变化。理解清楚这些 Shape,后面阅读完整 Transformer Attention 实现会容易很多。