三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

【人工智能】深入浅出 Transformer 架构:从 Self-Attention 到 PyTorch 完整代码实现

【人工智能】深入浅出 Transformer 架构:从 Self-Attention 到 PyTorch 完整代码实现

摘要:Transformer 架构是现代大语言模型(LLM)与 Transformer 系列模型(如 BERT、GPT、ViT)的基础基石。本文旨在帮助开发者与 AI 学习者彻底理解 Transformer 的核心机制。文章从传统 RNN/LSTM 的长距离依赖与并行化瓶颈出发,深入拆解自注意力机制(Self-Attention)、多头注意力(Multi-Head Attention)、位置编码(Positional Encoding)及残差连接等关键结构;随后基于 PyTorch 从零手写完整 Transformer 模块并进行前向传播验证;最后总结了 Mask 处理、LayerNorm 选择及训练 Warmup 等实战踩坑经验,助力快速掌握大模型底层原理。


深入浅出 Transformer 架构:从 Self-Attention 到 PyTorch 完整代码实现

1. 背景原理:为什么需要 Transformer?

在 2017 年 Google 提出论文 《Attention Is All You Need》 之前,自然语言处理(NLP)领域的主流模型是基于循环神经网络(RNN)及其变体 LSTM / GRU 的 Encoder-Decoder 结构。然而,RNN 结构存在两个致命的瓶颈:

  1. 无法并行计算(Sequential Computation):RNN 的隐状态hth_tht必须依赖上一个时刻的隐状态ht−1h_{t-1}ht1,这导致计算必须按时间步顺序执行,无法充分利用 GPU 的大规模并行能力。
  2. 长距离依赖衰减(Long-range Dependency Loss):尽管 LSTM/GRU 引入了门控机制,但在处理超长序列时,序列早期的信息依然容易丢失或梯度消失。

为了打破上述限制,Transformer 抛弃了传统的循环与卷积结构,完全依赖自注意力机制(Self-Attention)来进行序列全局关联度的建模,实现了高效的矩阵并行计算与跨长距离语义信息的直接捕捉。


2. 问题分析:序列建模的核心挑战

在进行序列到序列(Seq2Seq)的任务(如机器翻译、文本生成)时,我们需要解决以下核心问题:

  • 语义关联建立:如何判断句中某个词与上下文其他词的关系(例如“苹果”在“苹果手机”和“吃苹果”中的语义差异)?
  • 时序位置感知:注意力机制本身是无序的(置换不变性),如何让模型区分“张三打李四”与“李四打张三”?
  • 训练稳定性与深层拟合:如何保证深层网络在训练时不发生梯度爆炸或消失?

3. 方案思路:Transformer 核心架构拆解

Transformer 整体采用Encoder-Decoder(编码器-解码器)结构。下面我们逐一拆解其核心组件。

模块组件核心作用关键细节
Scaled Dot-Product Attention计算 Query 与 Key 的相关性,加权聚合 Value引入dk\sqrt{d_k}dk缩放,防止点积过大导致 Softmax 梯度消失
Multi-Head Attention允许模型在不同投影子空间中同时关注多维度语义信息dmodeld_{\text{model}}dmodel拆分为hhh个头并行计算后拼接
Positional Encoding注入位置信息采用正余弦周期函数交替计算
Feed-Forward Network (FFN)非线性特征转换两层全连接层 + ReLU/GELU 激活函数
Add & Norm稳定深层网络训练Residual Connection(残差连接) + Layer Normalization(层归一化)

3.1 缩放点积注意力(Scaled Dot-Product Attention)

注意力机制将输入特征映射为三个向量矩阵:Query (QQQ)Key (KKK)Value (VVV)。计算公式如下:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)VAttention(Q,K,V)=softmax(dkQKT)V

其中dk\sqrt{d_k}dk是缩放因子。当维度dkd_kdk较时,点积数值可能非常大,导致 Softmax 函数进入梯度极小的饱和区,缩放因子能够有效稳定梯度。

