通过拉格朗日乘数每层独立约束,PC-ALM在保持逐层本地更新同时在线性网络中恢复了精确反向传播梯度,MNIST上训练1000层残差MLP,MIT许可JAX代码已公开。
反向传播是一种全局算法:前向传播、反向传播、权重更新,三步顺序锁定、环环相扣。大脑目前没有发现任何能实现这种网络级相位锁定的机制,这也是局部学习替代方案(如预测编码 PC)持续吸引研究关注的原因。Sakana AI 研究人员提出了增广拉格朗日预测编码(PC-ALM),这是 PC 的一种变体,它让每一次权重更新都停留在层本地,同时恢复与反向传播对齐的信用信号。研究团队报告称,在 MNIST 数据集上,PC-ALM 能够训练残差 MLP 达到 1000 层,与反向传播的差距仅在 2 个百分点以内。
是否可以部署?答案是肯定的,但目前仅作为研究代码:MIT 许可的 JAX 参考实现可以在 CPU 上运行,并复现了论文中的宽度-深度网格搜索。它是一种训练方法,而非模型,且只在小型图像基准上进行了测试。
PC 将每个隐藏激活都视为优化变量,并对每层激活与其从下一层接收的预测之间的平方误差进行惩罚。推理是对该能量的梯度下降;学习则类似赫布规则的权重步。关键问题在于:监督信号从输出端输入,必须通过网络各层逐步扩散,在深层窄网络中,信用信号在到达输入之前就已经衰减殆尽。Innocenti 等人将这种 PC-BP 差距描述为宽度和深度的函数,当宽度小于深度时情况最为严重。
PC-ALM 从约束优化的视角看待训练:在每一层满足 hi=σ(Wihi−1) 的条件下最小化监督损失。PC 是该问题的二次惩罚松弛方法。PC-ALM 则使用增广拉格朗日方法,在保持 PC 惩罚项的同时,为每层约束附加一个拉格朗日乘子 λi∈ℝdi 且 dim(λi)=dim(hi)。将 λ 设为 0 即可精确恢复 PC。
推理交替执行 2 个局部步:激活的原版梯度步,以及拉格朗日乘子的对偶步 λi←λi+αρi,用于累积该层的预测误差。配方平方后可以看出,每个原版步实际上是标准 PC 步,只是预测目标被平移了 −λi/ρ。经过 T 步后,权重更新作用于复合信号 λi+ρri。研究团队将其解读为每层的 PI 控制器:预测误差是比例项,乘子是积分项。α=0 时得到 PC;α=ρ 且内问题精确求解时得到经典增广拉格朗日方法。
LeCun 在 1988 年观察到,约束网络的拉格朗日乘子在 KKT 点处等于反向传播伴随量。团队证明了在线性 PC 网络中,在谱半径稳定性条件下,PC-ALM 收敛到该 KKT 点:激活值恢复到前向传播时的值,同时每个 λi 积分至精确的 BP 伴随量。每模式的稳定性边界为 ηhσi2(2ρ+α)<4,在 α=0 时退化为 PC 的条件。与 PC 的单调梯度流不同,PC-ALM 的迭代矩阵具有复特征值,能产生阻尼振荡;α 设置振荡频率但不设置衰减率。
研究团队在 Fashion-MNIST 和 MNIST 上对残差 MLP 的宽度和深度(从 8 到 128)进行了扫参实验,采用 Innocenti 等人的平均场参数化,训练 1 个 epoch。在推理预算 T=2L 下,PC-ALM 在每种宽度、深度和激活函数(identity、tanh、ReLU)上都与反向传播相匹配,而 PC 在深层窄单元中急剧下降。仓库的参考单元(宽度 32、深度 32、ReLU、Fashion-MNIST)报告:BP 测试准确率 78.66%、PC 68.13%、PC-ALM 77.75%,与 BP 的梯度余弦相似度从 0.604 升至 0.909。
研究扩展了图景:MNIST 上 1000 层残差 MLP(宽度 32、ReLU、5 个 epoch)与 BP 的差距保持在约 2 个百分点内;PC-ALM 在所有尝试的基准上均优于 PC,包括 CIFAR-10 和 Tiny ImageNet 上的 ResNet-18。
PC-ALM 在预测编码的每层添加一个拉格朗日乘子;每次更新都保持在层本地。
在线性网络中,乘子收敛到精确反向传播梯度。
在 T=2L 下匹配所有宽度-深度网格(8 到 128)的 BP;PC 在深层窄单元中失效。
在 MNIST 上训练 1000 层残差 MLP,与 BP 差距约 2 个百分点。
MIT 许可的 JAX 代码在 CPU 上可复现结果。