本站此前发过概念篇《一篇读懂 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 省了重复计算」这类结论时,你能对应上它们压缩的正是这篇文章里的哪一步、哪块内存。
参考资料
- Vaswani et al., Attention Is All You Need(缩放论证与 mask 细节出处,3.2.1 节)— https://arxiv.org/abs/1706.03762
- karpathy/nanoGPT,model.py(对照的手写 MHA 实现)— https://github.com/karpathy/nanoGPT/blob/master/model.py
- HuggingFace Transformers 文档:Attention backends — https://huggingface.co/docs/transformers/attention_interface
- PyTorch 文档:torch.nn.functional.scaled_dot_product_attention — https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
- Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — https://arxiv.org/abs/2205.14135
- 本站概念篇:一篇读懂 Attention:Q、K、V 到底在算什么 — https://alishangtian.com/article/user-f06d713e1c
读者留言
COMMENTS 暂无还没有留言,来说第一句?