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 是目标序列长度。

  1. 初始嵌入与位置编码
    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}}} PRT×dmodel 是位置编码。

    💡 在大多数实现中(包括 Pre-LN 架构),位置编码直接加到词嵌入上,作为解码器第一层的输入。

  2. 逐层传播
    对于 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(l1),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}}} HencRS×dmodel 经线性变换后的 Key 和 Value,维度为 S × d k S \times d_k S×dk S S S 为源序列长度)。

  3. 最终输出
    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} yt 的自回归上下文
    • 与源序列全局对齐后的语义信息

✅ 此时我们停止。不进行 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}}} XRT×dmodel:上一层输出(或初始嵌入)

步骤详解

  1. 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 的核心标志。

  2. 线性投影生成 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,WVRdmodel×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}}} WOselfRdmodel×dmodel 投影回原维度。
  3. 因果掩码(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={0if ijif i<j

    • 应用于注意力分数:
      A = softmax ( Q K ⊤ d k + M ) A = \text{softmax}\left( \frac{QK^\top}{\sqrt{d_k}} + M \right) A=softmax(dk QK+M)

    • 效果:位置 i i i 只能关注 j ≤ i j \leq i ji,保证自回归性质。

  4. 注意力输出与投影
    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

  5. 残差连接
    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}}} HencRS×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

步骤详解

  1. 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)

  2. 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 是固定的(来自编码器)
  3. 无掩码注意力(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(dk QcrossKenc)

    • 解码器每个位置可关注源序列所有位置(无因果限制)
  4. 输出与投影
    CrossAttnOut = A cross V enc W O cross \text{CrossAttnOut} = A_{\text{cross}} V_{\text{enc}} W_O^{\text{cross}} CrossAttnOut=AcrossVencWOcross

  5. 残差连接
    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,通常具有“扩张-压缩”结构:

  1. 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)

  2. 第一线性层(扩张)
    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),W1Rdmodel×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 等
  3. 第二线性层(压缩)
    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,W2Rdff×dmodel

  4. 残差连接
    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-LNPost-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×24.72M
  • LayerNorm: 3 × 2 × 768 ≈ 4.6   K 3 \times 2 \times 768 \approx 4.6\,\text{K} 3×2×7684.6K
  • 单层总计 ≈ 8.26M 参数
  • 12 层解码器 ≈ 99M 参数(不含嵌入和输出层)

七、总结:Pre-LN 解码器主体的核心要点

  1. 堆叠结构 N N N 个 identical decoder layers 串联,每层独立参数。
  2. 三层流水线:掩码自注意力 → 交叉注意力 → FFN。
  3. Pre-LN 范式:每个子模块前先 LayerNorm,再计算,最后残差加原始输入。
  4. 信息流
    • 自注意力捕获目标序列历史
    • 交叉注意力注入源序列上下文
    • FFN 提供非线性表达与知识存储
  5. 训练优势:Pre-LN 支持更深网络、更大 LR、无需 warmup,已成为工业标准。
  6. 输出:得到 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

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