从零手写注意力机制:PyTorch实现Transformer核心
2026/8/31 8:06:32 网站建设 项目流程

注意力机制是 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_sizeh一起参与批量并行计算。

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 但未加 maskpadding 位置参与注意力检查 mask 传入是否生效创建 padding mask 并传入所有注意力层
因果生成时看到未来信息causal mask 未正确应用打印 mask 矩阵torch.tril构建下三角掩码
多头输出不对头维度和序列维度混在一起固定 batch_size=1 对比手算结果严格按照(batch, h, seq, d_k)组织张量
CPU/GPU 结果不一致存在未固定种子或非确定性算子设置torch.manual_seedtorch.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 架构最值得花的时间。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询