ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

别再只会调 API 了:跟着 ai-engineering-from-scratch 从零手写自注意力机制(Self-Attention)

别再只会调 API 了:跟着 ai-engineering-from-scratch 从零手写自注意力机制(Self-Attention) 别再只会调 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→ 从零 LLMPhase 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 卷积核大小结构每层计算复杂度顺序操作数并行瓶颈最大路径长度长距离依赖能力循环 RNNO(n · d²)O(n)O(n)差易梯度消失卷积 CNNO(k · n · d²)O(1)O(logₖ n)中需堆叠多层Self-AttentionO(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(dk​​QK⊤​)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}}dk​​1​缩放这是关键一步。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,maskNone): 缩放点积注意力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_kQ.shape[-1]# 步骤 1点积打分得到相关性矩阵scoresQ K.T# (T_q, T_k)# 步骤 2缩放防止 softmax 进入饱和区scoresscores/np.sqrt(d_k)# 步骤 3可选掩码把被屏蔽位置打到 -1e9softmax 后趋近 0ifmaskisnotNone:scoresnp.where(mask,scores,-1e9)# 步骤 4数值稳定的 softmax每行减去该行最大值scoresscores-scores.max(axis-1,keepdimsTrue)exp_scoresnp.exp(scores)attn_weightsexp_scores/exp_scores.sum(axis-1,keepdimsTrue)# 步骤 5用注意力权重加权求和所有 Valueoutputattn_weights V# (T_q, d_v)returnoutput,attn_weights下面用一个 4 个 token 的小例子跑通它并验证注意力权重每行和为 1np.random.seed(42)T,d_k,d_v4,8,8Qnp.random.randn(T,d_k)Knp.random.randn(T,d_k)Vnp.random.randn(T,d_v)output,attnscaled_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,dropout0.1):super().__init__()assertd_model%n_heads0,d_model 必须能被 n_heads 整除self.d_modeld_model self.n_headsn_heads self.d_kd_model//n_heads# 每个头的维度# 四个线性投影Q/K/V 输入投影 输出投影self.W_qnn.Linear(d_model,d_model,biasFalse)self.W_knn.Linear(d_model,d_model,biasFalse)self.W_vnn.Linear(d_model,d_model,biasFalse)self.W_onn.Linear(d_model,d_model,biasFalse)self.dropoutnn.Dropout(dropout)defforward(self,x,maskNone):B,T,_x.shape# batch, 序列长度, 维度# 1) 线性投影后切成多头并交换维度便于批量点积# 形状从 (B, T, d_model) - (B, n_heads, T, d_k)Qself.W_q(x).view(B,T,self.n_heads,self.d_k).transpose(1,2)Kself.W_k(x).view(B,T,self.n_heads,self.d_k).transpose(1,2)Vself.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:scoresscores.masked_fill(mask0,float(-inf))# 4) softmax 归一化 dropoutattntorch.softmax(scores,dim-1)attnself.dropout(attn)# 5) 加权求和并拼接多头、做输出投影outattn V# (B, n_heads, T, d_k)outout.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屏蔽未来 tokenreturntorch.tril(torch.ones(T,T)).bool()# 示例T4 时的可见性矩阵# [[1,0,0,0],# [1,1,0,0],# [1,1,1,0],# [1,1,1,1]]再把上面两点串起来验证一次完整的多头前向torch.manual_seed(0)mhaMultiHeadAttention(d_model512,n_heads8)xtorch.randn(2,10,512)# batch2, 序列长度10, 维度512maskcausal_mask(10).unsqueeze(0)# (1, 10, 10)广播到 batch 和 headout,attnmha(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)petorch.zeros(max_len,d_model)postorch.arange(0,max_len).unsqueeze(1).float()# (max_len, 1)itorch.arange(0,d_model,2).float()# 偶数维索引divtorch.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-productQ K ⊤ d k \frac{Q K^\top}{\sqrt{d_k}}dk​​QK⊤​Vaswani et al., 2017[7]Transformer 默认缩放保证训练稳定7.2 单头 vs 多头维度单头多头Multi-Head子空间数1h并行低维子空间关注模式只能表达一种同时捕捉语法/语义/位置等多种关系参数量基准相同d_model 被拆分总参数不变计算复杂度O(n²d)O(n²d)量级一致7.3 三种注意力变体按 Q/K/V 来源与掩码分变体Q 来源K/V 来源掩码典型用途自注意力Self本序列本序列无Transformer Encoder、BERT因果自注意力Causal本序列本序列下三角GPT 类 Decoder交叉注意力CrossDecoderEncoder无seq2seq 翻译的 Decoder7.4 工程延伸O(n²) 注意力的优化方向方案核心思路代表文献FlashAttentionIO 感知、分块 重计算不改变数学结果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多头注意力 前馈网络 残差 LayerNormBlock 堆叠成 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-scratchGitHub 主仓库ROADMAP.md课程路线与进度phases/07-transformers-deep-dive/02-self-attention-from-scratch/docs/en.md本文剖析对象DeepWikiSelf-Attention Transformer ArchitectureDeepWikiCurriculum 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.03762Bahdanau D., Cho K., Bengio Y.Neural Machine Translation by Jointly Learning to Align and Translate.ICLR, 2015. https://arxiv.org/abs/1409.0473Luong M.-T., Pham H., Manning C. D.Effective Approaches to Attention-based Neural Machine Translation.EMNLP, 2015. https://arxiv.org/abs/1508.04025Lin Z., Feng M., Santos C. N., et al.A Structured Self-Attentive Sentence Embedding.ICLR, 2017. https://arxiv.org/abs/1703.03130Devlin 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.04805Brown T., Mann B., Ryder N., et al.Language Models are Few-Shot Learners.NeurIPS, 2020. https://arxiv.org/abs/2005.14165Dao 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.14135Katharopoulos A., Vyas A., Pappas N., Fleuret F.Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention.ICML, 2020. https://arxiv.org/abs/2006.16236Child 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 加速。如有疏漏欢迎指正交流。
返回列表