目标序列为何既要看过去又要看源文?Transformer Decoder 的三条信息流
从条件生成的信息边界出发,组装因果自注意力、交叉注意力与 FFN,并拆解并行训练和逐 token 生成。
上一篇把自注意力、逐位置 FFN、残差和 LayerNorm 组成了 Transformer Encoder block。它能让源序列的每个 token 双向交换信息,却不能直接解决翻译这类问题:生成第 个目标 token 时,模型既要读取已生成的目标前缀,又要回到整条源序列取证。
Transformer Decoder 用三个子层分开这三种职责:因果自注意力读目标前缀,交叉注意力(Cross-Attention)读编码器记忆,FFN 在每个目标位置内加工特征。本文只追踪这三条信息流以及训练—推理边界;下一篇再专门解决逐 token 生成的重复计算。
01 只复制 Encoder block,会泄漏什么?#
设源 token 长度为 ,目标 token 长度为 ,模型宽度为 :
编码器允许源序列内双向注意,输出记忆(Encoder Memory):
若目标侧也用无 mask 的双向自注意力,位置 就能读到标签中的 。训练损失会虚假地降低,而真正生成时这些未来 token 根本不存在。
因此解码器的第一条边界是:
而不是 。因果 mask 不是可选的正则化,而是把训练信息集合限制成推理时真实可见信息的契约。
02 三个子层分别读什么?#
以 Pre-LN 解码块为例,记子层输入为 :
目标前缀 Y [N,T,D]
│
├─ LN1 ─► causal self-attention ─► 残差相加 ─► U [N,T,D]
│ 只读目标位置 <= t
│
├─ LN2 ─► Query [N,T,D] ─┐
│ ├─ cross-attention ─► 残差相加 ─► V
源记忆 M [N,S,D] ─► Key/Value ─┘ 可读所有真实源位置
│
└─ LN3 ─► FFN: D → F → D ─► 残差相加 ─► Z [N,T,D]text最容易混淆的是 Query、Key、Value 的来源:
| 子层 | Query 来源 | Key 来源 | Value 来源 | 关系矩阵 |
|---|---|---|---|---|
| 目标 causal self-attention | 目标表示 | 目标表示 | 目标表示 | [N,H,T,T] |
| 源—目标 cross-attention | 目标表示 | 编码器记忆 | 编码器记忆 | [N,H,T,S] |
| 逐位置 FFN | 无 Q/K/V | 无 | 无 | 不混合位置 |
交叉注意力的输出长度是 Query 长度 ,不是源长度 。它为每个目标位置从 个源 Value 中取回一个 维向量。
03 两种 mask 分别屏蔽哪条边?#
解码器通常同时需要两类布尔 mask:
tgt_causal_mask [T,T]:第 行禁止读取所有 的目标 Key。memory_key_padding_mask [N,S]:每个样本禁止读取源序列的 padding Key。
目标因果 mask(✓=可读,×=禁止)
Query\Key <bos> y1 y2 y3
<bos> ✓ × × ×
y1 ✓ ✓ × ×
y2 ✓ ✓ ✓ ×
y3 ✓ ✓ ✓ ✓
源 padding mask(对每个目标 Query 都屏蔽 PAD 列)
Source key x1 x2 <pad>
任意目标位置 ✓ ✓ ×texttgt_key_padding_mask [N,T] 还可用来阻止真实目标 Query 读取目标 padding 列。它不会自动清空 padding 查询行;计算 token 损失时仍要用 ignore_index 或等价 mask 排除 padding 标签。
04 用一个 Query 手算交叉注意力#
只看一个头,令 。某目标位置的 Query 与两个源 Key 为:
缩放分数均为 ,因此 softmax 权重是 [0.5,0.5]。再令:
取回的源信息为:
若第二个源位置是 padding,将它的分数改为 ,权重变为 [1,0],输出就是 [2,0]。这个例子也说明:交叉注意力可以同时取回多个源位置的混合证据,它不是必然选中单个对齐词。
05 为什么训练要把目标序列错开一位?#
设真实目标为:
target: [我, 喜欢, 机器, 学习, <eos>]
decoder_input: [<bos>, 我, 喜欢, 机器, 学习]
prediction_for: [我, 喜欢, 机器, 学习, <eos>]text解码器输入右移一位(Shift Right)后,位置 使用真实前缀预测下一个 token。这是 Teacher Forcing:它允许训练时一次并行计算所有位置,但因果 mask 仍保证每个位置看不见右侧答案。
若不错开一位,把同一 token 同时作为该位置输入和标签,模型可通过残差路径轻易复制当前 token,学到的不是下一 token 分布。
06 不调用 Decoder 封装,组装一个 Pre-LN 解码块#
import torch
from torch import nn
class PreNormDecoderBlock(nn.Module):
def __init__(self, model_dim, num_heads, ffn_dim, dropout=0.1):
super().__init__()
assert model_dim % num_heads == 0
self.norm1 = nn.LayerNorm(model_dim)
self.norm2 = nn.LayerNorm(model_dim)
self.norm3 = nn.LayerNorm(model_dim)
self.self_attn = nn.MultiheadAttention(
model_dim, num_heads, dropout=dropout, batch_first=True
)
self.cross_attn = nn.MultiheadAttention(
model_dim, num_heads, dropout=dropout, batch_first=True
)
self.ffn = nn.Sequential(
nn.Linear(model_dim, ffn_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(ffn_dim, model_dim),
)
self.drop1 = nn.Dropout(dropout)
self.drop2 = nn.Dropout(dropout)
self.drop3 = nn.Dropout(dropout)
def forward(
self,
target, # [N,T,D]
memory, # [N,S,D]
target_valid, # [N,T], True = 真实 token
memory_valid, # [N,S], True = 真实 token
):
n, target_len, width = target.shape
assert memory.shape[0] == n and memory.shape[2] == width
assert target_valid.shape == (n, target_len)
assert memory_valid.shape == memory.shape[:2]
# MHA 的布尔 attn_mask=True 表示“禁止读取”。
causal_block = torch.ones(
target_len, target_len, dtype=torch.bool, device=target.device
).triu(diagonal=1)
qkv = self.norm1(target)
self_out, _ = self.self_attn(
qkv, qkv, qkv,
attn_mask=causal_block,
key_padding_mask=~target_valid,
need_weights=False,
)
target = target + self.drop1(self_out)
query = self.norm2(target)
cross_out, _ = self.cross_attn(
query, memory, memory,
key_padding_mask=~memory_valid,
need_weights=False,
)
target = target + self.drop2(cross_out)
target = target + self.drop3(self.ffn(self.norm3(target)))
return target
block = PreNormDecoderBlock(32, 4, 128, dropout=0.0)
target = torch.randn(2, 5, 32)
memory = torch.randn(2, 7, 32)
target_valid = torch.ones(2, 5, dtype=torch.bool)
memory_valid = torch.tensor([
[True, True, True, True, True, False, False],
[True, True, True, True, True, True, True],
])
decoded = block(target, memory, target_valid, memory_valid)
assert decoded.shape == (2, 5, 32)python三次残差相加都要求输出为 [N,T,D]。交叉注意力内部虽然建立 [T,S] 关系,最终仍为每个目标 Query 返回一个 维结果。
07 与 PyTorch 2.13 当前官方层对齐#
PyTorch 2.13 的 nn.TransformerDecoderLayer ↗ 是用于理解原始架构的参考实现:
official = nn.TransformerDecoderLayer(
d_model=32,
nhead=4,
dim_feedforward=128,
dropout=0.1,
activation="gelu",
batch_first=True,
norm_first=True,
)
target_block = torch.ones(5, 5, dtype=torch.bool).triu(diagonal=1)
out = official(
tgt=target, # [N,T,D]
memory=memory, # [N,S,D]
tgt_mask=target_block, # [T,T], True = 禁止
tgt_key_padding_mask=~target_valid, # [N,T], True = 忽略
memory_key_padding_mask=~memory_valid,
)
assert out.shape == target.shapepython当前 API 的关键契约是:
batch_first=True才使用[N,T,D]与[N,S,D];默认仍是序列维在前。norm_first=True对应 Pre-LN;默认False是 Post-LN。tgt_mask管目标位置之间的边,memory_key_padding_mask管源 Key 列,它们不能互相替代。tgt_is_causal和memory_is_causal是因果性提示;官方文档警告,错误提示可导致前向、反向或版本兼容性错误。- 标准的编码器记忆不是因果序列,因此不应随手设置
memory_is_causal=True。 - 官方层定位为基础参考实现,只提供有限的现代 Transformer 特性;生产推理的缓存、调度与融合内核需要另行设计。
08 一次并行训练如何流动?#
source_ids [N,S] ─► source embedding + position ─► Encoder ─► memory [N,S,D]
target_ids [N,T+1]
├─ [:,:-1] ─► decoder_input [N,T] ─► embedding + position ─┐
└─ [:, 1:] ─► labels [N,T] │
▼
memory [N,S,D] ────────────────────────► K 个 Decoder blocks ─► hidden [N,T,D]
│
Linear(D,Vocab)
▼
logits [N,T,V]textfrom torch.nn import functional as F
decoder_input_ids = target_ids[:, :-1]
labels = target_ids[:, 1:]
logits = model(source_ids, decoder_input_ids) # [N,T,V]
loss = F.cross_entropy(
logits.reshape(-1, logits.size(-1)),
labels.reshape(-1),
ignore_index=pad_id,
)
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0, error_if_nonfinite=True)
optimizer.step()pythoncross_entropy 接收未经 softmax 的 logits。将 [N,T,V] 展平为 [N*T,V] 时,标签也必须用相同的 batch-major 顺序展平;否则形状合法,token 却全部错位。
09 为什么生成不能像训练一样并行?#
训练时,完整真实目标已知,因果 mask 使所有位置在一次前向中各自只看到合法前缀。开放式生成时, 必须先被选出,才能成为预测 的输入:
model.eval()
generated = torch.full((n, 1), bos_id, device=device, dtype=torch.long)
with torch.inference_mode():
memory, memory_valid = model.encode(source_ids)
for _ in range(max_new_tokens):
# 透明但低效:每轮重算整个已生成前缀。
logits = model.decode(generated, memory, memory_valid)
next_id = logits[:, -1].argmax(dim=-1, keepdim=True)
generated = torch.cat([generated, next_id], dim=1)
if (next_id == eos_id).all():
breakpython贪心选择 argmax 只是一种解码策略。采样、top-k、nucleus sampling 和 beam search 改变的是如何从当前 logits 选 token,不改变解码块的三条信息流。
上面的循环每次把完整 generated 重新送入所有解码层,正是下一篇 KV Cache 要修复的工程瓶颈。
10 一条可执行的调试路径#
- 先过拟合一个极小 batch。 关闭 dropout,用两个短序列反复训练,确认 token 损失能接近 0。
- 做未来泄漏测试。 固定目标前缀至位置 ,任意替换右侧 token;位置 的 hidden state 和 logits 必须不变。
- 做源信息依赖测试。 保持目标前缀不变,交换两条差异很大的源序列;若 logits 完全不变,检查 cross-attention 是否被旁路或 mask 掉。
- 打印两张关系矩阵的形状。 自注意力应是
[N,H,T,T],交叉注意力应是[N,H,T,S]。 - 检查错位输入。 直接打印
decoder_input_ids[0]和labels[0],逐位验证左边比右边早一个 token。 - 分别测试两种 padding。 替换源 padding token 不应改变任何真实目标位置;目标 padding 标签不应进入损失。
- 在
dropout=0下对齐手写层和官方层。 复制参数后逐子层assert_close,别只比较最终 loss。 - 分开 train/eval 与 autograd 开关。
model.eval()关闭模块式 dropout,inference_mode()关闭梯度记录,两者不互相替代。
11 最常见的“形状正确,语义却错了”#
- 自注意力没有 causal mask。 训练 loss 异常好看,生成时却失效。
- 目标输入和标签没有错开。 模型通过残差路径复制当前 token。
- 把 cross-attention 写成又一次 self-attention。 Q/K/V 都来自目标侧,源序列从未进入解码器。
- 把编码器记忆当 Query。 输出长度变成 ,不再与目标位置一一对应。
- 对 memory 错用 causal mask。 普通翻译解码的每个目标位置都应能读取全部真实源 token。
- 只屏蔽目标 padding,漏掉源 padding。 cross-attention 会把补齐位置当成证据。
- 只 mask 注意力,不 mask 损失。 padding 标签仍会改变梯度。
- 把
True在 SDPA 和 MHA 中当成同一语义。 两个接口的布尔attn_mask方向相反。 - 将 Teacher Forcing 误解为推理算法。 生成时没有真实目标前缀可供喂入。
- 并行训练速度外推为并行生成速度。 逐 token 依赖仍然是串行的。
12 它会在哪些场景失败?#
- 暴露偏差(Exposure Bias)。 训练总看到正确前缀,推理却要接着自己的错误继续生成,小错可能累积。
- 长序列成本。 目标自注意力为 ,交叉注意力为 ;训练显存和生成延迟都会增长。
- 源序列含噪或过长。 全局 cross-attention 不保证精确检索,相似证据可能被混合。
- 训练与业务解码目标不一致。 token 交叉熵不直接优化事实性、全局结构或人类偏好。
- 无条件生成。 Decoder-only 模型通常只保留 causal self-attention 和 FFN,并没有可读的编码器 memory;不应强行塞入空 cross-attention。
- 非自回归任务。 序列标注、双向理解或固定输出的回归问题未必需要 causal Decoder。
13 与相近结构的边界#
| 结构 | 目标侧可见性 | 如何读源信息 | 典型用途 |
|---|---|---|---|
| Transformer Encoder | 通常双向 | 输入本身就是源序列 | 理解、表示、分类 |
| Encoder–Decoder Transformer | 目标侧因果 | 专门的 cross-attention | 翻译、摘要、条件生成 |
| Decoder-only Transformer | 整个拼接序列因果 | 条件也作为左侧 token 前缀 | 语言建模、通用生成 |
| RNN Encoder–Decoder + 注意力 | 逐步递推 | 每步用隐状态查询源状态 | 经典序列转换 |
Decoder-only 模型把 prompt 和输出放进同一条因果序列,不等于它内部隐藏了一个 cross-attention。Encoder–Decoder 的源记忆可以一次编码、被每个解码层重复查询;两者的 mask、缓存和部署契约不同。
14 今天真正需要记住什么?#
- Transformer Decoder block 有三条不同的信息流:目标前缀的 causal self-attention、目标查询源记忆的 cross-attention,以及逐位置 FFN。
- 交叉注意力中 Query 来自目标侧,Key/Value 来自编码器 memory,权重形状是
[N,H,T,S]。 - 目标 causal mask 禁止读未来,源 padding mask 禁止读补齐 Key,token 损失还要单独忽略 padding 标签。
- 训练通过右移目标和 Teacher Forcing 并行计算所有位置;推理的下一 token 依赖已生成前缀,仍然逐步进行。
- 最小未来泄漏测试和源扰动测试比单纯的形状断言更能捕捉解码器语义错误。
15 思考题与小练习#
- 设 ,写出目标自注意力和交叉注意力的 Q/K/V、分数与输出形状。解释为什么交叉注意力输出长度是 3 而不是 4。
- 为
PreNormDecoderBlock写未来泄漏测试:固定位置 0–2,替换位置 3 以后的 target,验证前三个输出不变;再去掉causal_block观察失败。 - 把一个长度为 5 的 target 写成
decoder_input和labels,人为在末尾加两个 padding,列出因果 mask、目标 padding mask 和损失 mask 各自作用的位置。
相关工作#
- Vaswani et al. (2017), Attention Is All You Need ↗:提出由 masked self-attention、encoder–decoder attention 和 FFN 组成的 Transformer Decoder。
- Bahdanau, Cho & Bengio (2015), Neural Machine Translation by Jointly Learning to Align and Translate ↗:在循环 Encoder–Decoder 中引入可学习对齐,是 cross-attention 信息流的重要前身。
- Sutskever, Vinyals & Le (2014), Sequence to Sequence Learning with Neural Networks ↗:展示了早期无注意力序列转换的定长语境瓶颈。
- Williams & Zipser (1989), A Learning Algorithm for Continually Running Fully Recurrent Neural Networks ↗:早期 Teacher Forcing 与循环网络训练语境的经典工作。
- Bengio et al. (2015), Scheduled Sampling for Sequence Prediction with Recurrent Neural Networks ↗:尝试缓解训练真实前缀与推理模型前缀之间的暴露偏差。
16 下一篇预告#
解码器的信息边界已经正确,但上面的推理循环每生成一个 token 都重算整个前缀。下一篇将拆解 KV Cache:为什么历史 Key/Value 可以复用、Query 通常不缓存,以及缓存如何改变每步张量形状、计算量与显存占用。