清华「有效阶」:把神经网络有多简单算成一个数,预测泛化胜过 sharpness

Quantifying and Optimizing Simplicity via Polynomial Representations

Tianren Zhang, Xiangxin Li, Minghao Xiao, Guanyu Chen, Feng Chen

ICML 2026

cs.AI

2026-05-28

清华用插值路径上拟合多项式的「有效阶」当简单性度量,与泛化差距的相关性超过 sharpness;作正则项在 CIFAR-10 ViT-Tiny 上比基线高 3.02 个点。

这篇在解决什么

深度网络为什么能泛化,主流解释是「简单性偏置」:模型倾向于找简单的解。可「简单」到底怎么算,一直没有一个大家都认的、能量化的答案。这不是纯学术洁癖。训练完一堆 checkpoint,你想挑一个最可能泛化的;在两种训练配方之间,你想选更稳的那个。今天最常见的做法是看 sharpness(锐度),也就是损失面有多陡,但它对参数重标定极度敏感,换个实现细节数值就变。

论文把一个合格的简单性度量拆成三条要求:跨任务跨架构通用、能在训练好的大模型上算出来、(近似)可微从而能反过来当训练目标。现有候选各有短板。能证明的隐式偏置(最大间隔、最小范数解)只在特殊情形成立,推不到深网络;信息论那一路(压缩、描述长度)原则普适,但给神经函数算不出来,更别提当损失;几何和容量指标(样条、线性区计数)依赖具体架构,大规模算不动;参数空间的代理(范数、锐度)又受重参数化摆布。三条同时满足的度量,目前缺位。

方法

核心是把高维输入压成一维,再看函数有多「弯」。

取两个真实样本 x₁、x₂,它们之间的直线 x(α)=αx₁+(1−α)x₂,α∈[0,1],是一条插值路径。把网络在这条线上求值,g(α)=f(x(α)),高维函数就此退化成一条一维曲线。接着用切比雪夫正交多项式基逼近 g:P(α)≈Σck Tk(2α−1)。切比雪夫基数值稳定;采样点用一种分层切比雪夫节点取,把点聚到两端,压住多项式拟合在边缘震荡的 Runge 现象。输出维度高时,每条路径先做一次 PCA,只在前 m 个主成分上拟合。

度量本身很直接:有效阶 ED(P)=Σ|ck|·k,把每个系数的绝对值乘以对应阶数再求和。曲线越接近低阶多项式,ED 越小。再对 PCA 维度和采样到的路径对取平均,得到整网的 ED。定理 3.1 保证:只要采样足够,随机路径平均后保留多项式的阶数序,所以「路径上的阶」是「函数阶」的合法代理。

关键在于端点用的是真实样本,路径贴着数据流形走。这不是细节:换成随机像素,整个方法就失效。

ED 可微,作者给了闭式梯度(命题 5.1),用阻尼最小二乘(TᵀT+εI)避开发病求逆,能直接当正则项:L=Ltask+λ·ED。还有一招叫「标签锚定」:把路径两端(α=0、α=1)的网络输出替换成真实标签,免得 ED 项在数据点上和交叉熵打架。

结果

ED 先后被当两样东西用:一个测量仪器,一个正则项。

当测量仪器,它和泛化差距的相关性压过所有对手。CIFAR-10 上 ResNet18 和 ViT-Tiny,ED 与泛化差距的皮尔逊相关最强,sharpness 一族明显更弱,参数 L2 范数甚至是负相关或几乎无关。CLIP ViT-B/32 在 ImageNet 上微调结论一样:ED 正相关,sharpness 反而负相关。Grokking 实验最能说明问题:在 Z97 模除法(30% 训练集)上,模型先死记后突然泛化,只有 ED 准确捕捉到这个相变,记忆阶段 ED 上升、验证损失骤降时达峰、之后回落,说明最终泛化的解确实更简单;参数范数一路单调上升,sharpness 来回抖,都给不出清晰转折信号。

当正则项,提升在多种设置上一致。CIFAR-10 ViT-Tiny:

方法Top-1 (%)
Baseline87.80
Mixup88.83
SAM87.85
ASAM87.85
Jacobian reg87.81
ED(本文)90.82

ED 比基线高 3.02 个点,SAM、ASAM、Jacobian 基本原地踏步。ImageNet 从头训 ViT-S/16,原始配方 71.37→72.76,强配方 74.42→75.01。CLIP 微调:ViT-B/32 在 ImageNet 上 76.20→77.14,五个 OOD 集平均 44.04→45.31;ViT-B/16 更进一步 81.35→82.19。文本侧 GLUE 的 BERT-base,RTE 70.28→71.12、MRPC 86.74→87.66、CoLA 62.31→62.45,而 embedding-mixup 在文本上并不稳定、有时还掉点,ED 没有。强化学习这边,Procgen 上的 CNN-PPO 把 ED 加到 actor 网络,四个环境在未见关卡上的泛化全部提升。

为什么重要

ED 作为一个跨架构、跨任务的泛化诊断数,比那 3 个点更值钱。今天从业者要比较训练配方或挑 checkpoint,手里没有好用的泛化预测器,sharpness 又不靠谱。ED 给了一个能用一个数字说话的工具,而且这个数字还能反过来当损失去优化,测量和优化合二为一。正则项本身架构无关,ViT、ResNet、BERT、CLIP、PPO 的 actor 都直接挂得上,和现有技巧正交。诚实讲:ImageNet 上绝对提升只有 1 个点上下,这不是新 SOTA,卖点是「换哪儿都涨一点」的一致性。

局限与存疑

「简单」不等于「鲁棒」,这是 ED 最硬的天花板。作者自己给了一个失败案例:把 MNIST 的 0、1 和 CIFAR 的汽车、卡车拼成一个二分类(用简单 MNIST 特征或复杂 CIFAR 特征都能解),结果 ED 正则后的模型(99.90)和基线(99.85)一样,几乎全靠那个更简单也更脆弱的 MNIST 信号;一旦把 MNIST 部分随机化,准确率崩到 48%。ED 会把模型推向更简单的解,但当更简单的特征恰好更脆弱时,它帮倒忙。

ED 强依赖端点是真实样本。消融里把插值端点换成均匀随机像素,90.82 直接掉回 87.31,几乎等于没用 ED。这说明它是个分布感知的度量,不是纯函数空间的,而真实采样本身要成本。

计算开销也得算进账。附录 G 实测:CIFAR-10 每个 epoch 从 6.14 秒涨到 9.75 秒,CLIP 微调每步从 0.44 秒到 0.90 秒,接近 2 倍。作者称「可接受」,但预训练规模下这是实打实的成本,λ 和采样预算都得调。

理论层面,「路径上的多项式阶」何时能忠实代表「函数的简单性」,作者也只是部分形式化,留作开放问题。

术语

原文与代码

社区讨论

相关论文

全部论文解读