🎯 课程主题
因果掩码注意力机制在推理过程中的应用——逐步生成序列时的掩码处理。
📝 核心知识点
1. 推理过程的特点
- 概念说明:推理时输入来自前一时间步的输出,序列逐步增长。
- 关键细节:
- 第 1 步:输入起始符
<S> - 第 2 步:输入
<S>+ 第 1 步预测结果 - 第 3 步:输入
<S>+ 前两步预测结果 - 以此类推,直到输出结束符
- 第 1 步:输入起始符
2. 第 1 步推理(无掩码)
- 概念说明:第一步只有起始符,无需因果掩码。
- 关键细节:
- 输入:
<S>→ 词嵌入 → $X$ - $X \times W_Q, W_K, W_V$ → $Q, K, V$
- $Q \times K^T$ → $1 \times 1$ 矩阵(无未来需屏蔽)
- Softmax → 注意力得分为 1(只关注自己)
- 与 $V$ 相乘 → 输出
<S>的信息融合 - 经后续网络 → 输出第一个词(如 "I")
- 输入:
3. 第 2 步推理(2×2 掩码)
- 概念说明:第二步序列长度为 2,需 2×2 因果掩码。
- 关键细节:
- 输入:
<S> I→ 序列长度 2 - $Q \times K^T$ → $2 \times 2$ 注意力得分矩阵
- 因果掩码:右上角为 $-\infty$
- Softmax 后:
- 位置 1(
<S>):只看<S>,看不到I - 位置 2(
I):看<S>和I
- 位置 1(
- 与 $V$ 相乘:
- 第 1 行:只含
<S>信息 - 第 2 行:含
<S>和I的信息融合
- 第 1 行:只含
- 输入:
4. 第 N 步推理(N×N 掩码)
- 概念说明:序列长度逐步增长,掩码矩阵相应增大。
- 关键细节:
- 序列长度 $N$ → $N \times N$ 注意力矩阵
- 上三角为 $-\infty$,经 Softmax 后为 0
- 每个位置只关注当前位置及之前的信息
- 最后一行包含所有已生成 token 的信息融合
5. 推理与训练的对比
- 概念说明:推理是串行逐步生成,训练是并行计算。
- 关键细节:
- 训练:一次性输入所有标签,并行计算所有位置
- 推理:逐步输入,每步序列长度增加 1
- 本质都是倒三角掩码,上三角为 $-\infty$
- 推理实际实现中会使用 KV Cache 优化,避免重复计算
🧮 核心公式与推导
第 $t$ 步推理的因果掩码($t \times t$ 矩阵):
$$\text{Mask}_t = \begin{bmatrix} 0 & -\infty & \cdots & -\infty \ 0 & 0 & \cdots & -\infty \ \vdots & \vdots & \ddots & -\infty \ 0 & 0 & \cdots & 0 \end{bmatrix} \in \mathbb{R}^{t \times t}$$
第 $t$ 步注意力计算:
$$A_t = \text{Softmax}\left(\frac{Q_t K_t^T}{\sqrt{d_k}} + \text{Mask}_t\right)$$
$$B_t = A_t V_t$$
其中 $Q_t, K_t, V_t$ 是序列长度为 $t$ 时的 QKV 矩阵。
物理意义:推理时每步序列增长,掩码矩阵相应扩大,但始终保证上三角为 $-\infty$,确保不看到未来。
🏗️ 模型架构与数据流向
💻 代码实战
import torch
import torch.nn.functional as F
import math
class CausalAttention:
"""推理时的因果掩码注意力"""
def __init__(self, d_model, d_k):
self.W_Q = torch.nn.Linear(d_model, d_k, bias=False)
self.W_K = torch.nn.Linear(d_model, d_k, bias=False)
self.W_V = torch.nn.Linear(d_model, d_k, bias=False)
self.d_k = d_k
def step(self, x):
"""
推理单步计算
x: [1, t, d_model] 当前所有已生成的token
return: [1, t, d_k] 输出, 取最后一行作为预测
"""
t = x.size(1)
Q = self.W_Q(x) # [1, t, d_k]
K = self.W_K(x)
V = self.W_V(x)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
# 因果掩码: 上三角为-inf
mask = torch.triu(torch.ones(t, t), diagonal=1).bool()
scores = scores.masked_fill(mask, float('-inf'))
attn = F.softmax(scores, dim=-1)
output = torch.matmul(attn, V)
return output, attn
# 推理过程模拟
d_model = 512
d_k = 64
attn = CausalAttention(d_model, d_k)
# 模拟逐步生成
tokens = [torch.randn(1, 1, d_model)] # 起始符
for step in range(1, 5):
x = torch.cat(tokens, dim=1) # [1, step, d_model]
output, attn_weights = attn.step(x)
print(f"Step {step}: attention matrix shape = {attn_weights.shape}")
# 生成下一个token (这里用随机模拟)
tokens.append(torch.randn(1, 1, d_model))
⚠️ 常见问题与避坑指南
- 推理时每步序列长度增加,掩码矩阵相应增大
- 第 1 步无需掩码($1 \times 1$ 矩阵)
- 实际推理中用 KV Cache 优化,避免每步重新计算所有 K、V
- 推理比训练慢,因为需要串行生成
💡 个人总结与延伸
推理时的因果掩码机制与训练本质相同,都是通过上三角 $-\infty$ 屏蔽未来信息。但推理需要逐步生成,效率较低。现代大模型通过 KV Cache(缓存已计算的 K、V,避免重复计算)大幅加速推理,这是当前大模型推理优化的核心技术。