ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

【ICML 2024】Medusa:多解码头驱动的 LLM 推理加速框架|从LLM推理系统优化视角

【ICML 2024】Medusa:多解码头驱动的 LLM 推理加速框架|从LLM推理系统优化视角 摘要本文解读 ICML 2024 论文《Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads》。该论文提出Medusa 多解码头推理加速框架通过融合多个单层解码头、树注意力并行验证与典型接受采样在不引入任何独立草稿模型的前提下并行推进多个 token其特别之处在于解码头直接长在主干模型上、训练参数极少且与主干分布天然对齐。实验表明Medusa-1 冻结主干即可无损加速 2.2 倍Medusa-2 联合微调进一步达到 2.3-2.8 倍MT-Bench 生成质量基本持平±0.14 分以内为 LLM 推理系统优化提供了零架构改动的即插即用方案。视频讲解点击观看 B 站视频摘要论文基本信息背景与动机为什么 LLM 推理慢投机解码又为什么难用研究主线从问题到结论基准/方法设计Medusa 的三个核心组件分类全景Medusa 框架的四大组成方法细节两级训练配方与两个实用扩展实验设计与结果四个模型、两级训练、全线 2.3 倍以上加速结果对比总结关键发现局限性常见问题FAQMedusa 与投机解码的本质区别是什么Medusa-1 和 Medusa-2 怎么选树注意力为什么能一次验证多条候选没有训练数据时 Medusa 还能用吗Medusa 的加速上限在哪里参考链接论文基本信息项目内容标题英文Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads标题中文Medusa多解码头驱动的 LLM 推理加速框架作者Tianle Cai*, Yuhong Li*共同一作, Zhengyang Geng, Hongwu Peng, Jason D. Lee, Deming Chen, Tri Dao机构Princeton University · UIUC · CMU · UConn · Together AI会议ICML 2024arXivhttps://arxiv.org/abs/2401.10774项目网站https://github.com/FasterDecoding/Medusa背景与动机为什么 LLM 推理慢投机解码又为什么难用LLM 自回归解码是内存带宽受限的每一步都要把全部模型参数从高带宽内存HBM搬进计算单元却只推进 1 个 token算力远未用满。投机解码Speculative DecodingLeviathan et al., ICML 2023Chen et al., 2023用一个小草稿模型先草拟若干 token、再由大模型并行验证能把步数压缩到 $\approx 1/\gamma$。但草稿模型的获取与维护代价高昂需要专门预训练如 SpecInfer 消耗 275 个 A100 GPU 小时、存在与主模型之间的分布偏移且难以集成进分布式服务系统。论文实测发现投机解码在 Vicuna 系列上最终加速只有 1.47-1.60 倍且最优草稿配置随模型规模变化7B 用 Llama-68M 草稿 $\gamma4$33B 需要 Tiny-Vicuna 1B 草稿 $\gamma3$系统调参负担很重。核心问题由此产生能否不引入任何额外模型仅靠主干自身的预测能力并行推进多个 token研究主线从问题到结论图 9Medusa 论文研究主线——问题、动机、设计、方法、实验到结论Mermaid 流程图基准/方法设计Medusa 的三个核心组件Medusa 沿用投机解码生成-处理-接受三步框架但把三步全部收编进主干模型自身。第一步生成候选由 Medusa heads 完成在最后隐状态 $h_t$ 上附加 $K$ 个单层前馈头第 $k$ 个头预测第 $tk1$ 个 token输出 $p_t^{(k)} \mathrm{softmax}(W_2^{(k)}(\mathrm{SiLU}(W_1^{(k)} h_t) h_t))$输出投影 $W_2$ 用原语言模型头初始化、$W_1$ 置零保证初始预测与主干完全对齐。第二步并行处理由树注意力完成各头的 top 预测组合成候选树通过只允许 token 回看自身分支前驱的注意力掩码一次前向验证全部候选。第三步接受由典型接受完成以熵相关阈值 $\min(\epsilon,\ \delta\exp(-H(p)))$ 判定候选是否合理首 token 贪婪无条件接受取最长合法前缀进入下一轮。图 1Medusa 总览——解码头、树注意力与典型接受构成闭环生成-验证循环分类全景Medusa 框架的四大组成图 10Medusa 框架分类全景——四大组件分工Mermaid 结构图方法细节两级训练配方与两个实用扩展Medusa-1冻结主干只训练解码头损失为 $L_1 -\sum_{k1}^{K}\lambda_k \log p_t^{(k)}(y_{tk1})$权重 $\lambda_k 0.8^k$ 平衡远距离头的不确定性主干可量化QLoRA 风格Vicuna-7B 在单张 A100 上约 5 小时训完 6 万条 ShareGPT 样本实现无损加速。Medusa-2联合微调把损失改为 $L_2 L_{\mathrm{LM}} \lambda_0 L_1$配合三条保护主干的策略差分学习率head 学习率为 backbone 的 4 倍、两阶段 warmup先训 head 再联合训练、LoRA 低秩适配rank 32、$\alpha$ 16。直接微调主干会掉质量MT-Bench 5.925 vs 基线 6.17而 Medusa-2 的配方把质量保住在 6.18。两个扩展覆盖工程落地场景典型接受在温度越高时接受越长温度 0 退化为贪婪解码替代拒绝采样后加速更大自蒸馏在无训练数据如 RLHF 模型时用 ShareGPT/UltraChat 种子提示让模型自答生成约 10 万条样本主干损失换成与原模型分布的 KL 散度利用 LoRA 适配器开关实现近乎零额外显存的蒸馏。稀疏树构造则用校准集估计各头 top-$i$ 准确率贪心选择期望接受长度增益最大的节点64 节点稀疏树即可胜过 256 节点稠密树。图 2树注意力机制——把多候选组织成树结构一次前向并行验证实验设计与结果四个模型、两级训练、全线 2.3 倍以上加速评测使用 MT-Bench 多轮对话基准GPT-4 打分 0-10覆盖 Vicuna-7B/13BShareGPT 公开数据、Vicuna-33B私有数据与 Zephyr-7BSFTRLHF。三个核心指标加速率每步解码 token 数、开销每步延迟比、Speedup 加速率 / 开销。基线为 HuggingFace 默认实现与多种草稿模型的投机解码。主表Medusa-2 在四个模型上的表现Medusa-2 模型加速率开销MT-Bench 质量SpeedupVicuna-7B3.471.226.18 (0.01)2.83×Zephyr-7B3.141.187.25 (-0.07)2.66×Vicuna-13B3.511.236.43 (-0.14)2.83×Vicuna-33B3.011.277.18 (0.05)2.35×图 3Medusa-1 冻结主干即超 2 倍加速Medusa-2 进一步提升Extraction 类最高 3.62×、Coding 类 3.29×技术贡献分解印证每个组件的价值仅解码头约 1.5×加树注意力约 1.9×优化树结构约 2.2×Medusa-2 联合训练约 2.8×。AlpacaEval 上结果保持一致Vicuna-7B 2.88×、13B 3.16×、33B 2.26×、Zephyr-7B 2.91×说明加速不依赖单一基准。图 4自蒸馏训练的 Zephyr-7B、Vicuna-13B/33B 加速略弱但全部超过 2.2 倍Roofline 硬件分析解释了加速机理解码阶段所有算子都沿 HBM 带宽线运行带宽受限Medusa 增加候选 token 数后 FLOP/s 与操作强度同步上升batch16、seq1024、64 候选时达 44× FLOP/s 与 41× 操作强度但 batch 超过 32 后线性层转向计算受限加速回落——这正是候选数存在最优甜点区的原因。图 5Medusa 把带宽受限的算子推向更高计算强度区域稀疏树与典型接受消融64 节点稀疏树优于 256 节点稠密树加速率更高、速度衰减更小典型接受阈值 $\epsilon$ 从 0.01 增到 0.25 时质量上升、加速率下降与随机采样质量相当但加速更高。图 6稀疏树左倾结构反映算法偏好高概率节点图 7典型采样在相当质量下取得比随机采样更高的加速与投机解码的系统对比Medusa-2 加速 2.35-2.83×全面优于投机解码的 1.47-1.60×且省去了为每个模型规模挑选与预训练草稿模型的全套调参负担。图 8草稿模型最优配置随模型规模变化Medusa 无需任何配置即全面占优结果对比总结图 11结果对比总结——Medusa 全面优于投机解码Mermaid 流程图关键发现无损加速成立Medusa-1 冻结主干在 Vicuna-7B 上获 2.18× 加速MT-Bench 质量 6.23 高于基线 6.17。两级训练分工明确Medusa-1 适合资源有限/不可动主干场景Medusa-2 通过组合损失、差分学习率与 warmup 把加速推到 2.83×同时保住质量6.18 vs 直接微调的 5.925。树注意力是关键增量加树注意力把加速从约 1.5× 提到约 1.9×稀疏树优化再到约 2.2×。自蒸馏可行无训练数据时用约 10 万条自生成样本Vicuna-33B 加速 2.35× 且质量 0.05、Zephyr-7B 加速 2.66× 质量 -0.07。典型接受优于拒绝采样温度越高接受越长质量与随机采样相当而加速更高。生态落地已验证论文发表后 TensorRT-LLM 与 HuggingFace TGI 等推理库原生支持 Medusabatch 化推广到服务端吞吐场景。局限性候选数收益递减加速率随候选数近似为 $acc 0.477\log(#candidates)$超过 64 个候选后 speedup 开始回落。大 batch 失效风险batch 超过 32 时线性层从带宽受限转为计算受限speedup 下降甚至为负——论文实验以 batch1 为主。长序列注意力开销序列变长时注意力矩阵乘开销增大整体性能下降需要进一步优化注意力机制。自蒸馏质量折中Vicuna-33B 加速率偏低3.01存在隐藏训练集与自蒸馏分布错配的问题。常见问题FAQMedusa 与投机解码的本质区别是什么投机解码需要额外预训练一个小草稿模型来草拟 tokenMedusa 直接把多个解码头长在主干最后隐状态上解码头与主干分布天然对齐无需草稿模型、无分布偏移且对分布式服务系统零架构改动。Medusa-1 和 Medusa-2 怎么选Medusa-1 冻结主干、只训解码头单张 A100 约 5 小时即可完成 7B 模型训练实现 2.2 倍无损加速Medusa-2 联合微调主干加速可达 2.8 倍但需要特殊训练配方组合损失、差分学习率、warmup保护主干能力。树注意力为什么能一次验证多条候选传统因果注意力只允许 token 看全部历史树注意力把候选组织成树掩码只允许每个 token 回看自身分支的前驱位置编码按树结构调整于是整棵候选树可以在一次前向中并行验证候选间的公共前缀计算被复用。没有训练数据时 Medusa 还能用吗可以。自蒸馏扩展用 ShareGPT/UltraChat 种子提示让模型自答生成约 10 万条样本主干损失改为与原模型预测分布的 KL 散度配合 LoRA 适配器开关实现几乎零额外显存Vicuna-33B 与 Zephyr-7B 均验证可行。Medusa 的加速上限在哪里加速率随候选数对数增长受两个硬件拐点约束候选过多时线性层从带宽受限转为计算受限batch 相关序列过长时注意力矩阵乘成为新瓶颈论文通过 Roofline 分析给出最优候选数甜点区。参考链接arXiv 论文页https://arxiv.org/abs/2401.10774官方开源仓库https://github.com/FasterDecoding/MedusaEAGLE特征层草稿后续扩展https://arxiv.org/abs/2401.15077Speculative DecodingLeviathan et al.https://arxiv.org/abs/2211.17192Accelerating LLM Decoding with Speculative SamplingChen et al.https://arxiv.org/abs/2302.01318给大家推荐一款自用写文献综述、无虚构文献的 AI复旦大学 FudanNLP 团队自研 切问学术官网qiewenpaper.com覆盖3.6 亿篇可溯源真实中英文文献能自动整合文献观点生成规范综述还能挖掘研究创新点、复现实验配合视频教学新手快速上手文献综述写作后记博客的关键词集中在编程、算法、机器人、人工智能、数学等等持续高质量输出中。讨论QQ群白拾的小屋 (750365700)⭐B站账号白拾的物理AI组会活跃于知识区和动画区✨GitHub主页YhbCode000工程文件
返回列表