MINI-TABPFN COURSE · LESSON 02

第 2 课讲义:Attention——模型怎样读取 token 之间的关系?

这是一节可直接阅读、实现和验收的课程单元。页面正文由 canonical Markdown lesson source 投影生成。

第 02 / 08 课course-v0.5 / owner-review完整讲义 + 动手验收包
课程讲义 · 基础知识层canonical source: course/lectures/02-attention.md · 先建立概念、符号、公式与边界,再进入工程实现。

讲义定位

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. 学习目标与先修知识

学完后你应该能够:

  1. 用「查询—匹配—读取」解释 Q/K/V;
  2. 写出 scaled dot-product attention 的每个中间 shape;
  3. 解释为什么 softmax 要沿 key 轴归一化;
  4. 区分 self-attention、cross-attention 与普通 MLP;
  5. 写出 key padding mask 和 pairwise attention mask 的语义;
  6. 通过手算、PyTorch parity、row-sum 和 mask audit 检查实现。

先修是矩阵乘法、softmax 和 batch tensor。

2. Attention 之前:一个可学习的加权平均

假设有 value 向量 \(v_1,\ldots,v_{N_k}\),希望根据查询决定如何加权读取:

\[\operatorname{output}=\sum_{j=1}^{N_k}a_j v_j\]

若 \(a_j\ge 0\) 且 \(\sum_j a_j=1\),则 output 是 value 的凸组合。Attention 的贡献是:权重不是固定 规则,而是由 query 与 key 的匹配关系学习出来。

3. Q、K、V 的三种角色

对输入 token \(x\),经三个不同线性投影:

\[q=W_Q x,\qquad k=W_K x,\qquad v=W_V x\]

\(q\):「我正在寻找什么」;\(k\):「我可以被怎样匹配」;\(v\):「匹配后真正读取的内容」。

权威形式(canon §3.2;单头、无 head 轴):

\[S=\frac{QK^{\mathsf T}}{\sqrt{d_k}},\quad A=\operatorname{softmax}_{\mathrm{key}}(S),\quad O=AV\]

形状契约:

\[\begin{aligned} Q&\in\mathbb{R}^{B\times N_q\times d_k},& K&\in\mathbb{R}^{B\times N_k\times d_k},& V&\in\mathbb{R}^{B\times N_k\times d_v},\\ S,A&\in\mathbb{R}^{B\times N_q\times N_k},& O&\in\mathbb{R}^{B\times N_q\times d_v}. \end{aligned}\]

每个 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\) 增长:

\[\operatorname{Var}(q^{\mathsf T}k)=d_k \quad\Rightarrow\quad \operatorname{Var}\!\Bigl(\tfrac{q^{\mathsf T}k}{\sqrt{d_k}}\Bigr)=1\]

未经缩放的 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:

\[X\in\mathbb{R}^{B\times N\times d_{\mathrm{in}}} \longrightarrow (Q,K,V) \longrightarrow O\in\mathbb{R}^{B\times N\times d_{\mathrm{out}}}\]

让一组 token 彼此交换信息。若输入是同一行的 feature states,self-attention 可让不同 feature 相互读取。

6.2 Cross-attention

Q 来自一组,K/V 来自另一组:

\[Q\leftarrow X_{\mathrm{query}},\qquad (K,V)\leftarrow X_{\mathrm{source}},\qquad O\in\mathbb{R}^{B\times N_q\times d_{\mathrm{out}}}\]

这正是 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 的差别

\[o_i=\operatorname{MLP}(x_i) \qquad\text{vs}\qquad o_i=\sum_j a_{ij} v_j\]

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. 知识检查

  1. 为什么 output 长度跟 \(N_q\) 一致,而不是跟 \(N_k\) 一致?
  2. 为什么 softmax 必须沿 \(N_k\)?
  3. self-attention 和 cross-attention 的 Q/K/V 来源分别是什么?
  4. 为什么 mask 要在 softmax 前应用?
  5. 为什么 attention weight 不能自动被称为 feature importance 或 causal effect?
  6. 当 \(N_q\neq N_k\) 时,哪些 shape assertion 可以最早发现 bug?

13. 本课的硬结论

Attention 是带可见性约束的、数据依赖的加权读取。实现是否正确,首先由 shape、softmax 轴、mask 语义和 reference parity 决定;漂亮的 heatmap 或较低的 loss 都不能替代这些基础证据。

HAND-ON LAB · 工程实现与验收包

第 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」,而是:

  1. 谁在发起读取(Query);
  2. 去哪里匹配(Key);
  3. 读取什么内容(Value);
  4. 哪些位置被 mask 禁止读取;
  5. softmax 沿哪一条轴归一化。

3 小时课堂流程

时间内容证据
00:00–00:25Q/K/V 与输出长度手算 2 个 token 的权重
00:25–00:55scaled dot-product 逐步推导每个中间 tensor shape 写出
00:55–01:35手写 PyTorch moduleforward + weights 可运行
01:35–01:50休息与数值稳定性找出未减 max 的 overflow
01:50–02:20self/cross attention 对照改变 query 数量
02:20–02:50mask 与 reference parityrow sums、误差、禁读位置
02:50–03:00review 与作业提交 attention audit

1. 形状与公式(对齐 canon §3.2)

\[S=\frac{QK^{\mathsf T}}{\sqrt{d_k}},\quad A=\operatorname{softmax}_{\mathrm{key}}(S),\quad O=AV\]
\[\begin{aligned} Q&\in\mathbb{R}^{B\times N_q\times d_k},& K&\in\mathbb{R}^{B\times N_k\times d_k},\\ V&\in\mathbb{R}^{B\times N_k\times d_v},& O&\in\mathbb{R}^{B\times N_q\times d_v}. \end{aligned}\]

多头实现可在中间插入 \(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=1Nq=2Nk=3dk=2 的小整数 tensor,要求手算:

\[S\in\mathbb{R}^{1\times1\times2\times3},\qquad \sum_{k=1}^{3}A_{0,0,0,k}=1,\qquad O\in\mathbb{R}^{1\times1\times2\times d_v}\]

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。

退出条件与评分

  1. self/cross 两种调用都通过 shape tests;
  2. 每行权重和在 float32 下误差小于 1e-6
  3. 一个 key 被 mask 后,所有 query 对其权重为 0;
  4. 手写 output 与 reference 在固定输入上满足约定误差;
  5. 提交一张图和一段文字,解释图中能证明什么、不能证明什么。

评分建议:公式和 shape 25%,实现 30%,mask/parity tests 30%,审计解释 15%。

课后作业

为 attention 模块增加 head_dim 与多头包装,但保持单头模块的 public contract 不变。下一课只允许 调用这个模块,不得重新复制一份 softmax 逻辑。