衔接前序:第 36 篇完成了编码器的代码实战——我们逐步骤手动验证了位置编码、多头自注意力、前馈网络的数值正确性。本文是"代码实战三部曲"的第二篇,在前文编码器的基础上,实现解码器的三大核心机制:因果自注意力(确保序列生成的因果性)、交叉注意力(编码器-解码器的信息融合)、自回归推理(逐步生成目标序列)。

建议运行配套 notebook 37_transformer_decoder.ipynb 边看边读,效果最佳。


一、概述

本文的目标是:在编码器基础上,完整实现 Transformer 解码器,并通过手动验证理解交叉注意力如何融合编码信息。

我们将依次完成:

  1. 数据准备与编码器前向:加载平行语料,运行编码器获取编码表示
  2. 因果自注意力(核心):手动验证因果掩码如何阻止未来信息的泄露,逐行对比掩码前后的得分矩阵
  3. 交叉注意力(核心):演示 Q 来自解码器、K/V 来自编码器的"非对称"注意力机制
  4. 完整解码器前向:3 层解码器堆叠,逐层对比自注意力和交叉注意力的模式
  5. 自回归推理:实现贪心解码,逐步观察解码器如何"一个 token 接一个 token"地生成序列

模型配置

与第 36 篇保持一致:

参数
d_model 32
注意力头数 h 4
每头维度 d_k 8
FFN 隐藏层 d_ff 128
编码器/解码器层数 N 3

代码组织

本文在 transformer_modules.py 基础上新增解码器相关组件:

组件 类/函数 所在文件
解码器层(含交叉注意力) DecoderLayer transformer_modules.py
完整解码器 Decoder transformer_modules.py
完整 Transformer Transformer transformer_modules.py
因果多头注意力 MultiHeadAttention(causal=True) transformer_modules.py

导入依赖:

import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from data_utils import build_vocab, load_parallel_data, encode_line
from transformer_modules import (
    PositionalEncoding, MultiHeadAttention, FeedForward,
    EncoderLayer, Encoder, DecoderLayer, Decoder
)

二、数据准备与编码器前向

2.1 加载语料、构建词表

沿用第 36 篇的数据流程,加载 500 个中英平行句对:

ZH vocab size: 858      # 中文字符(含标点)
EN vocab size: 53       # 英文字母 + 空格 + 标点

编码测试句子"我爱深度学习":

Source: 我爱深度学习
Token IDs: [2, 15, 1, 374, 153, 306, 422, 3]
Tokens: ['<sos>', '我', '<unk>', '深', '度', '学', '习', '<eos>']

注意:由于随机种子不同,这里的 token IDs 与第 36 篇不完全一致,但含义相同。

2.2 编码器前向

将"我爱深度学习"送入第 36 篇实现的编码器,得到编码表示:

Encoder output shape: torch.Size([1, 8, 32])
Number of encoder layers: 3

enc_output 的形状为 (batch=1, seq_len=8, d_model=32),这就是解码器交叉注意力中 K 和 V 的来源。


三、因果自注意力 ⭐

3.1 为什么需要因果掩码?

解码器在生成序列时,有一个基本原则:在预测第 t 个 token 时,不能看到第 t+1 个及之后的 token。这是因为解码器是以自回归方式工作的——当前步的输出会成为下一步的输入。如果解码器能"偷看"未来的 token,训练就会退化。

因果掩码(Causal Mask)通过在注意力得分矩阵的上三角区域填入 -∞ 来实现这一约束:

MaskedAttention(Q,K,V)=softmax(QK⊤dk+M)V\text{MaskedAttention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right) VMaskedAttention(Q,K,V)=softmax(dk QK+M)V

其中 Mij=0M_{ij} = 0Mij=0i≥ji \geq jij(关注自身和之前的位置),Mij=−∞M_{ij} = -\inftyMij=i<ji < ji<j(屏蔽未来位置)。

3.2 因果掩码可视化

下图展示了目标序列"i love deep learning"的因果掩码(序列长度 19):
在这里插入图片描述

左图:数值形式的掩码,黑色 = -∞(屏蔽),白色 = 0(允许)。

右图:蓝色表示被屏蔽的位置。可以清楚看到所有 (i,j)(i, j)(i,j)i<ji < ji<j 的位置都被屏蔽了——query 位置 0 只能看 key 位置 0,query 位置 1 只能看 key 位置 0-1,以此类推。

