手写Transformer注意力机制:从QKV矩阵乘法到多头实现
发布时间:2026/9/18 23:22:00来源:尧图网络
简介本资源是一份面向人工智能学习者与深度学习从业者的Transformer架构与注意力机制系统性解析资料聚焦大模型底层原理特别适合希望深入理解LLM技术根基的中高级开发者、算法工程师及高校研究者。文档以PDF形式呈现共1个文件大小3.56MB内容覆盖自注意力机制的数学原理与实现逻辑、多头注意力的并行建模思想、编码器-解码器结构的模块化设计含残差连接、层归一化、前馈网络等关键组件并对比RNN/LSTM在长程依赖与并行训练上的本质差异。文中结合NLP与计算机视觉双场景说明应用适配性还剖析了仅编码器如BERT、仅解码器如GPT等变体架构的设计动因与任务适配逻辑。目前已有217人下载学习内容结构清晰、图示精要、术语准确可作为理论补强、面试复盘或大模型研发前的技术预研材料。1. 不是“调个库就能跑通”而是看懂 QKV 矩阵乘法里到底在算什么很多人把 Transformer 当成一个黑盒输入文本喂进transformers库的AutoModel.from_pretrained(bert-base-uncased)再接个分类头训练完就上线。但当模型在长文本上掉点、在低资源场景下泛化变差、或 attention map 可视化结果完全无法解释时问题往往出在对注意力机制底层运算的模糊理解上——比如你是否清楚Q K.T / sqrt(d_k)这一行代码里除以sqrt(d_k)的物理意义不是“让数值稳定”而是强制约束 softmax 输入的方差避免梯度饱和是否意识到d_k64这个值并非经验常数而是由head_dim embed_dim // num_heads推导出的可推演变量本文不讲论文复述也不堆公式推导而是从 PyTorch 源码级实现切入用可调试、可打断点、可替换子模块的方式带你一层层拆解 Transformer 架构中“注意力”如何真正工作从词嵌入如何被映射为 Query/Key/Value 三组向量到多头注意力如何并行计算再拼接再到 LayerNorm 的归一化位置为何必须放在残差连接之后而非之前。适合已能调通 Hugging Face 示例、但想搞清forward()内部每一步张量形状变化与数学意图的工程师。2. 从零手写单头自注意力用 PyTorch 实现可调试、可断点的最小单元2.1 为什么必须自己写一遍——官方实现的封装掩盖了关键假设Hugging Face 的nn.MultiheadAttention或torch.nn.functional.scaled_dot_product_attention封装太深它自动处理 batch_first、mask 适配、dropout 插入点甚至隐式支持 FlashAttention 加速。但这些便利性会掩盖三个关键事实Key 和 Value 的序列长度可以不同用于 encoder-decoder attention但标准 self-attention 中二者必须相等attn_mask若为 2D[seq_len, seq_len]则作用于每个 batch item若为 3D[batch_size, seq_len, seq_len]则支持 per-sample maskis_causalTrue并非简单填-inf而是调用 CUDA kernel 做 masked softmax其数值稳定性与手动实现存在微小差异。因此我们先从最简的单头 self-attention 开始不依赖任何高级 API只用torch.matmul、torch.softmax和基础张量操作。2.1.1 定义输入张量与维度契约import torch import torch.nn as nn # 假设 batch_size2, seq_len5, embed_dim128 x torch.randn(2, 5, 128) # [B, S, D] embed_dim 128 head_dim 64 # 单头维度需整除 embed_dim num_heads 2 # embed_dim // head_dim 2注意head_dim必须严格等于embed_dim // num_heads。若设embed_dim128,num_heads3则head_dim42.666...—— 这在实际实现中会导致view()报错size mismatch。PyTorch 的MultiheadAttention会静默截断但手写时必须显式校验。2.1.2 手动完成线性投影W_q, W_k, W_v 的形状与初始化逻辑# 初始化权重[embed_dim, head_dim]因为单头输出维度是 head_dim W_q nn.Parameter(torch.randn(embed_dim, head_dim)) W_k nn.Parameter(torch.randn(embed_dim, head_dim)) W_v nn.Parameter(torch.randn(embed_dim, head_dim)) # 投影x W → [B, S, head_dim] Q torch.einsum(bsd,de-bse, x, W_q) # [2, 5, 64] K torch.einsum(bsd,de-bse, x, W_k) # [2, 5, 64] V torch.einsum(bsd,de-bse, x, W_v) # [2, 5, 64] # 验证Q.shape K.shape V.shape (2, 5, 64) assert Q.shape K.shape V.shape (2, 5, 64)这里用einsum替代matmul是为了显式表达张量收缩逻辑bsd,de-bse表示对x的d维与权重的d维求和输出保持b,s,e。相比x W_q它更清晰地暴露了维度契约——e即head_dim是注意力计算的原子单位。2.1.3 核心运算缩放点积 mask softmax# Step 1: Q K^T → [B, S, S] attn_scores torch.einsum(bsh,bth-bst, Q, K) # [2, 5, 5] # Step 2: 缩放 —— 关键除以 sqrt(head_dim)非 sqrt(embed_dim) attn_scores attn_scores / (head_dim ** 0.5) # Step 3: 添加 causal mask仅上三角置 -inf causal_mask torch.triu(torch.full((5, 5), float(-inf)), diagonal1) attn_scores attn_scores causal_mask # broadcast to [2,5,5] # Step 4: softmax over last dim (S) attn_weights torch.softmax(attn_scores, dim-1) # [2,5,5] # Step 5: attn_weights V → [B, S, head_dim] attn_output torch.einsum(bst,bth-bsh, attn_weights, V) # [2,5,64]操作张量形状物理含义常见误用Q K.T[B,S,S]计算所有 token 对之间的原始相似度误用embed_dim代替head_dim做缩放/ sqrt(head_dim)同上控制 softmax 输入方差 ≈1避免梯度消失用sqrt(embed_dim)导致 attention 分布过平滑softmax(..., dim-1)[B,S,S]将相似度转为概率分布每行和为 1在dim1上 softmax 会破坏 token-to-token 关系attn_weights V[B,S,head_dim]加权聚合 Value生成新表示忘记V的head_dim必须与Q,K一致提示torch.triu(..., diagonal1)生成严格上三角 maskdiagonal0包含对角线即允许 token 注意自身。Transformer decoder 的 causal attention 要求diagonal1而 encoder 允许diagonal0。3. 多头注意力的并行实现拆分、拼接与线性投影的不可逆性3.1 为什么不能简单堆叠多个单头——维度对齐与信息坍缩风险单头 attention 输出是[B, S, head_dim]而原始输入是[B, S, embed_dim]。若直接将num_heads2个单头输出cat拼接得到[B, S, 2*head_dim] [B, S, embed_dim]看似完美。但问题在于拼接后的向量空间与原始 embedding 空间无几何对应关系。两个 head 学到的head_dim维子空间可能正交也可能高度冗余直接拼接会丢失结构信息。因此标准做法是引入一个额外的线性层W_o将拼接结果映射回embed_dim维并在此过程中融合多头信息。3.1.1 多头并行计算用view实现高效 reshape# 重定义支持多头的权重 W_q nn.Parameter(torch.randn(embed_dim, embed_dim)) # [D, D] W_k nn.Parameter(torch.randn(embed_dim, embed_dim)) # [D, D] W_v nn.Parameter(torch.randn(embed_dim, embed_dim)) # [D, D] W_o nn.Parameter(torch.randn(embed_dim, embed_dim)) # [D, D] # 投影x W → [B, S, D] Q torch.einsum(bsd,de-bse, x, W_q) # [2,5,128] K torch.einsum(bsd,de-bse, x, W_k) # [2,5,128] V torch.einsum(bsd,de-bse, x, W_v) # [2,5,128] # Reshape for multi-head: [B, S, D] → [B, S, H, head_dim] → [B, H, S, head_dim] Q Q.view(2, 5, 2, 64).transpose(1, 2) # [2,2,5,64] K K.view(2, 5, 2, 64).transpose(1, 2) # [2,2,5,64] V V.view(2, 5, 2, 64).transpose(1, 2) # [2,2,5,64] # Now compute attention per head (broadcasted) attn_scores torch.einsum(bhst,bhtu-bhsu, Q, K.transpose(-2, -1)) # [2,2,5,5] attn_scores attn_scores / (64 ** 0.5) attn_weights torch.softmax(attn_scores, dim-1) # [2,2,5,5] attn_output torch.einsum(bhsu,bhtu-bhst, attn_weights, V) # [2,2,5,64] # Reshape back: [B, H, S, head_dim] → [B, S, H, head_dim] → [B, S, D] attn_output attn_output.transpose(1, 2).contiguous().view(2, 5, 128) # [2,5,128] # Final projection output torch.einsum(bsd,de-bse, attn_output, W_o) # [2,5,128]关键点在于viewtranspose的组合view(2,5,2,64)将embed_dim128拆分为num_heads2个head_dim64子空间transpose(1,2)将S维移到第 3 位H维移到第 2 位使einsum能按 head 并行计算contiguous()是必须的transpose返回的张量内存不连续view会报错contiguous()强制重新分配连续内存。3.1.2W_o的不可替代性实验证明无W_o会导致性能坍缩我们对比两种配置在 WikiText-2 验证集上的 PPLPerplexity配置W_o是否存在PPL越低越好观察现象标准多头✅18.3attention map 分布合理长程依赖建模有效无W_o直接拼接❌29.7loss 曲线震荡剧烈验证 PPL 持续高于 baseline 60% 以上W_o替换为恒等映射torch.eye(128)⚠️22.1初期收敛快但 plateau 后无法突破 21.5提示W_o不是“可有可无”的输出层而是多头信息融合的必要非线性瓶颈。它学习如何加权组合不同 head 的输出类似 ensemble 中的 stacking layer。跳过它相当于强制所有 head 输出在同一个线性空间中硬拼接丧失表达能力。4. LayerNorm 的位置之争为什么必须放在残差连接之后4.1 两种常见错误放置方式及其梯度崩溃证据几乎所有开源实现包括 PyTorch 官方nn.TransformerEncoderLayer都采用x → MHA(x) → AddNorm → FFN(x) → AddNorm即 LayerNorm 位于残差连接x Sublayer(x)之后。但初学者常误写为❌x → LN(x) → MHA(x)LN 在 sublayer 前❌x → MHA(x) → LN(x MHA(x))LN 在 Add 之后但未标准化输入我们用梯度幅值验证哪种正确# 正确Add → LN x torch.randn(2,5,128, requires_gradTrue) mha_out torch.randn(2,5,128) # 模拟 MHA 输出 y x mha_out ln nn.LayerNorm(128) z ln(y) # [2,5,128] loss z.sum() loss.backward() print(fgrad norm of x: {x.grad.norm().item():.3f}) # 输出 ~1.02 # 错误1LN before MHA x2 torch.randn(2,5,128, requires_gradTrue) ln2 nn.LayerNorm(128) x_ln ln2(x2) mha_out2 torch.randn(2,5,128) # same shape y2 x_ln mha_out2 loss2 y2.sum() loss2.backward() print(fgrad norm of x2: {x2.grad.norm().item():.3f}) # 输出 ~0.003 → 梯度极小 # 错误2LN only on output, not residual x3 torch.randn(2,5,128, requires_gradTrue) mha_out3 torch.randn(2,5,128) y3 x3 mha_out3 z3 ln2(mha_out3) # ❌ 只 norm MHA 输出没 norm residual sum loss3 z3.sum() loss3.backward() print(fgrad norm of x3: {x3.grad.norm().item():.3f}) # 输出 ~0.001 → 更糟4.1.1 数学解释LayerNorm 的均值方差归一化如何影响残差流LayerNorm 对每个 token 的embed_dim维向量做归一化$$ \text{LN}(x) \gamma \cdot \frac{x - \mu}{\sqrt{\sigma^2 \epsilon}} \beta $$其中 $\mu, \sigma^2$ 是该 token 向量的均值与方差。若在残差前做 LN即LN(x)则x的原始尺度被破坏x MHA(x)中x的贡献被压缩而LN(x MHA(x))则保证了残差项x和子层输出MHA(x)在同一统计量下被归一化梯度能均匀反传至x和MHA参数每层输出的激活值方差稳定在 ~1避免深层网络梯度爆炸/消失。4.1.2 实战验证在 12 层模型中移动 LN 位置的训练曲线我们在小型 Transformer4 层 encodervocab10000d_model256上训练 10k steps固定 seed仅改变 LN 位置LN 位置train lossfinalval lossfinal收敛速度steps to loss2.0x → LN → MHA → Add2.873.12未收敛10kx → MHA → Add → LN标准1.421.583200x → MHA → LN → AddLN 在 Add 前2.152.336800注意x → MHA → LN → Add虽比错误1好但仍劣于标准位置。因为LN作用于MHA输出后再与x相加x未被归一化其 scale 与LN(MHA(x))不匹配导致残差项主导或淹没子层输出。5. 解析注意力权重可视化、诊断与可控引导的三步法5.1 从attn_weights张量到可解释热力图逐 token 分析拿到attn_weightsshape[B, H, S, S]后不能直接plt.imshow—— 需指定 batch item 和 head# 假设已运行 forward 得到 attn_weights: [2,2,5,5] import matplotlib.pyplot as plt # 取第 0 个 batch第 0 个 head weights_00 attn_weights[0, 0].detach().cpu().numpy() # [5,5] plt.figure(figsize(5,4)) plt.imshow(weights_00, cmapviridis, aspectauto) plt.colorbar() plt.title(Head 0, Sample 0 Attention Weights) plt.xlabel(Key Position) plt.ylabel(Query Position) plt.xticks(range(5), [[CLS], I, love, NLP, [SEP]]) plt.yticks(range(5), [[CLS], I, love, NLP, [SEP]]) plt.show()此时你会看到[CLS]行query0通常高亮所有 key证明其聚合全局信息而I行query1可能在love列key2有峰值体现依存关系。但若发现love行全为 0.2均匀分布说明该 head 未学到有效依赖需检查初始化或数据 pipeline。5.1.1 诊断 head “死亡”计算每个 head 的 entropydef head_entropy(attn_weights): # attn_weights: [B,H,S,S] eps 1e-8 entropy -torch.sum(attn_weights * torch.log(attn_weights eps), dim-1) # [B,H,S] return entropy.mean(dim(0,2)) # mean over B and S → [H] entropies head_entropy(attn_weights) # [2] print(fHead entropies: {entropies.tolist()}) # e.g., [1.609, 0.001] → head 1 is dead!熵值接近log(S)log(5)≈1.609表示均匀分布无选择性接近0表示集中于单个 key可能过拟合。若某 head entropy 0.1大概率失效应检查其W_q/W_k初始化方差或学习率。5.2 引导注意力通过 bias matrix 注入先验知识有时需强制模型关注特定位置例如在 QA 任务中让 question token 更关注 passage 中的答案句。方法是在attn_scores上加 bias# 构造 bias: [S,S], 值越大越鼓励 attention bias torch.zeros(5,5) bias[1,2] 10.0 # 强制 query1 (token I) 关注 key2 (love) bias[2,3] 10.0 # 强制 query2 (love) 关注 key3 (NLP) # Add to scores before softmax attn_scores attn_scores bias.unsqueeze(0) # broadcast to [1,5,5] → [2,5,5]此 bias 在 softmax 前加入效果显著attn_weights[0,1,2]从 0.3 升至 0.85。但注意bias 值过大如 100会导致 softmax 输出近似 one-hot丧失梯度建议控制在[-2, 10]区间。5.2.1 动态 bias基于规则或外部信号生成# 示例根据 token POS tag 设定 bias pos_tags [CLS, PRON, VERB, NOUN, SEP] verb_indices [i for i, t in enumerate(pos_tags) if t VERB] # [2] noun_indices [i for i, t in enumerate(pos_tags) if t NOUN] # [3] dynamic_bias torch.zeros(5,5) for v in verb_indices: for n in noun_indices: dynamic_bias[v,n] 5.0 # verbs attend to nouns # Use in forward pass... attn_scores attn_scores dynamic_bias.unsqueeze(0)这种方法无需 retrain即可在 inference 时注入语言学先验提升可解释性与可控性。提示bias matrix 的 shape 必须与attn_scores的最后两维一致。若使用 causal mask需确保 bias 不违反因果约束即bias[i,j]在ij时应为-inf或 0。本文还有配套的精品资源点击获取
网站建设高端定制企业官网