Moonwalk:逆 - 前向微分
机器学习
2026-05-25 v4 人工智能
摘要
反向传播的主要局限性在于其需要在正向传播期间存储中间激活值(残差),这限制了可训练网络的深度。这引发了一个根本性问题:我们能否避免存储这些激活值?我们通过重新审视梯度计算的结构来解决这一问题。反向传播通过一系列向量 - 雅可比乘积来计算梯度,该操作通常是不可逆的。丢失的信息位于每一层雅可比矩阵的余核中。我们定义了浸没网络(submersive networks)——即其层雅可比矩阵具有平凡余核的网络——在此类网络中,梯度可以在不存储激活值的情况下通过正向扫描精确重构。对于非浸没层,我们引入了片段梯度检查点法(fragmental gradient checkpointing),该方法仅记录恢复被雅可比矩阵擦除的余切向量所需的最小残差子集。我们方法的核心是一个新颖的算子:向量逆雅可比乘积(vector-inverse-Jacobian product, vijp),它在余核之外反转梯度流。我们的混合模式算法首先通过内存高效的反向传递计算输入梯度,然后利用 vijp 在正向扫描中重构参数梯度,从而消除了存储激活值的需求。我们在 Moonwalk 中实现了该方法,并表明其在匹配反向传播运行时间的同时,能在相同的内存预算下训练深度超过两倍的网络。
引用
@article{arxiv.2402.14212,
title = {Moonwalk: Inverse-Forward Differentiation},
author = {Dmitrii Krylov and Armin Karamzade and Roy Fox},
journal= {arXiv preprint arXiv:2402.14212},
year = {2026}
}