3.3 手动验证:掩码前后的得分变化

我们取 Head 0 的前 3 个 query,分别输出掩码前后的得分矩阵:

掩码前(原始得分,前 3 query × 全部 19 key):

[[-0.2195  0.1671 -0.0728  0.0439 -0.2528 -0.1236  0.0291 ... -0.2743 -0.1134]
 [-0.1086  0.0517  0.2660 -0.1176  0.0343  0.0568 -0.1378 ... -0.1317 -0.1844]
 [-0.3401 -0.1102  0.7379 -0.5971 -0.1663  0.3420 -0.5252 ... -0.6013 -0.2787]]

掩码后(前 3 query × 全部 19 key):

[[-0.2195    -inf    -inf    -inf    -inf    -inf    -inf ...    -inf    -inf]
 [-0.1086  0.0517    -inf    -inf    -inf    -inf    -inf ...    -inf    -inf]
 [-0.3401 -0.1102  0.7379    -inf    -inf    -inf    -inf ...    -inf    -inf]]

注意 query 1 只能看到 key 0-1,query 2 只能看到 key 0-2,其余位置全部被 -inf 替代。经过 softmax 后,-inf 位置的权重变为 0。

3.4 Query 位置 3 的注意力分布

具体来看 query 位置 3(对应 token “o”)的注意力权重:

Key 位置 Token 注意力权重 状态
0 <sos> 0.1444 有效
1 i 0.2955 有效
2 l 0.2723 有效
3 o 0.2878 有效
4-18 v, e, d, … 0.0000 被因果掩码阻断

三个有效位置的权重和为 0.1444+0.2955+0.2723+0.2878=1.00.1444 + 0.2955 + 0.2723 + 0.2878 = 1.00.1444+0.2955+0.2723+0.2878=1.0,而位置 4-18 的权重全部为 0。

手动计算的注意力权重与 MultiHeadAttention(return_attention=True) 的输出对比:

Max diff between manual and module: 0.00e+00
Row sums (should all be 1.0): [1.0, 1.0, 1.0, 1.0, 1.0]...
Upper triangle (blocked) entries: 171 / 361
Lower triangle (valid) entries: 190

差异为 0——因为这里没有 softmax 的浮点误差累积(被屏蔽的位置精确为 0),手动计算和模块输出完全一致。

3.5 四头因果自注意力可视化

在这里插入图片描述

上图展示了 4 个头的因果自注意力权重矩阵(19×19=361 entries)。所有头的上三角区域全为 0(纯白),这是因果掩码的直接效果。

每个头的下三角区域中,对角线附近的权重相对较高(较深的蓝色),说明在随机初始化状态下,每个 token 倾向于关注自身和邻近位置。不同头的"亮点"分布略有不同,体现了多头分化的雏形。


四、交叉注意力 ⭐

4.1 核心概念

交叉注意力(Cross-Attention)是编码器-解码器架构的核心创新,也是 Transformer 能够完成"序列到序列"任务的关键。

与自注意力不同,交叉注意力的 Q 来自解码器(目标语言),而 K 和 V 来自编码器(源语言):

CrossAttention(Qdec,Kenc,Venc)=softmax(QdecKenc⊤dk)Venc\text{CrossAttention}(Q_{\text{dec}}, K_{\text{enc}}, V_{\text{enc}}) = \text{softmax}\left(\frac{Q_{\text{dec}} K_{\text{enc}}^\top}{\sqrt{d_k}}\right) V_{\text{enc}}CrossAttention(Qdec,Kenc,Venc)=softmax(dk QdecKenc)Venc

MultiHeadAttention 的实现中,自注意力和交叉注意力使用的是同一个类——区别仅在于传入的 Q/K/V 参数不同:

# 自注意力:Q=K=V=x
self_attn_out = self_attn(x, x, x)

# 交叉注意力:Q=decoder, K=V=encoder
cross_attn_out = cross_attn(dec_x, enc_output, enc_output)

交叉注意力不使用因果掩码——解码器的每个位置都可以关注编码器的所有位置。

4.2 维度追踪

Q (decoder) shape:           torch.Size([1, 19, 32])   # 目标语言,19个字符
K, V (encoder output) shape: torch.Size([1, 8, 32])    # 源语言,8个字符
Cross attention output shape: torch.Size([1, 19, 32])   # 输出维度与Q一致
Cross attention weights:     torch.Size([1, 4, 19, 8]) # 4个头,19×8的注意力矩阵

