均匀精度学 ReLU 网络,样本数随深度与输入维度指数爆炸

Learning ReLU networks to high uniform accuracy is intractable

Julius Berner, Philipp Grohs, Felix Voigtlaender

ICLR 2023

cs.LG, stat.ML

2022-05-27

ICLR 论文证明:以均匀范数学给定架构 ReLU 网络,样本随深度与维度指数增长;同样网络的 L2 误差只需多项式样本。

这篇在解决什么

统计学习理论通常保证的是平均误差:在某个数据分布下,期望损失够小就算学会了。安全关键系统、科学计算、以及测试分布会漂移的场景,平均好不够。你要的是对每个输入都接近,也就是均匀范数 L∞,取输入域上点误差的最大值。

经验上这件事很难。SGD 训出来的网络,L1 或 L2 误差可以到千分之几,L∞ 仍有明显尖峰。Adcock 与 Dexter 已经在实验里看到这条缝。这篇把它钉成一个信息论问题:如果目标类包含给定架构的 ReLU 网络,要把均匀误差压到 ε,最少需要多少个点样本。结论与算法无关,SGD、主动学习、随机化采样都算在内。

方法

工具来自信息基复杂度:只数用了多少个点样本,不问算法跑多慢。算法可以按已看到的函数值自适应选下一个点,也可以随机化,只统计平均用了 m 个点。目标类是系数受 ℓ^q 球约束的 ReLU 网络,输入维度 d、深度 L、隐层宽度最多 B,权重和偏置的 ℓ^q 范数不超过 c。

下界靠一簇尖帽函数。ReLU 网可以表示支撑在边长约 1/M 的小立方体上的 bump,高度由正则参数 c、q 和深度 L 决定。把 M 取成 m^{1/d} 量级时,体积论证保证:对任意 m 个采样点,网格上至少一半 bump 完全躲开这些点。算法分不清这些 bump 和零函数,均匀误差就被卡住。正则越弱(q 越大、c 越大),bump 可以越高,下界越狠。

上界更粗:先证这类网络的 Lipschitz 常数有显式上界,再对 Lipschitz 函数做分片常数插值。所以下界不是无穷大,只是指数级。自适应选点帮不上忙,因为最坏 bump 可以藏在任何尚未采样的小立方体里。

结果

主定理覆盖 L≥3 和所有 Lp 范数。任何用 m 个点样本的算法,最小最大误差至少是 c0·Ω/(32s)^{1+s/p}·m^{-1/p-1/s},其中 s≤min{B/3,d}。取 p=∞、宽度至少 3d,均匀精度 ε 需要

m ≥ (Ω/(32d))^d · ε^{-d}。

一个具体数字:ε=1/1024、宽度最多 3d 时,m ≥ 2^d · c^{dL} · (3d)^{d(L-2)}。d=15、c=2、L=7,这个下界已经超过宇宙原子数的常用估计。同一类网络,L2 误差的样本复杂度对 d 只是多项式,来自伪维数有限、经验风险最小化的标准论证。

数值上,student-teacher 用 Adam 拟合随机教师网络,d=1 和 d=3,m 从 10^2 到 10^5,共 8640 次实验。L1/L2 随 m 下降,L∞ 明显更差;一维里 L∞ 几乎停滞。m=100 的最差教师上,L∞ 误差 2.7×10^{-3},大约是 L2(3.9×10^{-4})的七倍,来源就是样本缝里的尖峰。引言那条演示曲线在 m=1000 时 L1=2.8×10^{-3}、L∞=0.19。

为什么重要

这是信息论硬,不是算力硬。就算优化器能完美拟合已有样本,点样本也不够把 ReLU 网钉到高均匀精度。对抗样本、幻觉、模型窃取共享一个几何图景:网络表达能力里藏着采样点看不见的 bump。

如果你真的需要处处准,PDE 求解器、认证、安全场景都属此类,别指望「再多采一点数据加 SGD」自动过关。要么把更强的目标类先验写进算法,比如光滑或低维结构,要么接受平均误差。L2 能学、L∞ 不能学,这篇把这条缝量化了。

局限与存疑

只对 ReLU。别的激活函数作者自己说需要全新方法。上下界在高维有缺口,两边都可能没顶到最优。这是最坏情况分析:对每个算法存在至少一个难函数,难函数是否「典型」没有证明。目标类是「包含整个架构的所有实现」;若额外结构被写进算法,下界可以绕开,但那就不是现在这种端到端深度学习。实验只到 d=3,指数墙在真正高维还没被数值撞上。

术语

原文与代码

社区讨论

相关论文

全部论文解读