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

资讯详情

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

从零实现Transformer注意力机制:多头、掩码与位置编码全解析

从零实现Transformer注意力机制:多头、掩码与位置编码全解析 注意力机制是 Transformer 架构中最核心的组件。无论看 BERT、GPT、ViT 还是 Swin Transformer底层反复运行的几乎都是同一个模块多头自注意力。很多初学者把注意力机制理解成“给重要的词更大的权重”这句话方向没错但距离真正在 PyTorch 里写出可运行代码还有很长的路Q、K、V 是怎么生成的为什么要除以根号 d_kpadding mask 和 causal mask 加在哪个位置多头是怎么拆开又合并的只看论文很难把这些细节彻底搞清楚。这篇文章的目标很直接从零实现 Transformer 中的注意力机制覆盖 Scaled Dot-Product Attention、多头注意力、两种掩码机制、位置编码并组装成一个最小可运行的 Transformer Block。代码会逐段解释最后给出维度验证、注意力可视化、常见报错排查和生产环境建议。读完以后你至少能回答上面几个问题也能看懂常见开源实现里的核心代码包括各类“手撕 Transformer”教程中反复出现的那几段关键逻辑。1. 先搞清楚自注意力、多头和 QKV 之间的关系1.1 自注意力解决的核心问题RNN 处理序列时输出是逐步产生的当前时刻只能看到左侧信息长距离信息需要靠隐藏状态一步步传递。这个方式有两个明显短板一是长距离依赖容易衰减早期信息经过多个时间步后变得很弱二是无法并行序列必须严格按顺序计算。自注意力做的事情完全不同它让序列中任意两个位置之间可以直接交互。处理第 i 个位置时模型可以同时“看到”序列中所有位置并根据它们与当前位置的相关程度决定从每个位置取多少信息。对应到 Transformer 里自注意力的输入是一整段序列输出是长度不变、但每个位置都聚合了全序列信息的新序列。这种机制带来两个直接好处长距离依赖不再依赖梯度传播链整个序列可以并行计算。这也是 Transformer 取代 RNN 成为主流序列建模工具的根本原因之一。1.2 QKV注意力机制的三种角色注意力计算依赖三个矩阵Query查询、Key键、Value值。可以打一个直观的比方Query 表示“我在找什么”。Key 表示“我身上有什么标签”。Value 表示“我实际能提供什么内容”。计算过程就是拿每个 Query 去和所有 Key 做匹配匹配度越高对应 Value 的权重就越大最后按权重把所有 Value 加权求和。整个过程可以用一句话概括输出是值的加权平均权重由查询和键的相似度决定。在自注意力场景里Q、K、V 都来自同一个输入序列只是经过三个不同的线性变换所以叫 self-attention。为什么要做线性变换因为原始输入向量如果直接互相对比表达能力太有限。经过可学习的 W_q、W_k、W_v 之后模型可以学到不同的投影空间让“查询”和“键”的匹配更适合当前任务。这三组权重是整个注意力机制里真正的可学习参数。1.3 多头注意力是自注意力的工程化扩展单头注意力虽然能建模两两关系但只有一个投影空间难以同时关注多种关系。比如一个词在语法层面和语义层面可能关联到不同的词单靠一个注意力头很难兼顾。多头注意力把特征维度切成 h 份每一份独立做一次注意力最后拼接后再经过一个线性层。每个头可以学习不同的关系类型模型表达能力因此增强。这里有一个容易混淆的点多头注意力不是把序列切成多段而是把特征维度切成了多份。序列长度从头到尾没有变化变化的是每个位置的向量被拆分到多个头里并行计算。理解这一点后面看维度变换时就不会被绕晕。2. 从零实现 Scaled Dot-Product Attention2.1 环境准备实现注意力机制只需要 PyTorch不需要额外复杂的依赖。推荐环境如下组件推荐配置说明Python3.9 或更高使用 f-string 和类型标注更顺手PyTorch2.x本文代码基于 PyTorch 2.x 验证NumPy随 PyTorch 自动安装主要用于调试Matplotlib可选用于注意力权重可视化如果使用 GPU需要按官方文档选择对应 CUDA 版本的安装命令。如果只是学习注意力机制的代码结构CPU 版本完全足够本文示例的 batch size 和序列长度都很小。验证环境是否可用python -c import torch; print(torch.__version__)能正常输出版本号说明环境没问题。2.2 最小注意力函数实现注意力机制的数学表达如下Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V其中 Q 的形状是 (batch, seq_len, d_k)K 和 V 的形状是 (batch, source_len, d_k)。先实现最核心的函数import math import torch import torch.nn as nn import torch.nn.functional as F def scaled_dot_product_attention(query, key, value, maskNone, dropoutNone): d_k query.size(-1) # scores 形状: (batch, ..., seq_len_q, seq_len_k) scores torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) if dropout is not None: attn_weights dropout(attn_weights) output torch.matmul(attn_weights, value) return output, attn_weights这段代码需要解释几个关键点。第一key.transpose(-2, -1)是把 Key 的最后两个维度转置因为矩阵乘法要求 Q 的最后一维和 K 的倒数第二维匹配。Q 的形状是 (batch, seq_len_q, d_k)转置后的 K 形状是 (batch, d_k, seq_len_k)乘出来的 scores 形状是 (batch, seq_len_q, seq_len_k)。第二除以sqrt(d_k)是缩放操作。当 d_k 较大时点积结果会偏大进入 softmax 的饱和区间梯度会非常小。除以根号 d_k 可以把点积的方差压回 1 附近让 softmax 的梯度更稳定。这是原论文特意强调的设计不能省。这里用math.sqrt(d_k)而不是torch.sqrt(torch.tensor(d_k))是为了避免额外创建 CPU 标量张量再和 GPU 张量混用。第三mask 的加法和 softmax 之间的顺序很关键。mask 加在 softmax 之前被 mask 掉的位置被替换成负无穷经过 softmax 后权重会趋近于 0。如果加在 softmax 之后权重的归一化效果就被破坏了。后文会单独讲两种 mask 的形状。2.3 两种掩码padding mask 和 causal mask实际训练时一个 batch 里的序列长度通常不一致需要把短序列填充到相同长度。填充位置是无效信息不应该参与注意力计算于是需要 padding mask。它的形状是 (batch, 1, 1, seq_len_k)值为 1 表示有效位置0 表示填充位置def make_padding_mask(seq, pad_idx0): # seq 形状: (batch, seq_len) mask (seq ! pad_idx).unsqueeze(1).unsqueeze(2) return mask生成的 mask 形状是 (batch, 1, 1, seq_len)在广播机制下会作用到所有 query 位置上。解码器里还有另一种 maskcausal mask也叫 look-ahead mask。生成第 i 个位置的输出时不能看到 i 之后的位置否则训练时模型会“偷看答案”。causal mask 是一个上三角为 0、下三角为 1 的矩阵def make_causal_mask(seq_len): mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask两个 mask 可以同时使用比如解码器训练时先做 causal mask再做 padding mask。推荐把两个 mask 提前用逻辑与合并减少重复计算。掩码最容易出错的地方是维度形状后面排查章节会详细讲。注意编码器只需要 padding mask解码器需要同时使用 causal mask 和 padding mask。两者的作用不同不要混用。3. 实现多头注意力view、transpose 和维度检查3.1 多头拆分的维度变换先约定参数d_model 是整个模型的隐藏维度num_heads 是头数每个头的维度 d_k d_model / num_heads。原论文使用 d_model512、num_heads8小 demo 里可以用 d_model64、num_heads4。多头拆分的核心是 view transpose。假设输入 x 的形状是 (batch, seq_len, d_model)经过线性变换得到 Q 后需要先 view 成 (batch, seq_len, num_heads, d_k)再 transpose 成 (batch, num_heads, seq_len, d_k)。这一步是初学者最容易写错的地方# 正确流程先 view 拆维度再 transpose 换轴 batch_size query.size(0) Q self.w_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)为什么要先 view 再 transpose因为 Linear 输出的最后一个维度是 d_modelview 负责把它按顺序切成 num_heads 块transpose 再把头维度从第 2 位换到第 1 位让注意力计算时 batch 和 head 都在最外层方便批量并行。注意transpose 之后内存布局不连续直接 view 会报错或得到错误结果。正确做法是先调用.contiguous()再 view。这个坑在串联多个头时非常常见。3.2 完整 MultiHeadAttention 类接下来是完整的多头注意力模块。把缩放点积注意力抽成函数多头模块负责线性变换、维度拆分和最后拼接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) 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.dropout nn.Dropout(dropout) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 线性投影后拆头 Q self.w_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K self.w_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V self.w_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) attn_output, attn_weights scaled_dot_product_attention( Q, K, V, maskmask, dropoutself.dropout ) # 拼接多头先换回轴顺序再用 view 合并最后两维 attn_output attn_output.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) output self.w_o(attn_output) return output, attn_weightsforward 里 query 和 key 通常来自同一序列也就是自注意力。但交叉注意力cross-attention也使用同一个类query 来自解码器key 和 value 来自编码器输出只需要在调用时传入不同的张量即可。这就是把 query、key、value 分开传入的意义。3.3 参数的维度检查写完模块后先跑一个最小用例确认输出形状符合预期torch.manual_seed(42) batch_size, seq_len, d_model, num_heads 2, 10, 64, 4 x torch.randn(batch_size, seq_len, d_model) mha MultiHeadAttention(d_model, num_heads) out, attn mha(x, x, x) print(输入形状:, x.shape) print(输出形状:, out.shape) print(注意力权重形状:, attn.shape)预期输出如下输入形状: torch.Size([2, 10, 64]) 输出形状: torch.Size([2, 10, 64]) 注意力权重形状: torch.Size([2, 4, 10, 10])看到注意力权重是 (batch, num_heads, seq_len, seq_len)说明多头拆分和拼接是正确的。如果这里维度不对后面的 Transformer Block 一定跑不通。4. 从注意力拼出最小 Transformer Block4.1 位置编码自注意力缺少顺序信息自注意力对序列中每个位置的计算方式完全相同位置 A 到位置 B 的注意力分数和位置 B 到位置 A 的注意力分数只取决于内容的相似度与它们的先后顺序无关。也就是说自注意力本身是置换不变的。为了让模型感知顺序必须把位置信息加到输入里。最简单也最稳定的方式是原文的正弦位置编码。每个位置的向量和词向量相加后送入编码器。实现如下def sinusoidal_positional_encoding(seq_len, d_model): pe torch.zeros(seq_len, d_model) position torch.arange(0, seq_len, dtypetorch.float).unsqueeze(1) div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) return pe.unsqueeze(0) # (1, seq_len, d_model)现在也有不少实现使用可学习位置编码直接声明nn.Parameter(torch.randn(1, max_len, d_model))。两者没有绝对优劣正弦编码泛化到更长序列的能力更好可学习编码在固定长度上更灵活。小项目里选哪一种都行关键是记得加别漏。4.2 前馈网络、残差连接和 LayerNorm多头注意力输出之后Transformer 的每个 Block 还要做三件事前馈网络、残差连接、LayerNorm。前馈网络是一个两层的 MLP先升维再降维中间用 ReLU 或 GELU 激活。原论文里隐藏维度是 d_model 的 4 倍。残差连接解决网络加深后梯度消失的问题LayerNorm 的作用是稳定训练。三者组合成标准的 TransformerBlockclass FeedForward(nn.Module): def __init__(self, d_model, d_ff2048, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): x F.gelu(self.linear1(x)) return self.linear2(self.dropout(x)) class TransformerBlock(nn.Module): def __init__(self, d_model, num_heads, d_ff512, dropout0.1): super().__init__() self.attention 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.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 第一个子层多头注意力 残差 LayerNorm attn_out, _ self.attention(x, x, x, mask) x self.norm1(x self.dropout(attn_out)) # 第二个子层前馈网络 残差 LayerNorm ffn_out self.ffn(x) x self.norm2(x self.dropout(ffn_out)) return x注意这里使用的是 Post-LN 写法先做子层计算再加残差最后归一化。原论文就是这个结构。实践中不少模型改用 Pre-LN也就是先归一化再进子层训练更稳定。两种写法在代码层面的区别只是norm放在子层之前还是之后理解这一点有助于阅读不同开源项目。4.3 组装一个两层编码器把位置编码和 TransformerBlock 组合起来就是一个最小的编码器骨架。下面的小型模型输入 token id
返回列表