3.2 多头注意力机制(Multi-Head Attention)

多头注意力将高维度的Q,K,VQ, K, VQ,K,V通过线性变换投影到hhh个不同的子空间,分别计算注意力后再进行拼接(Concat)与最终线性映射。

MultiHead(Q,K,V)=Concat(head1,…,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \dots, \text{head}_h)W^OMultiHead(Q,K,V)=Concat(head1,,headh)WO

headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)headi=Attention(QWiQ,KWiK,VWiV)


4. 代码实操:PyTorch 从零手写 Transformer 核心模块

下面我们使用 PyTorch 逐步实现各核心组件,代码均附带详细注释。

4.1 缩放点积注意力实现

importtorchimporttorch.nnasnnimporttorch.nn.functionalasFimportmathclassScaledDotProductAttention(nn.Module):""" 缩放点积注意力模块 """def__init__(self,d_k):super(ScaledDotProductAttention,self).__init__()self.d_k=d_kdefforward(self,Q,K,V,mask=None):# Q, K, V shape: [batch_size, n_heads, seq_len, d_k]# 计算 Q 和 K^T 的点积相似度评分scores=torch.matmul(Q,K.transpose(-2,-1))/math.sqrt(self.d_k)# 如果存在 mask(例如 Padding Mask 或 Causal Mask),将掩码位置设为极大负数ifmaskisnotNone:scores=scores.masked_fill(mask==0,-1e9)# Softmax 归一化得到注意力权重attn_weights=F.softmax(scores,dim=-1)# 与 V 相乘得到加权特征表示output=torch.matmul(attn_weights,V)returnoutput,attn_weights

4.2 多头注意力(Multi-Head Attention)实现

classMultiHeadAttention(nn.Module):""" 多头注意力模块 """def__init__(self,d_model,n_heads):super(MultiHeadAttention,self).__init__()assertd_model%n_heads==0,"d_model 必须能被 n_heads 整除"self.d_model=d_model self.n_heads=n_heads self.d_k=d_model//n_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.attention=ScaledDotProductAttention(self.d_k)defforward(self,Q,K,V,mask=None):batch_size=Q.size(0)# 1. 线性投影并维度变换: [batch_size, seq_len, d_model] -> [batch_size, n_heads, seq_len, d_k]q_s=self.W_Q(Q).view(batch_size,-1,self.n_heads,self.d_k).transpose(1,2)k_s=self.W_K(K).view(batch_size,-1,self.n_heads,self.d_k).transpose(1,2)v_s=self.W_V(V).view(batch_size,-1,self.n_heads,self.d_k).transpose(1,2)ifmaskisnotNone:# 扩展 mask 维度以适配多头形状 [batch_size, 1, seq_len, seq_len]mask=mask.unsqueeze(1)# 2. 计算 Scaled Dot-Product Attentioncontext,attn_weights=self.attention(q_s,k_s,v_s,mask=mask)# 3. 拼接多个头的输出: [batch_size, n_heads, seq_len, d_k] -> [batch_size, seq_len, d_model]context=context.transpose(1,2).contiguous().view(batch_size,-1,self.d_model)# 4. 最终线性变换output=self.W_O(context)returnoutput,attn_weights

4.3 位置编码(Positional Encoding)实现

classPositionalEncoding(nn.Module):""" 正弦/余弦位置编码 """def__init__(self,d_model,max_len=5000):super(PositionalEncoding,self).__init__()pe=torch.zeros(max_len,d_model)position=torch.arange(0,max_len,dtype=torch.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)pe=pe.unsqueeze(0)# [1, max_len, d_model]self.register_buffer('pe',pe)defforward(self,x):# x: [batch_size, seq_len, d_model]x=x+self.pe[:,:x.size(1)]returnx

4.4 Encoder Block 搭建

