树形交叉注意力
机器学习
2024-03-04 v2
摘要
交叉注意力(Cross Attention)是一种从上下文 token 集合中检索信息以进行预测的流行方法。在推理时,对于每个预测,交叉注意力会扫描全部 个 token。然而在实践中,通常仅需一小部分 token 即可获得良好性能。诸如 Perceiver IO 等方法在推理时开销较低,因为它们将信息提炼为更小的潜 token 集合 ,再对其应用交叉注意力,从而仅产生 复杂度。但在实践中,随着输入 token 数量与待提炼信息量的增加,所需潜 token 数量也显著增长。本工作中,我们提出树形交叉注意力(TCA)——一种基于交叉注意力、在推理时仅从对数级 个 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