@reprompting: 今天读到关于切片级激活重叠的文章 https://arxiv.org/pdf/2607.02521
摘要
本文介绍了基于 CUTLASS 的内核,将 SwiGLU 激活与 GeMM 在切片级进行融合,在 NVIDIA H100 上实现了高达 2.47 倍的加速,以用于高效 LLM 推理。
查看缓存全文
缓存时间: 2026/08/16 16:01
阅读关于层级激活重叠的资料
https://t.co/TQmsC1xJgI https://t.co/KBaJuO9gfU
用于高效大语言模型推理的层级激活重叠
来源:https://arxiv.org/html/2607.02521
摘要
SwiGLU是现代大型语言模型中主导的MLP激活函数,但其张量中间表示消耗了MLP执行时间的9–37%。我们提出两种基于CUTLASS的SM90互补内核,在层级将SwiGLU融合到GeMM中。内核1通过乒乓(Pingpong)调度策略在Gate累加器上进行Swish计算,同时加载Up层级数据;内核2通过自定义Epilogue访客树将SwiGLU与层级存储操作交织执行。在NVIDIA H100上对Qwen-2.5模型(0.5B–72B)的评估显示,我们的内核相比PyTorch最高可实现2.47倍加速,将负载从内存受限型转变为计算受限型,达到79.5%的BF16峰值利用率。我们证明torch.compile无法复现这种融合(比我们的内核慢3–7倍),验证了手工层级设计的必要性。我们的融合内核在数值精度上也表现更优,与cuBLAS相比实现了零错误匹配。
1引言
SwiGLU(https://arxiv.org/html/2607.02521#bib.bib1)已成为现代大型语言模型的主流激活函数。Qwen-2.5(https://arxiv.org/html/2607.02521#bib.bib10)、LLaMA(https://arxiv.org/html/2607.02521#bib.bib9)、Mistral和Gemma均采用门控MLP结构:\text{Gate}=A\times W_1,\text{Up}=A\times W_2,Y=\text{SiLU}(\text{Gate})\odot\text{Up}。该模式需要执行两次独立的矩阵乘法,然后进行逐元素门控激活,并在GeMM和激活阶段之间将两个完整的中间张量(Gate和Up)物化到高带宽存储器中。
随着通过量化(FP8、INT4)和架构改进提升张量核计算密度,内存受限操作的相对成本不断增长。我们在NVIDIA H100上对SwiGLU MLP进行性能分析,发现激活计算及其相关的中间张量物化消耗了MLP总执行时间的9–37%(取决于模型大小)。对于边缘部署模型(Qwen-2.5 0.5B),SwiGLU占据了MLP时间的30%以上——这是一个显著开销,且随着GeMM算术相对于内存传输成本变得更低,情况只会恶化。
现有编译器基础设施无法解决此瓶颈。PyTorch的torch.compile在最大优化下无法融合两个具有不同权重矩阵的独立GeMM——这是图级融合传递的根本限制。我们的实验表明,torch.compile针对此模式仅达到急切PyTorch性能的34–94%,且显式融合提示未带来显著提升(变化远低于4%)。这验证了针对特定硬件手工设计内核的必要性。
受FlashAttention(https://arxiv.org/html/2607.02521#bib.bib2)成功消除注意力机制中间表示的启发,我们将IO感知内核设计应用于MLP模块。但MLP融合挑战本质不同:它涉及两个独立的GeMM与需要协调的不同权重矩阵,而非单次注意力计算。
我们提出两种基于CUTLASS的SM90互补内核,在层级将SwiGLU融合到GeMM中:
- 首次在层级使用调度策略实现细粒度的GeMM-SwiGLU融合——内核1在乒乓调度的消费者阶段,将Swish计算与Up MMA重叠执行,创建优化大批次的[M,N]线程块网格。
- 互补的双内核方法——内核2使用自定义Epilogue访客树(PairMulStore)将SwiGLU与层级存储交织执行,创建[M,2N]线程块网格,为小批次提供2倍更佳的占用率。
- 系统性实验评估(4种模型大小×5种批次规模)显示相比PyTorch最高实现2.47倍加速,性能模型分析阐明融合如何将负载从内存受限转变为计算受限(达到79.5%的BF16峰值利用率)。
- 证明编译器基础设施无法复现此效果——torch.compile结合所有融合提示仍比我们的内核慢3–7倍,且我们的融合内核数值精度更优(零错误 vs. cuBLAS的4.5–11%错误率)。
2相关工作
2.1 SwiGLU与门控激活
Shazeer(https://arxiv.org/html/2607.02521#bib.bib1)提出了门控线性单元变体(GEGLU、SwiGLU、ReGLU)作为标准FFN激活的替代方案,展示了稳定的质量提升。SwiGLU计算\text{SiLU}(xW_1)\odot(xW_2),其中xW_1为门控投影,xW_2为上投影,此后成为LLaMA(https://arxiv.org/html/2607.02521#bib.bib9)、Qwen-2.5(https://arxiv.org/html/2607.02521#bib.bib10)、Mistral和Gemma的默认激活函数。使SwiGLU有效的双输入门控结构——需要通过独立权重矩阵同时进行门控和上投影——正是我们工作要解决的融合挑战的核心。先前的优化将SwiGLU视为GeMM之后启动的轻量级逐元素内核;我们证明对于中小模型,该逐元素操作及其中间表示可占MLP时间的37%。
2.2 IO感知内核设计
FlashAttention(https://arxiv.org/html/2607.02521#bib.bib2;https://arxiv.org/html/2607.02521#bib.bib8)确立了面向Transformer的IO感知GPU内核设计范式,通过将计算分块以适应SRAM、从不物化完整注意力矩阵,将注意力的HBM访问量从O(N^2)降至O(N^2/M)。我们的工作将此IO感知原则从注意力模块扩展到MLP模块。但融合挑战本质不同:FlashAttention融合单一计算流内的操作(QK^T→softmax→V),而我们的内核必须协调来自两个独立GeMM(具有不同权重矩阵)的输出,然后应用门控激活。这需要创新的调度策略(乒乓重叠、双线程块网格),而注意力机制中无需此类策略。
2.3 GPU内核优化框架
CUTLASS 3.x(https://arxiv.org/html/2607.02521#bib.bib3)为我们的实现提供了基础,提供调度策略(乒乓、协作)、用于可组合GeMM后操作的Epilogue访客树,以及用于SM90上硬件加速异步数据传输的TMA。我们的工作展示了CUTLASS的两个新用法:(1)在Pingpong消费者循环的两个GeMM阶段间插入激活计算(内核1),(2)用自定义PairMulStore节点扩展EVT,在存储阶段融合门控激活(内核2)。
Triton(https://arxiv.org/html/2607.02521#bib.bib4)提供了更高层级的内核开发替代方案,通过编译器管理的调度运行在层级抽象上。虽然Triton能快速原型化融合内核,但其对warp级调度的控制不足以实现我们细粒度的MMA-激活重叠。基于Triton的SwiGLU融合将局限于基础的epilogue融合,无法实现我们实现最大加速所需的层级时间重叠。
2.4 LLM服务与推理系统
生产级LLM服务系统采用不同级别的内核融合。TensorRT-LLM(https://arxiv.org/html/2607.02521#bib.bib5)使用图级模式匹配与预编译融合内核处理常见操作,但其MLP融合的粒度比我们的层级方法更粗。vLLM(https://arxiv.org/html/2607.02521#bib.bib6)和SGLang使用PyTorch默认执行路径(cuBLAS用于GeMM,独立SwiGLU内核),这正是我们的内核要改进的基准线。Megatron-LM(https://arxiv.org/html/2607.02521#bib.bib7)实现了融合偏置+GeLU,但未实现完整的SwiGLU+GeMM融合,因其聚焦于分布式训练而非单GPU推理优化。
我们的融合内核设计为即用替代方案:它们接受相同的输入张量(A、W_1、W_2)并产生与未融合基准线相同的输出(Y),可集成到任何服务框架中而无需架构变更。
3方法
我们提出两种基于CUTLASS的SM90互补内核,在层级将SwiGLU激活融合到GeMM计算中。两种内核都消除了中间张量到HBM的物化,但采用不同策略将激活计算与矩阵算术重叠。
3.1 背景:SwiGLU MLP结构
使用SwiGLU(https://arxiv.org/html/2607.02521#bib.bib1)的现代LLM中的MLP模块计算:
Gate=A\times W_1\in\mathbb{R}^{M\times N} \quad (1)\ Up=A\times W_2\in\mathbb{R}^{M\times N} \quad (2)\ Y=\text{SiLU}(\text{Gate})\odot\text{Up} \quad (3)
其中A\in\mathbb{R}^{M\times K}为输入激活,W_1、W_2\in\mathbb{R}^{K\times N}为独立权重矩阵,\text{SiLU}(x)=x\cdot\sigma(x)为Swish激活,\odot表示逐元素乘法。
内存流量问题
在标准(未融合)实现中,计算需要对中间张量执行8次HBM操作:
- 将Gate写入HBM(M\times N\times 2字节,BF16)
- 将Up写入HBM(M\times N\times 2字节)
- 从HBM读取Gate用于SwiGLU
- 从HBM读取Up用于SwiGLU
- 将Y写入HBM(M\times N\times 2字节)
加上读取A、W_1和W_2。四个中间操作(第1-4项)传输4\times M\times N\times 2字节数据,若在层级保留在寄存器或共享内存中时计算SwiGLU则可消除这些传输。
未融合:A → GeMM1 → Gate(HBM)→ GeMM2 → Up(HBM)→ SwiGLU → Y(HBM)
写入/读取/读取/读取/写入
融合:A → 融合GeMM+SwiGLU → Y(HBM)
写入
图2:内存访问模式对比。左:未融合SwiGLU需要通过HBM写入/读取中间Gate和Up张量(共8次HBM操作)。右:融合内核在寄存器中计算SwiGLU,消除4次中间HBM操作。
瓶颈量化分析
我们的性能分析(第5节(https://arxiv.org/html/2607.02521#S5))显示,在H100上SwiGLU及其相关内存流量消耗MLP总执行时间的9–37%,小模型占比更高:Qwen-2.5 0.5B为30%,Qwen-2.5 72B为9%。随着量化(FP8、INT4)提升张量核算术强度,此内存受限的激活开销将成为更大的相对瓶颈。
算术强度分析
对于未融合基准线,算术强度为:
AI_{\text{未融合}}=\frac{2\cdot 2MKN+2MN}{2(MK+2KN+4MN+MN)\cdot 2} \quad (4)
其中分子为FLOPs(两次各2MKN的GeMM加上2MN的SwiGLU),分母为传输字节数(输入A、两个权重、四次中间传输和输出Y,均为BF16)。通过融合,我们消除4MN个中间传输元素:
AI_{\text{融合}}=\frac{4MKN+2MN}{2(MK+2KN+MN)\cdot 2} \quad (5)
这使算术强度提升6–247%,取决于M/K比值(第5.3节(https://arxiv.org/html/2607.02521#S5.SS3)表3)。
3.2 内核1:通过乒乓调度重叠同步SwiGLU
乒乓调度策略
CUTLASS 3.x的SM90乒乓调度(https://arxiv.org/html/2607.02521#bib.bib3)将线程块内的warp划分为生产者和消费者。生产者异步发出TMA(张量存储加速器)加载填充共享内存缓冲区,而消费者对已加载的层级执行MMA(矩阵乘累加)指令。该调度在两个共享内存缓冲区间交替(“乒乓”),将下一层级的加载与当前层级的计算重叠。
关键洞察:消费者空闲期进行Swish计算
在乒乓调度中,当消费者warp完成层级MMA后、下一数据到达前,存在短暂窗口期消费者warp处于空闲状态(等待生产者TMA加载)。我们利用此窗口在寄存器中对累积的Gate层级计算Swish激活\text{SiLU}(\text{Gate}{\text{tile}})=\text{Gate}{\text{tile}}\cdot\sigma(\text{Gate}_{\text{tile}})。
时间线:
生产者:TMA Gate → TMA Up → TMA Gate → TMA Up
消费者:MMA Gate → Swish → MMA Up → 存储Y
重叠部分:绿色Swish与蓝色TMA Up重叠
图3:内核1乒乓调度时间线。在Gate累加器上进行的Swish计算(绿色)与生产者的Up层级TMA加载重叠,将激活延迟隐藏在加载-计算流水线中。
线程块设计
内核1创建[M,N]线程块网格。每个线程块顺序计算:
- Gate层级:A_{\text{tile}}\times W_{1,\text{tile}}的MMA,跨K维度累积
- Swish计算:在Gate和Up阶段间的同步屏障期间,在寄存器中计算\text{SiLU}(\text{Gate}_{\text{tile}})
- Up层级:A_{\text{tile}}\times W_{2,\text{tile}}的MMA,类似累积
- Epilogue:逐元素乘法\text{SiLU}(\text{Gate}{\text{tile}})\odot\text{Up}{\text{tile}}后存储至HBM
关键优势在于Swish计算(步骤2)在时间上与生产者warp加载Up权重层级的TMA操作重叠。由于Swish仅涉及寄存器上的逐元素操作(sigmoid近似和乘法),它在加载延迟内完成,不延长关键路径。
Epilogue效率
由于Gate(Swish后)和Up累加器结果都位于同一线程块的寄存器中,最终的逐元素乘法和HBM存储达到最大效率——单个融合存储操作写入最终输出Y,无需额外全局内存读取。这使内核1在大批次大小(M\geq2048)时特别有效,此时[M,N]网格提供足够线程块以饱和GPU的流多处理器(SM)。
实现
我们将内核1实现为自定义GemmSwiGLU构建器,扩展CUTLASS的CollectiveMainloop并修改消费者循环。构建器在两个GeMM阶段间插入Swish计算,通过cuda::cp_async_fence屏障协调,确保在应用Swish前Gate累加器完成,且Swish完成前Up MMA存储不覆盖共享内存缓冲区。
3.3 内核2:通过自定义Epilogue访客树交织SwiGLU
Epilogue访客树(EVT)
CUTLASS 3.x提供Epilogue访客树(https://arxiv.org/html/2607.02521#bib.bib3),可组合式构建复杂post-GeMM操作。每个EVT节点代表单个操作(乘法、加法、激活等),树结构定义操作顺序。传统EVT应用于单次GeMM输出;我们扩展EVT以处理SwiGLU所需的双输入门控操作。
PairMulStore节点
我们设计自定义PairMulStore节点,将两个层级乘法(\text{SiLU}(Gate)\odot Up)融合为单一原子操作。该节点在EVT树的叶节点接收两个输入:(1)Gate层级(经Swish激活后)和(2)Up层级,计算它们的逐元素乘积,并将结果直接存储至HBM。此设计确保逐元素乘法与内存存储操作完全融合,无中间全局内存读写。
双层级网格设计
内核2创建[M,2N]线程块网格,每个线程块计算两个相邻的N列输出。线程块内工作分配:
- 并行执行两个独立的GeMM:计算Gate和Up的[M,N]输出
- 应用激活:对Gate累加器执行Swish激活
- 融合后处理:通过PairMulStore节点同时将两个输出的逐元素乘积写入HBM
该设计使内核2在小批次大小(M<2048)时实现2倍更高的占用率,因为双层级设计利用更多线程块填充GPU的SM,而内核1的[M,N]网格在此情况下可能资源利用不足。
数值精度保障
通过将Swish激活保留在寄存器中(内核1)或通过PairMulStore原子化(内核2),我们的融合避免了中间值到HBM的往返写入/读取,消除了cuBLAS中存在的4.5–11%数值错误匹配。
相似文章
GPU上的无畏并发:在Rust中进行安全的GPU推理,与vLLM/SGLang竞争 [R]
cuTile Rust 引入了一种基于块(tile)的编程模型,利用 Rust 的所有权机制来保证 GPU 内核的内存安全和无数据竞争,基于该模型构建的 Grout 推理引擎在 Qwen3 模型上实现了与 vLLM/SGLang 相当的吞吐量。
TileMix:用于LLM推理加速的以分块为中心的混合精度注意力
TileMix引入了一种以分块为中心的混合精度注意力机制,以加速大型语言模型中的长上下文预填充,通过将分数分块组路由到FP16或INT8路径来平衡准确性和效率。
@hardmaru: 人脑极其高效,因为它只激活特定思维所需的神经元。现代LLM…
本文介绍了TwELL和Hybrid稀疏格式,配合自定义CUDA内核,有效利用LLM中的非结构化稀疏性,在H100 GPU上实现了训练和推理速度提升超过20%,同时降低了能耗和内存使用。
@leloykun:[进行中] 关于 Lean4-to-TileLang 张量程序超级优化器的博文:
一篇技术博文介绍了一种 Lean4-to-TileLang 张量程序超级优化器,能自动生成优化的 GPU/TPU 内核与超参数缩放规律,展示了相较 torch.compile 的性能提升。
利用适度非结构化稀疏权重矩阵加速大语言模型的GPU推理
本文提出了一种针对具有适度非结构化稀疏性的大语言模型的高效GPU推理方法。引入了一种三层矩阵存储格式和一个联合利用稀疏张量核心与CUDA核心的SpMM内核,实现了相比SpInfer最高1.64倍的内核级加速,以及相比FlashLLM最高1.41倍的端到端加速。