尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

手写Transformer:从零实现注意力机制,彻底搞懂QKV原理

手写Transformer:从零实现注意力机制,彻底搞懂QKV原理 手写 Transformer从零实现注意力机制这次终于把 QKV 搞明白了如果你和我一样第一次接触 Transformer 时就被“注意力机制”四个字绕得云里雾里网上的代码一搜一大把但真正敢说自己“手写过”的并不多。大多数项目里我们直接调用nn.MultiheadAttention或者model.forward()一把梭等面试官问起“Q、K、V 到底怎么来的”“为什么要除以根号 d_k”时还是会卡壳。这篇文章就用一个完整的实战过程带你在 PyTorch 里手写注意力机制并在此基础上搭建一个可运行的 Transformer 编码器。整个过程不依赖高级封装你只需要了解基本的张量操作跟着代码一步步走就能彻底看清注意力机制从输入到输出的完整计算链路。文章适合谁适合刚学完 PyTorch 基础、想深入理解 Transformer 原理的初学者也适合已经在用nn.Transformer但想搞清楚内部机制的开发者。全文包含核心公式推导、完整代码、可视化思路、常见报错排查和工程建议建议收藏后边看边敲。1. 背景与核心概念注意力机制到底在做什么1.1 从 RNN 和 CNN 的痛点说起在 Transformer 之前处理序列数据的主流方案是 RNN 和 LSTM。RNN 的核心思路是“按顺序处理”也就是把序列看成一个时间步一个时间步地往后传递h_1 f(x_1, h_0) h_2 f(x_2, h_1) h_3 f(x_3, h_2) ...这种方式有两个明显问题长距离依赖难捕捉如果句子很长最前面的信息要经过很多时间步才能传到后面梯度容易消失位置靠前的单词对最终结果的影响会越来越弱。无法并行第 t 步必须等第 t-1 步算完计算效率很低这让大模型训练变得非常吃力。CNN 虽然可以并行但它更擅长捕捉局部特征需要通过堆叠很多层才能扩大感受野对序列中“跨位置的全局依赖”建模能力有限。注意力机制的出现就是为了解决“让模型直接关注序列中任意两个位置之间的关系”不需要按序传递也不需要一层层堆叠就能看到全局。1.2 注意力机制的本质加权求和从直觉上讲注意力机制就是一句话对信息做加权求和权重代表了当前关注点与不同信息之间的相关程度。举个例子翻译句子 “The animal didnt cross the street because it was too tired”我们要确定 “it” 指代的是 “animal” 还是 “street”就需要让模型在计算 “it” 这个位置的时候同时“看到”句子中其他单词并给 “animal” 一个更高的重要性权重。这个“重要性权重”就是注意力分数。从数学上看给定一个查询向量Query和一组键值对Key-Value注意力输出可以表示为Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V其中QQuery当前要查询的内容代表“我想找什么”。KKey被查询内容的索引代表“我有什么”。VValue实际被提取的信息代表“我能提供什么”。d_kQ和K的向量维度缩放因子防止点积数值过大。Q、K、V 这三个词很抽象下面我们用一张清晰的流程串起来。1.3 从 QKV 到自注意力自注意力Self-Attention是 Transformer 中最核心的注意力类型。所谓“自”指的是 Q、K、V 都来自同一个输入序列。每个 token 既当“查询者”也当“被查者”这样计算出来的注意力就表达了序列内部每个位置与其他所有位置之间的关系。具体流程如下输入序列X形状为[batch_size, seq_len, d_model]分别乘以三个权重矩阵W_Q、W_K、W_V得到Q、K、V。计算Q与所有K的点积得到原始注意力分数。将分数除以sqrt(d_k)进行缩放防止数值过大进入 softmax 的饱和区。对最后维度做 softmax归一化成概率分布。将归一化后的注意力权重与V相乘并求和得到加权后的输出。我把这个过程画成文字示意输入 X │ ├──乘 W_Q── Q ├──乘 W_K── K └──乘 W_V── V │ Q·K^T ──缩放── softmax ──× V── 输出理解了这一步Transformer 的“心脏”你就已经掌握了。2. 环境准备与项目结构2.1 环境版本说明本文示例使用 PyTorch 实现核心只需要torch和numpy。版本方面PyTorch 1.13 以上都可以正常运行示例代码不依赖最新 API。如果还没有安装可以用以下命令创建虚拟环境并安装依赖conda create -n transformer_diy python3.9 -y conda activate transformer_diy pip install torch numpy说明如果你的机器有 NVIDIA GPU可以按官网命令安装对应 CUDA 版本的 PyTorch。本文示例数据量极小CPU 上即可运行。2.2 项目文件结构为了让代码便于阅读和复现本次实战按以下结构组织文件transformer_diy/ ├── attention.py # 自注意力、多头注意力实现 ├── transformer_encoder.py # 位置编码、前馈网络、编码器层 ├── main.py # 最小可运行示例下面的内容中每个完整代码块上方都会标注对应的文件路径。3. 自注意力机制的从零实现3.1 不带掩码的缩放点积注意力我们先写最核心的缩放点积注意力函数。它接收 Q、K、V输出注意力权重和加权后的结果。# 文件路径attention.py import torch import torch.nn as nn import torch.nn.functional as F def scaled_dot_product_attention(query, key, value, maskNone): 缩放点积注意力。 参数 query: [batch_size, ..., seq_len_q, d_k] key: [batch_size, ..., seq_len_k, d_k] value: [batch_size, ..., seq_len_v, d_v] mask: [batch_size, 1, seq_len_q, seq_len_k] 或广播兼容形状 返回 output: [batch_size, ..., seq_len_q, d_v] attention_weights: [batch_size, ..., seq_len_q, seq_len_k] d_k query.size(-1) # 1. Q 和 K 做点积得到原始注意力分数 scores torch.matmul(query, key.transpose(-2, -1)) # 2. 缩放 scores scores / torch.sqrt(torch.tensor(d_k, dtypetorch.float32)) # 3. 掩码将被掩码位置的分数设为极小的负数 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # 4. softmax 归一化 attention_weights F.softmax(scores, dim-1) # 5. 与 V 加权求和 output torch.matmul(attention_weights, value) return output, attention_weights这里的mask参数很关键。在实际应用中有两个典型场景需要掩码Padding Mask把填充位置对应的注意力分数设为-inf让模型忽略填充符。Causal Mask因果掩码在解码器里禁止当前位置看到后面位置的信息。为什么用-inf而不是0因为softmax中-inf经过指数运算会变成0这样对应位置的注意力权重就是 0等价于完全忽略该位置。3.2 为什么除以 sqrt(d_k)这是一个高频面试点。在Q·K^T中如果d_k很大点积的数值会变得很大。假设Q和K的元素是独立随机变量且均值为 0、方差为 1那么点积的均值是 0方差是d_k。方差大意味着什么意味着某些维度的点积结果会异常大进入 softmax 后这些异常大的值对应位置的权重会接近 1而其他位置的权重会趋近于 0。这样一来梯度会变得非常小模型难以学习。除以sqrt(d_k)后点积的方差被控制在 1 左右softmax 的输入不至于过大梯度更稳定。这也是论文原文《Attention Is All You Need》里的标准做法。3.3 自注意力的完整封装上面的函数是通用计算核心我们需要把它封装成带可学习参数的模块。自注意力模块中Q、K、V 来自同一个输入X所以需要三组可学习的线性变换。# 文件路径attention.py续上方代码 class SelfAttention(nn.Module): 单头自注意力模块。 def __init__(self, d_model, d_kNone, d_vNone, dropout0.1): super().__init__() # 默认 Q、K、V 的维度都是 d_model self.d_k d_k if d_k is not None else d_model self.d_v d_v if d_v is not None else d_model self.w_q nn.Linear(d_model, self.d_k, biasFalse) self.w_k nn.Linear(d_model, self.d_k, biasFalse) self.w_v nn.Linear(d_model, self.d_v, biasFalse) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): 参数 x: [batch_size, seq_len, d_model] mask: 广播兼容形状 返回 output: [batch_size, seq_len, d_v] attention_weights: [batch_size, seq_len, seq_len] query self.w_q(x) key self.w_k(x) value self.w_v(x) output, attention_weights scaled_dot_product_attention( query, key, value, mask ) output self.dropout(output) return output, attention_weights这样我们就有了一个完整的单头自注意力模块。可以试一下前向传播import torch from attention import SelfAttention x torch.randn(2, 10, 512) # batch_size2, seq_len10, d_model512 attn SelfAttention(d_model512) out, weights attn(x) print(out.shape) # torch.Size([2, 10, 512]) print(weights.shape) # torch.Size([2, 10, 10])输出形状符合预期每个位置的输出融合了所有位置的信息注意力矩阵是 10×10。4. 多头注意力机制的实现4.1 为什么要多头单头注意力虽然能建模全局依赖但它只有一个“注意力模式”。例如某个头关注语法依赖另一个头关注语义相似还有一个头关注位置邻近。多个头各司其职模型表达能力更强。用专业的话说多头注意力把 Q、K、V 投影到多个低维子空间中在每个子空间独立计算注意力最后拼接起来再线性变换。这样模型可以在不同表示子空间上关注不同维度的信息。需要注意多头注意力并不是让每个头算出完整结果再平均而是把d_model维度切分成h份每个头处理d_k d_model / h维度的子向量。4.2 多头注意力完整实现实现多头注意力的最优雅方式是利用 PyTorch 张量 reshape transpose把[batch_size, seq_len, d_model]变成[batch_size, seq_len, h, d_k]再转置成[batch_size, h, seq_len, d_k]这样就可以把“多头”当成批次维度并行计算。# 文件路径attention.py续上方代码 class MultiHeadAttention(nn.Module): 多头注意力模块。 def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model 必须能被 num_heads 整除 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads self.w_q nn.Linear(d_model, d_model, biasFalse) self.w_k nn.Linear(d_model, d_model, biasFalse) self.w_v nn.Linear(d_model, d_model, biasFalse) self.out_proj nn.Linear(d_model, d_model, biasFalse) self.dropout nn.Dropout(dropout) def split_heads(self, x): 把最后一个维度拆成 num_heads 份并转置。 输入: [batch_size, seq_len, d_model] 输出: [batch_size, num_heads, seq_len, d_k] batch_size, seq_len, _ x.size() x x.view(batch_size, seq_len, self.num_heads, self.d_k) x x.transpose(1, 2) # [batch_size, num_heads, seq_len, d_k] return x def combine_heads(self, x): 将多头结果拼接回完整维度。 输入: [batch_size, num_heads, seq_len, d_k] 输出: [batch_size, seq_len, d_model] batch_size, _, seq_len, _ x.size() x x.transpose(1, 2).contiguous() x x.view(batch_size, seq_len, self.d_model) return x def forward(self, x, maskNone): # 1. 生成 Q、K、V 并拆成多头 query self.split_heads(self.w_q(x)) # [B, H, T, d_k] key self.split_heads(self.w_k(x)) value self.split_heads(self.w_v(x)) # 2. 调用缩放点积注意力 # mask 需要扩展成能和 [B, H, T, T] 广播的形状 if mask is not None: # 假设传入 mask 是 [B, T, T]需要变成 [B, 1, T, T] mask mask.unsqueeze(1) output, attention_weights scaled_dot_product_attention( query, key, value, mask ) # 3. 合并多头结果做输出投影 output self.combine_heads(output) output self.out_proj(output) output self.dropout(output) return output, attention_weights这里的核心是split_heads和combine_heads。很多初学者第一次接触多头注意力时会写循环去遍历每个头虽然也能实现但效率很低。用 reshape 的方式可以把所有头的计算一次性交给矩阵乘法完成这也体现出了 PyTorch 张量操作的优势。4.3 自测与对比写完后可以做个简单测试mha MultiHeadAttention(d_model512, num_heads8) out, weights mha(torch.randn(2, 10, 512)) print(out.shape) # torch.Size([2, 10, 512]) print(weights.shape) # torch.Size([2, 8, 10, 10])注意这里weights的形状比单头多了一维第 2 维是num_heads说明每个头都有独立的注意力矩阵。为了验证实现正确性我们可以和 PyTorch 官方的nn.MultiheadAttention做一次输出维度对比需要注意的是官方模块默认batch_firstFalse且输出顺序不同这里仅验证形状official nn.MultiheadAttention(embed_dim512, num_heads8, batch_firstTrue) official_out, official_weights official( torch.randn(2, 10, 512), torch.randn(2, 10, 512), torch.randn(2, 10, 512), ) print(official_out.shape) # torch.Size([2, 10, 512])从形状上看我们的实现与官方一致说明整体流程没有偏差。5. 从零搭建 Transformer 编码器有了多头注意力我们距离完整的 Transformer 只差三个组件位置编码、前馈神经网络、残差连接与层归一化。5.1 位置编码给序列引入顺序信息注意力机制本身对 token 的位置不敏感。如果你把两个 token 交换顺序计算出的 Q、K、V 点积结果是一样的。也就是说纯注意力模型根本不知道“谁在前、谁在后”这显然不行。Transformer 原文使用了正弦位置编码Sinusoidal Positional Encoding公式如下PE(pos, 2i) sin(pos / 10000^(2i / d_model)) PE(pos, 2i1) cos(pos / 10000^(2i / d_model))其中pos是位置下标i是维度下标。这个设计的巧妙之处在于不同位置的编码向量不同编码值在[-1, 1]之间利于训练稳定可以用正弦/余弦的加法性质表达相对位置关系。代码实现如下# 文件路径transformer_encoder.py import math import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(pdropout) # 生成 [max_len, d_model] 的位置编码矩阵 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # 维度下标序列 div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) # 偶数维用 sin奇数维用 cos pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) # 注册为 buffer不参与梯度更新但会随模型保存 pe pe.unsqueeze(0) # [1, max_len, d_model] self.register_buffer(pe, pe) def forward(self, x): 参数 x: [batch_size, seq_len, d_model] x x self.pe[:, :x.size(1), :] return self.dropout(x)使用位置编码时直接将它加在词向量/输入向量的上方。这个加法本身就是一个“注入位置信息”的操作位置编码的数值相对于输入嵌入来说不能太大否则会破坏原有的语义信息。5.2 前馈神经网络Transformer 编码器中的注意力输出会进入一个前馈网络Feed-Forward Network, FFN。它包含两个线性层和一个 ReLU 激活函数常被称为 Position-wise FFN因为它对每个位置独立操作FFN(x) max(0, x·W1 b1)·W2 b2中间隐藏层维度一般扩大 4 倍例如d_model512时d_ff2048。# 文件路径transformer_encoder.py续上方代码 class FeedForward(nn.Module): def __init__(self, d_model, d_ff2048, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.relu nn.ReLU() self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): x self.linear1(x) x self.relu(x) x self.dropout(x) x self.linear2(x) return x为什么每个注意力层后面要接 FFN一种解释是注意力机制本质上是“加权求和”它的变换是线性的准确说是凸组合。如果没有非线性激活层多个注意力层堆叠起来仍然近似线性变换模型的表达能力会受到很大限制。FFN 中的 ReLU 引入了非线性让模型能学到更复杂的特征。5.3 残差连接与层归一化Transformer 编码器中的每个子层都包含“残差连接 层归一化”。公式如下x LayerNorm(x Sublayer(x))残差连接让梯度可以跨层直接传播训练深层的 Transformer 时更加稳定。层归一化LayerNorm则作用于每个 token 的所有特征维度把数据分布拉回均值为 0、方差为 1 的状态加速收敛。5.4 完整 Transformer 编码器层把以上组件组合起来就是一个 Transformer Encoder Block# 文件路径transformer_encoder.py续上方代码 from attention import MultiHeadAttention class TransformerEncoderLayer(nn.Module): 单个 Transformer 编码器层。 def __init__(self, d_model, num_heads, d_ff2048, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.ffn FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x, maskNone): # 1. 子层1多头自注意力 残差 attn_output, _ self.self_attn(x, mask) x self.norm1(x self.dropout1(attn_output)) # 2. 子层2前馈网络 残差 ffn_output self.ffn(x) x self.norm2(x self.dropout2(ffn_output)) return x注意这里使用的是 Post-Norm 结构也就是“残差之后再做 LayerNorm”这是原始 Transformer 论文的做法。现在很多新模型用的是 Pre-Norm即“先归一化再进子层”训练更稳定。这个差异在后来被大量实验验证过不过作为入门我们先以原版为准。5.5 组合成完整的编码器并运行最后把位置编码和多个编码器层堆叠成完整编码器# 文件路径transformer_encoder.py续上方代码 class TransformerEncoder(nn.Module): def __init__(self, vocab_size, d_model, num_heads, num_layers, d_ff2048, dropout0.1, max_len5000): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.pos_encoding PositionalEncoding(d_model, max_len, dropout) self.layers nn.ModuleList([ TransformerEncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.d_model d_model def forward(self, token_ids, maskNone): 参数 token_ids: [batch_size, seq_len] mask: 扩展后可用于注意力的掩码 返回 x: [batch_size, seq_len, d_model] x self.embedding(token_ids) x x * math.sqrt(self.d_model) # 论文中的缩放操作 x self.pos_encoding(x) for layer in self.layers: x layer(x, mask) return x写一个最小 demo# 文件路径main.py import torch from transformer_encoder import TransformerEncoder # 构造一个简单的词表 vocab_size 100 d_model 64 num_heads 4 num_layers 2 model TransformerEncoder( vocab_sizevocab_size, d_modeld_model, num_headsnum_heads, num_layersnum_layers, d_ff128, dropout0.1, ) # 模拟一个 batchbatch_size2seq_len8 token_ids torch.randint(0, vocab_size, (2, 8)) output model(token_ids) print(输入形状:, token_ids.shape) # [2, 8] print(输出形状:, output.shape) # [2, 8, 64]运行后输出输入形状: torch.Size([2, 8]) 输出形状: torch.Size([2, 8, 64])到这里我们其实已经实现了一个可用的 Transformer 编码器。如果你加上一个线性分类头就可以用它做文本分类如果接一个 LM Head就能训练一个小型的语言模型。6. 注意力权重的可视化手写注意力机制还有一个很大的好处可以非常方便地拿到每个头的注意力权重。可视化注意力能帮我们直观理解“模型在看什么”。下面给出一段简单的可视化代码随机选一个头画出注意力矩阵热力图。假设你已经跑通了main.py我们可以把编码器第一层的注意力权重导出来。# 文件路径main.py续上方代码 import matplotlib.pyplot as plt def extract_attention(model, token_ids, layer_idx0): 提取指定层的注意力权重。 model.eval() with torch.no_grad(): x model.embedding(token_ids) x x * (model.d_model ** 0.5) x model.pos_encoding(x) attention_weights None for idx, layer in enumerate(model.layers): if idx layer_idx: _, attention_weights layer.self_attn(x) return attention_weights x layer(x) return attention_weights token_ids torch.randint(0, vocab_size, (1, 8)) attn_weights extract_attention(model, token_ids, layer_idx0) # 形状 [batch_size, num_heads, seq_len, seq_len] # 展示第 0 个样本、第 0 个头 head_idx 0 plt.figure(figsize(6, 6)) plt.imshow(attn_weights[0, head_idx].cpu().numpy(), cmapBlues) plt.colorbar() plt.title(fLayer 0, Head {head_idx} Attention) plt.xlabel(Key Position) plt.ylabel(Query Position) plt.show()可视化后你会看到不同的头关注模式差异很大。有的头呈明显的对角分布当前位置主要关注邻近位置有的头分散到整句的少数几个关键位置。这就是多头机制带来“多种关注模式”的直观体现。7. 常见问题与排查思路手写 Transformer 过程中最常踩的坑大部分集中在维度不匹配和 mask 上。下面整理了一份排查清单。问题现象常见原因解决思路mat1 and mat2 shapes cannot be multiplied线性层输入维度与d_model不一致检查nn.Embedding输出维度、输入 token 的最后一个维度是否为d_modelThe size of tensor a must match the size of tensor b多头注意力中 head 拆分后维度不对或者位置编码和输入 seq_len 不匹配检查split_heads和combine_heads中的 reshape/transpose 顺序注意力权重全变成 0 或全变成 1mask 使用不当或-inf的位置设置错误检查 mask 的每一维是否和[B, H, T, T]匹配padding 位置应为False/0训练 loss 下降极慢或直接 NaN学习率过大、未做梯度裁剪、LayerNorm 缺失调小学习率加 warmup检查是否所有输出都经过了 LayerNorm多头结果和官方nn.MultiheadAttention不一致QKV 投影方式、是否 bias、输出投影顺序不同先对比形状再关掉 dropout、初始化相同权重后对比数值排查维度问题时最实用的方法是print(shape)大法print(x:, x.shape) print(query:, query.shape) print(key:, key.shape) print(score:, scores.shape) print(attention_weights:, attention_weights.shape) print(output:, output.shape)不要觉得这很初级实际调模型时这类逐步打印往往比看报错信息更高效尤其是面对多层结构时。还有一个常见问题是mask的维度。在scaled_dot_product_attention中scores的形状是[B, H, T, T]因此mask需要能广播成这个形状。很多人在做多头时忘记给 mask 扩展维度导致什么报错都出现了还是找不到原因。建议统一写成if mask is not None: mask mask.unsqueeze(1) # [B, T, T] - [B, 1, T, T]或者传入前就确保 mask 是[B, 1, T, T]。8. 最佳实践与工程建议写完一个可运行的注意力机制只是第一步。如果要把它真正用到项目里下面这些建议值得提前了解。8.1 模块设计建议代码结构上把scaled_dot_product_attention单独抽出来是一个好习惯。这个函数是纯计算逻辑与模型参数无关既能被单头注意力调用也能被多头注意力调用单元测试起来非常方便。而且后面如果你想尝试 flash-attention 等高效实现只需要替换这个函数不需要改动整个模块。8.2 训练稳定性Transformer 对训练超参数比较敏感。实际项目里建议注意以下几点学习率调度先用 warmup把学习率从 0 缓慢升到峰值再按余弦退火降低。Transformer 原论文对 512 维模型使用 4000 步 warmup。Dropout 设置小模型建议0.1左右数据量很大时 dropout 可以适当调低。梯度裁剪把梯度的范数裁剪到1.0或者5.0可以有效避免训练中期出现的 loss 突增。Pre-Norm 更稳如果训练深层编码器经常不收敛可以考虑把 LayerNorm 移到子层之前也就是 Pre-LN 结构。8.3 性能优化自注意力的复杂度是O(n^2)n是序列长度。序列一长显存和耗时都会迅速上升。如果项目里需要处理长文本可以这样优化使用 PyTorch 2.0 的F.scaled_dot_product_attention它会自动选择最高效的内存融合实现。尝试稀疏注意力或窗口注意力只让局部 token 互相注意力。在训练和推理时使用torch.compile加速模型执行。8.4 对比官方实现与扩展方向手写实现最大的价值是“消除黑盒感”。等你能把代码跑通、能画出注意力热力图之后我推荐你再打开 PyTorch 官方的nn.TransformerEncoderLayer源码做一次对比看看官方在细节上做了哪些差异处理比如官方默认batch_firstFalse第一维是序列长度。官方支持src_key_padding_mask和attention_mask两类掩码。官方还提供了解码器nn.TransformerDecoderLayer支持交叉注意力Cross-Attention即 Q 来自解码器输入、K/V 来自编码器输出。交叉注意力是多头注意力的重要变体在机器翻译、图像描述、语音识别等场景中非常关键。看懂了本文的多头注意力交叉注意力只需要把SelfAttention的 Q 改成外部输入、K/V 保持从编码器输出生成即可原理完全一致。8.5 不要盲目手写最后说句实在的如果你只是在业务项目里用 Transformer完全不需要手写直接使用 PyTorch 官方封装更高效、更稳定。手写价值体现在下面几个地方面试前理解原理避免用“调包侠”人设被问倒。做研究/发论文时需要对注意力机制做自定义改造。学习阶段用一个小项目彻底打通原理到实现的路径。把本文代码跑通后你的下一步可以尝试在文本分类任务上训练一个完整的 Transformer 分类器给代码加上因果掩码实现一个自回归语言模型用自己实现的编码器替换nn.TransformerEncoder做对比实验尝试加入相对位置编码如 RoPE、ALiBi感受位置编码的演化。手写一次注意力机制远比调包十次学到的多。希望这篇文章能帮你真正跨过 Transformer 入门的门槛。如果过程中遇到报错可以照着上面的排查表格一步步找原因也欢迎在评论区交流你踩到的坑。
返回列表