手撕 Multi-Head Attention:纯 Python 从零实现并跑通

本站此前发过概念篇《一篇读懂 Attention:Q、K、V 到底在算什么》,把注意力的直觉讲清了。本文是代码实现篇,目标只有一个:不借助任何深度学习框架,只用 numpy 把 Multi-Head Attention 的前向计算从零写出来、跑通、并且可验证。

先把约定说清楚:只实现前向,不做反向传播;参数随机初始化,不求「学到的含义」,只求数学正确;输入是形状 (batch, seq_len, d_model) 的三维张量;causal mask 作为开关传入。核心实现(类体)45 行——连注释和空行一起算——远低于 100 行的预算。文中全部代码在 Python 3.13 + numpy 2.1.3 下原样跑通,三个验证断言全部通过。

从一个公式到七步流水线

Attention 的核心公式只有一行(Transformer 论文 3.2.1 节):

Attention(Q, K, V) = softmax(QKᵀ / √d_k) · V

落到代码上,它就是一条七步流水线:

flowchart LR
    X["x (B, T, 512)"] --> A["QKV 线性投影"]
    A --> B["分头 (B, 8, T, 64)"]
    B --> C["QKᵀ / √64"]
    C --> D["causal mask"]
    D --> E["softmax"]
    E --> F["· V 加权求和"]
    F --> G["合头 (B, T, 512)"]
    G --> H["输出投影"]

第一步:QKV 线性投影。 同一份输入 x 分别乘三个权重矩阵加偏置,得到 Q(问询)、K(索引)、V(内容)。实现里用三个独立的 d_model × d_model 矩阵,比 nanoGPT 的「一个 c_attn 一次产出 3 倍宽度再切开」更直白,数学完全等价。

第二步:分头。 把 (B, T, d_model) 先 reshape 成 (B, T, h, d_k) 再 transpose 成 (B, h, T, d_k)。这一步容易看花眼,但它是「多头」的全部含义:不是把一次注意力切八份,而是让 8 个头各自在 64 维子空间里独立做注意力,最后再拼回来。HuggingFace transformers 的实现同样是 reshape + transpose,head_dim = embed_dim // num_attention_heads。这里有个容易踩的陷阱:transpose 之后张量在内存里不再连续,合头时必须先 transpose 回来再 reshape——顺序搞反的话 reshape 不会报错,但会悄悄把不同头的 token 打乱串位,这种 bug 输出形状全对、数值全错,极难排查。

第三步:缩放点积。 除以 √d_k 不是可选项。论文 3.2.1 节的脚注给了方差论证:设 q、k 的各分量独立、均值 0、方差 1,则点积 q·k = Σqᵢkᵢ 的均值为 0、方差为 d_k。d_k = 64 时点积的标准差是 8,这么大的数喂进 softmax,输出会逼近 one-hot、落在梯度极小的饱和区。除以 √d_k 恰好把方差拉回 1,让 softmax 保持在梯度健康的工作区。

第四步:causal mask。 训练 GPT 这类自回归模型时,位置 t 不许偷看 t 之后的内容。做法是构造下三角布尔矩阵,把上三角的打分置为 -inf。注意 mask 必须加在 softmax 之前:如果放到之后强行乘 0,被 mask 的位置仍会分走概率质量,每行和就不再是 1。mask 的效果是阶梯状的:第 0 个位置只看得见自己,第 1 个位置看得见前两个,到最后一行才放开整个序列——第 0 行的注意力分布恒为 1.0,输出就是它自己的 V。

第五步:softmax。 沿最后一维归一化,每个位置得到一个对 ≤t 全部位置的注意力分布。实现里先减去每行最大值再取指数——这是不改变结果的数值稳定技巧,防止指数溢出;有 -inf 参与时它同样安全,因为 causal mask 下每行至少有对角线一个合法位置,行最大值总是有限数。

第六步:加权求和。 attn @ v:每个位置的输出是 V 序列的凸组合,权重就是上一步的注意力分布。

第七步:合头 + 输出投影。 把 (B, h, T, d_k) 转回 (B, T, d_model),再过输出投影 Wo,多头信息重新混合进同一个表示空间。

形状账

手写注意力,九成的 bug 是形状 bug。把七步的形状变化列成表(B = batch,T = seq_len,h = n_heads,d_k = d_model / h):

步骤 操作 形状
输入 — (B, T, d_model)
1 QKV 投影 x @ W + b,共三次 (B, T, d_model)
2 分头 reshape → transpose (B, h, T, d_k)
3 打分 q @ kᵀ / √d_k (B, h, T, T)
4 mask 上三角置 -inf (B, h, T, T)
5 softmax 沿最后一维 (B, h, T, T)
6 加权求和 attn @ v (B, h, T, d_k)
7 合头 + 输出投影 transpose → reshape → @ Wo (B, T, d_model)

