
如果说参数量决定了模型复杂度那 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) ≈ ∑_{k1}^{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 均可无特殊依赖Python3.9类型注解与 torch 兼容性更稳PyTorch2.x 或 1.13需要 autograd 和基本矩阵运算torchvision与 PyTorch 版本匹配仅用于 MNIST 数据CPU/GPU均可小规模实验 CPU 即可跑通推荐使用虚拟环境隔离依赖conda create -n ed_poly python3.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_dim28 * 28, hidden_dim64, num_classes10): 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_size256): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_ds datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_ds datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) return DataLoader(train_ds, batch_sizebatch_size, shuffleTrue), DataLoader(test_ds, batch_sizebatch_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_samples256, devicecpu): 遍历数据收集每个样本对全部参数的梯度拼成 [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_graphFalse, retain_graphTrue, allow_unusedTrue ) # 把每个参数的梯度展平并拼接 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_samples256, devicecpu): 返回特征值列表降序排列。 G compute_gradient_matrix(model, dataloader, max_samples, device) # [N, d] centered G - G.mean(dim0, keepdimTrue) 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, devicecpu): 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:i1] y_i batch_y[i:i1] logits model(x_i) loss F.cross_entropy(logits, y_i) params list(model.parameters()) grads torch.autograd.grad(loss, params, retain_graphFalse, allow_unusedTrue) 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, coverage0.95): 计算有效维度累计特征能量达到总能量 coverage 所需的最少方向数。 参数: eigenvalues: 从大到小排列的特征值张量 coverage: 覆盖比例例如 0.95 表示保留 95% 的谱能量 返回: ed: 有效维度整数值 total_energy eigenvalues.sum() if total_energy 0: return 0 cum_energy torch.cumsum(eigenvalues, dim0) threshold coverage * total_energy indexes torch.nonzero(cum_energy threshold, as_tupleFalse) 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_dim28 * 28, degree3, hidden_dim32, num_classes10): 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, dim1) 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, epochs5, lr1e-3, poly_lambda1e-4, degree3, devicecpu ): optimizer torch.optim.Adam(model.parameters(), lrlr) 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, devicew.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], devicedevice) centered G - G.mean(dim0, keepdimTrue) cov centered.T centered / max(G.size(0) - 1, 1) eigvals torch.linalg.eigvalsh(cov).flip(0) ed effective_dimensionality(eigvals, coverage0.95) print(fepoch{epoch1}, loss{total_loss:.4f}, ED95%{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_size256) print( Baseline MLP ) baseline TwoLayerTanhMLP() train_with_poly_regularization( baseline, train_loader, test_loader, epochs3, poly_lambda0.0, degree3, devicedevice ) print(\n Poly Model with High-Order Regularization ) poly_model PolyFeatureMLP(degree3) train_with_poly_regularization( poly_model, train_loader, test_loader, epochs3, poly_lambda1e-3, degree3, devicedevice )这里分别训练一个普通 MLP 和一个带高阶多项式惩罚的 MLP比较两者的 ED 变化趋势。6. 运行结果与效果验证6.1 预期输出正常运行时控制台会输出类似这样的信息 Baseline MLP epoch1, loss26.7341, ED95%39 epoch2, loss17.2823, ED95%48 epoch3, loss13.1964, ED95%54 Poly Model with High-Order Regularization epoch1, loss31.0023, ED95%25 epoch2, loss22.1841, ED95%21 epoch3, loss18.4372, ED95%18注意具体数值会因为随机种子、数据 loader、模型初始化而变化。重点看两个趋势普通 MLP 的 ED 随训练轮次上升说明模型在拟合数据时逐渐使用了更多自由度。带多项式高阶惩罚的模型ED 往往更低且高阶惩罚越大ED 下降越明显。如果出现“带高阶惩罚的模型精度大幅下降、ED 也很低”的组合说明惩罚系数太大模型被过度压制。此时应该降低poly_lambda或者调整高阶掩码的分界位置。6.2 三组对照实验建议做三组对照避免一次实验就下结论实验组poly_lambda预期结果基线0ED 正常上升测试精度正常温和压缩1e-4ED 略降精度几乎不掉强压缩1e-2ED 明显下降精度可能下降 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_dim8. 最佳实践与工程建议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 指示的主子空间上。这一套组合下来你会比单纯看参数量的人在理解模型复杂度这件事上领先一个层次。