关键观察:交叉注意力的权重矩阵形状是 (19, 8),而不是 (8, 19)。这是因为 query 数量由解码器决定(19),key 数量由编码器决定(8)。每个解码器位置都有一个长度为 8 的注意力分布,表示它如何关注源语言的 8 个字符。

4.3 手动验证:解码器位置 1 的源语言关注

取 Head 0,解码器位置 1(第一个实词"i")在源语言 8 个字符上的注意力分布:

Cross-attention scores (Head 0, before softmax):
  Shape: [1, 19, 8]
  Range: [-0.7750, 0.9925]
  Mean: -0.0102

Cross-attention: Decoder position 1 -> Source positions (Head 0):
  "我": 0.1350 ██████
  "爱": 0.1086 █████
  "深": 0.1074 █████
  "度": 0.1253 ██████
  "学": 0.1239 ██████
  "习": 0.1132 █████
  <?>: 0.1165 █████
  <?>: 0.1700 ████████

这里有两个<?>是编码器端的 <sos><eos> token。由于模型是随机初始化的,注意力分布相对均匀(所有权重在 0.10-0.17 之间,没有明显聚焦),但这正好说明了交叉注意力的信息来源——解码器的每个位置都"看到了"编码器的所有位置

手动计算与模块输出的对比:

Max diff manual vs module: 0.00e+00
Cross-attention row sums (should all be 1.0): [1.0, 1.0, 1.0, 1.0, 1.0]...

完全一致。

4.4 四头交叉注意力可视化

在这里插入图片描述

上图中每个子图是一个 19×8 的矩阵,行对应解码器的 19 个目标 token(含 SOS/EOS),列对应编码器的 8 个源 token(含 SOS/EOS)。

与自注意力热力图不同,这里没有对角线结构——因为交叉注意力不是"自己注意自己",而是"目标注意源"。每个解码器位置都可以关注所有编码器位置。4 个头的注意力模式略有不同,体现了多头注意力从不同角度提取源语言信息的能力。

4.5 编码器自注意力 vs 解码器交叉注意力

在这里插入图片描述

上下两排形成鲜明对比:

  • 上排(编码器自注意力):8×8 的方阵,Q=K=V 都来自源语言。矩阵沿对角线对称(自注意力中,query i 对 key j 的注意力与 query j 对 key i 的注意力由不同的 query/key 投影计算,所以不完全对称,但倾向于对角线),对角线较亮。
  • 下排(解码器交叉注意力):19×8 的矩形,Q 来自目标语言,K=V 来自源语言。没有对角线结构,行与行之间相对独立。

这种对比直观地展示了两种注意力的本质区别:自注意力建模序列内部的依赖关系,交叉注意力建模序列之间的对齐关系


五、完整解码器

5.1 解码器架构

完整的 Transformer 解码器由 N 个解码器层堆叠而成,每层包含三个子层:

  1. 因果自注意力:带因果掩码的多头自注意力,确保因果性
  2. 残差连接 + LayerNorm
  3. 交叉注意力:Q 来自解码器,K/V 来自编码器,融合源语言信息
  4. 残差连接 + LayerNorm
  5. 前馈网络:升维 → ReLU → 降维
  6. 残差连接 + LayerNorm

最后通过线性投影将 d_model=32 维映射到词表大小(53),得到每个位置对所有 token 的 logits。

完整实现如下:

class DecoderLayer(nn.Module):
    """单个解码器层:因果自注意力 → 残差+LN → 交叉注意力 → 残差+LN → FFN → 残差+LN。"""
    def __init__(self, d_model, h, d_ff):
        super().__init__()
        self.self_attn = MultiHeadAttention(d_model, h, causal=True)
        self.cross_attn = MultiHeadAttention(d_model, h)
        self.ln1 = LayerNorm(d_model)
        self.ln2 = LayerNorm(d_model)
        self.ln3 = LayerNorm(d_model)
        self.ffn = FeedForward(d_model, d_ff)

    def forward(self, x, enc_output, src_mask=None, tgt_mask=None, return_attention=False):
        if return_attention:
            self_out, self_attn_w = self.self_attn(x, x, x, mask=tgt_mask, return_attention=True)
        else:
            self_out = self.self_attn(x, x, x, mask=tgt_mask)
        x = self.ln1(x + self_out)

        if return_attention:
            cross_out, cross_attn_w = self.cross_attn(
                x, enc_output, enc_output, mask=src_mask, return_attention=True
            )
        else:
            cross_out = self.cross_attn(x, enc_output, enc_output, mask=src_mask)
        x = self.ln2(x + cross_out)

        x = self.ln3(x + self.ffn(x))

        if return_attention:
            return x, self_attn_w, cross_attn_w
        return x


