从零实现《Attention Is All You Need》!Transformer源码深度剖析(纯PyTorch可运行)
我们是由枫哥组建的IT技术团队,成立于2017年,致力于帮助IT从业者提供实力,成功入职理想企业,我们提供一对一学习辅导,由知名大厂导师指导,分享Java技术、参与项目实战等服务,并为学员定制职业规划,全面提升竞争力,过去8年,我们已成功帮助数千名求职者拿到满意的Offer:IT枫斗者、IT枫斗者-Java面试突击。
一、开篇前言
2017 年,Google 研究团队发表里程碑式论文 《Attention Is All You Need》,正式提出 Transformer 架构。
该架构彻底摒弃了传统 RNN/LSTM 的递归时序结构,完全依托自注意力机制(Self-Attention)捕获文本长距离语义依赖,一举解决了递归模型长序列梯度消失、无法并行计算的痛点。
时至今日,所有主流大模型均基于 Transformer 衍生而来:
- BERT:仅使用 Transformer Encoder
- GPT、LLaMA、Qwen:仅使用 Transformer Decoder
- T5、BART:完整 Encoder-Decoder 架构
本文将通过350行纯PyTorch极简可运行代码,从零逐层拆解 Transformer 所有核心组件,公式+代码+原理三重解析,无冗余封装,适合新手入门、原理调试、二次开发学习。
二、Transformer 核心组件总览
Transformer 所有能力均由七大基础组件构成,各组件功能、时间复杂度一目了然:
| 核心组件 | 核心功能 | 时间复杂度 |
|---|---|---|
| Scaled Dot-Product Attention | Q/K/V 相似度计算与加权聚合 | O(n·dₖ) |
| Multi-Head Attention | 多表示空间并行注意力,捕捉多维语义 | O(n·d_model) |
| Position-wise FFN | 单位置非线性特征变换,提升模型容量 | O(d_model·d_ff) |
| Positional Encoding | 为序列注入位置时序信息 | O(max_len·d_model) |
| Layer Norm | 特征维度规范化,稳定训练 | O(d_model) |
| Encoder Layer | 自注意力编码全局语义,双向理解 | O(n²·d_model) |
| Decoder Layer | 掩码自注意力+交叉注意力,序列生成 | O(n²·d_model) |
三、核心组件从零实现与深度解析
1. Scaled Dot-Product Attention(缩放点积注意力)
1.1 核心公式与原理
缩放点积注意力是 Transformer 的最小核心计算单元,核心逻辑:通过 Query(查询)与 Key(键)的相似度,对 Value(值)进行加权聚合。
核心公式:
A t t e n t i o n ( Q , K , V ) = s o f t m a x ( Q K ⊤ d k ) V Attention(Q,K,V)=softmax(\frac{QK^\top}{\sqrt{d_k}})V Attention(Q,K,V)=softmax(dkQK⊤)V
1.2 为什么需要除以 d k \sqrt{d_k} dk?
当特征维度 d k d_k dk 较大时,Q、K 点积结果的方差会随维度增大而急剧升高,数值会变得极大。此时 softmax 函数会落入梯度饱和区,梯度趋近于0,导致训练停滞。
除以 d k \sqrt{d_k} dk 可将点积结果方差稳定在1左右,保证 softmax 梯度有效,让训练过程更稳定。
1.3 完整代码实现
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class ScaledDotProductAttention(nn.Module):
"""
缩放点积注意力
Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
"""
def __init__(self, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
def forward(self, Q, K, V, mask=None):
# 获取特征维度
d_k = Q.size(-1)
# 计算注意力分数并缩放
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
# 掩码处理:padding位置置为-∞,softmax后权重为0
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
# 计算注意力权重并dropout正则化
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# 加权聚合Value特征
output = torch.matmul(attn_weights, V)
return output, attn_weights
1.4 核心关键点总结
- masked_fill:将padding无效位置填充为负无穷,经过softmax后权重趋近于0,彻底屏蔽无效字符干扰
- 注意力后dropout:经典正则化方案,有效防止注意力过拟合
- 返回attn_weights:用于训练可视化、注意力权重调试分析
- 性能瓶颈: Q K ⊤ QK^\top QK⊤ 矩阵乘法时间复杂度为 O ( n 2 ⋅ d k ) O(n^2 \cdot d_k) O(n2⋅dk),序列长度n越大,计算量爆炸,是Transformer核心性能瓶颈
2. Multi-Head Attention(多头注意力)
2.1 核心原理
单头注意力仅能在单一特征空间计算语义关联,表征能力有限。多头注意力将 Q/K/V 投影至多个独立特征子空间,并行计算多组注意力,最后拼接融合所有子空间特征,捕捉丰富的多维语义关系。
核心公式:
M u l t i H e a d ( Q , K , V ) = C o n c a t ( h e a d 1 , h e a d 2 , . . . , h e a d h ) W O MultiHead(Q,K,V)=Concat(head_1,head_2,...,head_h)W_O MultiHead(Q,K,V)=Concat(head1,head2,...,headh)WO
其中单个头: h e a d i = A t t e n t i o n ( Q W Q i , K W K i , V W V i ) head_i = Attention(QW_{Q_i},KW_{K_i},VW_{V_i}) headi=Attention(QWQi,KWKi,VWVi)
关键设计优化:不单独定义3h个线性层,而是先整体投影再拆头,仅用4个线性层即可实现等价效果,大幅减少参数量与计算量。
2.2 完整代码实现
class MultiHeadAttention(nn.Module):
"""
多头注意力机制
MultiHead(Q, K, V) = Concat(head_1, ..., head_h) @ W_O
"""
def __init__(self, d_model: int, n_heads: int, dropout: float = 0.1):
super().__init__()
# 特征维度必须整除头数,保证每个头维度均等
assert d_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
# 4个核心投影矩阵
self.W_Q = nn.Linear(d_model, d_model, bias=False)
self.W_K = nn.Linear(d_model, d_model, bias=False)
self.W_V = nn.Linear(d_model, d_model, bias=False)
self.W_O = nn.Linear(d_model, d_model, bias=False)
# 复用缩放点积注意力
self.attention = ScaledDotProductAttention(dropout)
def forward(self, Q, K, V, mask=None):
batch_size = Q.size(0)
# 1. 整体线性投影
Q = self.W_Q(Q)
K = self.W_K(K)
V = self.W_V(V)
# 2. 拆分多头:(batch, seq_len, d_model) -> (batch, n_heads, seq_len, d_k)
Q = Q.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
K = K.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
V = V.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# 3. 并行计算多头注意力
attn_output, attn_weights = self.attention(Q, K, V, mask)
# 4. 拼接所有头特征,恢复维度
attn_output = attn_output.transpose(1, 2).contiguous().view(
batch_size, -1, self.d_model
)
# 5. 最终融合投影
output = self.W_O(attn_output)
return output
2.3 核心关键点总结
- 维度约束:d_model % n_heads == 0,保证每个注意力头分配均等维度
- view+transpose:拆分特征维度,将多头维度前置,实现并行计算
- contiguous():transpose仅修改视图、不改变内存布局,该方法重置内存,避免后续维度变换报错
- 广播掩码:mask格式为(batch, 1, 1, seq_len),可直接与注意力分数矩阵广播运算,无需额外维度扩展
3. Position-wise Feed-Forward Network(位置前馈网络)
3.1 核心原理
注意力层负责捕捉序列全局关联,而 FFN 负责对每个位置特征独立做非线性变换,等价于 kernel_size=1 的一维卷积,大幅提升模型特征表达容量。
核心公式:
F F N ( x ) = R e L U ( x W 1 + b 1 ) W 2 + b 2 FFN(x)=ReLU(xW_1+b_1)W_2+b_2 FFN(x)=ReLU(xW1+b1)W2+b2
3.2 完整代码实现
class PositionWiseFeedForward(nn.Module):
"""
位置式前馈网络
维度变换:d_model → d_ff(升维)→ d_model(降维)
"""
def __init__(self, d_model: int, d_ff: int, dropout: float = 0.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):
# 升维+非线性激活+dropout+降维
return self.linear2(self.dropout(F.relu(self.linear1(x))))
3.3 核心关键点总结
- 维度设计:d_ff 远大于 d_model(论文512→2048),通过升维提升非线性拟合能力
- dropout位置:置于ReLU激活之后,是工业界最优实践
- 激活函数迭代:原始论文用ReLU,后续GPT、LLaMA等模型升级为GELU,拟合效果更优
- 两层最优解:单层FFN表达能力不足,三层及以上收益极低、算力消耗剧增,两层是性能与算力的最优平衡
4. Positional Encoding(位置编码)
4.1 核心原理
Self-Attention 是排列不变性结构,无法感知序列顺序。为模型注入时序位置信息,论文采用正余弦位置编码,无需训练、可泛化超长序列。
核心公式:
P E ( p o s , 2 i ) = s i n ( p o s 10000 2 i / d m o d e l ) PE(pos,2i)=sin(\frac{pos}{10000^{2i/d_{model}}}) PE(pos,2i)=sin(100002i/dmodelpos)
P E ( p o s , 2 i + 1 ) = c o s ( p o s 10000 2 i / d m o d e l ) PE(pos,2i+1)=cos(\frac{pos}{10000^{2i/d_{model}}}) PE(pos,2i+1)=cos(100002i/dmodelpos)
4.2 正余弦编码优势
- 可外推性:支持推理时出现训练集未见过的超长序列
- 无参数:无需训练参数,节省内存与算力
- 相对位置感知:利用三角函数和差公式,可线性表达相对位置偏移
4.3 完整代码实现
class PositionalEncoding(nn.Module):
"""
正余弦位置编码
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
"""
def __init__(self, d_model: int, max_len: int = 5000, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
# 初始化位置编码矩阵
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)
)
# 偶数维度sin,奇数维度cos
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
# 注册为缓冲区:不参与梯度更新,随模型迁移设备
pe = pe.unsqueeze(0)
self.register_buffer('pe', pe)
def forward(self, x):
# 广播相加,注入位置信息
x = x + self.pe[:, :x.size(1), :]
return self.dropout(x)
4.4 核心关键点总结
- 数值稳定性:采用指数运算替代直接幂运算,避免数值溢出/下溢
- register_buffer:位置编码是固定常量,不参与参数更新,仅随模型迁移CPU/GPU
- 相加融合:特征嵌入与位置编码直接广播相加,轻量化融合位置信息
5. Layer Normalization(层归一化)
5.1 核心原理
层归一化对单个样本的所有特征维度做归一化,解决训练过程中特征分布偏移问题。与BatchNorm不同,LN不依赖batch大小,对变长NLP序列极其友好,是Transformer的标配归一化方案。
核心公式:
L a y e r N o r m ( x ) = γ ⋅ x − μ σ 2 + ϵ + β LayerNorm(x)=\gamma \cdot \frac{x-\mu}{\sqrt{\sigma^2+\epsilon}}+\beta LayerNorm(x)=γ⋅σ2+ϵx−μ+β
5.2 完整代码实现
class LayerNorm(nn.Module):
"""
手写层归一化,便于底层原理理解
生产环境可直接使用 nn.LayerNorm
"""
def __init__(self, d_model: int, eps: float = 1e-6):
super().__init__()
# 可学习缩放、偏移参数
self.gamma = nn.Parameter(torch.ones(d_model))
self.beta = nn.Parameter(torch.zeros(d_model))
self.eps = eps
def forward(self, x):
# 对最后一维(特征维)计算均值、方差
mean = x.mean(dim=-1, keepdim=True)
std = x.std(dim=-1, keepdim=True, unbiased=False)
# 归一化+缩放偏移
return self.gamma * (x - mean) / (std + self.eps) + self.beta
5.3 核心关键点总结
- unbiased=False:使用有偏样本标准差,与原始Transformer论文实现完全对齐
- eps极小值:防止分母为0,避免数值报错
- 可学习参数:gamma、beta让归一化后特征具备自适应拟合能力
6. Encoder Layer(编码器层)
6.1 核心结构
编码器用于双向语义理解,每层由两个子层组成,搭配残差连接+层归一化,堆叠形成深层特征提取网络:
输入 → 多头自注意力 → 残差连接+LN → FFN → 残差连接+LN → 输出
6.2 完整代码实现
class EncoderLayer(nn.Module):
"""Transformer 编码器单层结构"""
def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, n_heads, dropout)
self.ffn = PositionWiseFeedForward(d_model, d_ff, dropout)
self.norm1 = LayerNorm(d_model)
self.norm2 = LayerNorm(d_model)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
def forward(self, x, mask=None):
# 1. 自注意力子层 + 残差归一化
attn_output = self.self_attn(x, x, x, mask)
x = x + self.dropout1(attn_output)
x = self.norm1(x)
# 2. 前馈网络子层 + 残差归一化
ffn_output = self.ffn(x)
x = x + self.dropout2(ffn_output)
x = self.norm2(x)
return x
6.3 核心关键点总结
- 残差连接:
x + sublayer(x),打通梯度直通通道,彻底解决深层网络梯度消失问题 - Post-LN结构:先残差、后归一化,严格对齐原始论文实现
- 双向自注意力:Q/K/V均来自输入自身,可感知上下文双向语义
7. Decoder Layer(解码器层)
7.1 核心结构
解码器用于自回归序列生成,在编码器基础上新增交叉注意力,共三个子层,严格防止未来信息泄露:
输入 → 掩码自注意力 → 残差LN → 交叉注意力 → 残差LN → FFN → 残差LN → 输出
7.2 完整代码实现
class DecoderLayer(nn.Module):
"""Transformer 解码器单层结构"""
def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, n_heads, dropout)
self.cross_attn = MultiHeadAttention(d_model, n_heads, dropout)
self.ffn = PositionWiseFeedForward(d_model, d_ff, dropout)
self.norm1 = LayerNorm(d_model)
self.norm2 = LayerNorm(d_model)
self.norm3 = LayerNorm(d_model)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
self.dropout3 = nn.Dropout(dropout)
def forward(self, x, enc_output, src_mask=None, tgt_mask=None):
# 1. 掩码自注意力:屏蔽未来位置信息
attn_output = self.self_attn(x, x, x, tgt_mask)
x = x + self.dropout1(attn_output)
x = self.norm1(x)
# 2. 交叉注意力:Decoder查Encoder全局特征
attn_output = self.cross_attn(x, enc_output, enc_output, src_mask)
x = x + self.dropout2(attn_output)
x = self.norm2(x)
# 3. 前馈网络
ffn_output = self.ffn(x)
x = x + self.dropout3(ffn_output)
x = self.norm3(x)
return x
7.3 核心关键点总结
- Target掩码:下三角掩码,屏蔽当前token之后的所有位置,杜绝未来信息泄露,保障自回归生成逻辑
- 交叉注意力机制:Q来自解码器当前输出,K/V来自编码器全局编码结果,实现输入序列对生成序列的语义引导
- 三层残差归一化:相比编码器多一组子层结构,适配生成任务复杂度
8. 完整 Transformer 模型拼装
8.1 整体网络流程
编码器流程:源序列 → 嵌入层 → 位置编码 → N层编码器堆叠 → 全局语义特征
解码器流程:目标序列 → 嵌入层 → 位置编码 → N层解码器堆叠 → 全连接输出词表概率
8.2 完整模型代码
class Transformer(nn.Module):
"""
完整 Transformer 模型
适配机器翻译、序列生成等经典任务
"""
def __init__(self, src_vocab, tgt_vocab, d_model=512,
n_heads=8, d_ff=2048, n_layers=6,
dropout=0.1, max_len=5000):
super().__init__()
# 嵌入层
self.encoder_embed = nn.Embedding(src_vocab, d_model)
self.decoder_embed = nn.Embedding(tgt_vocab, d_model)
# 位置编码
self.pos_encoding = PositionalEncoding(d_model, max_len, dropout)
# 堆叠多层编码器、解码器
self.encoder_layers = nn.ModuleList([
EncoderLayer(d_model, n_heads, d_ff, dropout)
for _ in range(n_layers)
])
self.decoder_layers = nn.ModuleList([
DecoderLayer(d_model, n_heads, d_ff, dropout)
for _ in range(n_layers)
])
# 最终输出层:映射到词表维度
self.fc_out = nn.Linear(d_model, tgt_vocab)
def forward(self, src, tgt, src_mask=None, tgt_mask=None):
# 编码器前向传播
src_emb = self.pos_encoding(self.encoder_embed(src))
for layer in self.encoder_layers:
src_emb = layer(src_emb, src_mask)
# 解码器前向传播
tgt_emb = self.pos_encoding(self.decoder_embed(tgt))
for layer in self.decoder_layers:
tgt_emb = layer(tgt_emb, src_emb, src_mask, tgt_mask)
# 输出每个位置的词表概率分布
return self.fc_out(tgt_emb)
8.3 核心关键点总结
- ModuleList:正确注册多层网络参数,保证模型可训练、参数可更新
- 双嵌入层设计:编码器、解码器独立词嵌入,适配源域、目标域不同语义分布
- 输出映射:将512维特征映射至词表维度,用于token预测与序列生成
四、全文核心总结
- 核心本质:Transformer 彻底抛弃递归结构,完全依靠注意力机制实现全局语义建模,支持全并行训练
- 缩放注意力:除以 d k \sqrt{d_k} dk是训练稳定的关键,解决高维点积梯度消失问题
- 多头机制:多空间并行建模,丰富语义表征,是模型强大拟合能力的核心
- 位置编码:正余弦编码实现无参数、可外推的序列位置注入
- 编解码分工:Encoder双向理解、Decoder自回归生成,适配各类NLP任务
- 残差+LN:深层网络训练的基础保障,解决梯度消失、特征分布偏移问题
⭐️推荐:
更多推荐




所有评论(0)