MINI-TABPFN COURSE · LESSON 06

第 6 课讲义:Prediction Heads——为什么分类和回归不能共用同一个输出语义?

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

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

讲义定位

前面得到 query / context 的 ICL 表示。本课回答读出问题:怎样输出分类分布或回归分布?

核心直觉:分类=对 context one-hot 的软检索;回归要单独声明输出语义。「换一个 activation」不够。

与 TabPFN-3 对齐 / Mini 省略:官方分类用 多头 retrieval decoder,再 \(\operatorname{logits}=\log(\operatorname{clip}(p))\);\(C_{\max}=160\);decoder 几何为 \(H=6\)、\(D_h=64\)(「6×64」不是类别上限)。官方回归为 ICL 后 MLP → num_buckets=5000 的 bar-distribution。Mini 可先单头 retrieval,并实现 scalar/bar/quantile 三条教学对照,须标注。

1. 学习目标与先修知识

学完后你应该能够:

  1. 区分 logits、概率、点预测与分布预测;
  2. 写出官方多头 decoder 与 Mini 单头教学形,并说明 \(\log(\mathrm{clip}(p))\);
  3. 审计 \(\sum_n A=1\) 与 \(\sum_c P=1\) 不是同一件事;
  4. 解释 scalar / bar / quantile 假设与官方 bar 路径;
  5. 设计公平 comparison protocol;
  6. 不把 retrieval 升级为因果邻居或机制结论。

符号:\(N_c,N_q\);\(C\) 当前 episode 类别数;\(C_{\max}\) 预训练上限;\(H\) heads;\(D_h\) head dim。

2. 预测头承接什么

\[H_{\mathrm{context}}\in\mathbb{R}^{B\times N_c\times D},\qquad H_{\mathrm{query}}\in\mathbb{R}^{B\times N_q\times D}\]
\[\begin{aligned} \text{classification:}\quad& p\!\left(y\in\{0,\ldots,C-1\}\mid \text{query},\text{context}\right),\\ \text{regression:}\quad& p\!\left(y\in\mathbb{R}\mid \text{query},\text{context}\right). \end{aligned}\]

共享 representation ≠ 共享 output semantics。官方分类与回归用不同 decoder / checkpoint。

3. 官方多头 many-class decoder(canon §3.6,权威)

以最终层 train embeddings \(\{h^{\mathrm{train}}_n\}_{n=1}^{N}\) 为 keys,one-hot \(\mathbf{y}_n\in\{0,1\}^{C}\) 为 values,test \(h^{\mathrm{test}}_m\) 为 queries;经 \(W_Q,W_K\) 与多头拆分:

\[p_m =\frac{1}{H}\sum_{h=1}^{H}\sum_{n=1}^{N} \alpha^{(h)}_{m,n}\,\mathbf{y}_n, \qquad \alpha^{(h)}_{m,n} =\operatorname{softmax}_{n}\!\Bigl( \tfrac{q^{(h)}_m\cdot k^{(h)}_n}{\sqrt{D_h}}\Bigr)\]

随后:

\[\operatorname{logits}_m=\log\bigl(\operatorname{clip}(p_m)\bigr)\]

性质:

  1. 对类别下标置换等变(类不绑死在固定 MLP 输出位);
  2. 参数不随 \(C\) 增长(只依赖 \(D\) 与 \(H\));
  3. 但 checkpoint 仍有 \(C_{\max}=160\)(正交 label embedding 与 one-hot 张量)。

「6×64」= decoder_num_heads × decoder_head_dim不是类别上限;类别上限来自 \(C_{\max}\)。

4. Mini 单头教学形(须标 Mini 简化)

\[A=\operatorname{softmax}_{N_c}\!\Bigl( \tfrac{Q_{\mathrm{query}}K_{\mathrm{context}}^{\mathsf T}}{\sqrt{d_k}}\Bigr),\qquad P=A\,Y_{\mathrm{context}}^{\mathrm{onehot}}\]

形状:\(A\in\mathbb{R}^{B\times N_q\times N_c}\), \(Y^{\mathrm{onehot}}\in\mathbb{R}^{B\times N_c\times C}\), \(P\in\mathbb{R}^{B\times N_q\times C}\)。

审计两条和:

\[\sum_{n=1}^{N_c}A_{bqn}=1,\qquad \sum_{c=1}^{C}P_{bqc}=1\]

不是同一件事。课程实现可先单头,再注明与官方多头 + \(\log(\mathrm{clip}(p))\) 的差距;进阶作业可加 多头平均与 clip-log。

4.1 Task-local vocabulary

须明确:class id 是否映射到 episode 连续范围;是否按 \(C_{\max}\) 再 slicing;多 task batch 如何 pad class 维;未见类是否仍在输出空间。

5. 分类损失与校准

accuracy      top-1 是否正确
NLL / logloss 分布是否给正确答案足够质量
ECE            置信度与频率是否一致

