想象一场听写考试:老师念一段话,学生一个字一个字写下来。写到第 5 个字时,学生只能依据老师已经念过的前 4 个字加上正在念的第 5 个字,绝不能翻到老师讲义后面偷看答案——否则就是作弊。
因果掩码(Causal Mask,又叫下三角掩码)就是 GPT 这类生成模型里的"防作弊装置":它在注意力分数矩阵上盖一张下三角的遮布,让模型在生成第 t 个 token 时,只能看到当前及之前的 token,看不到未来的 token。
为什么生成任务必须"防偷看"
GPT 这类 Decoder(解码器)模型的工作方式是自回归(autoregressive):每次根据已生成的内容,预测下一个 token。推理(实际使用)时,下一个 token 根本还不存在,模型自然看不到。
但训练时,为了让模型并行学习所有位置(一次前向传播就同时预测整句每个位置的下一个 token,效率高),我们会把整句话一次性喂进去。这时如果不加干预,注意力机制会让每个位置都能看到整句所有 token——包括它本该预测的那个"未来答案"。模型直接抄答案就行,学不到东西,训练和推理就脱节了。
因果掩码就是来堵这个漏洞的。
图 1:自回归生成——每写一个词,可见的上下文就多一个;未来词此时还不存在,训练时必须用掩码制造同样的条件。
因果掩码怎么工作
注意力机制的核心是:每个 token 看一眼所有 token,算一个相关度分数,再用分数加权求和。这个分数矩阵长这样(3 个 token 为例):
、 、 :每个 token 经过三种线性变换得到的"查询 / 键 / 值"——粗略说就是"我要找什么""我有什么""我提供什么" :每对 token 的相关度(点积越大越相关),这就是上面那个分数矩阵 :缩放因子,防止点积数值过大导致训练不稳
因果掩码做的事,是在分数矩阵
:当前 token 的位置(行) :被看的 token 的位置(列) :当前 token 在被看 token 的后面或同一位置——允许看,掩码为 0 :当前 token 在被看 token 的前面——这是未来,禁止看,掩码为
为什么负无穷等于"看不见"
softmax 把分数归一成概率(每行加起来等于 1):
:第 个位置的分数 :自然对数的底(约 2.718),指数函数的底数 :指数化,把任意实数变成正数 - 分母
:所有位置指数值的和,用来归一化
关键在指数:
图 2:softmax 把被掩成 −∞ 的未来位置指数化后变成 0,归一化后权重仍是 0,等价于完全看不见。
小例子:3 个 token 走一遍
假设 3 个 token 的注意力分数矩阵(已经过缩放)是:
S = [[2, 1, 3],
[1, 2, 4],
[2, 0, 1]]加上因果掩码
S + M = [[2, -inf, -inf],
[1, 2, -inf],
[2, 0, 0 ]]每行做 softmax:
- 第 1 行(第 1 个 token 只能看自己):
- 第 2 行(第 2 个 token 看前两个):
- 第 3 行(第 3 个 token 看全部三个):
自检:每行和约等于 1、所有值非负、上三角(未来位置)全是 0——三样都对,掩码生效。
为什么 BERT 不需要它
BERT 用的是 Transformer 的编码器(Encoder),任务是理解:给一整句话,做分类、问答、命名实体识别这种"看完再判断"的活。它的注意力是双向的——每个 token 都能同时看左和看右,一次拿全部上下文。
BERT 的训练任务叫掩码语言模型(Masked Language Model,MLM):随机挖空几个 token 让模型填,挖空的位置是随机选的,不是"未来"。既然没有"按顺序生成"这回事,也就没有"未来"可言,因果掩码自然派不上用场。
一句话区分:
- GPT(Decoder + 因果掩码):一个字一个字往外蹦,写第 t 个字时只能看前面——为生成而生
- BERT(Encoder + 双向注意力):一眼看全文,理解整段意思——为理解而生
图 3:GPT 的每个词只能看到自己和左边(下三角扩张);BERT 的每个词都能看到全文(每行满格),所以 BERT 不需要因果掩码。
完整代码
下面是一个带因果掩码的最小自注意力模块,复制即可跑:
import torch
import torch.nn as nn
class CausalSelfAttention(nn.Module):
def __init__(self, d_model):
super().__init__()
# 三个线性层,把输入分别转成 Q / K / V(注意力三件套)
self.q_proj = nn.Linear(d_model, d_model)
self.k_proj = nn.Linear(d_model, d_model)
self.v_proj = nn.Linear(d_model, d_model)
self.d_k = d_model
def forward(self, x):
# x 形状: [batch, 序列长度 n, 维度 d_model]
n = x.size(1)
Q = self.q_proj(x)
K = self.k_proj(x)
V = self.v_proj(x)
# 注意力分数 S = QK^T / sqrt(d_k),对应上文公式
scores = Q @ K.transpose(-2, -1) / (self.d_k ** 0.5)
# 因果掩码:torch.triu 取上三角(diagonal=1 表示跳过对角线)
# 上三角 = 1 用来标记"未来位置"
mask = torch.triu(torch.ones(n, n), diagonal=1)
# masked_fill 把未来位置的分数改成 -inf,softmax 后权重变 0
scores = scores.masked_fill(mask.bool(), float('-inf'))
# softmax 把每行分数归一成概率(和为 1)
attn = torch.softmax(scores, dim=-1)
# 用注意力权重加权求和 V
return attn @ V
torch.manual_seed(0) # 固定随机种子,让每次跑结果一样(可复现)
# 假数据:batch=2,序列长度=4,维度=8(即 2 个句子、每句 4 个 token、每个 token 8 维向量)
x = torch.randn(2, 4, 8)
module = CausalSelfAttention(d_model=8)
out = module(x)
print(out.shape) # torch.Size([2, 4, 8]) —— 因果掩码不改输出形状,只禁止看未来核心就两行:torch.triu 造掩码、masked_fill 把未来位置改成
小结
因果掩码是一张下三角的遮布:它把注意力分数矩阵的上三角(未来位置)压成
参考资料
- Attention Is All You Need(原始 Transformer 论文,提出掩码注意力) - Vaswani et al., 2017 https://arxiv.org/abs/1706.03762
- What is causal attention, and why can GPT-style models train on next-word prediction - Sebastian Raschka https://sebastianraschka.com/faq/docs/causal-attention.html
- A Gentle Introduction to Attention Masking in Transformer Models - Machine Learning Mastery https://machinelearningmastery.com/a-gentle-introduction-to-attention-masking-in-transformer-models/