揭示SGD中大损失尖峰的根源,解释为何模型偏好平坦解。
Large Spikes in Stochastic Gradient Descent: A Large-Deviations View
- 用大偏差理论分析浅层网络,发现尖峰由对数漂移决定。
- 尖峰至少以多项式概率出现,是逃离尖锐极小值的主要机制。
- 结果适用于ReLU网络,对课程学习设计有启发意义。
通过严格的大型偏差分析,研究了在NTK标度下浅层全连接网络中随机梯度下降(SGD)的大损失尖峰现象。与全批量梯度下降不同,弹射阶段被证明可分裂为由显式对数漂移准则决定的膨胀和收缩两种状态。在这两种情况下,大尖峰均至少以多项式概率出现。此外,这些尖峰被证明是逃离尖锐极小值并降低曲率的主要机制,从而倾向于更平坦的解。相关结论还扩展至某些ReLU网络,并推导出对课程学习的启示。
原文摘要 · Abstract (English)
Large loss spikes in stochastic gradient descent are studied through a rigorous large-deviations analysis for a shallow, fully connected network in the NTK scaling. In contrast to full-batch gradient descent, the catapult phase is shown to split into inflationary and deflationary regimes, determined by an explicit log-drift criterion. In both cases, large spikes are shown to be at least polynomially likely. In addition, these spikes are shown to be the dominant mechanism by which sharp minima are escaped and curvature is reduced, thereby favouring flatter solutions. Corresponding results are also obtained for certain ReLU networks, and implications for curriculum learning are derived.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。