专栏算法工具链【大模型】Causal Attention与 causal_mask简介

【大模型】Causal Attention与 causal_mask简介

no_name2026-07-26
31
0

声明:本文主要参考开源资料进行学习整理,如有错漏,欢迎评论交流~

引言

在进行prompt长度对齐和 chunk prefill 时,反复出现一个动作:左 padding 之后,attention_mask 要在 pad 位置同步补 0。很多时候只一句话带过——"这些是 pad,不要 attend(关注) 它们"——但背后其实藏着 LLM 注意力机制的一块拼图:causal attention(因果注意力),以及它怎么被实现成一张 causal_mask(因果掩码)。我们先来看几个问题:
  • 为什么生成式 LLM 的注意力叫"因果"(causal)注意力?它和普通注意力差在哪?

  • "不能看未来"这件事,怎么用一张 mask 矩阵表达?

  • 当序列里混进了 pad token,pad 的屏蔽和"不能看未来"的屏蔽如何叠加成一张 mask?
  • 为什么屏蔽位有时候填的是 -32768 ,而不是 -inf?
  • prefill 和 decode 两个阶段,mask 的形状和构造有何不同?


1. 从普通注意力到因果注意力

普通注意力:谁都能看谁

标准 self-attention(自注意力)里,序列中每个位置都可以 attend(关注)到所有位置(包括自己)。对一条长度为 L 的序列,attention 是一个 L x L 的矩阵:第 i 行第 j 列表示"位置i 对位置 j 的注意力权重"。普通注意力下,这个矩阵全是有效值
这在理解类任务(如 BERT 做文本分类)里没问题——整句话都摆在那儿,互相看很自然。

因果注意力:不能看未来

生成式 LLM 是逐 token 生成的:预测第 i 个 token 时,只能依据它前面已经出现的 token(位置 0..i-1)和它自己(位置 i),绝不能偷看后面的 token——否则就是"用答案预测答案",训练时泄漏、推理时根本拿不到未来。
于是注意力矩阵必须变成下三角:位置 i 只能 attend(关注)位置 0..i,对角线以上的"未来"位置全部屏蔽。
这就是 causal attention(因果注意力),也叫 masked self-attention(带掩码的自注意力)。"因果"二字取自物理因果:结果(当前 token)只能由原因(过去的 token)产生,不能依赖尚未发生的未来

一个直觉:causal attention(因果注意力)把 attention 矩阵从"满矩阵"砍成"下三角矩阵"。这个形状上的约束,正是自回归生成能成立的前提。

2. causal_mask(因果掩码):把"不能看未来"写成矩阵

用上三角表达"未来"

代码里"下三角可见、上三角屏蔽"是这样实现的:

torch.triu(..., diagonal=1) 取的是严格上三角(不含对角线)——这些正是"未来"位置。把它置为 True(表示"要屏蔽"):
注意:这里 True 表示"屏蔽"。后面会用 torch.where(mask == 1, min_value, 0) 把 True 的位置填成屏蔽值、False 的位置填 0(即不干预 attention 分数)。

把它加进 attention 计算

屏蔽位置的权重归零,等价于"位置 i 完全不关注未来"。这就是 causal_mask(因果掩码)落到 attention 计算上的效果。

3. 当 pad 混进来:两种屏蔽的叠加

只屏蔽"未来"还不够。回想 左 padding——序列左端塞了一堆 pad token,它们的 attention_mask=0。这些 pad 既不该被别人 attend(关注),也不该 attend(关注)别人(它们是"假"token,没有语义)。于是 mask 里除了 causal(因果,屏蔽未来),还要再叠加一层 pad 屏蔽

从 1D attention_mask 到 2D 屏蔽矩阵

上一篇对齐后,attention_mask 是一个长度为 $$$$ 的一维向量(batch 维忽略),形如:
1 表示真实 token,0 表示 pad。要把这个一维信息变成"$$L \times $$ 的屏蔽矩阵",需要从 query 和 key 两个方向各广播一次:

直观理解:

  • k_mask(key 是 pad → 屏蔽整列):任何 query 都不该 attend(关注)到 pad key。只要第 j 列对应的 key 是 pad,整列屏蔽。
  • q_mask(query 是 pad → 屏蔽整行):pad query 本身就是假的,它 attend(关注)什么都不重要,但为了干净,整行屏蔽。这也保证 pad 位置算出的 attention 不影响后续(比如不会被错误地当作有效上下文)。

两者取 OR,就是"只要 query 或 key 任一是 pad,就屏蔽"。

三种屏蔽 OR 在一起

最终的屏蔽矩阵 = causal(因果)屏蔽(未来) OR pad 屏蔽(query 是 pad 或 key 是 pad)
用一个具体例子把三层叠加看清楚。设 L=5,其中前 2 个是 pad(attention_mask = [0,0,1,1,1]):
第 1 层:causal(因果)屏蔽(上三角 = 未来)
第 2 层:pad 屏蔽(key 是 pad 的列 0、1 全屏蔽;query 是 pad 的行 0、1 全屏蔽)
第 3 层:OR 叠加后的最终 mask(屏蔽 = ✗)
注意第 2 行(query 2,第一个真实 token):pad 屏蔽了 key0、1 两列,causal(因果)屏蔽了 key3、4 两个"未来",双重夹击下最终只剩 key2(它自己的过去与自己)。逐行往下看,causal 的下三角逐渐"打开"——q3 多见了 key3,q4 多见了 key3、key4——这正是因果注意力的形状:越靠后的 token 能看到的过去越多causal(因果)和 pad 两层屏蔽各司其职又完美叠加:causal 管"时间因果"(砍掉上三角未来),pad 管"真假 token"(砍掉 pad 行列)。

