中文

可证明地学习多头注意力层

机器学习 2024-02-07 v1 数据结构与算法 机器学习

摘要

多头注意力层是 Transformer 架构的关键组件之一,使其有别于传统的前馈模型。给定序列长度 kk、注意力矩阵 Θ1,,ΘmRd×d\mathbf{\Theta}_1,\ldots,\mathbf{\Theta}_m\in\mathbb{R}^{d\times d} 和投影矩阵 W1,,WmRd×d\mathbf{W}_1,\ldots,\mathbf{W}_m\in\mathbb{R}^{d\times d},相应的多头注意力层 F:Rk×dRk×dF: \mathbb{R}^{k\times d}\to \mathbb{R}^{k\times d} 通过 F(X)i=1msoftmax(XΘiX)XWiF(\mathbf{X}) \triangleq \sum^m_{i=1} \mathrm{softmax}(\mathbf{X}\mathbf{\Theta}_i\mathbf{X}^\top)\mathbf{X}\mathbf{W}_idd 维 token 的长度为 kk 的序列 XRk×d\mathbf{X}\in\mathbb{R}^{k\times d} 进行变换。在本工作中,我们开创性地从随机样本中可证明地学习多头注意力层,并给出了该问题的首个非平凡上下界:\n- 在 {Wi,Θi}\{\mathbf{W}_i, \mathbf{\Theta}_i\} 满足某些非退化条件的前提下,我们提出了一种时间复杂度为 (dk)O(m3)(dk)^{O(m^3)} 的算法,该算法在给定从 {±1}k×d\{\pm 1\}^{k\times d} 均匀抽取的随机标注样本时,能以极小误差学习 FF。\n- 我们证明了计算下界,表明在最坏情况下,对 mm 的指数依赖是不可避免的。我们聚焦于布尔型 X\mathbf{X} 以模拟大型语言模型中 token 的离散性质,尽管我们的技术自然地扩展到标准的连续设置(如高斯分布)。我们的算法以利用样本雕刻出包含未知参数的凸体为核心,这与现有的可证明的前馈网络学习算法有显著不同,后者主要利用高斯分布的代数和旋转不变性。相比之下,我们的分析更加灵活,因为它主要依赖于输入分布及其“切片”的各种上下尾界。

关键词

引用

@article{arxiv.2402.04084,
  title  = {Provably learning a multi-head attention layer},
  author = {Sitan Chen and Yuanzhi Li},
  journal= {arXiv preprint arXiv:2402.04084},
  year   = {2024}
}

备注

105 pages, comments welcome