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