PFN 风格更应记录分布质量,不只 top-1。官方后处理可有 temperature scaling——Mini 可作可选作业。

6. 回归三路径(canon §3.7;Mini 教学)

6.1 Scalar

\[\mu=\operatorname{MLP}(H_{\mathrm{query}}),\qquad \mathcal{L}_{\mathrm{MSE}}=(\mu-y)^2\]

稳、便宜;不表达多峰与区间。

6.2 Bar-distribution

\[p=\operatorname{softmax}_{m}(\operatorname{logits}),\qquad \widehat{y}=\sum_m p_m v_m,\qquad \mathcal{L}=-\log p_{m^\star}\]

官方 regressor:\(512\to 1024\to 5000\) buckets(App.C);任意分位数可由预测 CDF 反演,无需按 \(\tau\) 重训。Mini 可用更少 bins,须记录 range / width / 越界处理。

6.3 Quantile / pinball(教学对照)

\[\mathcal{L}_\tau(y,q)=\tau\max(y-q,0)+(1-\tau)\max(q-y,0)\]

可直接给区间;须记录 crossing。官方主路径是 bar,不是独立 pinball 头——Mini 三条头是对照,不是 官方一一对应。

Head输出优点风险
scalar点/均值简单无分布
bar离散质量直接给分布;对齐官方bins/range 敏感
quantile分位点区间直观crossing;非官方主路径

7. 共享骨干与独立头

shared H_query
  ├── classification head (retrieval / many-class)
  └── regression head (bar / scalar / quantile)

须记录 loss scaling、task kind、output width padding、梯度互压。不能把「同一套参数跑两种 task」写成 「统一解决了分类与回归」。

8. 公平比较与课堂实验

固定:task sampler、held-out split、context/query 规模、backbone/预算、后处理与 target 归一化; 分类报 accuracy+NLL,回归报 RMSE+分布/区间;多种子。

  • A:class axis audit(两条 sum)。
  • B:context 复制,观察 retrieval mass。
  • C:异方差任务上比较 scalar/bar/quantile。

9. 知识检查

  1. 官方 \(p_m\) 公式里 \(H\) 平均与 \(\alpha\) 的 softmax 轴各是什么?
  2. 为何要 \(\log(\mathrm{clip}(p))\),而不是直接把 \(p\) 当 logits?
  3. 「6×64」与 \(C_{\max}=160\) 各约束什么?
  4. Mini 单头与官方多头差在哪里?
  5. 官方回归为何强调 bar + CDF 反演分位数?
  6. 怎样比较 head 才不把不同 loss 混为一谈?

10. 本课的硬结论

Prediction head 的正确性来自输出语义、概率轴、target 边界、loss 与校准证据。官方分类权威形式是 多头 soft retrieval + \(\log(\mathrm{clip}(p))\);Mini 单头须标注简化。Retrieval 首先是预测机制, 不自动等于因果解释。

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

第 6 课:分类头与回归头

项目问题

同一套 ICL hidden 怎样分别回答分类与回归?分类需要概率质量;回归可做点预测、分布或分位数。 对应讲义:lectures/06-prediction-heads.md(canon §3.6–3.7)。核心直觉:分类=对 context one-hot 的软检索;回归须单独声明输出语义。

分类实现

官方权威形式(讲义必须出现;实现可二阶段)

\[p_m =\frac{1}{H}\sum_{h=1}^{H}\sum_{n=1}^{N} \alpha^{(h)}_{m,n}\,\mathbf{y}_n, \qquad \alpha^{(h)}_{m,n} =\operatorname{softmax}_{n}\!\Bigl( \tfrac{q^{(h)}_m\cdot k^{(h)}_n}{\sqrt{D_h}}\Bigr)\]
\[\operatorname{logits}_m=\log\bigl(\operatorname{clip}(p_m)\bigr)\]

官方:\(H=6\),\(D_h=64\)(「6×64」≠类别上限);\(C_{\max}=160\)。

Mini 单头教学形(须标 Mini 简化)

\[A=\operatorname{softmax}_{N_c}\!\Bigl( \tfrac{Q_{\mathrm{query}}K_{\mathrm{context}}^{\mathsf T}}{\sqrt{d_k}}\Bigr),\qquad P=A\,Y_{\mathrm{context}}^{\mathrm{onehot}}\]

实现 models/classification_head.py:至少单头;进阶加多头平均与 \(\log(\mathrm{clip}(p))\)。 可视化 query→context 读出权重。

回归实现

models/regression_heads.py 三条教学路径:

  • scalar:MSE;
  • bar-distribution:bucket softmax + NLL(对齐官方 num_buckets 思路,Mini 可用更少 bins);
  • quantile:pinball(教学对照;官方主路径是 bar + CDF 反演分位数)。

Mini 三条头是对照,不是官方一一对应。官方 regressor:ICL 后 MLP → 5000 buckets(App.C)。

