
注意力机制是 Transformer 的核心组件也是当前大语言模型、多模态模型推理性能的瓶颈所在。与其反复阅读别人画的 QKV 示意图不如亲自动手把注意力机制写出来。这篇文章我直接带你从零实现一个可运行的缩放点积注意力、多头注意力并把它接进一个简化版 Transformer Block 里跑通训练测试。你会看到三样东西完整的 PyTorch 代码、每一步的张量形状变化、以及实际训练时的收敛表现。读完你就知道为什么Q K.T之后要除以sqrt(d_k)为什么多头能提升效果以及因果掩码和 KV Cache 到底在干什么。1. 注意力机制核心能力速览维度说明核心公式Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V关键参数Q、K、V、d_k、d_v、多头数 h主要作用让模型动态关注输入序列中相关性更高的位置典型形态自注意力、交叉注意力、因果自注意力实现难点维度变换、掩码、数值稳定性、批量计算适用场景Transformer、BERT、GPT、ViT、TTS、OCR 等运行环境Python 3.8PyTorch 1.10CPU 即可学习批量任务可在 batch 维度并行处理多句序列注意力机制不是玄学它是一个确定性的张量运算。所有“关注”行为都来自softmax之后得到的权重矩阵。把这一步吃透Transformer 的其余部分基本就是“线性层 残差 归一化”的堆叠。2. 手写注意力之前的数学基础注意力机制的输入是三个向量组查询 QueryQ、键 KeyK、值 ValueV。假设输入序列长度为n每个位置的特征维度是d_model。经过投影后得到Q维度(n, d_k)K维度(n, d_k)V维度(n, d_v)注意力的第一步是计算查询和所有键的点积相似度scores[i][j] Q[i] · K[j]Q[i]表示第i个查询K[j]表示第j个键。点积越大说明这两个位置的语义相关性越高。然后除以sqrt(d_k)。原因是当d_k很大时点积的方差也会变大导致softmax的梯度进入饱和区训练不稳定。除以sqrt(d_k)可以让方差恢复到 1 附近。最后用softmax归一化成注意力权重再对V做加权求和Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V这就是缩放点积注意力Scaled Dot-Product Attention。理解这个公式时最直观的方式是把自己想象成一个检索系统Q你想查什么。K资料库里的每份文档标题。V资料库里的每份文档正文。先用标题匹配相关度再把相关内容取出来。3. 环境准备与项目结构这篇文章的代码不需要特殊硬件。CPU 完全可以跑通所有示例。如果你有 NVIDIA GPU 且安装了对应版本的 PyTorch代码会自动调用 GPU。环境建议# 创建虚拟环境可选 python -m venv attn_env source attn_env/bin/activate # Windows 使用 attn_env\Scripts\activate # 安装依赖 pip install torch numpy matplotlib然后建立如下项目结构transformer-attention/ ├── attention.py # 自注意力、多头注意力实现 ├── transformer_block.py # Transformer Block 与简单训练验证 ├── test_attention.py # 单元测试 └── visualize.py # 注意力权重可视化建议所有代码保留纯 PyTorch 写法不引入额外高级封装方便你打断点查看张量形状。4. 从零实现缩放点积自注意力机制先实现最基础的单头缩放点积注意力。4.1 基础版本代码import torch import torch.nn as nn import torch.nn.functional as F class ScaledDotProductAttention(nn.Module): 缩放点积注意力 输入 Q, K, V: Q: (batch_size, seq_len, d_k) K: (batch_size, seq_len, d_k) V: (batch_size, seq_len, d_v) mask: 可选形状 (batch_size, seq_len, seq_len) 或广播为相同形状 def __init__(self, d_k, dropout0.1): super().__init__() self.d_k d_k self.dropout nn.Dropout(dropout) def forward(self, q, k, v, maskNone): # q k^T - (batch_size, seq_len, seq_len) scores torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtypeq.dtype, deviceq.device)) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) output torch.matmul(attn_weights, v) # (batch_size, seq_len, d_v) return output, attn_weights关键点k.transpose(-2, -1)把 K 的最后两个维度交换才能让 Q 的序列维度和 K 的序列维度做点积。masked_fill(mask 0, float(-inf))让被掩码的位置在 softmax 后权重为 0。返回attn_weights后面可视化会用到。4.2 测试基础注意力def test_scaled_dot_product_attention(): batch_size 2 seq_len 4 d_k 8 d_v 8 q torch.randn(batch_size, seq_len, d_k) k torch.randn(batch_size, seq_len, d_k) v torch.randn(batch_size, seq_len, d_v) attention ScaledDotProductAttention(d_k) output, weights attention(q, k, v) print(output shape:, output.shape) print(weights shape:, weights.shape) # 权重每一行之和应该等于 1 assert torch.allclose(weights.sum(dim-1), torch.ones_like(weights.sum(dim-1)), atol1e-6) print(Attention weights normalized correctly.) if __name__ __main__: test_scaled_dot_product_attention()如果你看到output shape: torch.Size([2, 4, 8]) weights shape: torch.Size([2, 4, 4]) Attention weights normalized correctly.说明最基本的注意力计算已经跑通了。5. 实现多头注意力机制Multi-Head Attention单头注意力只能捕获一种“关注模式”。多头注意力把d_model维度的特征切分成h个子空间在每个子空间独立做注意力最后拼接起来。这样模型可以同时关注语法关系、指代关系、语义相似度等不同层面的信息。5.1 多头注意力代码class MultiHeadAttention(nn.Module): 多头注意力 d_model: 输入特征维度 h: 注意力头数 d_k, d_v 可以手动指定默认 d_model // h def __init__(self, d_model, h, dropout0.1): super().__init__() assert d_model % h 0, d_model must be divisible by h self.d_model d_model self.h h self.d_k d_model // h self.d_v d_model // h self.w_q nn.Linear(d_model, d_model) self.w_k nn.Linear(d_model, d_model) self.w_v nn.Linear(d_model, d_model) self.w_o nn.Linear(d_model, d_model) self.attention ScaledDotProductAttention(self.d_k, dropoutdropout) self.dropout nn.Dropout(dropout) def forward(self, q, k, v, maskNone): batch_size, seq_len, _ q.size() # 线性投影后分割成 h 个头 Q self.w_q(q).view(batch_size, seq_len, self.h, self.d_k) K self.w_k(k).view(batch_size, seq_len, self.h, self.d_k) V self.w_v(v).view(batch_size, seq_len, self.h, self.d_v) # 把 (batch, seq, h, d_k) 转换成 (batch, h, seq, d_k) Q Q.transpose(1, 2) K K.transpose(1, 2) V V.transpose(1, 2) # 多头注意力 attn_output, attn_weights self.attention(Q, K, V, mask) # attn_output: (batch_size, h, seq_len, d_v) # 把多头结果拼回去 attn_output attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) # 最后的输出投影 output self.w_o(attn_output) return output, attn_weights这里的维度变化是初学者最容易踩坑的地方。我们拆开看输入 Q: (batch_size, seq_len, d_model) 经过 w_q: (batch_size, seq_len, d_model) .view(batch_size, seq_len, h, d_k): 拆出“头”维度 .transpose(1, 2): (batch_size, h, seq_len, d_k)为什么要把头的维度放到第二维因为torch.matmul要求最后两个维度参与矩阵乘法把头的维度放前面就可以让batch_size和h一起参与批量并行计算。5.2 多头注意力测试def test_multi_head_attention(): batch_size 2 seq_len 6 d_model 12 h 4 x torch.randn(batch_size, seq_len, d_model) mha MultiHeadAttention(d_model, h) output, weights mha(x, x, x) print(MHA output shape:, output.shape) print(MHA weights shape:, weights.shape) assert output.shape (batch_size, seq_len, d_model) assert weights.shape (batch_size, h, seq_len, seq_len) print(Multi-Head Attention works correctly.) if __name__ __main__: test_multi_head_attention()运行后MHA output shape: torch.Size([2, 6, 12]) MHA weights shape: torch.Size([2, 4, 6, 6])这就说明多头注意力已经能够正确输出。6. 在 Transformer 中实现注意力机制完整 Block 搭建Transformer 的完整结构包含编码器和解码器。编码器里的核心是“多头自注意力 前馈网络”解码器里则是“掩码多头自注意力 交叉注意力 前馈网络”。这里我们用一个小型 Transformer Block 来验证注意力机制的真实效果。为了让代码可运行、可训练我直接写一个极简版本。6.1 位置编码注意力机制本身没有顺序信息所以必须把位置信息加到输入里。这里使用最常见的正弦位置编码。class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len512): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp(torch.arange(0, d_model, 2).float() * (-torch.log(torch.tensor(10000.0)) / d_model)) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) pe pe.unsqueeze(0) # (1, max_len, d_model) self.register_buffer(pe, pe) def forward(self, x): return x self.pe[:, :x.size(1)]6.2 Transformer Encoder Blockclass TransformerEncoderBlock(nn.Module): def __init__(self, d_model, h, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, h, dropout) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model), ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 自注意力子层 attn_out, _ self.self_attn(x, x, x, mask) x self.norm1(x self.dropout(attn_out)) # 前馈网络子层 ffn_out self.ffn(x) x self.norm2(x self.dropout(ffn_out)) return x这就是 Transformer 的基本残差结构注意力输出过 dropout再加到原输入上最后做 LayerNorm。6.3 带掩码的 Decoder Block含交叉注意力解码器比编码器多一个交叉注意力层并且第一个多头自注意力必须使用因果掩码。def build_causal_mask(seq_len): 生成下三角全 1 的因果掩码矩阵 mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask class TransformerDecoderBlock(nn.Module): def __init__(self, d_model, h, d_ff, dropout0.1): super().__init__() self.masked_self_attn MultiHeadAttention(d_model, h, dropout) self.cross_attn MultiHeadAttention(d_model, h, dropout) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model), ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.norm3 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, encoder_output, causal_maskNone, padding_maskNone): # 被掩码的自注意力 attn_out, _ self.masked_self_attn(x, x, x, causal_mask) x self.norm1(x self.dropout(attn_out)) # 交叉注意力Q 来自解码器K、V 来自编码器 cross_out, _ self.cross_attn(x, encoder_output, encoder_output, padding_mask) x self.norm2(x self.dropout(cross_out)) ffn_out self.ffn(x) x self.norm3(x self.dropout(ffn_out)) return x交叉注意力是理解机器翻译、摘要生成等任务的关键Q 来自当前解码器状态表示“现在需要参考哪个源语言信息”。K、V 来自编码器输出表示“所有源语言特征”。7. 在 Transformer 中实现注意力机制时的关键细节掩码与数值稳定性注意力机制的坑不在公式本身而在工程细节上。7.1 为什么需要因果掩码语言模型生成第i个词时不应该看到第i个词之后的内容。因果掩码把未来位置的 score 设为-inf这样 softmax 后这些位置的权重为 0。seq_len 5 mask build_causal_mask(seq_len) print(mask)输出tensor([[ True, False, False, False, False], [ True, True, False, False, False], [ True, True, True, False, False], [ True, True, True, True, False], [ True, True, True, True, True]])7.2 数值稳定性处理当d_k较大或者 score 整体偏大时softmax内部的exp可能溢出。PyTorch 的F.softmax已经做了减去最大值处理但如果你要自定义实现一定要用这种稳定版本def stable_softmax(scores, maskNone): if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) max_val scores.max(dim-1, keepdimTrue).values scores scores - max_val exp_scores torch.exp(scores) if mask is not None: exp_scores exp_scores.masked_fill(mask 0, 0) return exp_scores / exp_scores.sum(dim-1, keepdimTrue)7.3 mask 类型padding mask实际训练时序列会 padding 到相同长度。padding 位置没有意义只在计算注意力 score 之前把这些位置对应的score设为负无穷。def create_padding_mask(seq, pad_idx0): # seq: (batch_size, seq_len) return (seq ! pad_idx).unsqueeze(1).unsqueeze(2) # (batch, 1, 1, seq_len)这个 mask 会在多头注意力里自动广播到(batch, h, seq_len, seq_len)。8. 功能测试与效果验证训练一个极简情感分类器光实现还不够必须用一个可收敛的小任务验证。下面我构造一个简单数据集用刚才实现的 Transformer Encoder Block 做情感分类观察 loss 能否下降。8.1 数据构造torch.manual_seed(42) # 简单样本词向量用随机的 one-hot 索引代替 vocab_size 20 d_model 16 h 4 seq_len 8 num_samples 200 def generate_data(num_samples, seq_len): 构造一个可学习规律的数据 如果序列中有 token 3 和 token 7则标签为 1否则为 0。 这样模型必须学会根据特定 token 的位置和内容做判断。 xs [] ys [] for _ in range(num_samples): x torch.randint(1, vocab_size, (seq_len,)) label 1 if (3 in x and 7 in x) else 0 xs.append(x) ys.append(label) return torch.stack(xs), torch.tensor(ys, dtypetorch.long)8.2 模型定义class TinyTextClassifier(nn.Module): def __init__(self, vocab_size, d_model, h, d_ff, num_layers2, num_classes2, max_len32): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.pos_encoding PositionalEncoding(d_model, max_len) self.encoder_blocks nn.ModuleList([ TransformerEncoderBlock(d_model, h, d_ff) for _ in range(num_layers) ]) self.classifier nn.Linear(d_model, num_classes) def forward(self, x): x self.embedding(x) # (batch, seq, d_model) x self.pos_encoding(x) for block in self.encoder_blocks: x block(x) # 取序列第一个 token 的表示做分类 pooled x[:, 0, :] return self.classifier(pooled)8.3 训练并观察注意力机制的效果def train_classifier(): x_train, y_train generate_data(200, seq_len) x_val, y_val generate_data(50, seq_len) model TinyTextClassifier(vocab_sizevocab_size, d_modeld_model, hh, d_ff32, num_layers2) optimizer torch.optim.Adam(model.parameters(), lr1e-3) loss_fn nn.CrossEntropyLoss() model.train() for epoch in range(50): optimizer.zero_grad() logits model(x_train) loss loss_fn(logits, y_train) loss.backward() optimizer.step() if (epoch 1) % 10 0: pred torch.argmax(model(x_train), dim-1) acc (pred y_train).float().mean() val_pred torch.argmax(model(x_val), dim-1) val_acc (val_pred y_val).float().mean() print(fepoch {epoch1:3d} | loss {loss.item():.4f} | train acc {acc:.2f} | val acc {val_acc:.2f}) if __name__ __main__: train_classifier()输出大致如下epoch 10 | loss 0.6001 | train acc 0.69 | val acc 0.68 epoch 20 | loss 0.4020 | train acc 0.87 | val acc 0.84 epoch 30 | loss 0.2518 | train acc 0.95 | val acc 0.92 epoch 40 | loss 0.1634 | train acc 0.98 | val acc 0.94 epoch 50 | loss 0.1162 | train acc 1.00 | val acc 0.96这说明注意力机制确实学到了数据中的规律。你可以把num_layers增加到 4、把d_model增加到 32继续观察效果。9. 注意力权重的可视化与解读训练完成后我们可以把注意力权重打印成热度图直接观察模型关注了哪些 token。import matplotlib.pyplot as plt def visualize_attention(model, x, block_idx0, head_idx0): model.eval() x_emb model.embedding(x) x_emb model.pos_encoding(x_emb) # 手动执行 encoder block 并捕获 attention weights attn_weights None for i, block in enumerate(model.encoder_blocks): if i block_idx: _, attn_weights block.self_attn(x_emb, x_emb, x_emb) break else: x_emb block(x_emb) # attn_weights: (batch, h, seq_len, seq_len) attn attn_weights[0, head_idx].detach().numpy() plt.figure(figsize(6, 6)) plt.imshow(attn, cmapviridis) plt.colorbar() plt.title(fBlock {block_idx} Head {head_idx} Attention) plt.xlabel(Key Position) plt.ylabel(Query Position) plt.tight_layout() plt.show() # 使用一个样本可视化 x_sample, _ generate_data(1, seq_len) visualize_attention(model, x_sample, block_idx0, head_idx0)通过热度图可以看到某些 head 会明显关注序列中的固定 token这就验证了多头注意力捕捉不同模式的能力。10. 常见问题与排查方法问题现象可能原因排查方式解决方案输出形状与预期不符.view()前张量不连续打印每一步张量形状在transpose后加.contiguous()注意力权重行和不为 1手动实现 softmax 时未做稳定处理检查归一化代码使用F.softmax或减去最大值训练 loss 不下降未除以 sqrt(d_k) 或初始化不合适打印 loss 和梯度统计确认缩放因子、降低学习率序列过长时显存爆炸注意力矩阵 O(n^2) 太大观察显存占用使用稀疏注意力、窗口注意力或 FlashAttention有 padding 但未加 maskpadding 位置参与注意力检查 mask 传入是否生效创建 padding mask 并传入所有注意力层因果生成时看到未来信息causal mask 未正确应用打印 mask 矩阵用torch.tril构建下三角掩码多头输出不对头维度和序列维度混在一起固定 batch_size1 对比手算结果严格按照(batch, h, seq, d_k)组织张量CPU/GPU 结果不一致存在未固定种子或非确定性算子设置torch.manual_seedtorch.use_deterministic_algorithms(True)11. 在 Transformer 中实现注意力机制时的工程优化建议如果只是学习直接使用标准实现就够了。但如果要在项目中使用下面的优化方向很重要。11.1 使用 FlashAttention2022 年后的大模型训练基本都使用 FlashAttention。它通过分块计算和重计算减少显存占用同时避免显式构造完整的QK^T矩阵。PyTorch 2.0 内置了torch.nn.functional.scaled_dot_product_attention底层会自动选择 FlashAttention 或 memory-efficient attention。# PyTorch 2.x 内置实现 output F.scaled_dot_product_attention(q, k, v, attn_maskNone, dropout_p0.1, is_causalFalse)如果你在自己的项目里建议直接调用内置版本性能和显存表现都会好很多。11.2 引入 KV Cache 加速推理推理时每生成一个 token不需要重新计算所有历史 token 的 K、V。把 K、V 缓存下来只让新 token 的 Q 与历史 K、V 做注意力可以把自回归生成从 O(n^2) 降到接近 O(n)忽略缓存增长时的拷贝开销。class KVCache: def __init__(self): self.cache None def update(self, k, v): # k, v: (batch, head, seq, d_k) if self.cache is None: self.cache (k, v) else: k torch.cat([self.cache[0], k], dim2) v torch.cat([self.cache[1], v], dim2) self.cache (k, v) return self.cache这虽然是一个简化版缓存逻辑但已经是推理优化里最核心的部分。11.3 分批推理与批量任务注意力机制天然支持 batch 并行。如果要做批量任务例如一次性给一批句子生成回复可以凑成一个 batch 输入但要注意 padding 的 mask 处理。批量过大时也要留意显存占用因为(batch, head, seq, seq)的注意力矩阵会随 batch 线性增长。12. 最佳实践从实现到落地如果你已经把这篇文章里的代码亲手跑完注意力机制就不再是抽象的图。建议你在自己的项目里按下面的顺序推进第一遍对照代码把每个张量的形状标在注释里运行单元测试。第二遍改用F.scaled_dot_product_attention对比手写版本输出是否一致。第三遍实现 KV Cache测试自回归生成。第四遍加入 padding mask 和 causal mask处理真实 batch 训练。第五遍用真实文本数据替换 toy dataset接入小型 Transformer。无论你后续使用 BERT、GPT、ViT 还是各类多模态模型核心都离不开这篇文章实现的这套张量运算。下次看到“注意力机制”四个字你的第一反应应该是QK^T / sqrt(d_k)和softmax后的加权求和而不是一张看不懂的示意图。把注意力机制手写一遍是你理解整个 Transformer 架构最值得花的时间。