P31 · mathematical reading path

批量端到端 Deep Kernel 分类,为什么可以闭式训练?

最反直觉的一步是:多分类 kernel ridge classification 并不需要先训练一个独立的分类器。one-hot 标签的每一列只是一个 KRR 回归目标;同一个 Gram matrix 做一次分解,就能同时求出所有类别的 score functions。

阅读时间约 10 分钟 · 适合已经理解 KRR 回归、但对分类与端到端梯度仍有疑问的读者

先给结论:这里只有一个 persistent optimizer。

每个 batch 内的 coefficient matrix \(A_\theta\) 由 linear solve 临时算出,是 forward pass 的一部分;它不是由第二个 optimizer 维护的模型参数。query loss 通过这个可微求解回传,真正被 optimizer 更新的仍然只有感知网络参数 \(\theta\)。

1. 多分类 KRR = \(C\) 个共享 Gram matrix 的回归

设一个 support batch 有 \(q_s\) 个样本、\(C\) 个类别。把标签写成 one-hot matrix \(Y_s\in\mathbb R^{q_s\times C}\)。例如三分类时:

\[ Y_s=\begin{bmatrix}1&0&0\\0&1&0\\0&0&1\\1&0&0\end{bmatrix} =\big[y^{(1)},y^{(2)},y^{(3)}\big]. \]

第 \(c\) 列 \(y^{(c)}\) 回答的只是一个问题:这个样本是否属于第 \(c\) 类?因此每一列都可以作为普通 KRR 的数值目标:

\[ a^{(c)}_\theta=\left(K_{ss,\theta}+q_s\lambda I\right)^{-1}y^{(c)}, \qquad c=1,\ldots,C. \]

把 \(C\) 个右端项并排放回矩阵,一次求解就是:

\[ \boxed{A_\theta=\left(K_{ss,\theta}+q_s\lambda I\right)^{-1}Y_s} \qquad A_\theta\in\mathbb R^{q_s\times C}. \]
为什么计算上很自然?

左侧矩阵完全相同,只是右端项从一个 target vector 变成 \(C\) 列 targets。数值实现中可复用同一次 Cholesky / linear factorization。

得到的是概率吗?

不是。query 与 \(A_\theta\) 相乘得到的是 \(C\) 个 class scores。最终可用 argmax 解码;若需要概率,还要单独校准或采用 probabilistic loss。

2. 一个 batch 内,到底发生了什么?

把 batch 分成 support 与 query 两部分。support 提供本轮临时核分类器,query 提供真正训练感知网络的外层损失。

① 感知与 kernelZ_s=h_θ(X_s)
Z_q=h_θ(X_q)
K_ss,θ 与 K_qs,θ
② 闭式 kernel headA_θ = solve(K_ss,θ + q_s λI, Y_s)
③ Query scores 与 lossS_q,θ = K_qs,θ A_θ
L = MSE 或 CE
反向:query loss ← scores ← linear solve ← Gram matrices ← \(h_\theta\);optimizer 只更新 \(\theta\)

这里的“先解 \(A_\theta\),再预测”只是 forward pass 的计算顺序,不等于启动第二套训练过程。中间量由当前参数确定,梯度仍能穿过它。

3. 最后一层的分类损失是什么?

query scores 为

\[ S_{q,\theta}=K_{qs,\theta}A_\theta =K_{qs,\theta}\left(K_{ss,\theta}+q_s\lambda I\right)^{-1}Y_s. \]

训练感知网络时有两种直接选择:

One-hot squared error

与 KRR 回归最完全对称:

\[\mathcal L_{\mathrm{MSE}}(\theta)=\frac1{q_q}\lVert Y_q-S_{q,\theta}\rVert_{\mathrm F}^2.\]
Cross-entropy on scores

把 KRR scores 当作 logits 来优化类别区分:

\[\mathcal L_{\mathrm{CE}}(\theta)=-\frac1{q_q}\sum_i\log\operatorname{softmax}(S_{q,\theta})_{i,c_i}.\]
argmax 不进入训练图

\(\widehat c_i=\arg\max_c S_{ic}\) 只用于 validation accuracy 与最终推理。训练直接作用于连续 scores,所以不存在“argmax 不可微导致网络无法训练”的问题。

4. 梯度为什么能穿过 linear solve?

令 \(M_\theta=K_{ss,\theta}+q_s\lambda I\),则 \(A_\theta=M_\theta^{-1}Y_s\)。标签固定时:

\[ \mathrm dA_\theta=-M_\theta^{-1}(\mathrm dK_{ss,\theta})A_\theta. \]

query score 同时通过 query-support kernel 与 support solve 两条路径变化:

\[ \boxed{\mathrm dS_{q,\theta}=(\mathrm dK_{qs,\theta})A_\theta+K_{qs,\theta}\,\mathrm dA_\theta}. \]

因此外层 loss 的梯度可以继续传到 kernel 参数,再传到 representation \(h_\theta(x)\)。在 PyTorch 中,torch.linalg.solve 会为这条路径提供自动微分;实现时不能对 \(A_\theta\) 使用 detach()

Minimal PyTorch-shaped pseudocode
optimizer.zero_grad()

z_s = perception(x_support)
z_q = perception(x_query)
K_ss = kernel(z_s, z_s)
K_qs = kernel(z_q, z_s)

system = K_ss + q_s * ridge * eye(q_s)
A = torch.linalg.solve(system, one_hot(y_support))
scores = K_qs @ A

loss = cross_entropy(scores, y_query)   # or one-hot MSE
loss.backward()                          # gradient crosses solve
optimizer.step()                         # updates perception/kernel θ only

5. 为什么 batch 最好分 support 与 query?

如果同一批样本既建立 kernel classifier、又评估自己,那么对角相似度 \(k(x_i,x_i)\) 可能让目标过度依赖自我重构。support/query split 让外层 loss 真正询问:“由 support 建出的当前 kernel classifier,能否预测未参与求解的 query?”

6. 回归与分类,其实只差在哪里?

对象KRR 回归Kernel ridge classification
Support target\(y_s\in\mathbb R^{q_s}\)\(Y_s\in\mathbb R^{q_s\times C}\)
闭式系数\(\alpha_\theta=M_\theta^{-1}y_s\)\(A_\theta=M_\theta^{-1}Y_s\)
Query output一个连续预测值\(C\) 个连续 class scores
训练 lossscalar MSEone-hot MSE 或 score cross-entropy
推理解码直接输出数值对 scores 取 argmax
感知网络梯度穿过 KRR solve同样穿过多右端项 KRR solve

7. 这个闭式结论的边界

闭式的是 kernel ridge classification,不是所有 kernel classifiers。

Kernel logistic regression 使用 cross-entropy 直接拟合 RKHS functions,kernel SVM 使用 hinge loss;它们通常没有上述闭式 coefficient solve,需要联合优化、unrolled optimization 或 implicit differentiation。P31 若做第一个分类扩展,KRC 是与当前 KRR 回归架构最干净的对应。

分类实验已经完成,但仍是独立的探索性证据。

P31 的主实验与主结论仍是 scalar regression RMSE。另行完成的 classification-v1-exploratory screen 覆盖 Breast Cancer、Digits 与五个 paired seeds,证明该 Gram construction 可以执行二分类与多分类;它不支持“分类上普遍优于 RBF、MLP 或 trees”,也不改变回归结论。