PolySketchFormer:通过多项式核草图化实现快速 Transformer
机器学习
2024-03-19 v3
摘要
自注意力机制相对于序列长度固有的二次时间与内存复杂度,在大规模基于 Transformer 的语言模型训练与部署中构成了关键计算瓶颈。近期理论结果表明,在合理复杂度假设下,亚二次 softmax 注意力近似是不可行的。本文通过首先证明高次多项式注意力可在不损失模型质量的前提下有效替代 softmax 来应对该挑战。接着,我们基于数值线性代数开发多项式草图化技术,以近似保证实现线性时间多项式注意力。关键的是,我们的方法无需对注意力矩阵进行稀疏化即可获得此加速。我们还提出一种基于分块的算法来高效施加因果掩码。结合这些技术,我们给出了 \emph{PolySketchFormer},一种用于语言建模的实用线性时间 Transformer 架构,并提供可证明保证。我们通过训练能够处理长上下文的语言模型对 PolySketchFormer 进行实证验证。这些实验在 Google Cloud TPU 上使用合成与真实世界数据集(PG19、Wikipedia 与 C4)进行。对于 32k 上下文长度与 GPT-2 风格模型,我们的模型相比 FlashAttention 在训练中实现了 2.5-4 倍加速,且在所有实验中未观测到质量下降。
引用
@article{arxiv.2310.01655,
title = {PolySketchFormer: Fast Transformers via Sketching Polynomial Kernels},
author = {Praneeth Kacham and Vahab Mirrokni and Peilin Zhong},
journal= {arXiv preprint arXiv:2310.01655},
year = {2024}
}
备注
Added results of more experiments. Added a link to our JAX implementation of models