讲义定位
Attention 是 Mini-TabPFN 的通用读取原语。Cell token、inducing point、row token 和 context/query 之间的交互,最终都可以还原为一个问题:当前 token 应该从哪些 token 读取多少信息?
核心直觉:Attention = 带可见性约束的、数据依赖的加权读取。
本讲义从加权平均讲起,推导 Q/K/V、scaled dot-product、softmax 轴和 mask,再区分 self-attention 与 cross-attention。先理解「在计算什么」,再写成 PyTorch module。
本课符号
| 符号 | 含义 |
|---|---|
| \(N_q, N_k\) | query / key(及 value)的 token 数 |
| \(d_k, d_v\) | key 维与 value 维 |
| \(H\) | attention heads(多头时;本课先单头) |
| \(S, A, O\) | score、attention 权重、输出 |
1. 学习目标与先修知识
学完后你应该能够:
- 用「查询—匹配—读取」解释 Q/K/V;
- 写出 scaled dot-product attention 的每个中间 shape;
- 解释为什么 softmax 要沿 key 轴归一化;
- 区分 self-attention、cross-attention 与普通 MLP;
- 写出 key padding mask 和 pairwise attention mask 的语义;
- 通过手算、PyTorch parity、row-sum 和 mask audit 检查实现。
先修是矩阵乘法、softmax 和 batch tensor。
2. Attention 之前:一个可学习的加权平均
假设有 value 向量 \(v_1,\ldots,v_{N_k}\),希望根据查询决定如何加权读取:
若 \(a_j\ge 0\) 且 \(\sum_j a_j=1\),则 output 是 value 的凸组合。Attention 的贡献是:权重不是固定 规则,而是由 query 与 key 的匹配关系学习出来。
3. Q、K、V 的三种角色
对输入 token \(x\),经三个不同线性投影:
\(q\):「我正在寻找什么」;\(k\):「我可以被怎样匹配」;\(v\):「匹配后真正读取的内容」。
权威形式(canon §3.2;单头、无 head 轴):
形状契约:
每个 query 得到一行权重;一行描述「这个 query 如何分配它对所有 key 的读取预算」。 输出长度由 \(N_q\) 决定,读取范围由 \(N_k\) 决定——这会直接支撑下一课的 \(N\to K\to N\)。
多头时通常写成 \(Q\in\mathbb{R}^{B\times H\times N_q\times d_k}\);本课先把单头做对,再在作业中加 head 轴包装。
4. 为什么要除以 \(\sqrt{d_k}\)
若 \(q,k\) 各分量近似独立、均值近 0、方差近 1,则点积方差随 \(d_k\) 增长:
未经缩放的 score 容易极端,softmax 接近 one-hot,例如 \(\operatorname{softmax}(100,0,0)\approx(1,0,0)\)。这会使梯度不稳、过早只读一个位置。 除以 \(\sqrt{d_k}\) 是方差尺度修正,不是「让模型更有解释性」的操作。
5. softmax 轴是语义契约
对每一个 batch、每一个 query,沿所有 key 归一化:
weights = scores.softmax(dim=-1) # last axis = N_k
row_sums = weights.sum(dim=-1)
assert torch.allclose(row_sums, torch.ones_like(row_sums), atol=1e-6)
若写成 dim=-2,得到的是每个 key 在不同 query 之间的列归一化。代码可能仍能运行,但「每个 query 如何读取 key」的意义已经反了。
6. Self-attention 与 cross-attention
6.1 Self-attention
Q、K、V 都来自同一组 token:
让一组 token 彼此交换信息。若输入是同一行的 feature states,self-attention 可让不同 feature 相互读取。
6.2 Cross-attention
Q 来自一组,K/V 来自另一组:
这正是 inducing points 从 \(N\) 个 cell 读取、或 query row 从 context rows 读取时需要的结构。
7. Mask:不是「把某处设为零」
Mask 的正确位置是在 softmax 之前 的 score 上:禁止读取的位置填 \(-\infty\)(工程常用 torch.finfo(dtype).min):
scores = scores.masked_fill(blocked, torch.finfo(scores.dtype).min)
weights = scores.softmax(dim=-1)
这样被屏蔽位置在归一化后权重才是 0。若在 softmax 之后简单乘零却不重新归一化,剩余权重总和会小于 1,改变了读取预算。
课程 additive mask 约定:允许处加 0.0,禁止处加 -inf;与 bool blocked + masked_fill 等价, 二选一即可,但 polarity 必须写清。
7.1 Key padding mask
「哪些 key 只是 padding?」形状常为 \([B, N_k]\),广播到 \([B, N_q, N_k]\)。
7.2 Pairwise mask
「某个 query 是否允许读某个 key?」形状 \([B, N_q, N_k]\) 或 \([B, H, N_q, N_k]\)。
Mask 不是 feature,是计算图的可见性控制。
8. 最小实现与契约
def scaled_dot_product_attention(q, k, v, blocked=None, return_weights=False):
d_k = q.shape[-1]
scores = q @ k.transpose(-2, -1) / math.sqrt(d_k)
if blocked is not None:
scores = scores.masked_fill(blocked, torch.finfo(scores.dtype).min)
weights = scores.softmax(dim=-1)
output = weights @ v
return (output, weights) if return_weights else output
若把 \(N_q\) 与 \(N_k\) 写死成相同长度,self-attention 可能通过,cross-attention 会暴露接口错误。
9. Attention 与 MLP 的差别
Attention 的输出依赖其他位置,因此可以表达「这个 token 需要读取哪一个 row/column/feature」;但这 不自动等于因果解释、feature importance 或真实机制发现。权重首先是模型内部的读取系数。
10. 复杂度与设计直觉
当 \(N_q=N_k=N\) 时,主项近似 \(\mathcal{O}(N^2 d)\)。一般情形是 \(\mathcal{O}(N_q N_k d)\)。
这解释了下一课为什么用 inducing points:先把长序列读到较短的 \(K\) 个 latent,再从 \(K\) 读回, 用 \(NK\) 量级交互代替 \(N^2\)。
Mini 简化:本课使用标准 scaled softmax。官方 TabPFN-3 在 Stage 1/3 与 many-class decoder 使用 QASSMax(query-aware scalable softmax),按输入长度重标定 query,以改善大训练集上的长度泛化; Mini 可不实现,但知道边界。
11. 三个课堂实验
实验 A:手算两个 key
构造一个 query 和两个 key,计算 score、softmax 权重和 output。必须能解释每个数字。
实验 B:scale 对数值的影响
固定 seed,比较有无 \(\sqrt{d_k}\) 的 score、权重熵和梯度。不要把「权重更尖锐」写成「更正确」。
实验 C:padding key
给 source 末尾增加 padding key,确认其对所有 query 的权重为 0,且有效 key 的行和仍为 1。
12. 知识检查
- 为什么 output 长度跟 \(N_q\) 一致,而不是跟 \(N_k\) 一致?
- 为什么 softmax 必须沿 \(N_k\)?
- self-attention 和 cross-attention 的 Q/K/V 来源分别是什么?
- 为什么 mask 要在 softmax 前应用?
- 为什么 attention weight 不能自动被称为 feature importance 或 causal effect?
- 当 \(N_q\neq N_k\) 时,哪些 shape assertion 可以最早发现 bug?
13. 本课的硬结论
Attention 是带可见性约束的、数据依赖的加权读取。实现是否正确,首先由 shape、softmax 轴、mask 语义和 reference parity 决定;漂亮的 heatmap 或较低的 loss 都不能替代这些基础证据。
第 2 课:Attention 与 Cross-Attention
本课交付
本课把 attention 从框架调用还原成一个可以逐项检查的模块。结束时应有:
models/attention.py:手写 scaled dot-product attention;- self-attention 和 cross-attention 的同一接口;
return_weights=True的诊断输出;- key padding mask 与 additive mask 的测试;
- 与 PyTorch 参考实现的误差报告和一张 attention matrix 图。
核心直觉:Attention = 带可见性约束的、数据依赖的加权读取。
本课符号
| 符号 | 含义 |
|---|---|
| \(N_q, N_k\) | query / key-value token 数 |
| \(d_k, d_v\) | key / value 维 |
| \(S, A, O\) | score、权重、输出 |
项目问题
表格模型需要让一个 token 读取一组其他 token。Attention 的关键不是「用了 Transformer」,而是:
- 谁在发起读取(Query);
- 去哪里匹配(Key);
- 读取什么内容(Value);
- 哪些位置被 mask 禁止读取;
- softmax 沿哪一条轴归一化。
3 小时课堂流程
| 时间 | 内容 | 证据 |
|---|---|---|
| 00:00–00:25 | Q/K/V 与输出长度 | 手算 2 个 token 的权重 |
| 00:25–00:55 | scaled dot-product 逐步推导 | 每个中间 tensor shape 写出 |
| 00:55–01:35 | 手写 PyTorch module | forward + weights 可运行 |
| 01:35–01:50 | 休息与数值稳定性 | 找出未减 max 的 overflow |
| 01:50–02:20 | self/cross attention 对照 | 改变 query 数量 |
| 02:20–02:50 | mask 与 reference parity | row sums、误差、禁读位置 |
| 02:50–03:00 | review 与作业 | 提交 attention audit |
1. 形状与公式(对齐 canon §3.2)
多头实现可在中间插入 \(H\) 轴:\(Q\in\mathbb{R}^{B\times H\times N_q\times d_k}\)。本课先保证单头正确。
\(N_q\) 决定输出 token 数;\(N_k\) 只决定每个 query 可从多少位置读取。这个事实直接支撑下一课的 \(N\to K\to N\) inducing attention。
若 \(q,k\) 分量近似独立、方差近 1,则 \(\operatorname{Var}(q^{\mathsf T}k)=d_k\),除以 \(\sqrt{d_k}\) 把点积方差拉回常数尺度。
2. 手写实现
# models/attention.py
import math
import torch
from torch import Tensor, nn
class ScaledDotProductAttention(nn.Module):
def forward(self, q: Tensor, k: Tensor, v: Tensor,
additive_mask: Tensor | None = None,
*, return_weights: bool = False):
assert q.ndim == k.ndim == v.ndim == 4 # [B,H,N,*]
assert q.shape[-1] == k.shape[-1]
assert k.shape[-2] == v.shape[-2]
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(q.shape[-1])
if additive_mask is not None:
assert additive_mask.shape[-2:] == scores.shape[-2:]
scores = scores + additive_mask
weights = torch.softmax(scores, dim=-1)
output = torch.matmul(weights, v)
return (output, weights) if return_weights else output
additive mask 约定:允许读取加 0.0,禁止读取加 -inf。不要把未转换的 True/False mask 语义 直接传入;不同 PyTorch API 的 bool polarity 不相同。等价写法是 scores.masked_fill(blocked, finfo.min)(canon:mask 在 softmax 之前)。
3. 手算与可视化练习
用 B=H=1、Nq=2、Nk=3、dk=2 的小整数 tensor,要求手算:
把 weights[0, 0] 画成 \(2\times 3\) heatmap,标记 query row 与 key row。图用于审计路由,不是把 颜色深浅升级成「因果重要性」。
4. self-attention 与 cross-attention
x = torch.randn(2, 5, 16) # [B, N, d]
context = torch.randn(2, 7, 16) # [B, M, d]
# self: q/k/v 都来自 x,输出 [B, 5, d]
# cross: q 来自 x,k/v 来自 context,输出仍是 [B, 5, d]
必须改变 \(N_q\):query 从 5 改成 2,输出第一维从 5 变成 2;context 从 7 改成 11,输出第一维 不能随之变成 11。这是最便宜、却最能抓住 Q/K/V 方向错误的测试。
5. Reference parity 与 mask audit
在固定 seed 下,用相同投影后的 Q/K/V 比较手写模块和 torch.nn.functional.scaled_dot_product_attention。若参考函数接收 bool mask,先转成课程约定的 additive mask,再比较:
max_abs_error(output) < 1e-5 # float32 smoke test 的课程门槛
row_sum_error < 1e-6
再把一个 key 位置设为 -inf,确认所有 query 在该位置的权重为 0。不要只看 output;某些 mask 错误会被 Value 的偶然数值掩盖。
6. 必做实验
实验 A:scale 的作用
\(d_k\) 从 8 增到 128 时,比较有无 \(1/\sqrt{d_k}\) 的 attention entropy;无 scale 的 logits 会变大, softmax 过度尖锐。记录 entropy,而不是只展示一张图。
实验 B:padding key
在一个 batch 内给第二个样本增加 padding keys。mask 前后,第一样本的有效 key 权重总和都应为 1, padding 位置必须为 0。
实验 C:query 数量
固定 context,分别使用 1、4、16 个 query,确认计算图只沿 query 轴扩展,不依赖固定 query size。
Mini 简化:本课用标准 scaled softmax。官方 Stage 1/3 与 decoder 使用 QASSMax;本课不实现。
常见失败
- 把
softmax(dim=-2)写成 key 轴,导致每个 key 的列和为 1; K.transpose(-2, -1)错位,得到「能运行但语义反了」的矩阵;- mask shape 只适配
[N,N],一到 batch/head 就广播错误; - 直接在
float16上对极大 logits 做 softmax; - 把 attention heatmap 的高权重写成 feature importance 或 causal effect。
退出条件与评分
- self/cross 两种调用都通过 shape tests;
- 每行权重和在 float32 下误差小于
1e-6; - 一个 key 被 mask 后,所有 query 对其权重为 0;
- 手写 output 与 reference 在固定输入上满足约定误差;
- 提交一张图和一段文字,解释图中能证明什么、不能证明什么。
评分建议:公式和 shape 25%,实现 30%,mask/parity tests 30%,审计解释 15%。
课后作业
为 attention 模块增加 head_dim 与多头包装,但保持单头模块的 public contract 不变。下一课只允许 调用这个模块,不得重新复制一份 softmax 逻辑。