声明:本文主要参考开源资料进行学习整理,如有错漏,欢迎评论交流~
引言
为什么生成式 LLM 的注意力叫"因果"(causal)注意力?它和普通注意力差在哪?
"不能看未来"这件事,怎么用一张 mask 矩阵表达?
- 当序列里混进了 pad token,pad 的屏蔽和"不能看未来"的屏蔽如何叠加成一张 mask?
- 为什么屏蔽位有时候填的是 -32768 ,而不是 -inf?
prefill 和 decode 两个阶段,mask 的形状和构造有何不同?
1. 从普通注意力到因果注意力
普通注意力:谁都能看谁
因果注意力:不能看未来

一个直觉:causal attention(因果注意力)把 attention 矩阵从"满矩阵"砍成"下三角矩阵"。这个形状上的约束,正是自回归生成能成立的前提。
2. causal_mask(因果掩码):把"不能看未来"写成矩阵
用上三角表达"未来"
代码里"下三角可见、上三角屏蔽"是这样实现的:
把它加进 attention 计算

屏蔽位置的权重归零,等价于"位置 i 完全不关注未来"。这就是 causal_mask(因果掩码)落到 attention 计算上的效果。
3. 当 pad 混进来:两种屏蔽的叠加
从 1D attention_mask 到 2D 屏蔽矩阵
直观理解:
- 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 在一起

这就是第一篇里"attention_mask 同步补 0"的真正含义——补 0 不是终点,终点是它在这张叠加 mask 里把 pad 位置彻底屏蔽掉。
4. 为什么填 -32768,而不是 -inf
- -32768 正好是 int16 能表示的最小值,量化友好的"最负"。
softmax 时 exp(-32768) ≈0(相对于其他有限分数,指数后趋近 0),屏蔽位的权重实际归零。
又因为是有限值,int 算子能正常处理,不会溢出。
这就是端侧和云端推理的一个细节差异:云端 FP16/BF16 下大家习惯用 -inf 或大负数;端侧量化下统一用 min_value=-32768 这个 int16 下限。
5. prefill 与 decode:mask 形状的两副面孔
causal_mask(因果掩码)在 prefill 和 decode 两个阶段长得不一样,这也是容易糊涂的地方。
Prefill:L x C 的完整矩阵
两个细节:
- 因果屏蔽只在 LxL 范围内做(triu(seq_len, seq_len)),因为"未来"只存在于当前这批 query 内部。
- 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。
pad 历史位置屏蔽(左 padding 留下的 pad KV);
cache 未填充区屏蔽。
每生成一个 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。
后续
- 长度对齐决定 pad 在哪(左 padding 把真实 token 推到右端,attention_mask 在 pad 位补 0);
- causal attention(因果注意力)决定未来在哪(上三角屏蔽),并和 pad 屏蔽叠加成最终 mask;
- chunk prefill决定 query 怎么切(沿 query 维切成 chunk,key 维仍右对齐到 C)。
