深入理解 Transformer Decoder 与 Masked Attention——从因果掩码到 KV Cache 的原理、验证与 PyTorch 实战

📅 发布时间:2026/10/3 20:01:07
深入理解 Transformer Decoder 与 Masked Attention——从因果掩码到 KV Cache 的原理、验证与 PyTorch 实战
深入理解Transformer Decoder与Masked Attention深入理解Transformer Decoder与Masked Attention——从因果掩码到KV Cache的原理、验证与PyTorch实战#Transformer #Decoder #Masked Attention #Causal Mask #Self-Attention #GPT #KV Cache #PyTorch #大语言模型#深度学习Transformer 最反直觉的一点是训练时明明把完整目标序列一次送进 GPU模型却必须像真实生成一样“看不到未来”。这并不是一句“加个下三角 Mask”就能解释完整的问题。本文从标签泄露这一核心矛盾出发逐层拆解 Decoder、Masked Self-Attention、Causal Mask、目标序列右移与 Teacher Forcing并通过手算矩阵和 PyTorch 代码验证未来位置如何在 Softmax 前被严格排除随后继续追到 Decoder-only、非方阵因果对齐、Prefill/Decode、KV Cache、SDPA 与 FlashAttention建立从数学正确性、训练机制到现代大模型推理系统的完整知识链。读完后你不仅能画出 Decoder更能独立判断一个 Mask 或 Cache 实现究竟对不对。▌先抛出一个问题假设训练样本是“我 喜欢 吃 苹果”。GPU 明明同时拿到了“苹果”为什么模型在预测“苹果”时仍然不能偷看它如果你能从注意力矩阵、训练目标和代码三个层面回答这个问题Decoder 的核心就真正打通了。一、先看结论Decoder 的核心不是“解码”而是“受约束地预测未来”经典 Transformer Decoder 的任务可以概括为在已知源序列若存在和目标前缀的条件下预测下一个 Token。对自回归语言模型而言联合概率被分解为一连串条件概率▌自回归分解P(x₁,…,xₜ) ∏ P(xᵢ | x₁,…,xᵢ₋₁)这条公式直接规定了信息边界预测 xᵢ 时只能使用它之前的 Token。于是 Decoder 的所有关键机制——Causal Mask、目标右移、Teacher Forcing、逐 Token 生成、KV Cache——都可以从这条约束推导出来。图1经典Transformer Decoder Layer的核心信息流二、为什么普通 Self-Attention 会“作弊”标准缩放点积注意力为 Attention(Q,K,V)softmax(QKᵀ/√dₖ)V。若序列长度为 TQKᵀ 形成 T×T 分数矩阵。普通 Self-Attention 中第 i 个 Query 可以与全部 Key 建立联系。Encoder 做双向理解时这很合理Decoder 做下一 Token 预测时却会造成标签泄露。例如我们希望位置“吃”的隐藏状态预测“苹果”。如果“吃”可以直接关注右侧的“苹果”训练损失可能很好看但模型学到的是一条真实生成时根本不存在的捷径。▌核心矛盾训练希望整段序列并行计算生成逻辑又要求每个位置只能看到历史。Causal Mask 正是把“逻辑上的先后关系”编码进并行矩阵运算。三、Causal Mask为什么是下三角图2方形序列中的Causal Mask当行表示 Query、列表示 Key 时第 i 行只允许访问 j≤i 的列因此形成下三角可见区域。常见 additive mask 定义为合法位置加 0未来位置加 −∞。▌掩码定义Mᵢⱼ 0j ≤ iMᵢⱼ −∞j i工程里未必真的存储 IEEE 负无穷低精度实现可能使用数据类型的极小有限值或直接使用布尔/融合 kernel 表达相同语义。判断实现是否正确应看 Softmax 后非法位置是否为零概率而不是执着于某个常数。四、Mask 为什么必须在 Softmax 前图3 Masked Attention的计算顺序带掩码的注意力写成 Attention(Q,K,V)softmax((QKᵀ/√dₖ)M)V。Softmax 会把一行 logits 归一化成概率分布。只有先把非法位置压到 −∞才有 exp(−∞)0从而使它们不进入有效概率竞争。如果非法位置只加 0分数没有任何变化如果只加 −1也只是降低概率并不能保证为 0。这是很多手写 Mask 的第一个隐蔽错误。五、手算一个 4×4 示例未来权重怎样真正归零S [[2.1, 1.2, 0.7, 1.5],[0.8, 2.4, 1.1, 1.7],[1.2, 0.9, 2.8, 1.6],[1.4, 1.8, 0.7, 2.5]]S M [[2.1, -∞, -∞, -∞],[0.8, 2.4, -∞, -∞],[1.2, 0.9, 2.8, -∞],[1.4, 1.8, 0.7, 2.5]]逐行 Softmax 后约为A ≈[[1.0000, 0.0000, 0.0000, 0.0000],[0.1680, 0.8320, 0.0000, 0.0000],[0.1494, 0.1107, 0.7399, 0.0000],[0.1669, 0.2489, 0.0829, 0.5013]]这揭示了 Mask 的本质未来 Token 在训练张量里可以存在但它们对当前 Query 的注意力概率为 0。于是完整序列可以并行进入模型而每个位置仍保持合法因果视野。六、目标序列右移与 Teacher Forcing两个经常被混淆的概念位置Decoder输入监督目标1BOS我2我喜欢3喜欢吃4吃苹果5苹果EOS右移决定“这个位置要预测谁”Causal Mask 决定“这个位置允许看谁”。Teacher Forcing 则表示训练时使用真实历史 Token 作为条件而不是使用模型刚生成的错误 Token。▌不要误解Teacher Forcing 不等于看到未来答案。真实历史可以喂给模型但更右侧的未来位置仍由 Causal Mask 屏蔽。七、训练为什么并行生成为什么仍然串行图4训练并行与自回归生成串行RNN 的 hₜ 显式依赖 hₜ₋₁时间步之间存在计算图顺序依赖。Transformer 训练时可以一次得到所有位置的 Q/K/V再用 Mask 规定每行可见域因此多个预测位置能在矩阵运算中并行。推理则不同第 t1 个 Token 的输入包含第 t 步真实选出的 Token而这个 Token 在第 t 步之前并不存在。因此 Token 维度上的自回归依赖无法被普通并行矩阵运算直接消除。八、Causal Mask 与 Padding Mask名字相似职责完全不同图5 Causal Mask与Padding Mask比较项Causal MaskPadding Mask屏蔽对象未来 Token补齐产生的 PAD目的保证因果性、防止标签泄露避免无效填充参与注意力来源位置先后关系每个样本的真实长度能否组合可以可以在变长 Batch 的 Decoder 训练中两者经常叠加。不要把具体代码中的一个 mask 张量误认为它只表达一种语义框架可能已经把多种约束合并成统一 bias。九、经典 Decoder 与 Decoder-only为什么 GPT 没有独立 Encoder经典机器翻译 Transformer 的 Decoder 除了 Masked Self-Attention还有 Cross-AttentionQ 来自 DecoderK/V 来自 Encoder。源序列在生成前已经完整可用因此 Cross-Attention 通常不需要目标侧的 causal 下三角限制但可能需要源端 Padding Mask。Decoder-only 模型则把系统指令、用户输入和生成内容统一组织成一条因果 Token 序列没有独立 Encoder自然也通常没有经典 Encoder–Decoder Cross-Attention。它并非“只会生成不会理解”而是使用单向可见的多层 Self-Attention 对完整历史前缀建立条件表示。十、一个容易被教程省略的边界非方阵 Causal Mask图6 Query与Key/Value长度不相等时的因果对齐入门教程几乎总画 T×T 下三角矩阵但增量解码时 q_len 和 kv_len 往往不同。例如历史已有 8 个 Token本轮只产生 1 个新 Query而 K/V 覆盖 9 个位置。此时“因果”应该按完整时间轴对齐。当前 PyTorch 提供 upper-left 与 lower-right 两种 CausalBias 变体。方阵时二者等价非方阵时可见区域不同。因此不要把 causal 的定义简化成“无脑 torch.tril 一个 q_len×kv_len 矩阵”。▌完整定义Causal Mask 的本质是时间位置约束下三角只是 Query 与 KV 长度相等时最直观的几何表现。十一、Boolean Mask 与 Additive Mask同一个词两种数值语义Mask类型合法位置非法位置典型实现Boolean Mask由 API 的布尔约定表示允许相反布尔值masked_fill / 框架 biasAdditive Mask0−∞ 或极小值直接加到 attention logits不同 API 对布尔 True/False 的含义可能存在约定差异因此不要把某个框架的布尔语义机械迁移到另一个接口。最稳妥的方法是查当前 API 文档再用 3×3 小矩阵做单元测试。十二、从零实现 Causal Self-Attentionimport mathimport torchimport torch.nn.functional as Fdef causal_attention(q, k, v):# q/k/v: [B, T, D]d q.size(-1)scores q k.transpose(-2, -1)scores scores / math.sqrt(d)T q.size(-2)future torch.triu(torch.ones(T, T, dtypetorch.bool, deviceq.device),diagonal1)scores scores.masked_fill(future, float(-inf))attn F.softmax(scores, dim-1)output attn vreturn output, attn不要只验证“能运行”。至少检查未来区域的最大权重out, attn causal_attention(q, k, v)future_max attn[0].triu(diagonal1).max().item()print(未来位置最大注意力权重, future_max)# 正确结果应为 0或数值上等价于 0十三、现代 PyTorch用 SDPA 表达标准因果注意力import torchimport torch.nn.functional as F# [batch, heads, seq_len, head_dim]q torch.randn(2, 8, 128, 64, devicecuda, dtypetorch.float16)k torch.randn(2, 8, 128, 64, devicecuda, dtypetorch.float16)v torch.randn(2, 8, 128, 64, devicecuda, dtypetorch.float16)y F.scaled_dot_product_attention(q, k, v,dropout_p0.0,is_causalTrue)高层 API 的价值不仅是少写几行代码。框架可以依据设备、dtype、形状等条件选择合适的高效后端避免用户手工物化大量中间张量。十四、三组实验把“我觉得对”升级成“我证明它对”图7因果注意力的三层验证闭环实验 1——前缀不变性固定前缀只修改未来 Token。在关闭 dropout 的评估模式下更早位置的隐藏状态不应被未来变化影响。实验 2——显式 Mask 与 is_causal 对照同一组 Q/K/V 分别走手写 Mask 和框架 causal 路径输出应在允许的浮点误差内一致。实验 3——KV Cache 等价性一次性 full forward 与 prefill 单步 decode 对同一位置产生的 logits 应近似一致。with torch.no_grad():full model(input_ids).logits[:, -1, :]prefix input_ids[:, :-1]last input_ids[:, -1:]prefill model(prefix, use_cacheTrue)step model(last,past_key_valuesprefill.past_key_values,use_cacheTrue)cached step.logits[:, -1, :]print((full - cached).abs().max())十五、KV Cache为什么历史 K/V 可以直接复用图8 KV Cache减少历史K/V的重复计算在因果 Self-Attention 中过去 Token 在某一层得到的 K/V 不会因为未来新增 Token 而被改写。推理时因此可以保存历史 K/V下一步只计算新 Token 的 qₜ、kₜ、vₜ并让 qₜ 与缓存历史 K 及新 kₜ 做注意力。为什么通常缓存 K/V 而不是 Q因为未来步骤需要用新的 Query 去检索历史 Key并聚合历史 Value过去 Query 已经完成了自己的输出计算通常不需要再次参与当前 Token 的输出。▌代价交换KV Cache 用显存/内存换取计算复用。上下文越长、层数越多、KV heads 越多Cache 占用越显著。十六、Prefill 与 Decode同一次生成中的两种硬件特征图9 Prefill与DecodePrefill 一次处理完整 Prompt并建立初始 KV Cache矩阵规模较大、并行度较高。Decode 通常每步只新增一个 Token却不断读取越来越长的历史 Cache因此更容易受到内存带宽、缓存容量和批处理调度影响。这解释了为什么线上 LLM 系统会继续讨论 Dynamic/Static/Quantized Cache、Paged KV Cache、continuous batching 等机制它们是在不改变自回归因果语义的前提下优化真实系统成本。十七、FlashAttention它优化的是执行不是因果定义朴素标准注意力会产生大规模中间分数/概率矩阵长序列时显存读写昂贵。FlashAttention 的关键思想是 IO-aware通过分块等方式减少 GPU HBM 与片上 SRAM 之间的数据搬运同时计算精确注意力。因此 Causal Mask 与 FlashAttention 解决的是两个层次的问题前者定义“哪些位置允许互相看”后者研究“同样的注意力怎样在硬件上更高效地算”。▌一句话区分Mask 决定语义正确性高效 Attention kernel 决定执行效率。十八、从 MHA 到 MQA/GQA为什么现代模型开始减少 KV Heads经典 Multi-Head Attention 为多个 Query heads 配置对应 K/V heads。自回归推理中历史 K/V 长期驻留 Cache因此 KV heads 数量会直接影响缓存容量和读取带宽。MQA 让多个 Query heads 共享一组 K/VGQA 则让一组 Query heads 共享 K/V在表达能力与推理成本之间折中。它们并不改变 Causal Mask 的基本逻辑Mask 决定“看哪里”MQA/GQA 改变“用多少组 K/V 表示可见历史”。十九、6 个高频误区看似会了其实最容易写错图10 Masked Attention高频误区误区正确理解Mask 就是删除未来 Token未来 Token 可存在于训练张量中只是对当前 Query 的注意力概率被屏蔽非法位置加 0 就够additive mask 中 0 表示不改变分数训练也必须逐 Token完整序列可并行计算Mask 负责保持因果性Padding Mask 就是 Causal Mask一个处理无效填充一个处理未来信息Decoder 一定有 Cross-AttentionDecoder-only 通常没有独立 Cross-AttentionKV Cache 会改变生成逻辑正确实现应与无 Cache 路径保持语义等价二十、性能复杂度别用一句 O(T²) 解释所有推理现象阶段主要特征常见瓶颈训练 / Prefill多个 Query × 多个 Key算力、显存、中间张量无 Cache Decode每一步重复处理历史大量冗余计算有 Cache Decode新增 Q/K/V 读取历史 Cache内存带宽、Cache 容量、调度标准全注意力的分数矩阵规模随序列长度二次增长但在线生成速度还受 KV Cache、Batch、采样、模型并行和硬件带宽等共同影响。复杂度是理解成本的起点不是完整性能模型。二十一、把整个 Decoder 串成一条可执行链1. Tokenizer 把文本转换为 Token ID并映射为向量位置机制提供顺序信息。2. 隐藏状态投影成 Q、K、V。3. 计算缩放点积分数并施加 Causal Mask。4. Softmax 得到合法注意力分布对 V 加权聚合。5. 经输出投影、残差、归一化和 FFN/MLP 进入下一层。6. 经典 Seq2Seq Decoder 还通过 Cross-Attention 读取 Encoder 输出。7. 最后一层隐藏状态经 LM Head 投影到词表 logits。8. Greedy、Temperature、Top-k 或 Top-p 等策略选择新 Token。9. 推理时把新 K/V 追加到 Cache下一步复用历史。10. 重复生成直到 EOS、长度上限或其他停止条件。二十二、一张表建立最终知识地图层次关键问题机制语义层当前位置允许依赖谁Causal Mask数学层非法依赖如何变成零概率logits mask → Softmax结构层如何形成深层条件表示Masked Attention FFN Residual Norm训练层如何学习下一 Token右移 Teacher Forcing Cross Entropy生成层如何从分布选 TokenGreedy / Temperature / Top-k / Top-p推理层如何避免重复计算KV Cache系统层如何降低 Attention IOSDPA / FlashAttention / 高效 kernel二十三、最后总结真正理解 Decoder只需要抓住一条主线语言模型要预测下一个 Token → 当前预测不能读取未来 → Causal Mask 把因果约束写进 Self-Attention → 完整序列因此可以在训练时并行计算 → 推理时未来 Token 尚不存在所以仍然逐步生成 → 历史 K/V 不会被未来改写因此可以缓存 → Prefill/Decode 与高效 Attention 再把同一套数学语义映射到真实硬件。▌记住这四句话Mask 决定“能看什么”Attention 决定“重点看什么”Decoder 决定“如何把已知变成下一步预测”KV Cache 决定“如何避免重复计算”。参考资料1. Vaswani A. et al. Attention Is All You Need. 2017.2. PyTorch Documentation. scaled_dot_product_attention.3. PyTorch Documentation. torch.nn.attention.bias.CausalBias / CausalVariant.4. Hugging Face Transformers Documentation. KV cache strategies.5. Dao T. et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. 2022.