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
2
3
4
5
6
[x0, x1, x2, x3]



[[x0, x1],
[x2, x3]]

然后通过:

1
torch.view_as_complex(...)

把:

1
[a, b]

看成复数:

1
a + bi

再乘上:

1
cosθ + i sinθ

就相当于在复平面中旋转一个角度。

因此 RoPE 本质上就是:

把 hidden dimension 两两配对,然后根据 token 的位置进行旋转。


reshapeview_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 会把最后一个维度两两组成复数。

第二个坑是把 reshapeview_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
2
3
4
5
6
7
8
9
Q 和 K

计算相关程度

softmax 得到注意力权重

使用权重对 V 加权求和

得到输出

Attention 的核心公式是:

$$
\operatorname{Attention}(Q,K,V)

\operatorname{softmax}
\left(
\frac{QK^T}{\sqrt{D}}
\right)V
$$


Multi-Head Attention 的 Shape

假设:

1
2
3
4
B = batch size
S = sequence length
H = num_heads
D = head_dim

原始输入通常是:

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
2
3
Q: [B, H, Sq, D]

K: [B, H, Skv, D]

K 转置最后两个维度:

1
K^T: [B, H, D, Skv]

所以:

1
Q @ K^T

得到:

1
scores: [B, H, Sq, Skv]

例如:

1
2
3
Q: [2, 4, 3, 32]

K: [2, 4, 11, 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
2
3
4
5
scores: [B, H, Sq, Skv]



probs: [B, H, Sq, Skv]

shape 不会发生变化。

Softmax 做的只是让每一个 Query 对所有 Key 的注意力权重之和变成 1。


probs @ V 的 Shape

假设:

1
2
3
probs: [B, H, Sq, Skv]

V: [B, H, Skv, D]

矩阵乘法关注最后两个维度:

1
[Sq, Skv] @ [Skv, D]

中间的 Skv 会被消掉,所以结果是:

1
[Sq, D]

最终:

1
output: [B, H, Sq, D]

例如:

1
2
3
probs: [2, 4, 3, 11]

V: [2, 4, 11, 32]

得到:

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
2
4 个 Q Head
2 个 KV Head

那么:

1
2
Q0、Q1 → KV0
Q2、Q3 → KV1

真正进行 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
2
K
V

保存起来,也就是 KV Cache。

新的 K、V 只需要沿 sequence dimension 拼接:

1
2
3
4
5
6
7
旧 KV: [B, H, 10, D]

新 KV: [B, H, 1, D]



[B, H, 11, 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
2
3
4
5
6
7
[B, H, Sq, D]
@
[B, H, D, Skv]



[B, H, Sq, Skv]

以及:

1
probs @ V

的 shape:

1
2
3
4
5
6
7
[B, H, Sq, Skv]
@
[B, H, Skv, D]



[B, H, Sq, D]

GQA 则是在 MHA 的基础上减少 KV Head 数量,主要用于降低 KV Cache 开销。

这一部分代码真正容易出问题的地方并不是 Attention 公式,而是 reshapetranspose、broadcasting、矩阵乘法以及 KV Cache 带来的 Shape 变化。理解清楚这些 Shape,后面阅读完整 Transformer Attention 实现会容易很多。