BotOf TechAI / IoT / Full-Stack / 植物养护知识分享
W0201 / 夯实底层理论与架构

QKV 与多头注意力

从相似度、缩放项、掩码和张量形状推导注意力,并写出可与 PyTorch 对齐的多头实现。

建议时长
8–10 小时
难度
核心机制
课程位置
2 / 12
LEARNING OBJECTIVES

学完这一课,你应该能:

  • 能解释 Q、K、V 分别承担的匹配与聚合职责
  • 能推导除以 √dₖ 的方差原因
  • 能正确构造 causal mask 与 padding mask
  • 能实现多头注意力并验证前向和梯度

1. Attention 在解决什么问题

当前位置 (i) 要从所有可见位置 (j) 收集信息。过程分两步:

  1. 匹配:当前位置需要什么,历史位置能提供什么?
  2. 聚合:按匹配权重,对历史信息加权求和。
Q = XWq   当前 token 想找什么
K = XWk   每个 token 如何被匹配
V = XWv   被找到后交付的内容

Q 与 K 用来“算权重”,V 用来“传内容”。分开投影让模型可以独立学习检索标准与信息载荷。

2. 为什么要除以 (\sqrt{d_k})

单个 query 与 key 的相似度是点积:

sⱼ = q · kⱼ

若 q、k 各分量独立、均值为 0、方差约为 1,点积是 (d_k) 项乘积之和,方差随 (d_k) 增长。维度变大时 logits 绝对值变大,softmax 更容易饱和,梯度变得尖锐。缩放后:

Attention(Q,K,V) = softmax(QKᵀ / √dₖ + M)V

缩放使 logits 方差回到较稳定量级。(M) 是 mask:允许位置加 0,不允许位置加极大负数。必须在 softmax 之前加,因为被屏蔽位置不应参与归一化。

3. 一次写清所有张量形状

设 batch 为 B、序列长 T、模型维 D、头数 H、每头维度 Dh=D/H:

Q,K,V projection: [B,T,D]
split heads:       [B,H,T,Dh]
Q @ Kᵀ:            [B,H,T,T]
softmax(scores)@V: [B,H,T,Dh]
merge heads:       [B,T,D]

调试注意力时,先打印这六个形状,再看热力图。大多数错误来自 transpose 轴次序、mask 广播或 head 合并。

4. Causal mask 与 padding mask

自回归模型位置 (i) 不能看到未来位置 (j>i):

[[0, -∞, -∞, -∞],
 [0,  0, -∞, -∞],
 [0,  0,  0, -∞],
 [0,  0,  0,  0]]

Causal mask 常为 [1,1,T,T],沿 batch 和 head 广播。Padding mask 从 [B,T] 变成 [B,1,1,T],屏蔽作为 key/value 来源的补齐 token。注意不同 API 对布尔 mask 的定义可能相反,不能凭习惯猜 True 是保留还是屏蔽。

5. 最小多头注意力实现

import math
import torch
from torch import nn

class MultiHeadSelfAttention(nn.Module):
    def __init__(self, d_model: int, n_heads: int):
        super().__init__()
        assert d_model % n_heads == 0
        self.n_heads = n_heads
        self.d_head = d_model // n_heads
        self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
        self.out = nn.Linear(d_model, d_model, bias=False)

    def forward(self, x, padding_mask=None):
        b, t, d = x.shape
        q, k, v = self.qkv(x).chunk(3, dim=-1)

        def split(z):
            return z.view(b, t, self.n_heads, self.d_head).transpose(1, 2)

        q, k, v = map(split, (q, k, v))
        scores = q @ k.transpose(-2, -1) / math.sqrt(self.d_head)

        future = torch.triu(
            torch.ones(t, t, device=x.device, dtype=torch.bool),
            diagonal=1,
        )
        scores = scores.masked_fill(future[None, None], float("-inf"))

        if padding_mask is not None:
            scores = scores.masked_fill(
                padding_mask[:, None, None, :],
                float("-inf"),
            )

        weights = torch.softmax(scores, dim=-1)
        context = weights @ v
        context = context.transpose(1, 2).contiguous().view(b, t, d)
        return self.out(context), weights

transpose 通常只改变 stride。随后使用 view 前调用 contiguous(),是为了让底层布局和目标形状一致。使用 reshape 也可自动处理,但理解 stride 有助于排查性能问题。

6. 不止比较输出形状:对齐官方实现

复制相同权重,比较前向与输入梯度:

torch.manual_seed(17)
b, t, d, h = 2, 5, 16, 4
x1 = torch.randn(b, t, d, dtype=torch.float64, requires_grad=True)
x2 = x1.detach().clone().requires_grad_(True)

mine = MultiHeadSelfAttention(d, h).double()
ref = nn.MultiheadAttention(d, h, bias=False, batch_first=True).double()

with torch.no_grad():
    ref.in_proj_weight.copy_(mine.qkv.weight)
    ref.out_proj.weight.copy_(mine.out.weight)

mask = torch.triu(torch.ones(t, t, dtype=torch.bool), diagonal=1)
y1, _ = mine(x1)
y2, _ = ref(x2, x2, x2, attn_mask=mask, need_weights=False)

print("forward diff", (y1 - y2).abs().max().item())
y1.sum().backward()
y2.sum().backward()
print("gradient diff", (x1.grad - x2.grad).abs().max().item())

验收标准不是“能运行”,而是相同输入、权重和 mask 下,前向与梯度最大误差都足够小。

7. 复杂度与系统瓶颈

标准注意力分数矩阵大小为 [T,T],计算近似 (O(T^2D)),中间概率内存近似 (O(T^2))。但不要只背复杂度:

  • 短序列、大 D 时,QKV 与 FFN 线性层可能占更多 FLOPs。
  • 训练要保存反向中间量,Attention 概率显著占显存。
  • 自回归 decode 用 KV Cache 复用历史 K/V,但要持续读取增长的缓存。
  • FlashAttention 不改变精确注意力数学结果,主要通过分块与 IO-aware 算法减少显存读写。

8. 注意力权重不自动等于解释

热力图适合发现 mask、位置偏好与特殊 token 行为,但高权重不等于对最终输出因果重要:

  1. Value 与输出投影会继续改变影响。
  2. 多层、多头与残差共同作用。
  3. 相关性不能替代干预实验。

真正分析重要性应结合梯度、特征归因、attention rollout 或遮蔽干预。

9. 本课实验

变量对照必须观察
缩放项除以 √dₖ / 不缩放logits 方差、entropy、grad norm
causal mask正确 / 关闭是否偷看未来
head 数1 / 2 / 4 / 8吞吐、效果、单头维度
实现自研 / PyTorchforward diff、gradient diff

再注入一个错误:把 padding mask 写成 [B,T,1,1]。它可能可以广播,但屏蔽轴完全错误。用只有两个有效 token 的样本证明这一点。

10. 练习与答案提示

  1. D=768、H=12 时 Dh?**答:**64,分数最后两维是 [T,T]
  2. 为什么不能 softmax 后再乘 mask?**答:**概率和不再为 1;重新归一化等价于前置 mask。
  3. 总 D 不变时,头数增加会让 QKV 参数量增长吗?**答:**不会,矩阵只是被重排为多个头。
  4. causal LM 的第 0 个位置能看什么?**答:**只能看到自身和模型约定的前缀。
  5. 为什么点积适合作为相似度?**答:**能高效批量矩阵乘法,并与可学习投影自然组合。

11. 延伸阅读