中文

树形交叉注意力

机器学习 2024-03-04 v2

摘要

交叉注意力(Cross Attention)是一种从上下文 token 集合中检索信息以进行预测的流行方法。在推理时,对于每个预测,交叉注意力会扫描全部 O(N)\mathcal{O}(N) 个 token。然而在实践中,通常仅需一小部分 token 即可获得良好性能。诸如 Perceiver IO 等方法在推理时开销较低,因为它们将信息提炼为更小的潜 token 集合 L<NL < N,再对其应用交叉注意力,从而仅产生 O(L)\mathcal{O}(L) 复杂度。但在实践中,随着输入 token 数量与待提炼信息量的增加,所需潜 token 数量也显著增长。本工作中,我们提出树形交叉注意力(TCA)——一种基于交叉注意力、在推理时仅从对数级 O(log(N))\mathcal{O}(\log(N)) 个 token 中检索信息的模块。TCA 将数据组织为树结构,并在推理时执行树搜索以检索相关 token 用于预测。利用 TCA,我们引入了 ReTreever,一种用于 token 高效推理的灵活架构。我们通过实验表明,树形交叉注意力(TCA)在各种分类与不确定性回归任务中与交叉注意力性能相当,同时显著更省 token。此外,我们将 ReTreever 与 Perceiver IO 比较,在使用相同推理 token 数时显示出显著增益。

关键词

引用

@article{arxiv.2309.17388,
  title  = {Tree Cross Attention},
  author = {Leo Feng and Frederick Tung and Hossein Hajimirsadeghi and Yoshua Bengio and Mohamed Osama Ahmed},
  journal= {arXiv preprint arXiv:2309.17388},
  year   = {2024}
}

备注

Accepted by ICLR 2024