class Decoder(nn.Module):
    """完整解码器:嵌入 → 位置编码 → N 层解码器层 → 线性投影。"""
    def __init__(self, vocab_size, d_model, h, d_ff, N, max_len=5000):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, d_model)
        self.pe = PositionalEncoding(d_model, max_len)
        self.layers = nn.ModuleList([
            DecoderLayer(d_model, h, d_ff) for _ in range(N)
        ])
        self.output_proj = nn.Linear(d_model, vocab_size, bias=True)
        self.d_model = d_model

    def forward(self, x, enc_output, src_mask=None, tgt_mask=None, return_attention=False):
        x = self.embedding(x) * math.sqrt(self.d_model)
        x = self.pe(x)

        self_attentions = []
        cross_attentions = []
        for layer in self.layers:
            if return_attention:
                x, self_attn_w, cross_attn_w = layer(
                    x, enc_output, src_mask=src_mask, tgt_mask=tgt_mask,
                    return_attention=True
                )
                self_attentions.append(self_attn_w)
                cross_attentions.append(cross_attn_w)
            else:
                x = layer(x, enc_output, src_mask=src_mask, tgt_mask=tgt_mask)

        logits = self.output_proj(x)
        if return_attention:
            return logits, self_attentions, cross_attentions
        return logits

与编码器层相比,解码器层的区别在于:

  • 使用 causal=TrueMultiHeadAttention 做自注意力
  • 多了一个交叉注意力子层,接收 enc_output 作为 K/V
  • 3 个 LayerNorm(每个子层各一个),而非编码器的 2 个

5.2 维度追踪

将目标序列"i love deep learning"(已加 SOS/EOS)送入解码器:

Target sequence:               torch.Size([1, 19])
Encoder output:                torch.Size([1, 8, 32])
Decoder logits:                torch.Size([1, 19, 53])
Self-attention per layer:      torch.Size([1, 4, 19, 19])
Cross-attention per layer:     torch.Size([1, 4, 19, 8])
  • 输入:19 个 token(目标语言)
  • 编码器输出:8 个 token 的编码表示(源语言)
  • 解码器输出:19 × 53 的 logits(53 是目标语言词表大小)
  • 每层自注意力:19×19(含因果掩码)
  • 每层交叉注意力:19×8(无因果掩码)

注意 logits 的维度[1, 19, 53] 中的 53 是英文词表大小。这意味着每个解码器位置都输出了一个长度为 53 的向量,其中第 i 个分量表示第 i 个 token 的"得分"。通过 argmax 取最高分对应的 token,就是当前步的预测。

5.3 逐层注意力对比

在这里插入图片描述

上图展示了 3 层解码器中,第 0 个头的自注意力和交叉注意力随层数的变化:

自注意力列(左)

  • 第 1 层的自注意力分布相对多样,下三角区域中有明显的亮点分布
  • 第 2-3 层的分布略有变化,但整体模式与第 1 层相似(这也是随机初始化下的特征,训练后才会出现明显的层间分化)

交叉注意力列(右)

  • 所有 3 层的交叉注意力都是 19×8 的矩形
  • 各层之间模式有差异,但没有哪一层出现特别极端的分布(如某个源位置独占所有权重)
  • 这说明在随机初始化状态下,每层都在"平均地"关注源语言的所有位置

六、自回归推理

6.1 贪心解码

前面所有的解码器前向都是在 teacher forcing 模式——即我们同时输入了完整的目标序列(含正确答案)。但在实际推理时,我们需要自回归生成:从一个 SOS token 开始,每步预测一个 token,将上一步的输出作为下一步的输入。

贪心解码的实现如下:

