中文

HashAttention:用于加速推理的语义稀疏化

机器学习 2025-06-05 v2 人工智能

摘要

利用长上下文是高级AI系统的关键,但注意力计算在可扩展性方面存在挑战。虽然缩放点积注意力(SDPA) exhibits token稀疏性,即仅少数关键令牌显著贡献于输出,但利用这种稀疏性仍具有挑战性。现有方法要么导致质量下降,要么需要额外的大量资源。我们表明,识别关键令牌是一个最大内积搜索(MIPS)问题。然而,现有的MIPS解决方案不适用于SDPA,因为它们不GPU友好,且由于查询和键分布分离,往往表现不佳。本文引入HashAttention,将关键令牌识别建模为推荐问题。给定查询,HashAttention在汉明空间中对键和查询进行编码,通过学习的映射函数捕获所需的语义相似性。HashAttention使用按位操作高效识别给定查询的关键令牌,并仅使用这些令牌计算注意力,从而提高整体注意力效率。在通用数据上训练的HashAttention可将使用的令牌数最多降低16×16\times,且仅需每个令牌32位的辅助内存。通过任务特定微调,可进一步提升稀疏性至32×32\times。在A100 GPU上,在32×32\times稀疏性下,集成HashAttention可在GPT-FAST中将注意力延迟降低最高可达4.3×4.3\times,在FlashDecode中降低2.54×2.54\times,并为GPT-FAST实现最高可达3.12×3.12\times的吞吐量提升。

关键词

引用

@article{arxiv.2412.14468,
  title  = {HashAttention: Semantic Sparsity for Faster Inference},
  author = {Aditya Desai and Shuo Yang and Alejandro Cuadron and Matei Zaharia and Joseph E. Gonzalez and Ion Stoica},
  journal= {arXiv preprint arXiv:2412.14468},
  year   = {2025}
}

备注

Accepted at ICML'2025