如果说参数量决定了模型复杂度,那 100 亿参数的模型应该比 10 亿参数难解释 10 倍。但真正做过压测和部署的人都知道,这种直觉几乎不成立。很多大模型在高强度蒸馏之后,去掉大量冗余结构,精度几乎不掉;反过来,有些模型参数不多,却极度难收敛、难解释。原因在于:参数量衡量的只是“存储空间”,而不是“有效自由度”。真正决定模型简单还是复杂的,是它在数据分布上的有效自由度,也就是有效维度(Effective Dimension, ED)。正好,ICML 2026 投稿周期里有一批工作开始把这个问题和多项式表示绑定在一起,其中最直接的一个思路就是:给神经网络一个多项式表示,用 ED 量化它有多少个“真正有用的自由度”,然后把降低 ED 作为优化目标,让网络在同精度下变得更简单、更好解释、更容易压缩。
这篇文章不会假装已经拿到了完整论文代码。标题里的 ICML 2026 表明它很可能是当前投稿/预印本阶段的工作,很多细节还没完全公开。因此本文做三件事:第一,拆解 ED 和多项式表示到底在解决什么问题;第二,用 PyTorch 给出一个可以直接运行的最小实现,让你不需要等官方代码也能体会这个思路;第三,结合工程经验,聊聊这类方法在真实项目里的适用边界和常见坑。
1. 这篇文章真正要解决的问题
先回想一下,为什么神经网络普遍被认为“复杂”?最常见的回答是:层数多、参数多、非线性强。但这三个特征都只能描述模型的形态,不能描述模型的行为。同样一个 ResNet 结构,训练集不同、正则手段不同、BN 统计量不同,行为差异可能非常大。一个模型在推理时,如果很多神经元输出饱和、很多权重共线、很多残差分支趋近于零,那它在数据分布上就只用了很少一部分表达能力,也就是“低复杂度的行为”。
这里真正值得思考的是:能不能用一种不依赖具体结构、只依赖“网络在数据上如何响应”的方式,来量化网络的复杂度?
ED 就是为此出现的。它的基本思想很朴素:把网络的参数看作一个高维空间,训练过程在这一空间里会显著改变某些方向,而对另一些方向几乎不敏感。那些被显著影响的“方向”数量,就是网络在这个任务上的有效维度。把这个数字降下来,网络的行为就变得更可理解、更低秩、更接近一个简单的函数。
但 ED 只是一个数值,光告诉别人“这道模型的 ED 是 12.7”没有太大帮助。问题在于:这 12.7 个自由度到底对应什么?于是,多项式表示登场了。如果网络用了 GELU、Tanh、SiLU 这类光滑激活函数,它在一个有界区域内可以视为多项式的函数,因此网络输出可以被展开成输入的多项式形式。这时候,ED 不再是一个笼统的数字,而可以被细化成“在 1 阶项、2 阶项、3 阶项上分别占了多少自由度”。一旦复杂度有了这个“可归属的结构”,我们就可以真正去优化它:让高阶多项式项的谱能量尽量小,让低阶项保留主要表达力。
整篇文章的核心判断是:把“复杂”理解为“输入-输出关系中的高次项太多”,比理解为“网络参数太多”要更接近深度学习的真实情况。多项式表示拿来描述这种关系,ED 拿来量化它,两个工具拼在一起,就构成了一条“先测量、再压缩、后解释”的技术路线。
2. 核心概念:有效维度、简单性与多项式表示
2.1 有效维度:不是参数量,而是激活方向的个数
先给一组直觉。假设你有一个 100 个参数的模型,但训练之后,Fisher 信息矩阵的特征谱里只有 3 个特征值明显大于零,剩余 97 个方向无论怎么扰动,损失函数都不变。那么这个模型实际上只用了 3 个有效自由度,ED 约等于 3。训练过程在参数空间中只是沿着一条三维修正曲面在走,剩下的 97 维都是“死胡同”。
更严谨地说,ED 通常借助参数 θ 在样本上的 Fisher 信息矩阵 F 的特征谱来定义。F 的特征值衡量了参数各个方向对模型输出的影响程度。特征值大的方向,对预测结果影响明显,属于“重要自由度”;特征值接近于零的方向,属于“可忽略自由度”。假设 F 的特征值为 λ1 ≥ λ2 ≥ ... ≥ λd,那么有效维度可以理解成:保留绝大多数谱能量所需的最小方向数量。
这样做的好处是,ED 完全不关心网络层的具体结构,只关心网络在某个数据分布上产生了什么样的几何性质。你完全可以用同一个 ED 指标去对比 MLP、CNN、Transformer 在同一任务上的“行为复杂度”。
2.2 简单性如何定义
“简单性”不是哲学概念,在工程上至少可以操作化。一个人为把网络的简单性定义为“在保持任务精度不下降的前提下,网络输入-输出关系的多项式展开中有效项越少越简单”。
这个定义有三层含义:
- 简单不等于参数量少。一个参数量大但高度低秩的网络,行为上仍然很“简单”。
- 简单和任务精度绑定。不能为了简单而毁掉任务表现,否则再简单也没有意义。
- 简单要落在输入-输出关系上。模型的复杂度度量应该基于函数行为,而不是权重张量的尺寸。
有了这个定义,我们就可以把“优化简单性”转化为“在损失函数上增加一个复杂度惩罚”,或者“对网络结构施加某个谱约束”。ED 恰好能充当这个惩罚项,多项式表示则让这个惩罚有明确的结构指向。
2.3 多项式表示从哪来
为什么偏偏是多项式,而不是三角级数、小波基或者其他函数基底?原因是神经网络本身就有很强的多项式偏向。
对于光滑激活函数,比如 GELU、Tanh、SiLU,网络的前向传播可以看作一系列光滑函数的复合。根据泰勒展开,任意光滑函数在某一点附近都能用多项式逼近。网络层数越深,逼近中的最高阶项越高,而不是说网络天然就是多项式函数。ReLU 网络严格来说不是光滑的,但在线性区域内同样呈现分段多项式性质。因此,用多项式来表示网络的行为,在局部意义上是合理的。
更重要的是,多项式展开中的“阶数”有直观解释:
| 多项式阶数 | 含义 | 网络行为特征 |
|---|---|---|
| 1 阶(线性项) | 输入特征被线性放大/缩小 | 可解释性强,近似线性函数 |
| 2 阶(二次项) | 特征之间的两两交互 | 能表达简单非线性关系 |
| 3 阶及以上(高次项) | 复杂交互和高频变化 | 表达能力强,但容易过拟合、难解释 |
一个网络如果高阶项能量低,那它的输入-输出关系就接近简单多项式;如果高阶项能量很高,说明模型在利用非常复杂的特征交互。通常我们希望后者受到控制和约束。
3. 方法框架拆解:ED 如何对齐到多项式结构
3.1 从网络到多项式展开的抽象
不妨假设经过某种变换后,网络输出的第 j 个分量可以写成关于输入 x 的多项式形式:
f_j(x) ≈ ∑_{k=1}^{p} ∑_{i_1,...,i_k} W_{j,i_1,...,i_k}^{(k)} x_{i_1} ... x_{i_k}
这里 W^{(k)} 就是 k 阶多项式项的系数张量。实际上这个展开可能并不是网络前向传播的精确等价,但是在光滑激活和局部采样条件下,它提供了一个可供分析的行为代理。
用多项式表示来重写网络,最大的价值在于:参数 W 在原始权重空间里没有明确的“阶数”归属,但在多项式系数空间里,每一项都带着明确的复杂度级别。高阶系数越大,表示模型越依赖复杂函数关系;高阶系数越小,表示模型行为越简单。
3.2 Fisher 信息与多项式项的关联
接下来把 ED 落实到多项式系数上。假设网络输出是 y = f_θ(x),我们在参数 θ 上定义 Fisher 信息矩阵:
F = E_x[ ∇_θ log p(y|x,θ) ∇_θ log p(y|x,θ)^T ]
在实际计算中,通常用训练集样本的梯度外积平均来逼近。F 的特征谱刻画了参数空间里的重要方向。
当网络用多项式表示后,不同阶数对应的系数会分布在 F 的不同特征方向上。于是 ED 的计算结果可以分解成:
ED = ED_1 + ED_2 + ... + ED_p
其中 ED_k 表示 k 阶多项式系数对特征谱的贡献。这样一来,ED 就不仅仅是“一个数”,而是“一张复杂度分布表”。比如,一个模型 ED 为 15,可能其中 8 个自由度来自 1 阶项、5 个来自 2 阶项、2 个来自 3 阶项。这对理解模型行为非常关键。
3.3 优化目标:精度与简单性的平衡
有了上述度量,优化“简单性”就变成约束 ED_3、ED_4 等高阶项的能量,或在训练损失上加入高阶谱惩罚:
L_total = L_task + λ * R_poly(θ)
其中 L_task 是原始任务损失,R_poly(θ) 是一个针对高阶多项式系数谱的惩罚项。这样做的好处是:它不像 L1 正则那样简单地把权重推到零,而是把“复杂度”从高阶项推向低阶项,实现“模型行为降阶”,而不是彻底杀死模型能力。
下面的第 5 节会给出一个可以运行的简化版本。
4. 环境准备与最小可运行实验
在动手实验前,先明确环境。本文示例偏教学与验证,不追求大规模跑分,因此版本约束很宽松:
| 组件 | 建议要求 | 说明 |
|---|---|---|
| 操作系统 | Linux/macOS/Windows 均可 | 无特殊依赖 |
| Python | 3.9+ | 类型注解与 torch 兼容性更稳 |
| PyTorch | 2.x 或 1.13+ | 需要 autograd 和基本矩阵运算 |
| torchvision | 与 PyTorch 版本匹配 | 仅用于 MNIST 数据 |
| CPU/GPU | 均可 | 小规模实验 CPU 即可跑通 |
推荐使用虚拟环境隔离依赖:
conda create -n ed_poly python=3.10 conda activate ed_poly pip install torch torchvision matplotlib这里不写死具体版本,因为以当前日期为准,PyTorch 版本迭代很快。如果你的项目已有版本约束,建议优先兼容现有环境。下面的代码会用到torch.autograd.grad、torch.linalg.eigvalsh和基本的nn.Module,这三个 API 在 1.13 以上都稳定存在。
5. 核心代码实现
5.1 准备小型模型与数据
先构造一个两层光滑激活的 MLP。注意,这里特意选择 Tanh 作为激活,因为 Tanh 在原点附近是光滑的,便于用多项式视角理解。
# 文件路径:demo_poly_ed/model.py import torch import torch.nn as nn import torch.nn.functional as F class TwoLayerTanhMLP(nn.Module): def __init__(self, input_dim=28 * 28, hidden_dim=64, num_classes=10): super().__init__() self.fc1 = nn.Linear(input_dim, hidden_dim) self.fc2 = nn.Linear(hidden_dim, num_classes) def forward(self, x): x = x.view(x.size(0), -1) h = torch.tanh(self.fc1(x)) out = self.fc2(h) return out使用 MNIST 的简化 loader:
# 文件路径:demo_poly_ed/data.py from torch.utils.data import DataLoader from torchvision import datasets, transforms def get_mnist_loaders(batch_size=256): transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_ds = datasets.MNIST(root="./data", train=True, download=True, transform=transform) test_ds = datasets.MNIST(root="./data", train=False, download=True, transform=transform) return DataLoader(train_ds, batch_size=batch_size, shuffle=True), DataLoader(test_ds, batch_size=batch_size)这段代码没有特别之处,作用是快速拿到一个可训练的数据流。如果你需要在自己的数据集上实验,只需要替换get_*_loaders的返回内容,保持 DataLoader 格式一致即可。
5.2 计算 Fisher 特征谱
这是 ED 的核心计算。思想是对测试集上的每个样本计算模型参数梯度,然后组成样本梯度矩阵,再计算其协方差矩阵的特征值。严格来说这并非完整 Fisher 信息矩阵,但在分类问题中,用损失函数的梯度外积近似 Fisher,已经是常见做法。
# 文件路径:demo_poly_ed/fisher.py import torch import torch.nn.functional as F def compute_gradient_matrix(model, dataloader, max_samples=256, device="cpu"): """遍历数据,收集每个样本对全部参数的梯度,拼成 [N, d] 矩阵。""" model.to(device) model.train() grads = [] total = 0 for x, y in dataloader: x, y = x.to(device), y.to(device) logits = model(x) loss = F.cross_entropy(logits, y) # 对全部参数求导 params = list(model.parameters()) grads_per_sample = torch.autograd.grad( loss, params, create_graph=False, retain_graph=True, allow_unused=True ) # 把每个参数的梯度展平并拼接 flat = [] for g in grads_per_sample: if g is None: continue flat.append(g.detach().reshape(-1)) flat = torch.cat(flat) # 此处得到的是整个 batch 的和,为了得到逐样本梯度, # 在实际工程中需要逐样本计算或使用批量 Jacobian。 # 为了演示,这里只有一个 batch 的整体梯度,后面的代码会做近似。 grads.append(flat) total += x.size(0) if total >= max_samples: break return torch.stack(grads) def estimate_fisher_spectrum(model, dataloader, max_samples=256, device="cpu"): """返回特征值列表,降序排列。""" G = compute_gradient_matrix(model, dataloader, max_samples, device) # [N, d] centered = G - G.mean(dim=0, keepdim=True) cov = centered.T @ centered / max(centered.size(0) - 1, 1) eigenvalues = torch.linalg.eigvalsh(cov) return eigenvalues.flip(0)需要注意,上面的compute_gradient_matrix为了保持代码简洁,实际得到的是 batch 梯度,而不是逐样本梯度。这是教学代码与论文实验之间的主要差距。如果要做严肃实验,应该使用下面的逐样本梯度实现:
def compute_per_sample_gradients(model, batch_x, batch_y, device="cpu"): model.to(device) batch_x, batch_y = batch_x.to(device), batch_y.to(device) per_sample_grads = [] for i in range(batch_x.size(0)): x_i = batch_x[i:i+1] y_i = batch_y[i:i+1] logits = model(x_i) loss = F.cross_entropy(logits, y_i) params = list(model.parameters()) grads = torch.autograd.grad(loss, params, retain_graph=False, allow_unused=True) flat = [] for g in grads: if g is None: continue flat.append(g.reshape(-1)) per_sample_grads.append(torch.cat(flat)) return torch.stack(per_sample_grads) # [batch_size, d]逐样本计算的代价是数据量大了之后很慢,但结果更接近 Fisher 信息矩阵的真实估计。你在做小规模论文复现时,优先用逐样本版本。生产环境要优化的话,可以改用 K-FAC 或 Lanczos 方法,避免显式构造完整的 d×d 矩阵。
5.3 计算有效维度 ED
拿到特征谱之后,ED 的计算就有多种口径。这里给一个最直观的“谱能量覆盖率”版本:
# 文件路径:demo_poly_ed/effective_dim.py import torch def effective_dimensionality(eigenvalues, coverage=0.95): """ 计算有效维度:累计特征能量达到总能量 coverage 所需的最少方向数。 参数: eigenvalues: 从大到小排列的特征值张量 coverage: 覆盖比例,例如 0.95 表示保留 95% 的谱能量 返回: ed: 有效维度整数值 """ total_energy = eigenvalues.sum() if total_energy <= 0: return 0 cum_energy = torch.cumsum(eigenvalues, dim=0) threshold = coverage * total_energy indexes = torch.nonzero(cum_energy >= threshold, as_tuple=False) if indexes.numel() == 0: return eigenvalues.numel() return int(indexes[0].item()) + 1这个定义简单、稳定、可解释。它衡量的是:在 Fisher 信息矩阵的特征谱中,需要多少个主要方向才能解释 95% 的参数敏感性。ED 越低,说明网络行为越集中在少数方向上。
如果你想和论文中常见的 ED 公式对齐,还需要仔细阅读最终公开版本的数学定义。但作为先跑通流程的实验,覆盖率版本已经足够说明问题。
5.4 多项式表示的简易化实现
为了演示“多项式表示”如何与 ED 结合,这里构造一个简单的多项式特征 MLP。它先对输入计算 1 到 p 次幂,再送入一个 MLP 分类器:
# 文件路径:demo_poly_ed/poly_model.py import torch import torch.nn as nn class PolyFeatureMLP(nn.Module): def __init__(self, input_dim=28 * 28, degree=3, hidden_dim=32, num_classes=10): super().__init__() self.input_dim = input_dim self.degree = degree # 多项式基的展开维数 self.poly_dim = input_dim * degree self.fc1 = nn.Linear(self.poly_dim, hidden_dim) self.fc2 = nn.Linear(hidden_dim, num_classes) def forward(self, x): x = x.view(x.size(0), -1) # 构造 1 阶、2 阶、...、p 阶多项式特征 poly_feats = [] for k in range(1, self.degree + 1): poly_feats.append(x ** k) poly_x = torch.cat(poly_feats, dim=1) h = torch.tanh(self.fc1(poly_x)) out = self.fc2(h) return out def get_poly_coefficient_layer(self): """返回负责多项式特征的第一层权重,形状为 [hidden, input_dim*degree]""" return self.fc1.weight在这个模型里,fc1.weight的第 k 个分块天然就对应 k 阶多项式的线性组合系数。我们完全可以在训练时只对高阶分块施加惩罚,实现“低阶保留、高阶压缩”。
5.5 训练循环:加入高阶结构调整
下面这个训练循环演示一个最简化的“优化简单性”方式:对高阶多项式系数施加额外惩罚,同时周期性打印当前 ED,观察压缩前后 ED 的变化。注意,真实论文的做法可能更复杂,比如通过 Fisher 谱的软阈值来反向传播梯度;这里保留核心思想即可。
# 文件路径:demo_poly_ed/train.py import torch import torch.nn.functional as F from .effective_dim import effective_dimensionality def train_with_poly_regularization( model, train_loader, test_loader, epochs=5, lr=1e-3, poly_lambda=1e-4, degree=3, device="cpu" ): optimizer = torch.optim.Adam(model.parameters(), lr=lr) model.to(device) for epoch in range(epochs): model.train() total_loss = 0.0 for x, y in train_loader: x, y = x.to(device), y.to(device) logits = model(x) ce_loss = F.cross_entropy(logits, y) # 取出 fc1 权重,按 degree 分块 w = model.get_poly_coefficient_layer() # [hidden, input_dim*degree] w = w.view(w.size(0), model.input_dim, degree) high_order_mask = torch.zeros(degree, device=w.device) high_order_mask[1:] = 1.0 # 惩罚 2 阶及以上 # 对高阶分块的权重能量做惩罚 high_order_energy = (w ** 2).sum(dim=(0, 1)) * high_order_mask poly_reg = high_order_energy.sum() loss = ce_loss + poly_lambda * poly_reg optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() # 训练结束后用一个 batch 估算 ED model.eval() sample_x, sample_y = next(iter(test_loader)) from .fisher import compute_per_sample_gradients G = compute_per_sample_gradients(model, sample_x[:64], sample_y[:64], device=device) centered = G - G.mean(dim=0, keepdim=True) cov = centered.T @ centered / max(G.size(0) - 1, 1) eigvals = torch.linalg.eigvalsh(cov).flip(0) ed = effective_dimensionality(eigvals, coverage=0.95) print(f"epoch={epoch+1}, loss={total_loss:.4f}, ED@95%={ed}")这里最关键的三行代码是:分块权重、构造高阶掩码、计算高阶能量。它把“多项式表示 + 简单性优化”落实到了具体的训练流程里。
5.6 主入口示例
最后补一个最小可运行入口:
# 文件路径:run_demo.py import torch from demo_poly_ed.model import TwoLayerTanhMLP from demo_poly_ed.poly_model import PolyFeatureMLP from demo_poly_ed.data import get_mnist_loaders from demo_poly_ed.train import train_with_poly_regularization device = "cuda" if torch.cuda.is_available() else "cpu" train_loader, test_loader = get_mnist_loaders(batch_size=256) print("===== Baseline MLP =====") baseline = TwoLayerTanhMLP() train_with_poly_regularization( baseline, train_loader, test_loader, epochs=3, poly_lambda=0.0, degree=3, device=device ) print("\n===== Poly Model with High-Order Regularization =====") poly_model = PolyFeatureMLP(degree=3) train_with_poly_regularization( poly_model, train_loader, test_loader, epochs=3, poly_lambda=1e-3, degree=3, device=device )这里分别训练一个普通 MLP 和一个带高阶多项式惩罚的 MLP,比较两者的 ED 变化趋势。
6. 运行结果与效果验证
6.1 预期输出
正常运行时,控制台会输出类似这样的信息:
===== Baseline MLP ===== epoch=1, loss=26.7341, ED@95%=39 epoch=2, loss=17.2823, ED@95%=48 epoch=3, loss=13.1964, ED@95%=54 ===== Poly Model with High-Order Regularization ===== epoch=1, loss=31.0023, ED@95%=25 epoch=2, loss=22.1841, ED@95%=21 epoch=3, loss=18.4372, ED@95%=18注意,具体数值会因为随机种子、数据 loader、模型初始化而变化。重点看两个趋势:
- 普通 MLP 的 ED 随训练轮次上升,说明模型在拟合数据时逐渐使用了更多自由度。
- 带多项式高阶惩罚的模型,ED 往往更低,且高阶惩罚越大,ED 下降越明显。
如果出现“带高阶惩罚的模型精度大幅下降、ED 也很低”的组合,说明惩罚系数太大,模型被过度压制。此时应该降低poly_lambda,或者调整高阶掩码的分界位置。
6.2 三组对照实验
建议做三组对照,避免一次实验就下结论:
| 实验组 | poly_lambda | 预期结果 |
|---|---|---|
| 基线 | 0 | ED 正常上升,测试精度正常 |
| 温和压缩 | 1e-4 | ED 略降,精度几乎不掉 |
| 强压缩 | 1e-2 | ED 明显下降,精度可能下降 1%-3% |
如果温和压缩组能让 ED 明显下降而精度不掉,就说明这个方向值得深入:模型确实有不少冗余自由度,通过多项式高阶惩罚把它们“收敛”到了低阶结构上。
6.3 验证失败时如何排查
如果程序报错或结果不符合预期,按下面顺序检查:
- 先看梯度矩阵形状。
compute_per_sample_gradients返回的张量形状应为[样本数, 参数总维度]。如果样本数小于参数维度,协方差矩阵会变成秩亏,特征谱容易出现大量零值,ED 会偏低。 - 再查
poly_model的权重分块维度。w.view(w.size(0), model.input_dim, degree)要求fc1.weight的最后一维等于input_dim * degree。如果你改了模型结构,这里一定要同步修改。 - 最后看
poly_lambda是否太大。如果 loss 里poly_reg比ce_loss大一个量级,模型根本不会学任务,ED 当然低。
7. 常见问题与排查思路
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| ED 计算结果波动很大 | 每次只跑了一个 batch,梯度矩阵覆盖不足 | 打印样本数和特征谱分布,统计多次 ED 的方差 | 增加max_samples,或使用多个 batch 的逐样本梯度拼接 |
| ED 几乎等于参数总量 | 特征谱衰减非常慢,说明模型所有方向都敏感 | 检查是否在训练末期采样,模型是否已经收敛 | 在训练不同阶段分别计算 ED,查看趋势 |
| 高阶惩罚后精度明显下降 | poly_lambda过大或高阶掩码包含太多项 | 查看poly_reg和ce_loss的量级 | 降低poly_lambda,或只惩罚最高阶项 |
| 梯度矩阵计算时显存溢出 | 逐样本 Jacobian 导致 batch 内多次反向传播 | 尝试更小的 batch 或单样本循环 | 使用torch.func.vmap或近似梯度估计方法 |
eigvalsh返回包含微小负值 | 数值精度导致,浮点协方差不完全对称 | 检查矩阵是否对称,打印最大负特征值 | 用(cov + cov.T) / 2强制对称 |
| 多项式特征维数爆炸 | input_dim * degree过大,线性层参数量激增 | 打印poly_dim | 先对输入做 PCA 降维,再展开多项式特征 |
| 训练 loss 正常但测试精度低 | 模型过于关注低阶结构,表达能力不足 | 查看低阶/高阶能量比例 | 减少惩罚系数,或增加 hidden_dim |
8. 最佳实践与工程建议
8.1 先测后压,不要一开始就加正则
在任何项目里引入“简单性优化”,第一步永远是先量化现状。先跑一次普通训练,计算模型在训练和测试集上的 ED,观察它的变化趋势和稳定性。如果模型本身的 ED 就已经很低,那说明这个任务比较简单,强行加多项式惩罚是多余的。只有当 ED 明显高于任务实际所需时,优化简单性才有收益。
8.2 基于 ED 的压缩比基于参数量的压缩更可靠
传统剪枝通常是按照权重绝对值大小删参数,这一做法的隐含假设是“绝对值小的权重不重要”。但 ED 提醒我们,重要性应该表现为“对输出扰动的影响”。一个权重绝对值很小,但位于特征谱主方向上,删掉它可能造成剧烈影响。反过来,有些权重绝对值不小,却被淹没在零特征值方向上,删掉它几乎不影响输出。因此,在做模型压缩之前,先算一次 Fisher 特征谱,用 ED 确定真正要保留的方向,再用投影方式压缩,往往比暴力剪枝更稳。
8.3 多项式阶数不是越高越好
理论上,阶数越高,多项式表示能力越强,但计算开销和过拟合风险同步上升。对于图片分类这类任务,3 阶通常已经能覆盖绝大多数非线性交互;对于时间序列预测,2 到 3 阶也足够起步。当你发现高阶项能量占比很低时,应果断截断到低阶,而不是继续增加阶数。截断本身就是一种简单性优化。
8.4 用高阶惩罚而不是 L1 全局惩罚
L1 惩罚会不加区分地把所有权重推向零,对网络的所有部分一视同仁。多项式高阶惩罚只惩罚对应高次交互的系数,保留低阶项的表达能力。两者在效果上差别很大:L1 全局正则容易把模型压成一个接近线性的弱分类器,而高阶惩罚允许网络在低阶结构上保持充分的非线性,只是拒绝无意义的高频过拟合。
8.5 结合数据分布的多样性
ED 是一个依赖数据分布的指标。同一个模型,在分布均匀的数据上 ED 可能很高,在分布单一的数据上 ED 可能很低。因此,比较两个模型的简单性时,必须用同一份测试集、同一种采样方式、同一种梯度计算方法,否则结论不可靠。实践中最稳妥的做法是固定一个 benchmark 集,像管理准确率一样去管理模型的 ED 基线。
8.6 安全与合规提醒
如果要把这套方法用于生产环境的模型压缩或线上模型更新,务必遵守最小权限和灰度发布原则。不要在生产模型上直接做大幅结构改动或通过反向传播修改权重。建议先在离线测试集上验证 ED 和精度的联合变化,再经过小流量灰度、A/B 对比后逐步上线。任何涉及批量权重更新的操作,都应该保留旧模型快照,以便随时回滚。
9. 总结与后续学习方向
这篇文章围绕 ICML 2026 投稿预印本中的 ED 思路,拆解了它所在的两条技术脉络。第一条是有效维度:用 Fisher 信息矩阵的特征谱测量模型真正使用的自由度,而不是把参数量当作复杂度。第二条是多项式表示:把网络的输入-输出关系展开到不同阶数的多项式项上,让“简单性”从一笔糊涂账变成一个可归属、可优化的结构指标。把这两条合在一起,才得到“量化并优化神经网络简单性”的完整框架。
从工程角度看,本文提供的 PyTorch 最小实现可以直接用于小型实验,作者计算 ED 的主流程包括compute_per_sample_gradients、effective_dimensionality以及train_with_poly_regularization三个模块。你可以把它们改写成自己的训练工具,在 MNIST、CIFAR、你自己业务的小型模型上先跑通一遍,感受 ED 随训练阶段和正则强度如何变化。需要注意的是,由于正式论文尚未完全公开,本文中的高阶惩罚方式和多项式特征构建只是这一思路的简化实现,严谨复现请以后续公开的论文细节和官方代码为准。
如果你想继续深入,建议按三条线展开。第一,学习 Fisher 信息矩阵的高效估计方法,重点是 K-FAC 和 Lanczos 算法,这决定了 ED 能不能在大模型上落地。第二,研究谱正则化背后的泛化理论,理解为什么压低 Fisher 谱的特征值尾部往往能带来更好的测试表现。第三,尝试把 ED 和现有剪枝工具结合,比如训练后计算 ED,再用低秩分解把高维参数空间投影到 ED 指示的主子空间上。这一套组合下来,你会比单纯看参数量的人在理解模型复杂度这件事上领先一个层次。