def greedy_decode(decoder, encoder, src_ids, en_i2t, max_len=20):
    decoder.eval()
    encoder.eval()
    with torch.no_grad():
        enc_out = encoder(src_ids)

        # 从 SOS token 开始
        tgt = torch.tensor([[en_spec['sos_idx']]])
        generated = ['<sos>']

        for step in range(max_len):
            logits = decoder(tgt, enc_out)
            next_token = logits[0, -1].argmax().item()
            generated.append(en_i2t.get(next_token, '<unk>'))

            if next_token == en_spec['eos_idx']:
                break

            # 将当前步的输出拼接到输入中
            tgt = torch.cat([tgt, torch.tensor([[next_token]])], dim=1)

    return ''.join(generated[1:-1])  # 去掉 SOS 和 EOS

关键点在于 tgt = torch.cat([tgt, torch.tensor([[next_token]])], dim=1)——每一步都将预测的 token 拼接到已有的序列末尾,作为下一步解码器的输入。

在未训练的模型上运行:

Source: 我爱深度学习
Decoded (untrained model, expect garbage): "lcyl3]r]r]r]r])u[j8"

输出是随机字符,这是预期行为——未训练的模型只是将随机权重映射到随机 token。这正是第 38 篇训练篇要解决的问题。

6.2 逐步骤解码 + 交叉注意力追踪

一个更有意义的观察是解码过程中交叉注意力的变化。我们实现逐步骤解码,并输出每一步的交叉注意力分布:

Source: 我<unk>深度学习

Step 0: predict "l" (prob=0.061)  [cross-attn: 我=0.13 | <unk>=0.10 | 深=0.11 | 度=0.13 | 学=0.10 | 习=0.08]
Step 1: predict "c" (prob=0.057)  [cross-attn: 我=0.07 | <unk>=0.14 | 深=0.10 | 度=0.13 | 学=0.07 | 习=0.09]
Step 2: predict "y" (prob=0.058)  [cross-attn: 我=0.12 | <unk>=0.11 | 深=0.15 | 度=0.11 | 学=0.14 | 习=0.14]
Step 3: predict "l" (prob=0.041)  [cross-attn: 我=0.17 | <unk>=0.12 | 深=0.12 | 度=0.12 | 学=0.14 | 习=0.13]

每一步输出的 cross-attn 是第 0 层 Head 0 在最后 query 位置对源语言 6 个有效 token 的注意力分布。虽然模型未训练,预测结果无意义,但我们可以观察到:

  1. 注意力动态变化:每一步的交叉注意力分布都不相同——即使概率很低,解码器在选择下一个 token 时"看"源语言的方式在变化
  2. 预测概率很低:所有步的 prob 在 0.04-0.06 之间,远低于 1.0——因为模型对任何 token 都没有信心,53 个 token 的均匀分布期望概率是 1/53 ≈ 0.019,实际略高于均匀分布
  3. 分布相对均匀:交叉注意力权重在 0.07-0.17 之间,没有某个源 token 被特别关注——随机初始化的特征

经过训练后(第 38 篇会展示),这些交叉注意力权重会演变为有意义的对齐模式,如"i"对应"我"、“love"对应"爱”、“learning"对应"学习”。


七、总结

通过本文的代码实战,我们完成了 Transformer 解码器的完整实现。核心发现如下:

1. 因果自注意力

  • 上三角掩码将未来位置的得分置为 -∞,softmax 后权重精确为 0
  • 手动验证与模块输出差异为 0.00e+00(完全一致)
  • 序列长度为 19 时,171/361 个条目被屏蔽,下三角 190 个有效条目

2. 交叉注意力

  • Q 的维度由解码器决定 (batch, tgt_seq, d_model),K/V 的维度由编码器决定 (batch, src_seq, d_model)
  • 注意力权重矩阵形状为 (tgt_seq, src_seq)——矩形而非方形
  • 不使用因果掩码,解码器的每个位置都可以关注编码器的所有位置
  • 与编码器自注意力形成鲜明对比:自注意力是"自己注意自己"的方阵,交叉注意力是"目标注意源"的矩形

3. 完整解码器

  • 每层包含 3 个子层(因果自注意力 + 交叉注意力 + FFN),比编码器多一个交叉注意力
  • 3 个 LayerNorm 确保每层输出的数值稳定
  • 最终输出为词表大小的 logits,通过 argmax 或采样得到预测 token

4. 自回归推理

  • 从 SOS 开始,每步预测一个 token 并拼接到输入中
  • 未训练模型输出随机字符——引出了第 38 篇训练篇的必要性
  • 交叉注意力在每步的动态变化展示了"解码器在生成过程中不断重新关注源语言"
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