手写Transformer:ai-engineering-from-scratch自注意力机制源码逐行讲解
在 ai-engineering-from-scratch 开源课程中,这一课教你用纯 NumPy 手写Transformer的核心——自注意力机制。不需要 GPU、不需要 PyTorch,只要 100 行 Python 代码,你就能看懂 GPT 等大模型最关键的"缩放点积注意力"是如何一步步算出来的。🔍
为什么要手写自注意力机制?
RNN 处理序列时是一个 token 一个 token 地"挤牙膏",到第 50 个词时,第 1 个词的信息已经被压缩了 50 次,长距离依赖被严重削弱。
而**自注意力(Self-Attention)**让序列中的每个位置一次性"看到"所有其他位置,完全并行——这正是 Transformer 快速、可扩展、并成为大模型主流架构的根本原因。
课程把这一课放在 Phase 7: Transformers Deep Dive,源码入口在这里:
- 核心源码:self_attention.py
- 配套讲解:en.md
- 动手练习:quiz.json
动手前先懂:Q、K、V 的"数据库查表"类比
注意力就像一次软性的数据库查询:
| 向量 | 通俗含义 |
|---|---|
| Query (Q) | "我在找什么信息?" |
| Key (K) | "我这里有什么信息?" |
| Value (V) | "如果被选中,我提供什么内容?" |
每个 token 通过三个可学习的权重矩阵 Wq、Wk、Wv 投影出 Q、K、V;Q 与所有 K 做点积得到"匹配分数",经 softmax 变成权重,再加权求和 V,得到输出。一行公式概括:
Attention(Q, K, V) = softmax(QKᵀ / √dk) · V
源码逐行讲解
整个实现分四步,代码都在 self_attention.py 中。
第一步:手写 softmax(数值稳定是关键)
def softmax(x):
shifted = x - np.max(x, axis=-1, keepdims=True) # 减去最大值,防止溢出
exp_x = np.exp(shifted)
return exp_x / np.sum(exp_x, axis=-1, keepdims=True)
逐行看(源码 L4-L7):
shifted = x - np.max(x, ...):先减去每行最大值。若输入是[100, 200, 300]这样的大数,直接exp会溢出;减完最大值后最大项变成exp(0)=1,数值绝对安全。np.exp(shifted):指数化,把分数变成"越大越突出"的正数。- 除以总和:归一化为一行权重和为 1 的概率分布。
💡 源码文件末尾还专门放了 [100, 200, 300] 的演示用例,跑一遍就能亲眼看到"大 logits 也不会溢出"。
第二步:缩放点积注意力——核心只有 4 行
def scaled_dot_product_attention(Q, K, V):
dk = Q.shape[-1]
scores = Q @ K.T / np.sqrt(dk) # 点积打分,并除以 √dk
weights = softmax(scores) # 分数 → 权重
output = weights @ V # 加权求和
return output, weights
逐行看(源码 L10-L15):
dk = Q.shape[-1]:取键向量维度,作为缩放因子。scores = Q @ K.T / np.sqrt(dk):一次矩阵乘法就算出所有 token 两两之间的注意力分数。为什么要除以√dk?高维随机向量的点积会随dk增大而变大,直接把分数喂给 softmax 会"饱和"(输出趋近 one-hot,梯度消失)。缩放后 softmax 输出平滑,梯度健康——这是 Transformer 论文里最经典的设计细节。weights = softmax(scores):每一行独立归一化,得到"这个 token 看其他 token 各花多少注意力"。output = weights @ V:按权重混合所有 Value 向量,得到每个 token 的新表示。
第三步:SelfAttention 类——加入可学习投影
class SelfAttention:
def __init__(self, d_model, dk, dv, seed=42):
rng = np.random.default_rng(seed)
scale_qk = np.sqrt(2.0 / (d_model + dk)) # Xavier 风格缩放
self.Wq = rng.normal(0, scale_qk, (d_model, dk))
self.Wk = rng.normal(0, scale_qk, (d_model, dk))
scale_v = np.sqrt(2.0 / (d_model + dv))
self.Wv = rng.normal(0, scale_v, (d_model, dv))
def forward(self, X):
Q = X @ self.Wq # 每个 token 的"查询"
K = X @ self.Wk # 每个 token 的"标签"
V = X @ self.Wv # 每个 token 的"内容"
return scaled_dot_product_attention(Q, K, V)
逐行看(源码 L18-L32):
- 权重初始化用
√(2/(d_model + dk))缩放:这是 Xavier 风格的初始化,控制方差,避免训练初期信号爆炸或消失——真实训练时同样重要。 Q = X @ self.Wq:同一份输入 X 投影三次,"自注意力的自"就体现在 Q、K、V 全部来自同一序列。forward只做投影 + 调用第二步:结构极简,职责单一,方便复用。
第四步:多头注意力 MultiHeadSelfAttention
class MultiHeadSelfAttention:
def __init__(self, d_model, n_heads, seed=42):
assert d_model % n_heads == 0
self.dk = d_model // n_heads
self.dv = d_model // n_heads
self.heads = [SelfAttention(d_model, self.dk, self.dv, seed=seed + i)
for i in range(n_heads)]
# 输出投影矩阵 Wo
self.Wo = rng.normal(0, scale, (n_heads * self.dv, d_model))
def forward(self, X):
head_outputs, all_weights = [], []
for head in self.heads: # 每个头独立算一遍注意力
out, w = head.forward(X)
head_outputs.append(out)
concatenated = np.concatenate(head_outputs, axis=-1) # 拼接各头
output = concatenated @ self.Wo # 输出投影
return output, all_weights
逐行看(源码 L35-L58):
d_model // n_heads:模型维度均分给每个头。比如d_model=512、8 个头,每头 64 维——总计算量和单头满维度几乎相同,却获得了多个"视角"。- 多头并行:不同头可以同时关注语法关系、指代关系、位置关系等不同类型,表达能力远超单头。
concatenate + Wo:各头结果拼接后过输出投影矩阵,融合成最终表示。
⚡ 为什么是"多头"而不是"一个大注意力"?因为单一 QKᵀ 容易让 softmax 收敛到相似的分布;拆成多个小空间各自竞争,模型才能学到多样的注意力模式。
运行验证:用一句话生成注意力矩阵
课程自带完整可运行示例,对句子 "The cat sat on the mat" 构造随机嵌入,输出单头与多头的注意力权重表,还画了 ASCII 热力图(源码 L88-L146)。
运行方式(在仓库根目录):
python3 phases/07-transformers-deep-dive/02-self-attention-from-scratch/code/self_attention.py
你会看到每个 token 对全句的注意力分布——哪一行"更亮",就说明这个 token 更依赖那些词。同一课程还提供 Julia 版(main.jl)和 Rust 版(main.rs)实现,适合对比不同语言的数值实现细节。
常见疑问与延伸练习
结合 quiz.json 里的自测题,新手最容易卡住的三个点:
- 为什么要除以 √dk? 防止高维点积过大把 softmax 推入梯度消失区,不是为了归一化向量长度。
- 因果掩码怎么加? 在 softmax 之前把"未来位置"的分数置为负无穷,权重即归零——这是 decoder 自回归生成的关键。
- 为什么复杂度是 O(n²)?
Q @ K.T产生 n×n 的注意力矩阵,计算和显存都随序列长平方增长。这也是后续课程 FlashAttention、KV Cache 要解决的问题,可继续看 KV Cache 与 FlashAttention 一课。
课程推荐的三个动手练习(来自 en.md):
- 给
scaled_dot_product_attention加可选的 mask 参数,实现因果掩码; - 从零实现多头注意力(本文件已给出参考,可先自己写一遍再对照);
- 用同一实例跑两句不同的话,对比注意力模式:什么变了?什么没变?
接下来学什么?
学完本课,建议按 Phase 7 的路径继续推进:
| 主题 | 路径 |
|---|---|
| 多头注意力深入 | 03-multi-head-attention/ |
| 位置编码 | 04-positional-encoding/ |
| 完整 Transformer | 05-full-transformer/ |
| BERT / GPT 两条路线 | 06-bert-masked-language-modeling/、07-gpt-causal-language-modeling/ |
| 前置:注意力机制初探 | 10-attention-mechanism/ |
本课程的口号是"Learn it. Build it. Ship it."——把自注意力亲手写一遍之后,再看任何 Transformer 论文,你会发现它们都只是在这 100 行代码上的变体。🚀
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考




