中文

超越 NTK:使用原始梯度下降对多项式宽度、样本数与时间的神经网络进行平均场分析

机器学习 2023-10-10 v2

摘要

尽管近期关于两层神经网络非凸优化的理论取得了进展,但未经非自然修改的神经网络上的梯度下降能否比核方法实现更好的样本复杂度,仍是一个开放问题。本文对多项式宽度的两层神经网络上的投影梯度流给出了清晰的平均场分析。与先前工作不同,我们的分析不需要对优化算法进行非自然修改。我们证明,当样本量 n=O(d3.1)n = O(d^{3.1})(其中 dd 为输入维度)时,用投影梯度流训练的网络在 poly(d)\text{poly}(d) 时间内收敛到核方法使用 nd4n \ll d^4 样本无法达到的非平凡误差,从而展示了未修改梯度下降与 NTK 之间的清晰分离。作为推论,我们表明具有正学习率和多项式次迭代的投影梯度下降以相同样本复杂度收敛到低误差。

关键词

引用

@article{arxiv.2306.16361,
  title  = {Beyond NTK with Vanilla Gradient Descent: A Mean-Field Analysis of Neural Networks with Polynomial Width, Samples, and Time},
  author = {Arvind Mahankali and Jeff Z. Haochen and Kefan Dong and Margalit Glasgow and Tengyu Ma},
  journal= {arXiv preprint arXiv:2306.16361},
  year   = {2023}
}

备注

Added result on projected gradient descent with inverse-polynomial learning rate