Neural Networks can Learn Representations with Gradient Descent

Neural Networks can Learn Representations with Gradient Descent
复制标题

DOI:
10.48550/arxiv.2206.15144
复制
发表时间:
2022-06
期刊:
ArXiv
影响因子:
--
通讯作者:
Alexandru Damian;Jason D. Lee;M. Soltanolkotabi
Alexandru Damian;Jason D. Lee;M. Soltanolkotabi
中科院分区:
其他
文献类型:
--
作者:
Alexandru Damian;Jason D. Lee;M. Soltanolkotabi

文献摘要

相似文献

重要的理论工作已经确定,在特定的制度下,通过梯度下降训练的神经网络表现得像核方法。然而,在实践中,众所周知,神经网络的性能远远超过其相关的内核。在这项工作中,我们解释了这一差距,证明有一大类功能,不能有效地学习内核方法,但可以很容易地学习与梯度下降的内核制度以外的两层神经网络通过学习表示,是相关的目标任务。我们还证明了这些表示允许有效的迁移学习,这在内核机制中是不可能的。具体地说,我们考虑学习多项式的问题,它只依赖于几个相关的方向,即形式$f^\星星(x)= g(Ux)$其中$U:\R^d \to \R^r$与$d \gg r$。当$f^\星星$的次数为$p$时,已知在内核机制中学习$f^\星星$需要$n \asymp d^p$样本。我们的主要结果是,梯度下降学习的数据表示只依赖于相关的方向$f^\星星$。这导致改进的样本复杂度为$n\asymp d^2 r + dr^p$。此外,在迁移学习设置中,源域和目标域中的数据分布共享相同的表示$U$,但具有不同的多项式头,我们证明了迁移学习的流行启发式算法具有独立于$d$的目标样本复杂度。
Significant theoretical work has established that in specific regimes, neural networks trained by gradient descent behave like kernel methods. However, in practice, it is known that neural networks strongly outperform their associated kernels. In this work, we explain this gap by demonstrating that there is a large class of functions which cannot be efficiently learned by kernel methods but can be easily learned with gradient descent on a two layer neural network outside the kernel regime by learning representations that are relevant to the target task. We also demonstrate that these representations allow for efficient transfer learning, which is impossible in the kernel regime. Specifically, we consider the problem of learning polynomials which depend on only a few relevant directions, i.e. of the form $f^\star(x) = g(Ux)$ where $U: \R^d \to \R^r$ with $d \gg r$. When the degree of $f^\star$ is $p$, it is known that $n \asymp d^p$ samples are necessary to learn $f^\star$ in the kernel regime. Our primary result is that gradient descent learns a representation of the data which depends only on the directions relevant to $f^\star$. This results in an improved sample complexity of $n\asymp d^2 r + dr^p$. Furthermore, in a transfer learning setup where the data distributions in the source and target domain share the same representation $U$ but have different polynomial heads we show that a popular heuristic for transfer learning has a target sample complexity independent of $d$.