这就是第一篇里"attention_mask 同步补 0"的真正含义——补 0 不是终点,终点是它在这张叠加 mask 里把 pad 位置彻底屏蔽掉。

4. 为什么填 -32768,而不是 -inf

到这里 mask 还是布尔矩阵(True=屏蔽)。最后一步要把它变成 attention 能用的数值:
屏蔽位填 min_value,可见位填 0(不改动 attention 分数)。代码里 min_value=-32768,不是 -inf。为什么?
因为端侧推理是量化(int)运算。 端侧推理芯片跑的是 int16/int8 算子,-inf 这种浮点特殊值根本无法在整数里表示,强行用会导致溢出或未定义行为。于是改用一个足够大的有限负数
  • -32768 正好是 int16 能表示的最小值,量化友好的"最负"。
  • softmax 时 exp(-32768) ≈0(相对于其他有限分数,指数后趋近 0),屏蔽位的权重实际归零。

  • 又因为是有限值,int 算子能正常处理,不会溢出。

这就是端侧和云端推理的一个细节差异:云端 FP16/BF16 下大家习惯用 -inf 或大负数;端侧量化下统一用 min_value=-32768 这个 int16 下限。
同理,KV cache 右对齐时左侧补的 pad_mask,也填的是这个 min_value——所有"不该被 attend(关注)"的位置,无论是未来、pad、还是 KV 未填充区,统一用同一个大负数屏蔽。

5. prefill 与 decode:mask 形状的两副面孔

causal_mask(因果掩码)在 prefill 和 decode 两个阶段长得不一样,这也是容易糊涂的地方。

Prefill:L x C 的完整矩阵

prefill 一次处理 L 个 query token,要 attend(关注)整个长度为 C(max_kvcache_len)的 KV cache。所以 mask 形状是 (L, C)——行是 L 个 query,列是 C 个 key。

两个细节:

  1. 因果屏蔽只在 LxL 范围内做(triu(seq_len, seq_len)),因为"未来"只存在于当前这批 query 内部。
  2. KV 侧左侧补 pad_tokens:C - L 是 cache 里还没被填满的空位,它们在 key 维的左端(因为有效 KV 右对齐),统一用 pad_mask 屏蔽。补完之后 mask 从 (L, L) 扩成 (L, C),key 维和 cache 对齐——这正是强调的"右对齐"的源头。

Decode:1 x C 的单行 mask

decode 一次只生成 1 个新 token,它只有 1 个 query,但要 attend(关注)整个 cache(所有已写入的历史 KV)。所以 mask 形状退化成 (1, C)——一行,C 个 key。

此时没有"未来"可言(只有 1 个 query,它就是最新的,没有比自己更晚的),所以因果屏蔽在 decode 里退化为空——新 token 能看到 cache 里所有有效历史。decode 的 mask 主要管两件事:
  • pad 历史位置屏蔽(左 padding 留下的 pad KV);

  • cache 未填充区屏蔽。

代码里 decode 的 mask 是从 prefill 的 attention_mask 演化来的(get_decoder_mask 每生成一个 token 就把窗口右滑一格):

每生成一个 token,把 mask 沿 key 维左移一格(丢掉最左、右端补 0),相当于"有效历史窗口"右滑一位,把刚生成的新 token 纳入可见范围。这个左移动作配合滚动 KV cache,共同维持"cache 右对齐、mask 同步右对齐"的不变式。

对比一览

总结

  • Causal attention(因果注意力):生成式 LLM 的注意力,每个位置只能 attend(关注)自己及之前的 token,不能看未来。这是自回归生成的前提。
  • causal_mask(因果掩码):用 torch.triu(ones, diagonal=1) 构造上三角矩阵表达"未来不可见",在 softmax(归一化指数函数)前把未来位置的分数压成极小值,使其权重归零。
  • 两层屏蔽叠加:causal(因果)屏蔽(未来)OR pad 屏蔽(query 或 key 是 pad)。attention_mask 一维的 0/1 标记,经 q_mask/k_mask 广播成二维,与 causal(因果)矩阵 OR 起来,就是最终屏蔽矩阵。
  • min_value=-32768:端侧 int 量化无法表示 -inf,改用 int16 下限 -32768,softmax 后权重实际归零,又不会让 int 算子溢出。
  • prefill vs decode:prefill mask 是 (L, C) 完整矩阵(有因果屏蔽、左补 KV 未填充区);decode mask 退化为 (1, C) 单行(无因果屏蔽,每步左移一格纳新 token),始终右对齐到C。

后续

至此,长度对齐、chunk prefill、causal attention(因果注意力)三块拼图合拢。三者其实围绕同一个数据结构打转——一张形状随阶段变化、但始终右对齐到 KV cache 容量 C 的屏蔽矩阵
  • 长度对齐决定 pad 在哪(左 padding 把真实 token 推到右端,attention_mask 在 pad 位补 0);
  • causal attention(因果注意力)决定未来在哪(上三角屏蔽),并和 pad 屏蔽叠加成最终 mask;
  • chunk prefill决定 query 怎么切(沿 query 维切成 chunk,key 维仍右对齐到 C)。
算法工具链
社区征文杂谈
评论0
0/600