注意力机制是 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.8+,PyTorch 1.10+,CPU 即可学习 |
| 批量任务 | 可在 batch 维度并行处理多句序列 |
注意力机制不是玄学,它是一个确定性的张量运算。所有“关注”行为都来自softmax之后得到的权重矩阵。把这一步吃透,Transformer 的其余部分基本就是“线性层 + 残差 + 归一化”的堆叠。
2. 手写注意力之前的数学基础
注意力机制的输入是三个向量组:查询 Query(Q)、键 Key(K)、值 Value(V)。
假设输入序列长度为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, dropout=0.1): super().__init__() self.d_k = d_k self.dropout = nn.Dropout(dropout) def forward(self, q, k, v, mask=None): # q @ k^T -> (batch_size, seq_len, seq_len) scores = torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtype=q.dtype, device=q.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)), atol=1e-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, dropout=0.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, dropout=dropout) self.dropout = nn.Dropout(dropout) def forward(self, q, k, v, mask=None): 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_len=512): 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 Block
class TransformerEncoderBlock(nn.Module): def __init__(self, d_model, h, d_ff, dropout=0.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, mask=None): # 自注意力子层 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, dropout=0.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_mask=None, padding_mask=None): # 被掩码的自注意力 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, mask=None): if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) max_val = scores.max(dim=-1, keepdim=True).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, keepdim=True)7.3 mask 类型:padding mask
实际训练时,序列会 padding 到相同长度。padding 位置没有意义,只在计算注意力 score 之前把这些位置对应的score设为负无穷。
def create_padding_mask(seq, pad_idx=0): # 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, dtype=torch.long)8.2 模型定义
class TinyTextClassifier(nn.Module): def __init__(self, vocab_size, d_model, h, d_ff, num_layers=2, num_classes=2, max_len=32): 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_size=vocab_size, d_model=d_model, h=h, d_ff=32, num_layers=2) optimizer = torch.optim.Adam(model.parameters(), lr=1e-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(f"epoch {epoch+1: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_idx=0, head_idx=0): 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, cmap='viridis') plt.colorbar() plt.title(f"Block {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_idx=0, head_idx=0)通过热度图可以看到,某些 head 会明显关注序列中的固定 token,这就验证了多头注意力捕捉不同模式的能力。
10. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 输出形状与预期不符 | .view()前张量不连续 | 打印每一步张量形状 | 在transpose后加.contiguous() |
| 注意力权重行和不为 1 | 手动实现 softmax 时未做稳定处理 | 检查归一化代码 | 使用F.softmax或减去最大值 |
| 训练 loss 不下降 | 未除以 sqrt(d_k) 或初始化不合适 | 打印 loss 和梯度统计 | 确认缩放因子、降低学习率 |
| 序列过长时显存爆炸 | 注意力矩阵 O(n^2) 太大 | 观察显存占用 | 使用稀疏注意力、窗口注意力或 FlashAttention |
| 有 padding 但未加 mask | padding 位置参与注意力 | 检查 mask 传入是否生效 | 创建 padding mask 并传入所有注意力层 |
| 因果生成时看到未来信息 | causal mask 未正确应用 | 打印 mask 矩阵 | 用torch.tril构建下三角掩码 |
| 多头输出不对 | 头维度和序列维度混在一起 | 固定 batch_size=1 对比手算结果 | 严格按照(batch, h, seq, d_k)组织张量 |
| CPU/GPU 结果不一致 | 存在未固定种子或非确定性算子 | 设置torch.manual_seed | torch.use_deterministic_algorithms(True) |
11. 在 Transformer 中实现注意力机制时的工程优化建议
如果只是学习,直接使用标准实现就够了。但如果要在项目中使用,下面的优化方向很重要。
11.1 使用 FlashAttention
2022 年后的大模型训练基本都使用 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_mask=None, dropout_p=0.1, is_causal=False)如果你在自己的项目里,建议直接调用内置版本,性能和显存表现都会好很多。
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], dim=2) v = torch.cat([self.cache[1], v], dim=2) 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 架构最值得花的时间。