验证实验

  • 分类:\(\sum_n A=1\) 与 \(\sum_c P=1\) 分开断言;
  • retrieval vs 普通 MLP head;
  • 回归:RMSE、区间、分布指标;
  • 异常值 / 异方差失败模式。

边界

更丰富的 predictive distribution 仍是预测对象,不自动等于 Unit belief 或因果不确定性。

交付

classification_head.pyregression_heads.py、分类/回归对照报告。

本课的具体教学包

时间内容证据
00:00–00:25预测对象与指标logits / prob / sample
00:25–00:55retrieval(先单头,对照官方多头公式)Q/K/V 与 class 轴
00:55–01:25scalar regressionMSE / RMSE
01:25–01:40休息target normalization
01:40–02:10bar / quantileNLL / coverage / pinball
02:10–02:45对照实验同 split、不同 head
02:45–03:00review预测边界说明

1. Classification retrieval head

Mini 单头:

\[\begin{aligned} Q&=W_QH_{\mathrm{query}}\in\mathbb{R}^{B\times N_q\times d_k},& K&=W_KH_{\mathrm{context}}\in\mathbb{R}^{B\times N_c\times d_k},\\ V&=\operatorname{one\_hot}(Y_{\mathrm{context}},C) \in\mathbb{R}^{B\times N_c\times C},& A&=\operatorname{softmax}_{N_c}\!\left(\frac{QK^{\mathsf T}}{\sqrt{d_k}}\right) \in\mathbb{R}^{B\times N_q\times N_c},\\ P&=AV\in\mathbb{R}^{B\times N_q\times C}. \end{aligned}\]

进阶对齐官方:多头拆分 → 头平均得 \(p_m\) → \(\log(\mathrm{clip}(p_m))\) 作 logits。

A 最后轴是 context row;P 最后轴是 class。必须分别断言 row sum。

class RetrievalClassificationHead(nn.Module):
    def forward(self, query_hidden, context_hidden, context_labels,
                n_classes: int, context_mask=None):
        # probs [B,Nq,C], diagnostics {"read_weights": [B,Nq,Nc]}
        ...

另写 LinearClassificationHead 作 control。未见类可为 0——这是结构结果,不是校准保证。

2. 三条回归路径

Scalar

\[\mu=\operatorname{MLP}(H_{\mathrm{query}}),\qquad \mathcal{L}_{\mathrm{MSE}}=(\mu-y)^2\]

Bar-distribution

\[p=\operatorname{softmax}_{m}(\operatorname{logits}),\qquad \widehat{y}=\sum_m p_m v_m,\qquad \mathcal{L}=-\log p_{m^\star}\]

记录 clipping、bin width、越界处理。可选:由 CDF 反演若干分位数,对照 pinball head。

Quantile

\[\mathcal{L}_\tau(y,q)=\tau\max(y-q,0)+(1-\tau)\max(q-y,0)\]

报告 median、区间 coverage / width;记录 crossing rate。

3. Guided lab:固定 comparison protocol

同 task seeds、context/query split、upstream hidden、优化步数、target 归一化、eval 列表。 分类:accuracy + NLL(+ 可选 ECE);回归:RMSE + 分布/区间 + cost。

4. 必做实验

A:class axis audit

\[\sum_{n}A_{bqn}=1,\qquad \sum_{c}P_{bqc}=1\]

形状示例:B=2,Nq=4,Nc=7,C=3

B:context duplication

复制一 context row,观察 retrieval mass 与预测变化;报告现象,不自动称 bias/causal。

C:异方差回归

scalar / bar / quantile 比中心误差与区间;画 x vs interval。

5. 常见失败

  • 把 context 轴 softmax 结果直接当 class probability;
  • one-hot 宽度不随 task \(C\) 调整;
  • 把「6×64」当成类别上限;
  • 实现只有单头却声称已复现官方 many-class decoder(未写 Mini 简化 / 未做 \(\log(\mathrm{clip}(p))\));
  • bar 的 bin 边界 train/eval 不一致;
  • quantile crossing 未记录;
  • 用 test target 做 normalization;
  • 把预测分布写成 Unit belief。

退出条件与评分

  1. 两条 softmax 轴通过 shape/sum;
  2. linear vs retrieval 对照可跑;
  3. scalar / bar / quantile 至少两种可训练出指标;
  4. 报告含中心、分布/区间、成本;
  5. 写明官方多头 + \(\log(\mathrm{clip}(p))\) 与 Mini 单头的差距;
  6. 未自动赋予因果/epistemic 边界。

评分:classification 30%,regression 30%,metric 25%,边界 15%。

课后作业

(1) 实现多头平均 + \(\log(\mathrm{clip}(p))\),与单头对照。 (2) temperature calibration:只在 held-out calibration tasks 上拟合;不用最终 test query labels 调参。