别再只会调 API 了:跟着 ai-engineering-from-scratch 从零手写自注意力机制(Self-Attention)
本文是 rohitg00/ai-engineering-from-scratch 这一开源项目的一篇定点深度剖析。全文只聚焦一个点——Transformer 的心脏:自注意力机制(Self-Attention)。我会把 Q/K/V 的数学含义、缩放点积注意力的逐步实现、多头注意力的完整代码、以及它与 RNN/CNN 的对比表格一次讲透,最后附上规范的参考文献方便溯源。
摘要
ai-engineering-from-scratch是一个 “Learn it. Build it. Ship it.” 的开源 AI 工程课程,覆盖从线性代数到 Transformer、LLM、RAG、Agent 再到 MCP 的完整链路,据社区文章统计课程规模在435 课 / 10k+ Star量级[1][6]。它最反主流的一点是:不把大模型当黑盒,而是让你动手"从零实现"每一块积木。
本文选择其中Phase 7「Transformers Deep Dive」的02-self-attention-from-scratch一课[3] 作为唯一剖析对象,原因有二:
- 自注意力是 GPT、BERT、几乎所有现代 LLM 的共同底层构件,吃透它就等于拿到理解后续 RAG、微调、Agent 的钥匙;
- 它的核心算法只有"点积 → 缩放 → softmax → 加权求和"四步,几十行 NumPy 就能跑通,是"从零实现"性价比最高的一块。
读完你会得到:一份可运行的 NumPy 版缩放点积注意力、一份完整的多头注意力 PyTorch 实现、四张对比表,以及一页可直接引用的文献清单。
1. 这个项目在解决什么问题
先把背景铺平,方便理解我为什么挑"自注意力"这一个点来深挖。
ai-engineering-from-scratch的结构大致是:数学基础 → 机器学习 → NLP 基础 → Transformers 深潜(Phase 7)→ 从零 LLM(Phase 10)→ LLM 工程(Phase 11)→ 多模态(Phase 12)→ Agent / MCP,并用ROADMAP.md逐课记录进度,每节课都带docs/en.md讲义和可运行的构建代码[1][2][5]。它刻意强调多语言(Python / TypeScript / Rust / Julia)与"亲手构建",目的就是对抗一种普遍现象:很多人会用chat.completions.create(),却说不清输入一个 token 之后模型内部到底发生了什么。
而所有"内部到底发生了什么"的问题,最终都会收敛到一句话:
注意力机制是 LLM 对 token 之间依赖关系建模的核心算子。
所以本文不面面俱到地罗列 435 节课,而是只拆这一个算子。这正是"选一个点深入剖析"的价值:把一层拆到分子,胜过把十层各看一眼。
2. 为什么需要自注意力:它到底解决了什么
在 Transformer 之前,序列建模基本靠两大家族:
- RNN(含 LSTM/GRU):按时间步顺序地读,第t步的隐藏状态依赖第t-1步。问题是无法并行,且长距离依赖会梯度消失;
- CNN(如带膨胀卷积的 seq2seq):可以并行,但感受野有限,要堆很多层才能让距离远的 token 相互看见,长距离建模"间接"且昂贵。
自注意力的核心洞察是:让序列里任意两个位置之间都有一条"直达通道",并让模型自己学出"谁该关注谁"的权重。这一思想最早以独立的 self-attention 形式出现在 Lin 等人的句子嵌入工作中[10],随后被 Vaswani 等人整合进 Transformer 并彻底放大[7]。
用一张表对比三者(基于 Vaswani 论文 Table 1 简化,n= 序列长度,d= 特征维度,k= 卷积核大小):
| 结构 | 每层计算复杂度 | 顺序操作数(并行瓶颈) | 最大路径长度 | 长距离依赖能力 |
|---|---|---|---|---|
| 循环 RNN | O(n · d²) | O(n) | O(n) | 差(易梯度消失) |
| 卷积 CNN | O(k · n · d²) | O(1) | O(logₖ n) | 中(需堆叠多层) |
| Self-Attention | O(n² · d) | O(1) | O(1) | 强(任意两位置直达) |
一眼能看出 trade-off:自注意力用O(n²)的显存/计算代价,换来了常数级的最大路径长度和完全并行——这就是它为什么值得被单独拿出来"从零实现"。
顺带一提:那个 O(n²) 正是后来 FlashAttention[13]、线性注意力等一堆优化的起点,本文第 7 节会作为延伸点到为止。
3. 数学拆解:Q、K、V 到底在干什么
很多人卡在"Q/K/V 是什么"这一步。一个最直观的类比是检索/查字典:
- Query(查询):当前 token 发出的"我想找什么";
- Key(键):序列里每个 token 的"我有什么可以被检索的标签";
- Value(值):每个 token 真正要传递出去的"内容"。
注意力分数就是"查询与键的匹配程度",然后用这个匹配度去加权平均所有 Value。公式如下:
Attention ( Q , K , V ) = softmax ( Q K ⊤ d k ) V \text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)VAttention(Q,K,V)=softmax(dkQK⊤)V
其中d k d_kdk是 Key/Query 的维度。逐项拆开:
- Q K ⊤ QK^\topQK⊤:点积打分,得到一个T q × T k T_q \times T_kTq×Tk的"相关性矩阵",数值越大表示这两个位置越"互相相关";
- 1 d k \frac{1}{\sqrt{d_k}}dk1(缩放):这是关键一步。d k d_kdk增大时,点积的方差会线性增大,softmax 输入会被推到饱和区,梯度趋近 0。除以d k \sqrt{d_k}dk把方差拉回常数级,保证训练稳定(这也是它叫scaleddot-product 的原因)[7];
- softmax:把每一行归一化成概率分布(和为 1),即"注意力权重";
- 乘V VV:用注意力权重对所有位置的 Value 加权求和,得到当前 Query 位置的输出表示。
一个容易踩的坑:softmax 在数值上要做"减去行最大值"的稳定化处理,否则大分数会溢出。第 4 节的代码里我会显式写出来。
4. 从零实现(一):NumPy 版缩放点积注意力
先给一份不依赖任何框架、纯 NumPy的实现,把上面四步翻译成代码。这里用二维张量(T, d)方便读懂,多 batch/多头的情况放到第 5 节用 PyTorch 处理。
importnumpyasnpdefscaled_dot_product_attention(Q,K,V,mask=None):""" 缩放点积注意力(Scaled Dot-Product Attention)—— NumPy 版 Q: (T_q, d_k) 查询矩阵 K: (T_k, d_k) 键矩阵 V: (T_k, d_v) 值矩阵 mask: (T_q, T_k) 可选掩码,True 表示允许参与,False 表示屏蔽 返回: (输出 (T_q, d_v), 注意力权重 (T_q, T_k)) """d_k=Q.shape[-1]# 步骤 1:点积打分,得到相关性矩阵scores=Q @ K.T# (T_q, T_k)# 步骤 2:缩放,防止 softmax 进入饱和区scores=scores/np.sqrt(d_k)# 步骤 3(可选):掩码,把被屏蔽位置打到 -1e9(softmax 后趋近 0)ifmaskisnotNone:scores=np.where(mask,scores,-1e9)# 步骤 4:数值稳定的 softmax(每行减去该行最大值)scores=scores-scores.max(axis=-1,keepdims=True)exp_scores=np.exp(scores)attn_weights=exp_scores/exp_scores.sum(axis=-1,keepdims=True)# 步骤 5:用注意力权重加权求和所有 Valueoutput=attn_weights @ V# (T_q, d_v)returnoutput,attn_weights下面用一个 4 个 token 的小例子跑通它,并验证"注意力权重每行和为 1":
np.random.seed(42)T,d_k,d_v=4,8,8Q=np.random.randn(T,d_k)K=np.random.randn(T,d_k)V=np.random.randn(T,d_v)output,attn=scaled_dot_product_attention(Q,K,V)print("注意力权重矩阵(每行和为 1):")print(np.round(attn,3))# 4x4 矩阵,具体数值随 seed 而定print("行和:",attn.sum(axis=-1))# 应输出 [1. 1. 1. 1.]print("输出形状:",output.shape)# (4, 8)要点复盘:整个"自注意力"就浓缩在Q @ K.T这一个矩阵乘法里——它一次性算出了所有 token 两两之间的相关性,这就是"常数级最大路径长度"的由来:不需要一步步传递,一个矩阵乘法就完成了全局信息交换。
5. 从零实现(二):PyTorch 版多头注意力 + 因果掩码
真实 Transformer 用的是多头注意力(Multi-Head Attention)。单头只能学一种"关注模式",多头则把d m o d e l d_{model}dmodel拆成h hh份,各自在不同的低维子空间里学习,最后拼接——相当于让模型同时捕捉"语法关系、指代关系、语义关系"等不同维度的依赖[7]。
importmathimporttorchimporttorch.nnasnnclassMultiHeadAttention(nn.Module):def__init__(self,d_model,n_heads,dropout=0.1):super().__init__()assertd_model%n_heads==0,"d_model 必须能被 n_heads 整除"self.d_model=d_model self.n_heads=n_heads self.d_k=d_model//n_heads# 每个头的维度# 四个线性投影:Q/K/V 输入投影 + 输出投影self.W_q=nn.Linear(d_model,d_model,bias=False)self.W_k=nn.Linear(d_model,d_model,bias=False)self.W_v=nn.Linear(d_model,d_model,bias=False)self.W_o=nn.Linear(d_model,d_model,bias=False)self.dropout=nn.Dropout(dropout)defforward(self,x,mask=None):B,T,_=x.shape# batch, 序列长度, 维度# 1) 线性投影后切成多头,并交换维度便于批量点积# 形状从 (B, T, d_model) -> (B, n_heads, T, d_k)Q=self.W_q(x).view(B,T,self.n_heads,self.d_k).transpose(1,2)K=self.W_k(x).view(B,T,self.n_heads,self.d_k).transpose(1,2)V=self.W_v(x).view(B,T,self.n_heads,self.d_k).transpose(1,2)# 2) 缩放点积注意力(对最后两维做矩阵乘)scores=(Q @ K.transpose(-2,-1))/math.sqrt(self.d_k)# 3) 掩码(如因果掩码),masked_fill 把屏蔽位置置为 -infifmaskisnotNone:scores=scores.masked_fill(mask==0,float("-inf"))# 4) softmax 归一化 + dropoutattn=torch.softmax(scores,dim=-1)attn=self.dropout(attn)# 5) 加权求和并拼接多头、做输出投影out=attn @ V# (B, n_heads, T, d_k)out=out.transpose(1,2).contiguous().view(B,T,self.d_model)returnself.W_o(out),attn因果掩码(causal mask)是 GPT 这类自回归模型的关键:当前位置只能看到它自己以及它之前的位置,不能"偷看"未来,否则推理时就泄露了答案。实现上就是一个下三角矩阵:
defcausal_mask(T):"""返回 (T, T) 的下三角布尔掩码,True=可见,False=屏蔽未来 token"""returntorch.tril(torch.ones(T,T)).bool()# 示例:T=4 时的可见性矩阵# [[1,0,0,0],# [1,1,0,0],# [1,1,1,0],# [1,1,1,1]]再把上面两点串起来,验证一次完整的多头前向:
torch.manual_seed(0)mha=MultiHeadAttention(d_model=512,n_heads=8)x=torch.randn(2,10,512)# batch=2, 序列长度=10, 维度=512mask=causal_mask(10).unsqueeze(0)# (1, 10, 10),广播到 batch 和 headout,attn=mha(x,mask)print(out.shape)# torch.Size([2, 10, 512]),与输入同形print(attn.shape)# torch.Size([2, 8, 10, 10])补充:
MultiHeadAttention里总参数量与单头相同——因为d_model被拆成n_heads份再并行投影,这是多头"不额外加参数却换来更强表达力"的巧妙之处。
6. 另一个绕不开的配角:位置编码(Positional Encoding)
自注意力本身是**"无序"的——把["我", "爱", "你"]任意打乱,注意力矩阵只是行/列跟着换位,模型感受不到顺序。所以 Transformer 会在输入上叠加位置编码,最经典的是 Vaswani 的正弦位置编码**[7]:
defsinusoidal_positional_encoding(max_len,d_model):"""正弦/余弦位置编码,返回 (max_len, d_model)"""pe=torch.zeros(max_len,d_model)pos=torch.arange(0,max_len).unsqueeze(1).float()# (max_len, 1)i=torch.arange(0,d_model,2).float()# 偶数维索引div=torch.exp(i*(-math.log(10000.0)/d_model))# 频率按指数衰减pe[:,0::2]=torch.sin(pos*div)# 偶数维用 sinpe[:,1::2]=torch.cos(pos*div)# 奇数维用 cosreturnpe它用不同频率的正弦波给每个位置一个唯一指纹,并让模型能通过相对位置关系泛化。这一块单独拎出来就能再写一篇,这里作为自注意力的"配套零件"点到为止。
7. 对比表格合集
为方便速查,我把本文涉及的几组对比集中在这里。
7.1 注意力打分函数对比(历史演进)
| 打分函数 | 公式 | 提出文献 | 特点 |
|---|---|---|---|
| Additive(拼接式) | v ⊤ tanh ( W [ Q ; K ] ) v^\top \tanh(W[Q; K])v⊤tanh(W[Q;K]) | Bahdanau et al., 2015[8] | 早期主流,表达灵活,需额外参数 |
| Dot(普通点积) | Q K ⊤ Q K^\topQK⊤ | Luong et al., 2015[9] | 无参数、快,但维度大时数值不稳 |
| General(乘法) | Q W K ⊤ Q W K^\topQWK⊤ | Luong et al., 2015[9] | 学习一个对齐矩阵,介于两者之间 |
| Scaled dot-product | Q K ⊤ d k \frac{Q K^\top}{\sqrt{d_k}}dkQK⊤ | Vaswani et al., 2017[7] | Transformer 默认,缩放保证训练稳定 |
7.2 单头 vs 多头
| 维度 | 单头 | 多头(Multi-Head) |
|---|---|---|
| 子空间数 | 1 | h(并行低维子空间) |
| 关注模式 | 只能表达一种 | 同时捕捉语法/语义/位置等多种关系 |
| 参数量 | 基准 | 相同(d_model 被拆分,总参数不变) |
| 计算复杂度 | O(n²d) | O(n²d)(量级一致) |
7.3 三种注意力变体(按 Q/K/V 来源与掩码分)
| 变体 | Q 来源 | K/V 来源 | 掩码 | 典型用途 |
|---|---|---|---|---|
| 自注意力(Self) | 本序列 | 本序列 | 无 | Transformer Encoder、BERT |
| 因果自注意力(Causal) | 本序列 | 本序列 | 下三角 | GPT 类 Decoder |
| 交叉注意力(Cross) | Decoder | Encoder | 无 | seq2seq 翻译的 Decoder |
7.4 工程延伸:O(n²) 注意力的优化方向
| 方案 | 核心思路 | 代表文献 |
|---|---|---|
| FlashAttention | IO 感知、分块 + 重计算,不改变数学结果 | Dao et al., 2022[13] |
| 线性注意力 | 用核函数把 QK^T 重排,降到 O(n) | Katharopoulos et al., 2020[14] |
| 稀疏/局部注意力 | 限制每个 token 只关注局部窗口 | Child et al., 2019[15] |
| KV Cache | 推理时缓存历史 K/V,避免重复计算 | 推理标配工程手段 |
之所以值得了解,是因为自注意力在 LLM 落地里的头号成本就是 O(n²):上下文一长,显存和延迟都会爆。理解了底层的 O(n²),你才能理解为什么大家拼命做长上下文优化。
8. 从"自注意力"到"AI 工程全貌":这一个点如何串起整个项目
最后把视角拉回项目本身,说明为什么挑这一个点是有"全局意义"的:
- 向上:自注意力拼装成 Transformer Block(多头注意力 + 前馈网络 + 残差 + LayerNorm),Block 堆叠成 Encoder/Decoder,再堆成 GPT/BERT[11][12]——项目 Phase 7 和 Phase 10 就是沿这条路从零搭 LLM[2][5];
- 向下:自注意力的打分矩阵Q K ⊤ QK^\topQK⊤,本质就是"向量相似度检索",这与项目后续RAG 的向量检索 / 相似度计算 / chunking一脉相承——在 Phase 5 的 chunking 策略和 Phase 11 的 LLM 工程里,你会反复用到同一套"向量 + 相似度"的直觉;
- 向右:理解了注意力,再看Agent 的工具调用、MCP 的上下文注入,无非是"如何组织进入注意力窗口的 token 序列",而不只是玄学。
一句话总结这个项目的价值:它把"调 API 的黑盒"拆回一个个可动手实现的零件,而自注意力是所有零件里杠杆最高、也最该第一个亲手写的那块。
参考文献
项目源码与讲义
- rohitg00/ai-engineering-from-scratch(GitHub 主仓库)
- ROADMAP.md(课程路线与进度)
- phases/07-transformers-deep-dive/02-self-attention-from-scratch/docs/en.md(本文剖析对象)
- DeepWiki:Self-Attention & Transformer Architecture
- DeepWiki:Curriculum Roadmap & Progress Tracking
- 社区解读:从模型、Agent 到 MCP,把 AI 工程学习路线重新铺了一遍
学术文献
- Vaswani A., Shazeer N., Parmar N., et al.Attention Is All You Need.NeurIPS, 2017. https://arxiv.org/abs/1706.03762
- Bahdanau D., Cho K., Bengio Y.Neural Machine Translation by Jointly Learning to Align and Translate.ICLR, 2015. https://arxiv.org/abs/1409.0473
- Luong M.-T., Pham H., Manning C. D.Effective Approaches to Attention-based Neural Machine Translation.EMNLP, 2015. https://arxiv.org/abs/1508.04025
- Lin Z., Feng M., Santos C. N., et al.A Structured Self-Attentive Sentence Embedding.ICLR, 2017. https://arxiv.org/abs/1703.03130
- Devlin J., Chang M.-W., Lee K., Toutanova K.BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding.NAACL, 2019. https://arxiv.org/abs/1810.04805
- Brown T., Mann B., Ryder N., et al.Language Models are Few-Shot Learners.NeurIPS, 2020. https://arxiv.org/abs/2005.14165
- Dao T., Fu D. Y., Ermon S., Ré C., Rudra A.FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.NeurIPS, 2022. https://arxiv.org/abs/2205.14135
- Katharopoulos A., Vyas A., Pappas N., Fleuret F.Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention.ICML, 2020. https://arxiv.org/abs/2006.16236
- Child R., Gray S., Radford A., Sutskever I.Generating Long Sequences with Sparse Transformers.2019. https://arxiv.org/abs/1904.10509
本文代码为演示"从零实现"的教学实现,侧重于可读性;生产环境建议直接使用 PyTorch 内置的nn.MultiheadAttention或torch.nn.functional.scaled_dot_product_attention(后者已内置 FlashAttention 加速)。如有疏漏,欢迎指正交流。