classEncoderLayer(nn.Module):""" 编码器单层模块(Encoder Block) """def__init__(self,d_model,n_heads,d_ff,dropout=0.1):super(EncoderLayer,self).__init__()self.mha=MultiHeadAttention(d_model,n_heads)self.norm1=nn.LayerNorm(d_model)self.norm2=nn.LayerNorm(d_model)# 前馈神经网络 Feed-Forward Networkself.ffn=nn.Sequential(nn.Linear(d_model,d_ff),nn.ReLU(),nn.Dropout(dropout),nn.Linear(d_ff,d_model))self.dropout=nn.Dropout(dropout)defforward(self,x,mask=None):# 1. 自注意力 + 残差连接 + LayerNormattn_out,_=self.mha(x,x,x,mask)x=self.norm1(x+self.dropout(attn_out))# 2. 前馈网络 + 残差连接 + LayerNormffn_out=self.ffn(x)x=self.norm2(x+self.dropout(ffn_out))returnx

5. 结果验证:模型前向传播测试

下面编写验证脚本,模拟 Batch 大小为 2、序列长度为 5 的文本数据输入,验证模块前向传播的正确性与张量维度变化。

if__name__=="__main__":# 超参数定义batch_size=2seq_len=5d_model=64n_heads=8d_ff=256# 模拟输入嵌入 Tensordummy_input=torch.randn(batch_size,seq_len,d_model)# 1. 测试位置编码pos_encoder=PositionalEncoding(d_model)x_pe=pos_encoder(dummy_input)print(f"输入形状:{dummy_input.shape}")print(f"位置编码后形状:{x_pe.shape}")# 2. 测试 Encoder Layer 前向传播encoder_layer=EncoderLayer(d_model,n_heads,d_ff)output=encoder_layer(x_pe)print(f"Encoder Layer 输出形状:{output.shape}")# 验证输出维度是否严格保持不变 [batch_size, seq_len, d_model]assertoutput.shape==(batch_size,seq_len,d_model),"维度校验失败!"print("\n✅ 前向传播测试成功!张量维度转换完全符合预期。")

运行效果输出:

输入形状: torch.Size([2, 5, 64]) 位置编码后形状: torch.Size([2, 5, 64]) Encoder Layer 输出形状: torch.Size([2, 5, 64]) ✅ 前向传播测试成功!张量维度转换完全符合预期。

6. 实战踩坑与优化总结

在实际开发与训练 Transformer 模型时,有以下几点需要特别注意:

1. Mask 机制的区别与使用

  • Padding Mask:用于处理变长文本。对填充位置(Padding Token,如 ID=0)进行掩码,防止无意义的填充项参与注意力计算。
  • Causal Mask(因果掩码 / 下三角掩码):用于 Decoder 或 GPT 等自回归生成模型,确保ttt时刻只能看到ttt及之前的词,防止“预测未来文本”。

2. LayerNorm 与 BatchNorm 的选择

  • NLP 任务中严禁使用 BatchNorm:因为文本序列长度动态可变,Batch 内不同样本在同一位置的统计量差异巨大。
  • 使用 LayerNorm:在单个样本内部的特征维度dmodeld_{\text{model}}dmodel上计算均值与方差,不受序列长度与 Batch 大小的波动影响。

3. Pre-LN 与 Post-LN 结构

  • 原始论文采用Post-LN(在残差相加后进行 LayerNorm),深层网络容易梯度消失,训练需要非常精细的 Warmup 策略。
  • 现代大模型(如 LLaMA、GPT-3)普遍改用Pre-LN(在输入 Self-Attention / FFN 前先进行 LayerNorm),训练收敛更加稳定。

7. 总结

Transformer 架构通过自注意力机制打破了传统 RNN 的串行瓶颈,配合多头注意力残差归一化结构,不仅具备强大的长距离上下文拟合能力,更拥有极佳的硬件并行效率。深刻理解 Transformer 的底层张量变换与 PyTorch 实现,是掌握当下 LLM(大语言模型)与 Agent 技术的必经之路。

← 返回列表