中文

Medusa: 具有多解码头的简单 LLM 推理加速框架

机器学习 2024-06-18 v3 计算与语言

摘要

大型语言模型(LLM)采用自回归解码,需要顺序计算,每一步都依赖于前一步的输出。这造成了瓶颈,因为每一步都需要将完整的模型参数从高带宽存储器(HBM)移动到加速器缓存中。虽然已经提出了推测解码等方法来解决此问题,但由于获取和维护单独的草稿模型相关的挑战,其实现受到阻碍。在本文中,我们提出了 Medusa,这是一种高效的方法,通过添加额外的解码头来并行预测多个后续 token,从而增强 LLM 推理。使用基于树的注意力机制,Medusa 在每个解码步骤中构建多个候选续写并同时对其进行验证。通过利用并行处理,Medusa 大幅减少了所需的解码步骤数。我们提出了两个级别的 Medusa 微调程序以满足不同用例的需求:Medusa-1:在冻结的骨干 LLM 之上直接微调 Medusa,实现无损推理加速。Medusa-2:将 Medusa 与骨干 LLM 一起微调,实现更高的 Medusa 头预测精度和更大的加速,但需要保留骨干模型能力的特殊训练方案。此外,我们提出了几个改进或扩展 Medusa 效用的扩展,包括用于处理无训练数据情况的自蒸馏,以及在保持生成质量的同时提高接受率的典型接受方案。我们在各种规模和训练程序的模型上评估了 Medusa。我们的实验表明,Medusa-1 可以在不影响生成质量的情况下实现 2.2 倍以上的加速,而 Medusa-2 将加速进一步提高到 2.3-3.6 倍。

关键词

引用

@article{arxiv.2401.10774,
  title  = {Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads},
  author = {Tianle Cai and Yuhong Li and Zhengyang Geng and Hongwu Peng and Jason D. Lee and Deming Chen and Tri Dao},
  journal= {arXiv preprint arXiv:2401.10774},
  year   = {2024}
}

备注

The code for this implementation is available at https://github.com/FasterDecoding/Medusa