深度学习的数学原理(三十七)—— Transformer 解码器代码实战
衔接前序:第 36 篇完成了编码器的代码实战——我们逐步骤手动验证了位置编码、多头自注意力、前馈网络的数值正确性。本文是"代码实战三部曲"的第二篇,在前文编码器的基础上,实现解码器的三大核心机制:因果自注意力(确保序列生成的因果性)、交叉注意力(编码器-解码器的信息融合)、自回归推理(逐步生成目标序列)。
建议运行配套 notebook
37_transformer_decoder.ipynb边看边读,效果最佳。
一、概述
本文的目标是:在编码器基础上,完整实现 Transformer 解码器,并通过手动验证理解交叉注意力如何融合编码信息。
我们将依次完成:
- 数据准备与编码器前向:加载平行语料,运行编码器获取编码表示
- 因果自注意力(核心):手动验证因果掩码如何阻止未来信息的泄露,逐行对比掩码前后的得分矩阵
- 交叉注意力(核心):演示 Q 来自解码器、K/V 来自编码器的"非对称"注意力机制
- 完整解码器前向:3 层解码器堆叠,逐层对比自注意力和交叉注意力的模式
- 自回归推理:实现贪心解码,逐步观察解码器如何"一个 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(dkQK⊤+M)V
其中 Mij=0M_{ij} = 0Mij=0 当 i≥ji \geq ji≥j(关注自身和之前的位置),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(dkQdecKenc⊤)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 个解码器层堆叠而成,每层包含三个子层:
- 因果自注意力:带因果掩码的多头自注意力,确保因果性
- 残差连接 + LayerNorm
- 交叉注意力:Q 来自解码器,K/V 来自编码器,融合源语言信息
- 残差连接 + LayerNorm
- 前馈网络:升维 → ReLU → 降维
- 残差连接 + 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=True的MultiHeadAttention做自注意力 - 多了一个交叉注意力子层,接收
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 的注意力分布。虽然模型未训练,预测结果无意义,但我们可以观察到:
- 注意力动态变化:每一步的交叉注意力分布都不相同——即使概率很低,解码器在选择下一个 token 时"看"源语言的方式在变化
- 预测概率很低:所有步的
prob在 0.04-0.06 之间,远低于 1.0——因为模型对任何 token 都没有信心,53 个 token 的均匀分布期望概率是 1/53 ≈ 0.019,实际略高于均匀分布 - 分布相对均匀:交叉注意力权重在 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 篇训练篇的必要性
- 交叉注意力在每步的动态变化展示了"解码器在生成过程中不断重新关注源语言"
更多推荐




所有评论(0)