「Transformer核心必读」Pre‑LN 解码器架构详解与 PyTorch 代码实现
1、解码器架构详解
对 采用 Pre-Layer Normalization(Pre-LN)结构的 Transformer 解码器主体部分 进行一次全面、系统、深入且详尽的讲解,限定在 解码器本身 的结构范围内,不涉及最终输出层(如线性投影 + softmax),也不讨论训练策略、优化技巧或与编码器交互以外的上下文。
一、整体架构定位
在标准的 Transformer 架构中,解码器(Decoder) 是一个独立的子网络,其主要任务是:
基于已生成的部分目标序列(自回归输入)和来自编码器的上下文信息,逐步生成下一个 token 的表示。
在 仅解码器模型(如 GPT 系列)中,解码器也承担全部语言建模任务,但此时 没有交叉注意力模块(因为无编码器)。然而,在 编码器-解码器架构(如原始 Transformer、T5、BART)中,解码器必须包含交叉注意力以融合源序列信息。
聚焦于 完整的编码器-解码器架构中的 Pre-LN 解码器,即包含:
- 掩码自注意力(Masked Self-Attention)
- 交叉注意力(Cross-Attention)
- 前馈网络(FFN)
并采用 Pre-LN 配置(LayerNorm 在残差分支之前)。
二、解码器的宏观结构:堆叠式设计
2.1 层数与参数
- 解码器由 N N N 个完全相同的解码器层(Decoder Layer) 垂直堆叠而成。
- N N N 是超参数,典型值为 6(原始 Transformer)、12(BERT-style)、24 或更高(如 LLaMA 使用 32–80 层)。
- 每一层拥有 独立的可学习参数,包括:
- 自注意力模块的 W Q self , W K self , W V self , W O self W_Q^{\text{self}}, W_K^{\text{self}}, W_V^{\text{self}}, W_O^{\text{self}} WQself,WKself,WVself,WOself
- 交叉注意力模块的 W Q cross , W O cross W_Q^{\text{cross}}, W_O^{\text{cross}} WQcross,WOcross(注意:Key 和 Value 来自编码器,因此解码器不学习 W K enc W_K^{\text{enc}} WKenc 或 W V enc W_V^{\text{enc}} WVenc,但需学习自身的 Query 投影和输出投影)
- FFN 的两个线性变换矩阵
- 三个 LayerNorm 的缩放( γ \gamma γ)和平移( β \beta β)参数
⚠️ 注意:虽然结构相同,但各层参数不共享。深层通常学习更抽象的语义表示。
2.2 数据流概览
设输入序列为 Y = [ y 1 , y 2 , . . . , y T ] Y = [y_1, y_2, ..., y_T] Y=[y1,y2,...,yT],其中 T T T 是目标序列长度。
-
初始嵌入与位置编码:
H ( 0 ) = E ( Y ) + P H^{(0)} = E(Y) + P H(0)=E(Y)+P
其中 E ( ⋅ ) ∈ R V × d model E(\cdot) \in \mathbb{R}^{V \times d_{\text{model}}} E(⋅)∈RV×dmodel 是词嵌入矩阵( V V V 为词表大小), P ∈ R T × d model P \in \mathbb{R}^{T \times d_{\text{model}}} P∈RT×dmodel 是位置编码。💡 在大多数实现中(包括 Pre-LN 架构),位置编码直接加到词嵌入上,作为解码器第一层的输入。
-
逐层传播:
对于 l = 1 l = 1 l=1 到 N N N:
H ( l ) = DecoderLayer ( l ) ( H ( l − 1 ) , K enc , V enc ) H^{(l)} = \text{DecoderLayer}^{(l)}\left(H^{(l-1)}, K_{\text{enc}}, V_{\text{enc}}\right) H(l)=DecoderLayer(l)(H(l−1),Kenc,Venc)
其中 K enc = H enc W K enc K_{\text{enc}} = H_{\text{enc}} W_K^{\text{enc}} Kenc=HencWKenc, V enc = H enc W V enc V_{\text{enc}} = H_{\text{enc}} W_V^{\text{enc}} Venc=HencWVenc 是编码器最终输出 H enc ∈ R S × d model H_{\text{enc}} \in \mathbb{R}^{S \times d_{\text{model}}} Henc∈RS×dmodel 经线性变换后的 Key 和 Value,维度为 S × d k S \times d_k S×dk( S S S 为源序列长度)。 -
最终输出:
H ( N ) ∈ R T × d model H^{(N)} \in \mathbb{R}^{T \times d_{\text{model}}} H(N)∈RT×dmodel 是解码器主体的输出,每个位置 t t t 的向量 h t ( N ) h_t^{(N)} ht(N) 编码了:- 目标序列前缀 y ≤ t y_{\leq t} y≤t 的自回归上下文
- 与源序列全局对齐后的语义信息
✅ 此时我们停止。不进行 H ( N ) W out H^{(N)} W_{\text{out}} H(N)Wout 投影。
三、单个 Pre-LN 解码器层的微观结构
每个解码器层由 三个顺序执行的子模块 构成,每个子模块都遵循 “LayerNorm → 子操作 → 残差连接” 的 Pre-LN 范式。
我们逐一分解:
3.1 第一子模块:掩码自注意力(Masked Self-Attention)
输入
- X ∈ R T × d model X \in \mathbb{R}^{T \times d_{\text{model}}} X∈RT×dmodel:上一层输出(或初始嵌入)
步骤详解
-
Layer Normalization(Pre-LN)
X ~ = LayerNorm ( X ) \tilde{X} = \text{LayerNorm}(X) X~=LayerNorm(X)-
LayerNorm 对每个 token 的特征维度做归一化(均值为 0,方差为 1),再通过可学习参数 γ \gamma γ(缩放)和 β \beta β(偏移)恢复表达能力:
LayerNorm ( x ) = γ ⊙ x − μ σ + ϵ + β \text{LayerNorm}(x) = \gamma \odot \frac{x - \mu}{\sigma + \epsilon} + \beta LayerNorm(x)=γ⊙σ+ϵx−μ+β -
这一步发生在注意力计算之前,是 Pre-LN 的核心标志。
-
-
线性投影生成 Q, K, V
Q = X ~ W Q self , K = X ~ W K self , V = X ~ W V self Q = \tilde{X} W_Q^{\text{self}}, \quad K = \tilde{X} W_K^{\text{self}}, \quad V = \tilde{X} W_V^{\text{self}} Q=X~WQself,K=X~WKself,V=X~WVself- W Q , W K , W V ∈ R d model × d k W_Q, W_K, W_V \in \mathbb{R}^{d_{\text{model}} \times d_k} WQ,WK,WV∈Rdmodel×dk,通常 d k = d model / h d_k = d_{\text{model}} / h dk=dmodel/h, h h h 为注意力头数。
- 多头机制:将 Q/K/V 分成 h h h 个头,分别计算注意力,拼接后再通过 W O self ∈ R d model × d model W_O^{\text{self}} \in \mathbb{R}^{d_{\text{model}} \times d_{\text{model}}} WOself∈Rdmodel×dmodel 投影回原维度。
-
因果掩码(Causal Masking)
-
构造下三角掩码矩阵 M ∈ { 0 , − ∞ } T × T M \in \{0, -\infty\}^{T \times T} M∈{0,−∞}T×T,使得:
M i j = { 0 if i ≥ j − ∞ if i < j M_{ij} = \begin{cases} 0 & \text{if } i \geq j \\ -\infty & \text{if } i < j \end{cases} Mij={0−∞if i≥jif i<j -
应用于注意力分数:
A = softmax ( Q K ⊤ d k + M ) A = \text{softmax}\left( \frac{QK^\top}{\sqrt{d_k}} + M \right) A=softmax(dkQK⊤+M) -
效果:位置 i i i 只能关注 j ≤ i j \leq i j≤i,保证自回归性质。
-
-
注意力输出与投影
SelfAttnOut = Concat ( head 1 , . . . , head h ) W O self \text{SelfAttnOut} = \text{Concat}(\text{head}_1, ..., \text{head}_h) W_O^{\text{self}} SelfAttnOut=Concat(head1,...,headh)WOself -
残差连接
X after_self = X + SelfAttnOut X_{\text{after\_self}} = X + \text{SelfAttnOut} Xafter_self=X+SelfAttnOut- 注意:残差加的是 原始输入 X,不是 LayerNorm 后的 X ~ \tilde{X} X~。
🔁 关键点:Pre-LN 中,梯度直接流经 LayerNorm 到输入,避免深层梯度消失。
3.2 第二子模块:交叉注意力(Cross-Attention)
输入
- Query 来源: X after_self X_{\text{after\_self}} Xafter_self
- Key/Value 来源:编码器最终输出 H enc ∈ R S × d model H_{\text{enc}} \in \mathbb{R}^{S \times d_{\text{model}}} Henc∈RS×dmodel
📌 重要:在整个解码过程中, K enc = H enc W K enc K_{\text{enc}} = H_{\text{enc}} W_K^{\text{enc}} Kenc=HencWKenc, V enc = H enc W V enc V_{\text{enc}} = H_{\text{enc}} W_V^{\text{enc}} Venc=HencWVenc 只需在编码器完成后计算一次并缓存,供所有解码器层和所有生成步使用。解码器自身不学习 W K enc W_K^{\text{enc}} WKenc 或 W V enc W_V^{\text{enc}} WVenc。
步骤详解
-
Layer Normalization(Pre-LN)
X ~ query = LayerNorm ( X after_self ) \tilde{X}_{\text{query}} = \text{LayerNorm}(X_{\text{after\_self}}) X~query=LayerNorm(Xafter_self) -
Query 投影(Key/Value 已预计算)
Q cross = X ~ query W Q cross Q_{\text{cross}} = \tilde{X}_{\text{query}} W_Q^{\text{cross}} Qcross=X~queryWQcross- K enc , V enc K_{\text{enc}}, V_{\text{enc}} Kenc,Venc 是固定的(来自编码器)
-
无掩码注意力(Full Attention)
A cross = softmax ( Q cross K enc ⊤ d k ) A_{\text{cross}} = \text{softmax}\left( \frac{Q_{\text{cross}} K_{\text{enc}}^\top}{\sqrt{d_k}} \right) Across=softmax(dkQcrossKenc⊤)- 解码器每个位置可关注源序列所有位置(无因果限制)
-
输出与投影
CrossAttnOut = A cross V enc W O cross \text{CrossAttnOut} = A_{\text{cross}} V_{\text{enc}} W_O^{\text{cross}} CrossAttnOut=AcrossVencWOcross -
残差连接
X after_cross = X after_self + CrossAttnOut X_{\text{after\_cross}} = X_{\text{after\_self}} + \text{CrossAttnOut} Xafter_cross=Xafter_self+CrossAttnOut
💡 交叉注意力是解码器“理解源语义”的桥梁。若为纯解码器模型(如 GPT),此模块不存在。
3.3 第三子模块:前馈网络(Feed-Forward Network, FFN)
输入
- X after_cross X_{\text{after\_cross}} Xafter_cross
结构细节
FFN 是一个两层 MLP,通常具有“扩张-压缩”结构:
-
Layer Normalization(Pre-LN)
X ~ ffn = LayerNorm ( X after_cross ) \tilde{X}_{\text{ffn}} = \text{LayerNorm}(X_{\text{after\_cross}}) X~ffn=LayerNorm(Xafter_cross) -
第一线性层(扩张)
Z = GeLU ( X ~ ffn W 1 + b 1 ) , W 1 ∈ R d model × d ff Z = \text{GeLU}\left( \tilde{X}_{\text{ffn}} W_1 + b_1 \right), \quad W_1 \in \mathbb{R}^{d_{\text{model}} \times d_{\text{ff}}} Z=GeLU(X~ffnW1+b1),W1∈Rdmodel×dff- d ff = 4 × d model d_{\text{ff}} = 4 \times d_{\text{model}} dff=4×dmodel 是常见设置(如 768 → 3072)
- 激活函数常用 GeLU(Gaussian Error Linear Unit),也可用 ReLU、SwiGLU 等
-
第二线性层(压缩)
FFNOut = Z W 2 + b 2 , W 2 ∈ R d ff × d model \text{FFNOut} = Z W_2 + b_2, \quad W_2 \in \mathbb{R}^{d_{\text{ff}} \times d_{\text{model}}} FFNOut=ZW2+b2,W2∈Rdff×dmodel -
残差连接
X out = X after_cross + FFNOut X_{\text{out}} = X_{\text{after\_cross}} + \text{FFNOut} Xout=Xafter_cross+FFNOut
🧠 FFN 被认为是“记忆库”或“知识存储”,每个位置独立处理,无 token 间交互。
四、Pre-LN 的数学形式化总结(单层)
给定输入 X X X,以及编码器提供的 K enc , V enc K_{\text{enc}}, V_{\text{enc}} Kenc,Venc,一个 Pre-LN 解码器层的完整计算流程为:
(1) Masked Self-Attention: X ~ 1 = LayerNorm ( X ) X 1 = X + MaskedMultiHeadAttn ( X ~ 1 , X ~ 1 , X ~ 1 ) (2) Cross-Attention: X ~ 2 = LayerNorm ( X 1 ) X 2 = X 1 + MultiHeadAttn ( X ~ 2 , K enc , V enc ) (3) Feed-Forward Network: X ~ 3 = LayerNorm ( X 2 ) X out = X 2 + FFN ( X ~ 3 ) \begin{aligned} &\text{(1) Masked Self-Attention:} \\ &\quad \tilde{X}_1 = \text{LayerNorm}(X) \\ &\quad X_1 = X + \text{MaskedMultiHeadAttn}(\tilde{X}_1, \tilde{X}_1, \tilde{X}_1) \\ \\ &\text{(2) Cross-Attention:} \\ &\quad \tilde{X}_2 = \text{LayerNorm}(X_1) \\ &\quad X_2 = X_1 + \text{MultiHeadAttn}(\tilde{X}_2, K_{\text{enc}}, V_{\text{enc}}) \\ \\ &\text{(3) Feed-Forward Network:} \\ &\quad \tilde{X}_3 = \text{LayerNorm}(X_2) \\ &\quad X_{\text{out}} = X_2 + \text{FFN}(\tilde{X}_3) \end{aligned} (1) Masked Self-Attention:X~1=LayerNorm(X)X1=X+MaskedMultiHeadAttn(X~1,X~1,X~1)(2) Cross-Attention:X~2=LayerNorm(X1)X2=X1+MultiHeadAttn(X~2,Kenc,Venc)(3) Feed-Forward Network:X~3=LayerNorm(X2)Xout=X2+FFN(X~3)
输出 X out X_{\text{out}} Xout 作为下一层的输入。
五、Pre-LN 与 Post-LN 的深度对比
| 维度 | Pre-LN | Post-LN(Vaswani et al., 2017) |
|---|---|---|
| LayerNorm 位置 | 子模块输入前 | 子模块输出后 |
| 残差路径 | x + Submodule ( LayerNorm ( x ) ) x + \text{Submodule}(\text{LayerNorm}(x)) x+Submodule(LayerNorm(x)) | LayerNorm ( x + Submodule ( x ) ) \text{LayerNorm}(x + \text{Submodule}(x)) LayerNorm(x+Submodule(x)) |
| 梯度传播 | 更稳定,深层梯度不易消失 | 深层梯度可能衰减,需 warmup |
| 训练动态 | 可使用更大初始学习率 | 必须使用学习率预热(warmup) |
| 表示分布 | 各层输出尺度更一致 | 底层和顶层表示尺度差异大 |
| 现代应用 | GPT-2/3, LLaMA, Mistral 等主流模型 | 原始 Transformer、早期 BERT |
📚 理论支持:Pre-LN 被证明在深度网络中具有更好的优化 landscape(Xiong et al., 2020, “On Layer Normalization in the Transformer Architecture”)。
六、关键设计原则与工程考量
6.1 为何需要三个子模块?
- 掩码自注意力:建模目标序列内部依赖(自回归)
- 交叉注意力:对齐并融合源序列信息(跨模态/跨语言)
- FFN:非线性变换,增强模型容量,存储“事实知识”
6.2 为何 Pre-LN 更适合深层堆叠?
- 残差连接直接连接原始输入,梯度可无损回传
- LayerNorm 提前规范化输入,使注意力和 FFN 的输入分布更稳定
- 避免 Post-LN 中 LN 对残差和子模块输出求和后的“混合分布”进行归一化,导致训练不稳定
6.3 参数量估算(以 d model = 768 d_{\text{model}}=768 dmodel=768, h = 12 h=12 h=12, d ff = 3072 d_{\text{ff}}=3072 dff=3072 为例)
- 自注意力: 4 × ( 768 × 768 ) ≈ 2.36 M 4 \times (768 \times 768) \approx 2.36\,\text{M} 4×(768×768)≈2.36M
- 交叉注意力: 2 × ( 768 × 768 ) ≈ 1.18 M 2 \times (768 \times 768) \approx 1.18\,\text{M} 2×(768×768)≈1.18M(仅 W Q cross , W O cross W_Q^{\text{cross}}, W_O^{\text{cross}} WQcross,WOcross;K/V 投影由编码器提供)
- FFN: 768 × 3072 × 2 ≈ 4.72 M 768 \times 3072 \times 2 \approx 4.72\,\text{M} 768×3072×2≈4.72M
- LayerNorm: 3 × 2 × 768 ≈ 4.6 K 3 \times 2 \times 768 \approx 4.6\,\text{K} 3×2×768≈4.6K
- 单层总计 ≈ 8.26M 参数
- 12 层解码器 ≈ 99M 参数(不含嵌入和输出层)
七、总结:Pre-LN 解码器主体的核心要点
- 堆叠结构: N N N 个 identical decoder layers 串联,每层独立参数。
- 三层流水线:掩码自注意力 → 交叉注意力 → FFN。
- Pre-LN 范式:每个子模块前先 LayerNorm,再计算,最后残差加原始输入。
- 信息流:
- 自注意力捕获目标序列历史
- 交叉注意力注入源序列上下文
- FFN 提供非线性表达与知识存储
- 训练优势:Pre-LN 支持更深网络、更大 LR、无需 warmup,已成为工业标准。
- 输出:得到 H ( N ) ∈ R T × d model H^{(N)} \in \mathbb{R}^{T \times d_{\text{model}}} H(N)∈RT×dmodel,每个位置向量蕴含完整上下文,等待后续词汇表投影。
2、代码
# coding: utf-8
import torch
import torch.nn as nn
import math
import torch.nn.functional as F
import copy
# 注意是 Embeddings, 不是 Embedding, 和 nn.Embedding 区分
# 仅仅是 词嵌入, 没有 位置编码
class Embeddings(nn.Module):
def __init__(self, vocabulary_size: int, d_model: int):
super().__init__()
self.vocabulary_size = vocabulary_size # 词表大小
self.d_model = d_model # 词向量维度
self.embed = nn.Embedding(num_embeddings=vocabulary_size, embedding_dim=d_model)
def forward(self, x):
# x.shape = (N, T)
# nn.Embedding 要求输入的 indices(即 x)必须是整数类型(如 torch.long 或 torch.int)
# math.sqrt 放缩的作用
x = self.embed(x) * math.sqrt(self.d_model)
return x
# 位置编码 【 Positional adj.位置的 】
class PositionalEncoding(nn.Module):
def __init__(self, d_model: int, max_len: int = 1024):
super().__init__()
assert d_model % 2 == 0, 'd_model 必须为偶数'
self.d_model = d_model
self.max_len = max_len
# max_len # 最大 token 数量
# d_model # 词向量维度
# 创建大表格
pe = torch.zeros(size=(max_len, d_model))
# 创建位置索引
position = torch.arange(start=0, end=max_len, step=1, dtype=torch.float)
print('position.shape =', position.shape) # torch.Size([1024])
print(f'position: {position}')
# tensor([0.0000e+00, 1.0000e+00, 2.0000e+00, ..., 1.0210e+03, 1.0220e+03, 1.0230e+03])
# 目的:
# pos dim_0 dim_1 dim_2 dim_3 dim_4 ... dim_d_model
# 0 0*w0 0*w1 0*w2 0*w3 0*w4 ... 0*wd_model
# 1 1*w0 1*w1 1*w2 1*w3 1*w4 ... 1*wd_model
# 2 2*w0 2*w1 2*w2 2*w3 2*w4 ... 2*wd_model
# 3 3*w0 3*w1 3*w2 3*w3 3*w4 ... 3*wd_model
# 4 4*w0 4*w1 4*w2 4*w3 4*w4 ... 4*wd_model
# ...
# max_len m*w0 m*w1 m*w2 m*w3 m*w4 ... m*wd_model
# position 是 (max_len, ) 即 [0, 1, 2, ..., max_len]
# 每行都需要 d_model 哥 0, 要变成 (max_len, d_model) 个0
# 如果变成 (1, max_len), 那即使广播,也只能变成:
# [ [0, 1, 2, ..., max_len],
# [0, 1, 2, ..., max_len],
# ...
# [0, 1, 2, ..., max_len] ]
# 所以只能变成 (max_len, 1) -> 广播后:
# [ [0], [ [0, 0, 0, ..., 0],
# [1], [0, 0, 0, ..., 0],
# [2], [0, 0, 0, ..., 0],
# [3], [0, 0, 0, ..., 0],
# ..., ...,
# [max_len] ] [0, 0, 0, ..., 0] ]
position = position.unsqueeze(dim=1)
# 计算 w = 10000 ^ (- 2i/d) = exp( 2i * (-ln(10000) / d) )
div_term = torch.exp(
torch.arange(0, d_model, 2, dtype=torch.float) *
(-math.log(10000) / d_model)
)
print('div_term.shape =', div_term.shape) # torch.Size([34])
# div_term 需要变成 每行都是 [w0, w1, w2, w3, ..., w_d_model]
# 即:
# [ [w0, w1, w2, w3, ..., w_d_model],
# [w0, w1, w2, w3, ..., w_d_model],
# ...,
# [w0, w1, w2, w3, ..., w_d_model] ]
# 所以 只能变成 (1, d_model)
# 但是 position 已经变了, (max_len, 1) 与 (d_model, ) 运算
# 可以广播, (max_len, 1) -> (max_len, d_model)
# (1, d_model) -> (max_len, d_model)
# 所以 div_term 可以不进行升维,但也可以显示升维
# product 相当于算出来了 相位, 即 sin(x) 中的 x
product = position * div_term
print('product.shape =', product.shape) # torch.Size([1024, 34])
# 万事俱备,只差 sin、cos
# sin:
# 取所有行, 每行从0开始,取到末尾,步长为 2
pe[:, 0:: 2] = torch.sin(product)
# cos:
# 取所有行, 每行从1开始,取到末尾,步长为 2
pe[:, 1:: 2] = torch.cos(product)
print('pe.shape =', pe.shape) # torch.Size([1024, 68])
print(pe) # 只要 max_len 和 d_model 确定了,那 pe 就确定了额
# tensor([[ 0.0000e+00, 1.0000e+00, 0.0000e+00, ..., 1.0000e+00, 0.0000e+00, 1.0000e+00],
# [ 8.4147e-01, 5.4030e-01, 6.9087e-01, ..., 1.0000e+00, 1.3111e-04, 1.0000e+00],
# [ 9.0930e-01, -4.1615e-01, 9.9897e-01, ..., 1.0000e+00, 2.6223e-04, 1.0000e+00],
# ...,
# [ 9.9236e-01, -1.2340e-01, -1.1004e-02, ..., -6.7493e-01, 2.6875e-01, 9.6321e-01],
# [ 4.3234e-01, -9.0171e-01, 6.8216e-01, ..., -6.7505e-01, 2.6888e-01, 9.6317e-01],
# [-5.2517e-01, -8.5100e-01, 9.9852e-01, ..., -6.7518e-01, 2.6900e-01, 9.6314e-01]])
# 注册为 buffer
self.register_buffer(name='position_encoding', tensor=pe, persistent=True)
def forward(self, x):
T = x.size(dim=1) # 获取 T
print(f'T = {T}') # 3
# 为 True 则继续执行
assert T <= self.max_len, '序列长度超过最大序列长度, 位置编码失败!'
print(f'x.shape = {x.shape}') # torch.Size([2, 3, 68])
print(f'self.position_encoding.shape = {self.position_encoding.shape}') # torch.Size([1024, 68])
# 输入 x 可能只有 T=3 个 token(如 (2, 3, 68))。
# 所以不能直接加整个 pe,而应该只取前 T 行!
# x = x + self.position_encoding # 错误!
# 这里有讲究: position_encoding[: T] vs position_encoding[: T, : ]
# 在 《`pe[: T]` vs `pe[: T, : ]`》 里有详情
x = x + self.position_encoding[: T] # 这里有广播机制
return x
# 多头注意力
class MultiHeadAttn(nn.Module):
def __init__(
self,
d_model: int,
num_heads: int,
d_k: int = None, # query/key 的每个 head 维度
d_v: int = None, # value 的每个 head 维度
dropout: float = 0.1
):
super().__init__()
# 保证词嵌入维度 能整除 头数
# False 则报错,True 则正常运行
# 正确约束是:当 d_k 未指定时,d_model 必须能被 num_heads 整除(因为 d_k = d_model // num_heads)
# 一旦指定了 d_k,就没有任何整除要求!
# d_v 同理
if d_k is None or d_v is None:
assert d_model % num_heads == 0, f'词嵌入维度不能整除头数'
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_k if d_k is not None else d_model // num_heads
# 通常 d_v = d_k,但可以 d_v 不等于 d_k
self.d_v = d_v if d_v is not None else d_model // num_heads
print(f'd_k = {self.d_k}') # 64
print(f'd_v = {self.d_v}') # 64
# 把多个小矩阵拼接成大矩阵
self.W_q = nn.Linear(in_features=self.d_model, out_features=self.num_heads * self.d_k, bias=False)
self.W_k = nn.Linear(in_features=self.d_model, out_features=self.num_heads * self.d_k, bias=False)
self.W_v = nn.Linear(in_features=self.d_model, out_features=self.num_heads * self.d_v, bias=False)
# 在Softmax之后,在加权融合V之前
self.dropout = nn.Dropout(p=dropout)
# 加权融合V 得到的是 (token, h * d_v),所以这里的输入是 h * d_v
# 输出的线性投影, 拼接后的 (h * d_v) 需要映射回 d_model
self.W_o = nn.Linear(in_features=self.num_heads * self.d_v, out_features=self.d_model, bias=False)
def forward(
self,
query_input: torch.Tensor,
key_input: torch.Tensor,
value_input: torch.Tensor,
mask: torch.Tensor = None
):
B, T, D = query_input.shape # 需要保证 batch_first = True
# 为什么形参名字叫 query_input, key_input, value_input,在《关键:分清 X 到底是什么》有详情
# 线性变换, 矩阵是由多个heads小矩阵合并而成
Q = self.W_q(query_input) # (B, T, h * d_k)
K = self.W_k(key_input) # (B, T, h * d_k)
V = self.W_v(value_input) # (B, T, h * d_v)
# 拆分为多头 Q: (B, T, h * d_k) -> (B, T, h, d_k) -> (B, h, T, d_k)
# 为什么写 -1,详情在《拆分多头时:`Q.reshape(B, -1, self.num_heads, self.d_k)`》
Q_heads = Q.reshape(B, -1, self.num_heads, self.d_k).transpose(1, 2) # (B, h, T, d_k)
K_heads = K.reshape(B, -1, self.num_heads, self.d_k).transpose(1, 2) # (B, h, T, d_k)
V_heads = V.reshape(B, -1, self.num_heads, self.d_v).transpose(1, 2) # (B, h, T, d_v)
# 计算 缩放点击注意力(不是除以 h*d_k!) (B, h, T, T)
score = torch.matmul(Q_heads, K_heads.transpose(-1, -2)) / math.sqrt(self.d_k)
if mask is not None:
# masked_fill 返回新张量,不是 in_place 操作
score = score.masked_fill(mask == 0, -torch.inf)
# Softmax计算分布
attn_weight = F.softmax(score, dim=-1) # (B, h, T, T)
# 使用 dropout
attn_weight = self.dropout(attn_weight)
# 加权聚合 value 得到 上下文
context = torch.matmul(attn_weight, V_heads) # (B, h, T, d_v)
# 把所有 head 合并
# (B, h, T, d_v) --> (B, T, h, d_v) --> (B, T, h * d_v)
context = context.transpose(1, 2).reshape(B, -1, self.num_heads * self.d_v)
# 线性投影
output = self.W_o(context)
return output
# 前馈全连接层
class PositionWiseFeedForwardNetwork(nn.Module):
def __init__(
self,
d_model: int,
d_ff: int = None,
dropout: float = 0.1,
activation='relu'
):
super().__init__()
self.d_model = d_model
self.d_ff = d_ff if d_ff is not None else 4 * d_model # 默认 4 倍
self.linear1 = nn.Linear(in_features=d_model, out_features=self.d_ff)
if activation == 'relu':
self.activation = nn.ReLU()
else:
raise ValueError('还需要添加更多激活函数')
raise ValueError(f"Unsupported activation: '{activation}'. Currently only 'relu' is supported.")
self.dropout = nn.Dropout(p=dropout)
self.linear2 = nn.Linear(in_features=self.d_ff, out_features=self.d_model)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# 返回值预计为 张量(但Python解释器可不管什么类型,即使不返回张量也不会报错)
"""
:param x: x.shape = (B, T, d_model)
:return: (B, T, d_model)
"""
x = self.linear1(x)
x = self.activation(x)
x = self.dropout(x)
x = self.linear2(x)
# 可一行搞定
# return self.linear2(self.dropout(self.activation(self.linear1(x))))
return x
# 规范化层
class LayerNorm(nn.Module):
def __init__(self, d_model: int, eps=1e-5):
super().__init__()
self.d_model = d_model
self.eps = eps
# 仿射变换 γ,用于缩放(初始为1,表示不缩放)
self.gamma = nn.Parameter(torch.ones(size=(self.d_model,)))
# 仿射变换 β,用于偏移(初始为0,表示不偏移)
self.beta = nn.Parameter(torch.zeros(size=(self.d_model,)))
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
支持任意形状,只要最后一维是 d_model
例如: (B, d_model) 或 (B, T, d_model) 或 (B, T1, T2, d_model)
"""
# 计算均值
mean = x.mean(dim=-1, keepdim=True)
# 计算有偏方差
var = x.var(dim=-1, keepdim=True, unbiased=False)
# 标准化
x_norm = (x - mean) / (torch.sqrt(var + self.eps))
# 仿射变换
output = self.gamma * x_norm + self.beta
return output
# 子层连接结构
class SublayerConnection(nn.Module):
def __init__(self, d_model: int, dropout: float = 0.1):
super().__init__()
# 规范化层
self.norm = LayerNorm(d_model=d_model)
# Dropout
self.dropout = nn.Dropout(p=dropout)
def forward(self, x, sublayer_fn) -> torch.Tensor:
# 绝对不能写成 x = self.norm(x), 这样会覆盖原始的 x, 残差连接还要用到原始的 x
x_norm = self.norm(x) # 规范化层
sublayer_output = sublayer_fn(x_norm) # MHSA、FNN 等
output = x + self.dropout(sublayer_output) # 残差连接
# 一行搞定
# return x + self.dropout(sublayer_fn(self.norm(x)) )
return output
# 编码器层
class EncoderLayer(nn.Module):
def __init__(self, d_model: int, self_attn: nn.Module, ffn: nn.Module, dropout: float = 0.1):
"""
:param self_attn: 多头自注意力机制(MHSA)
:param ffn: 前馈全连接(FFN)
"""
super().__init__()
self.d_model = d_model
self.self_attn = self_attn
self.ffn = ffn
self.dropout = dropout
# 第一个子层 —— 多头自注意力(MHSA Sub-layer)
self.sublayer_attn = SublayerConnection(d_model=self.d_model, dropout=self.dropout)
# 第二个子层 —— 前馈网络(FFN Sub-layer)
self.sublayer_ffn = SublayerConnection(d_model=self.d_model, dropout=self.dropout)
def forward(self, x, mask=None):
"""
:param mask: 这里先不管 mask,后续会对 mask 进行详细的改进
:param x: x.shape = (B, T, D)
"""
# 第一个子层
x1 = self.sublayer_attn(x, lambda x_norm: self.self_attn(x_norm, x_norm, x_norm, mask))
# 第二个子层
# 如果 self.ffn 本身就是一个接受单个张量输入并返回张量的模块(如 nn.Sequential 或自定义 FFN),那么直接传 self.ffn 即可。
# 使用 lambda x_norm: self.ffn(x_norm) 是冗余的函数包装,增加调用开销,且无任何收益。
x2 = self.sublayer_ffn(x1, self.ffn)
return x2
class Encoder(nn.Module):
def __init__(self, encoder_layer: nn.Module, num_layers: int):
"""
:param encoder_layer: 已构建好的编码器层
:param num_layers: 编码器的层数
"""
super().__init__()
# num_layer 个编码器
self.layers = nn.ModuleList([copy.deepcopy(encoder_layer) for _ in range(num_layers)])
# 最终输出需通过规范化层
self.norm = LayerNorm(d_model=encoder_layer.d_model)
def forward(self, x, mask=None):
# 经过多个编码器层
for layer in self.layers:
x = layer(x, mask)
# 最终输出经过规范化层
output = self.norm(x)
return output
class DecoderLayer(nn.Module):
def __init__(
self,
d_model: int,
self_attn: nn.Module, # 掩码多头自注意力(用于第一个子层)
cross_attn: nn.Module, # 掩码多头交叉注意力(用于第二个子层)
ffn: nn.Module, # 前馈全连接(用于第三个子层)
dropout: float = 0.1
):
super().__init__()
self.d_model = d_model
self.dropout = dropout
self.self_attn = self_attn
self.cross_attn = cross_attn
self.ffn = ffn
# 第一个子层: Masked Multi-Head Self-Attention(掩码多头自注意力)
self.sublayer_self_attn = SublayerConnection(d_model=self.d_model, dropout=self.dropout)
# 第二个子层: Multi-Head Cross-Attention(编码器-解码器注意力)
self.sublayer_cross_attn = SublayerConnection(d_model=self.d_model, dropout=self.dropout)
# 第三个子层: Position-wise Feed-Forward Network(逐位置前馈网络)
self.sublayer_ffn = SublayerConnection(d_model=self.d_model, dropout=self.dropout)
def forward(
self,
x: torch.Tensor,
memory: torch.Tensor, # 编码器最终的输出
target_mask: torch.Tensor = None, # 第一个子层的掩码
memory_mask: torch.Tensor = None # 第二个子层的掩码
):
# 第一个子层
x1 = self.sublayer_self_attn(
x,
lambda x_norm: self.self_attn(x_norm, x_norm, x_norm, target_mask)
)
# 第二个子层
x2 = self.sublayer_cross_attn(
x1,
lambda x_norm: self.cross_attn(x_norm, memory, memory, memory_mask)
)
# 第三个子层
x3 = self.sublayer_ffn(
x2,
self.ffn # 只有一个参数,可以简写
)
return x3
class Decoder(nn.Module):
def __init__(self, decoder_layer: nn.Module, num_layers: int):
super().__init__()
# N个解码器层
self.layers = nn.ModuleList([
copy.deepcopy(decoder_layer) for _ in range(num_layers)
])
# 注意:Pre-LN 架构中
# 有的地方说在经过堆叠后,需要加 LayerNorm,有的地方说不需要加
# 所以具体加不加,看具体情况,如果遇到了,就去问领导
# 或者查看要复现的论文原文(比如 T5、BERT、GPT)
# 或者查看官方开源代码
self.norm = LayerNorm(d_model=decoder_layer.d_model)
def forward(self, x, memory, target_mask: torch.Tensor = None, memory_mask: torch.Tensor = None):
for layer in self.layers:
x = layer(x, memory, target_mask, memory_mask)
# 注意:Pre-LN 架构中
# 有的地方说在经过堆叠后,需要加 LayerNorm,有的地方说不需要加
# 所以具体加不加,看具体情况,如果遇到了,就去问领导
# 或者查看要复现的论文原文(比如 T5、BERT、GPT)
# 或者查看官方开源代码
x = self.norm(x)
return x
if __name__ == '__main__':
torch.manual_seed(66)
vocabulary_size = 500 # 500 词库
d_model = 512 # 词向量维度
max_len = 1000 # 预测序列最大长度为 1000
num_heads = 8 # 多头注意力的头数为 8
num_layers = 6 # 编码器 & 解码器 堆叠层数
# 编码器: 生成 [0, 500) 范围内的随机整数
x = torch.randint(low=0, high=vocabulary_size, size=(2, 4)) # 原始输入
# 转为词向量
my_embed = Embeddings(vocabulary_size=vocabulary_size, d_model=d_model)
x_embed = my_embed(x)
# 添加位置编码
my_position = PositionalEncoding(d_model=d_model, max_len=max_len)
x_position = my_position(x_embed)
# 第一个子层————多头自注意力层(MHSA Sub-layer)
encoder_my_attn = MultiHeadAttn(d_model=d_model, num_heads=num_heads)
# 第二个子层————前馈网络层(FFN Sub-layer)
encoder_my_ffn = PositionWiseFeedForwardNetwork(d_model=d_model)
# 构建编码器层
my_encoder_layer = EncoderLayer(d_model=d_model, self_attn=encoder_my_attn, ffn=encoder_my_ffn)
# 构建编码器
my_encoder = Encoder(encoder_layer=my_encoder_layer, num_layers=num_layers)
encoder_output = my_encoder(x=x_position)
print(f'encoder_output.shape = {encoder_output.shape}') # torch.Size([2, 4, 512])
# 解码器: 生成 [0, 500) 范围内的随机整数
y = torch.randint(low=0, high=vocabulary_size, size=(2, 4)) # 原始输入
# 转为词向量
y_embed = my_embed(y)
# 添加位置编码
y_position = my_position(y_embed)
# 解码器:第一个子层
decoder_self_attn = MultiHeadAttn(d_model=d_model, num_heads=num_heads)
# 解码器: 第二个子层
decoder_cross_attn = MultiHeadAttn(d_model=d_model, num_heads=num_heads)
# 解码器: 第三个子层
decoder_ffn = PositionWiseFeedForwardNetwork(d_model=d_model)
# 构建解码器层
my_decoder_layer = DecoderLayer(
d_model=d_model,
self_attn=decoder_self_attn,
cross_attn=decoder_cross_attn,
ffn=decoder_ffn
)
# 构建解码器
my_decoder = Decoder(decoder_layer=my_decoder_layer, num_layers=num_layers)
decoder_output = my_decoder(x=y_position, memory=encoder_output)
print(f'decoder_output.shape = {decoder_output.shape}') # torch.Size([2, 4, 512])
经典代码:
class Decoder(nn.Module):
def __init__(
self,
decoder_layer: nn.Module,
num_layers: int
):
"""
Pre-LN Transformer 解码器主体(不含嵌入、位置编码、输出投影)
输入应为已嵌入并加好位置编码的张量。
:param decoder_layer: 单个 DecoderLayer 实例
:param num_layers: 解码器层数
"""
super().__init__()
self.layers = nn.ModuleList([
copy.deepcopy(decoder_layer) for _ in range(num_layers)
])
# 注意:Pre-LN 架构中
# 有的地方说在经过堆叠后,需要加 LayerNorm,有的地方说不需要加
# 所以具体加不加,看具体情况,如果遇到了,就去问领导
# 或者查看要复现的论文原文(比如 T5、BERT、GPT)
# 或者查看官方开源代码
# self.norm = LayerNorm(d_model=decoder_layer.d_model)
def forward(
self,
x: torch.Tensor, # ← 已嵌入的输入 (B, T, d_model)
memory: torch.Tensor, # 编码器输出 (B, S, d_model)
tgt_mask: torch.Tensor = None,
memory_mask: torch.Tensor = None
):
"""
:param x: 目标序列的嵌入表示(已包含位置编码),shape = (B, T, d_model)
:param memory: 编码器输出, shape = (B, S, d_model)
:param tgt_mask: 自注意力掩码(如因果掩码)
:param memory_mask: 交叉注意力掩码(如源 padding 掩码)
:return: (B, T, d_model)
"""
for layer in self.layers:
x = layer(
x,
memory,
target_mask=tgt_mask,
memory_mask=memory_mask
)
# 注意:Pre-LN 架构中
# 有的地方说在经过堆叠后,需要加 LayerNorm,有的地方说不需要加
# 所以具体加不加,看具体情况,如果遇到了,就去问领导
# 或者查看要复现的论文原文(比如 T5、BERT、GPT)
# 或者查看官方开源代码
# x = self.norm(x)
return x
更多推荐


所有评论(0)