以 d_model=512、n_heads=8 为例,d_k=64。参数量可以手算:四个投影矩阵 4 × 512 × 512 = 1,048,576,四个偏置 4 × 512 = 2,048,合计 1,050,624。注意它跟 n_heads 无关:8 个头分的是同一个 512×512 投影矩阵的不同列块,切 1 个头还是 16 个头,参数量都一样——这是比形状更容易记错的一点。

实现

两段代码拼起来就是跑通验证的完整脚本,原样贴出(运行环境 Python 3.13.5 + numpy 2.1.3):

import numpy as np


class MultiHeadAttention:
    """纯 numpy 的 Multi-Head Attention 前向实现,不依赖任何深度学习框架。"""

    def __init__(self, d_model: int, n_heads: int, seed: int = 0):
        assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除"
        self.d_model, self.n_heads = d_model, n_heads
        self.d_k = d_model // n_heads                      # 每个头的维度
        rng = np.random.default_rng(seed)
        std = 1.0 / np.sqrt(d_model)                       # 与常见初始化量级对齐
        self.Wq = rng.normal(0, std, (d_model, d_model))   # Q 投影矩阵
        self.Wk = rng.normal(0, std, (d_model, d_model))   # K 投影矩阵
        self.Wv = rng.normal(0, std, (d_model, d_model))   # V 投影矩阵
        self.Wo = rng.normal(0, std, (d_model, d_model))   # 输出投影矩阵
        self.bq = np.zeros(d_model)                        # 四组偏置
        self.bk = np.zeros(d_model)
        self.bv = np.zeros(d_model)
        self.bo = np.zeros(d_model)

    def parameters(self):
        return [self.Wq, self.bq, self.Wk, self.bk,
                self.Wv, self.bv, self.Wo, self.bo]

    def _split_heads(self, x):                             # (B, T, d_model) -> (B, h, T, d_k)
        B, T, _ = x.shape
        return x.reshape(B, T, self.n_heads, self.d_k).transpose(0, 2, 1, 3)

    def forward(self, x, causal: bool = False):
        B, T, _ = x.shape
        q = self._split_heads(x @ self.Wq + self.bq)       # (B, h, T, d_k)
        k = self._split_heads(x @ self.Wk + self.bk)
        v = self._split_heads(x @ self.Wv + self.bv)

        scores = q @ k.transpose(0, 1, 3, 2) / np.sqrt(self.d_k)   # (B, h, T, T)
        if causal:                                         # 因果掩码:位置 i 只能看 <= i
            allowed = np.tril(np.ones((T, T), dtype=bool))
            scores = np.where(allowed, scores, -np.inf)

        scores -= scores.max(axis=-1, keepdims=True)       # 减行最大值,防 softmax 溢出
        weights = np.exp(scores)
        weights /= weights.sum(axis=-1, keepdims=True)     # 每行和为 1
        self.attn = weights                                # 留给验证与可视化

        out = weights @ v                                  # (B, h, T, d_k) 加权求和
        out = out.transpose(0, 2, 1, 3).reshape(B, T, self.d_model)  # 合头
        return out @ self.Wo + self.bo                     # (B, T, d_model)

初始化只写了一行却值得停下来看一眼:权重用标准差 1/√d_model 的高斯分布,偏置全零。这不是随手写的——上面的方差论证对初始化同样成立,std = 1/√d_model 保证了投影输出的各分量方差仍约为 1,于是逐层堆叠时激活值不会指数级放大或湮灭。Transformer 原始实现用的 Xavier 初始化、GPT 系列用的 0.02 * normal,都是同一思想的变体。

数值实验:三个断言

实现对不对,光看代码不够,得让数字说话。下面三个实验分别验证:causal mask 的因果性、softmax 的归一性、参数量与手算一致。

if __name__ == "__main__":
    B, T, d_model, n_heads = 2, 6, 512, 8
    mha = MultiHeadAttention(d_model, n_heads, seed=42)
    rng = np.random.default_rng(0)
    x = rng.normal(0, 1, (B, T, d_model))
    print("输出形状:", mha.forward(x, causal=True).shape)

    # 性质 a:causal mask 下,t 时刻的输出只依赖输入 <= t 的部分
    out = mha.forward(x, causal=True)
    t = 2
    x2 = x.copy()
    x2[:, t + 1:, :] = rng.normal(0, 1, (B, T - t - 1, d_model))   # 改掉所有 > t 的输入
    out2 = mha.forward(x2, causal=True)
    assert np.array_equal(out[:, : t + 1, :], out2[:, : t + 1, :]), "性质 a 失败"

    x3 = x.copy()
    x3[:, t, :] += 1.0                                             # 对照组:改 t 时刻本身
    out3 = mha.forward(x3, causal=True)
    assert not np.array_equal(out[:, t, :], out3[:, t, :]), "对照组失败:t 时刻输出竟然没变"

    # 性质 b:softmax 每行和为 1(causal 与非 causal 都成立)
    assert np.allclose(mha.attn.sum(axis=-1), 1.0), "性质 b 失败"
    mha.forward(x, causal=False)
    assert np.allclose(mha.attn.sum(axis=-1), 1.0), "性质 b 失败(非 causal)"

    # 性质 c:参数量与手算 4*d_model^2 + 4*d_model 一致
    n_params = sum(p.size for p in mha.parameters())
    assert n_params == 4 * d_model * d_model + 4 * d_model == 1_050_624, "性质 c 失败"

    print("全部断言通过:")
    print(f"  参数量 {n_params:,} = 4x512^2 + 4x512")
    print(f"  注意力权重形状 {mha.attn.shape},行和示例 {mha.attn.sum(axis=-1).ravel()[:3]}")

