神经 ODE 中用于梯度估计的自适应检查点伴随方法
机器学习
2020-12-07 v1 机器学习
摘要
神经常微分方程(NODEs)近来受到越来越多的关注;然而,它们在基准任务(如图像分类)上的经验性能明显逊于离散层模型。我们证明导致其较差性能的一个原因是现有梯度估计方法的不准确性:伴随方法在反向模式积分中存在数值误差;朴素方法直接通过 ODE 求解器反向传播,但在搜索最优步长时遭受冗余过深的计算图。我们提出自适应检查点伴随(ACA)方法:在自动微分中,ACA 采用轨迹检查点策略,将前向模式轨迹记录为反向模式轨迹以保证精度;ACA 删除冗余组件以获得浅层计算图;且 ACA 支持自适应求解器。在图像分类任务上,相较于伴随方法和朴素方法,ACA 以一半的训练时间取得了减半的错误率;用 ACA 训练的 NODE 在准确率和重测信度上均优于 ResNet。在时间序列建模上,ACA 优于对比方法。最后,在一个三体问题示例中,我们展示带 ACA 的 NODE 可融入物理知识以获得更好精度。我们提供 ACA 的 PyTorch 实现:\url{https://github.com/juntang-zhuang/torch-ACA}。
引用
@article{arxiv.2006.02493,
title = {Adaptive Checkpoint Adjoint Method for Gradient Estimation in Neural ODE},
author = {Juntang Zhuang and Nicha Dvornek and Xiaoxiao Li and Sekhar Tatikonda and Xenophon Papademetris and James Duncan},
journal= {arXiv preprint arXiv:2006.02493},
year = {2020}
}