通俗版
多头(Multi-Head):为什么单头不够
考虑这句话:"苹果公司发布了新的苹果手机。"
你在理解这句话时同时在做好几件事:要搞清楚"苹果公司"和"苹果手机"是两个不同概念(语义消歧),同时要搞清楚"发布了"的主语是"苹果公司"(语法关系),同时还要知道"新的"修饰的是"手机"(依存关系)。这些理解任务同时进行,互不干扰。
单头注意力是一个"注意力通道",每次只能学到一种权重分布——类比你只有一个注意力资源,却要同时盯着语义、语法、指代三件事,做不到。
多头注意力就是开 N 份并行的注意力计算。Head 1 可能专门在学语义关系,Head 2 在学语法依存,Head 3 在学指代关系。每个头独立计算自己的 Q/K/V,得到自己的注意力矩阵,最后把所有头的输出拼接起来再过一个线性变换。
有人会问计算量不就翻倍了吗?实际没有。多头的做法是把原来的向量维度"分头"——如果原来是 512 维用 8 个头,每个头只处理 64 维。总计算量和单头 512 维差不多,但模型能同时从多个角度理解语言。
位置编码:Self-Attention 为什么不感知顺序
Self-Attention 计算的核心是矩阵乘法,参与运算的向量是一个集合,不是有序数组。如果把输入序列的顺序打乱,矩阵乘法的结果只是对应行重排,计算本身完全不受影响——模型感知不到"你"排在第一位还是第三位。
所以 Transformer 必须人工把位置信息注入进去。
最原始的做法用正弦波和余弦波叠加来给位置编号。直觉类似二进制计数:低位快速变化,高位缓慢变化,多个不同频率的波叠加,让每个位置的"波形组合"都唯一。这和二进制里用不同权重的位来表示不同数字是一个道理——用 4 位二进制,0000 到 1111 能表示 16 个不同的数;用 N 个不同频率的波叠加,N 个位置都有唯一的"频率指纹"。问题是这是绝对位置编码——模型学到的是"位置 1 长这样,位置 2 长这样",推理时遇到没见过的超长位置,就会懵掉。
RoPE(旋转位置编码)的核心改进是:从绝对位置变成相对位置。它不告诉模型"你在第几位",而是告诉模型"你和另一个词之间差了几位"。具体操作是在做点积之前,根据位置对 Q 和 K 分别做旋转变换,旋转角度取决于位置。旋转的妙处在于:两个向量点积之后,结果只和它们的相对旋转角度(也就是相对位置差)有关,与各自绝对位置无关。
这对长上下文扩展很重要:模型学到的是相对位置关系,而不只是某个绝对位置编号。但“8K 训练直接推到 128K”并不会自动成立,通常还要配合 RoPE scaling、继续训练或其他长上下文方法,并重新验证长距离质量。
GQA:显存和质量之间的工程权衡
先算一笔账。假设你有一个模型:32 个注意力头,每个头维度 128,上下文 128K,FP16。KV Cache 每层需要:
32 头 × 128 维 × 128000 位置 × 2 字节 × 2(K 和 V)≈ 2 GB / 每层
如果模型有 32 层,总共 64 GB,还没算模型权重本身。长上下文推理为什么那么贵,答案就在这。
三种方案的对比:
- MHA(多头注意力):每个头有独立的 K 和 V,表达能力最强,但 KV Cache 最大
- MQA(多查询注意力):所有 Query 头共享一组 K/V;相对 32 个独立 KV 头的 MHA,理论 KV Cache 可缩小到约 1/32,但实际质量变化取决于模型和训练方式
- GQA(分组查询注意力):把 32 个头分成若干组,组内共享 KV,组间独立。折中方案——比 MHA 省显存,比 MQA 质量好
Llama 3、Qwen、Mistral 选 GQA 而不选 MQA,因为 MQA 在长对话和复杂推理上质量掉得明显。GQA 只掉一点质量,就能省下大量显存。
进阶版
多头计算是否真正并行
各头之间完全独立,可以真正并行计算。PyTorch 的底层 CUDA kernel 通常会把多头打包成一个大矩阵乘法,利用 GPU 的 Tensor Core 批量处理,比"启动 N 个独立线程"的并行更高效,是向量化批量计算。头与头之间唯一有顺序依赖的是最后的拼接和线性投影,计算量可以忽略不计。
RoPE 的工程实现
RoPE 的旋转变换是把每个向量的相邻两维当作一个复数平面,按位置做复数乘法(相当于在二维平面上旋转)。整个操作没有新增可学习参数,是纯计算变换,对推理速度的影响极小。
外推能力的来源在于:相对位置 m-n 对应的旋转角差值,在训练集内短序列上出现过。只要模型学会了"距离 k 的相对关系",即使绝对位置超出训练范围,相对关系依然成立。实践中还会用 NTK-aware scaling、YaRN 这类频率插值方法对高频分量做进一步调整,让更长的外推更稳定——Qwen、DeepSeek 的 128K/1M 长上下文都是这条路。