实际运行输出:

输出形状: (2, 6, 512)
全部断言通过:
  参数量 1,050,624 = 4x512^2 + 4x512
  注意力权重形状 (2, 8, 6, 6),行和示例 [1. 1. 1.]

几处值得说明的细节:

  • 性质 a 用的是 np.array_equal(逐位相等)而不是 np.allclose。 这不是碰运气:-inf 经 softmax 后是精确的 0.0,加权求和时被 mask 位置的贡献精确为零,浮点加 0 不改变部分和,所以前 t+1 个位置的输出应当逐位相同。断言能这么写,恰恰说明 mask 的实现是「真零」而非「近似零」。
  • 对照组防「虚过」。 只验证「改后面的输入输出不变」还不够——一个输出恒为常数的退化实现也能通过。所以补了反向对照:只改 t 时刻本身的输入,断言 t 时刻的输出必须变化。
  • 性质 b 对 causal 与非 causal 两种模式都断言:mask 打在 softmax 之前,所以即使一半位置是 -inf,每行和仍然严格为 1。
  • 性质 c 与上一节的手算闭环:sum(p.size ...) 数出的 1,050,624 与 4 × 512² + 4 × 512 精确相等。

玩具与生产的距离

这份 45 行实现在数学上与生产实现等价,工程上则差着一代人的优化:

  • 无融合 kernel。 这里 (B, h, T, T) 的注意力矩阵被完整物化在内存里,序列一长就是 O(T²) 的显存开销。FlashAttention 的思路是用分块重算让注意力矩阵从不离开 GPU 片上 SRAM——数学结果不变、访存量骤降。PyTorch 2.0 把它连同 memory-efficient 后端封装进 F.scaled_dot_product_attention;HuggingFace transformers 用 attn_implementation 在 eager / sdpa / flash_attention_2 之间切换;nanoGPT 里那段手写注意力也只在 PyTorch 较旧时作为回退,默认走 SDPA 快路径(截至 2026-10-07 的仓库现状)。细节本站此前写过《一篇读懂 FlashAttention:注意力为什么能又快又省显存》。
  • 无 KV Cache。 自回归生成时,本实现每生成一个 token 都重算全部历史的 K 和 V;生产实现会缓存 K、V,每步只算增量。原理同源的本站文章还有《vLLM 与 PagedAttention:把 GPU 显存当操作系统内存管》。
  • 只有前向。 训练需要反向传播(自动微分或手推四条梯度公式),还有 dropout、混合精度(fp16/bf16)等一整层工程细节,这里全部省略。

如果想在这个骨架上继续动手,有三条难度递增的练习路线:给 forward 加 padding mask,让变长 batch 里补齐的位置既不被看见也不参与提问(提示:-inf 的位置要多一处来源);手推并实现 Q、K、V、Wo 的四条梯度公式,然后与有限差分对拍;再把 RoPE 之类的位置编码挂进 Q、K——挂的位置恰好在打分之前,这也是本文流水线图能一路延伸下去的地方。

一句话收束:手撕的意义不是造轮子,而是把每一步摊开。之后再读到「FlashAttention 不改变数学」「KV Cache 省了重复计算」这类结论时,你能对应上它们压缩的正是这篇文章里的哪一步、哪块内存。

参考资料

  1. Vaswani et al., Attention Is All You Need(缩放论证与 mask 细节出处,3.2.1 节)— https://arxiv.org/abs/1706.03762
  2. karpathy/nanoGPT,model.py(对照的手写 MHA 实现)— https://github.com/karpathy/nanoGPT/blob/master/model.py
  3. HuggingFace Transformers 文档:Attention backends — https://huggingface.co/docs/transformers/attention_interface
  4. PyTorch 文档:torch.nn.functional.scaled_dot_product_attention — https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
  5. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — https://arxiv.org/abs/2205.14135
  6. 本站概念篇:一篇读懂 Attention:Q、K、V 到底在算什么 — https://alishangtian.com/article/user-f06d713e1c
← 返回资讯列表

读者留言

COMMENTS 暂无
仅本站原创文章开放留言 · 请勿留下手机号、邮箱等个人信息

还没有留言,来说第一句?