@PyTorch: https://bit.ly/4yawNqB..*

X AI KOLs Timeline 论文

摘要

这篇来自PyTorch的博客文章介绍了针对LayerNorm和RMSNorm等归一化操作的新型内核融合技术,通过减少内存IO开销实现了显著的加速。这些技术包括Lazy Pre-Norm和Multi-CTA Norm Fusion,在与GEMM融合时可隐藏高达90%的归一化延迟,而FlashNormAttention算法可实现高达35%的内核加速。

https://t.co/pZKSCPtZfc..*
查看原文
查看缓存全文

缓存时间: 2026/07/10 20:17

将归一化融合到 GEMM 与 Attention 内核中 – PyTorch 源: https://pytorch.org/blog/towards-free-normalization-fusing-normalization-into-gemm-and-attention-kernels/?_gl=11ljkcyh_up*MQ ### 精选项目 - Helion项目徽标 (https://pytorch.org/projects/helion/) 代码获取地址:https://github.com/facebookresearch/ads_model_kernel_library/tree/main/multi_cta_norm_fusion 以及 https://github.com/facebookresearch/ads_model_kernel_library/tree/main/gdpa_megakernel ## 太长不看 在这篇博客中,我们介绍了多种针对常见归一化操作(如 LayerNorm 和 RMSNorm)的新型内核融合技术。这些技术通过减少高度内存受限内核的内存-I/O 开销,实现了显著的加速。我们首先简要概述了在 LLM(大语言模型)和广告推荐模型中常见归一化操作的重要性及其性能挑战,然后介绍了应对性能瓶颈的新颖策略,包括延迟预归一化多 CTA 归一化融合。我们展示了这些技术可以将归一化内核的延迟隐藏高达 90%(通过融合到 GEMM 中)。最后,我们介绍了 FlashNormAttention 算法,该算法将多个归一化融合到如 GDPA [1] 这样的注意力内核周围,实现了高达 35% 的内核加速。这项工作主要使用了两种内核 DSL:TLX (https://arxiv.org/abs/2605.10905),这是一组具有更低级、硬件感知 GPU 执行控制支持的 Triton DSL 扩展;以及 Helion (https://pytorch.org/blog/helion/),这是一种高级 DSL,擅长开发者效率、可移植性和全面的自动调优。基准测试使用 bfloat16 数据类型,在 Meta 数据中心的 NVIDIA B200 GPU 上执行,功率上限为 750 W。 ## 引言 归一化技术因其在稳定训练和加速收敛方面的出色效果,已成为大多数深度学习架构中不可或缺的一部分。特别是,传统上在最内层嵌入维度(例如 LayerNorm、RMSNorm)上的归一化,在当代大语言模型以及像 Meta 广告模型这样的推荐系统模型(recsys)中,是最常见且无处不在的类型。例如,在部署于 Meta 最大 Recsys 训练基础模型上的 Kunlun 架构中 (https://arxiv.org/abs/2602.10016)[2](即生成式广告模型 GEM (https://engineering.fb.com/2025/11/10/ml-applications/metas-generative-ads-model-gem-the-central-brain-accelerating-ads-recommendation-ai-innovation/)[3]),LayerNorm/RMSNorm 几乎存在于所有关键组件中,如多头注意力、层级种子池化、以及 GDPA 增强的 PFFN [1]。然而,归一化的普遍性也带来了一个棘手的性能挑战:它高度受限于内存,且无法利用 TensorCore。这阻碍了我们在模型训练中充分利用硬件计算能力。以 Kunlun [2] 为例,归一化大约占总训练延迟的 20%。这意味着如果不进行优化,我们立即损失了 20% 的硬件计算吞吐量。在典型的、更受计算限制的 LLM 中,归一化仍可能占据总延迟的约 10%。为解决这个问题,我们必须以 IO 感知的方式设计归一化内核,通过内核融合仔细节省内存 IO 开销而不损失计算精度,并将这些内存/CUDACore 密集型操作与 TensorCore 密集型操作重叠。由于大多数归一化操作位于矩阵乘法操作(如 MLP、注意力)之前或之后,我们的工作专注于如何高效地将归一化与矩阵乘法融合。我们首先描述将归一化与单个 GEMM 高效融合的多种策略,最后介绍 FlashNormAttention 算法,该算法将 LayerNorm 和 RMSNorm 融合到注意力中。 注意:在以下基准测试结果中,除非另有说明,我们禁用了逐元素仿射 (https://docs.pytorch.org/docs/2.12/generated/torch.nn.LayerNorm.html),因为我们发现这些操作会带来显著的性能开销,而在我们的模型中,它们对模型质量的影响很小。另外,除非另有说明,本文介绍的核心优化和算法思想无论是否存在逐元素仿射都适用,尽管性能结果可能有所不同。 ## 1. 归一化融合的挑战 概述:在本节中,我们将讨论以典型方式将归一化与计算密集型内核(如 GEMM)融合所面临的挑战,其根本原因在于不同的分块策略。然后我们提出一种“朴素”融合方案,该方案强制 GEMM 算法遵循与归一化算法相同的分块方式,并观察到对于非常小的 N 性能良好,但随着 N 增加,由于分块约束和低效性,该方案变得次优甚至不可行。 与标准的激活函数融合(例如 GEMM+ReLU)相比,归一化融合的根本挑战在于分块方式的不同。归一化本质上是一种规约操作,需要访问整个维度上的数据才能计算出正确结果。特别是,对于 LayerNorm 和 RMSNorm,典型的内核会沿外部维度对输入进行分块,但不会沿内部维度分块,这意味着每个 CTA 始终需要加载完整行数据。相比之下,典型的 GEMM 在两个维度上都会分块,这意味着每个数据块不会跨越整个行,从而使得后续的逐行归一化变得不可能。 最直接的解决方法是扩展 GEMM 的数据块大小,使每个数据块跨越整个内部维度。对于一个典型的 (MxK) @ (KxN) GEMM,这意味着沿 N 维度的数据块大小必须大于 N(通常取 N 的下一个 2 的幂)。高级算法如下图所示: 这种方法主要有两个问题: - 它偏离了纯 GEMM 本应最优的分块策略,会因缓存行为、流水线行为等降低 GEMM 自身的性能。 - 它对输入形状设置了硬限制,尤其是 N 可以有多大。N 太大将无法容纳在共享内存中。 我们来做个粗略估算,看看 N 可以有多大。假设 Blackwell GPU 有 228 KB 共享内存,数据类型为 bfloat16,并且为了高效流水线/重叠,最少需要 2 个流水线阶段。进一步假设 M 维度和 K 维度的最小数据块大小为 32。那么我们有: 2 stages x 2 bytes / element x (tile_m x tile_k + tile_k x tile_n + tile_m x tile_n) < 228KB => 32 x 32 + 32 x tile_n + 32 x tile_n < 228KB / 4 => 512 < tile_n < 1024 由于数据块大小通常应为 2 的幂,这限制了 tile_n,进而限制 N 最多为 512,否则内核甚至无法运行。在后续章节中,我们将讨论解决这些限制的方法。但尽管如此,我们发现此类融合策略对于较小的 N 值仍能产生显著的收益。在此实验中,我们使用了 Helion (https://pytorch.org/blog/helion/),因为它具有高开发者效率和详尽的自动调优,在这种其中一个数据块大小受硬约束的非标准情况下特别有用。以下是广告模型中典型输入形状的基准测试结果。请注意,延迟节省按 torch inductor 归一化内核延迟的百分比计算。我们这样做是为了使指标独立于基础 GEMM 内核的延迟(以及它与归一化内核延迟的比较),并自然反映了此类融合尝试的提升空间(即 100% 是最佳情况,需要完全将归一化与 GEMM 重叠)。 对于像 64 和 128 这样的小形状,这种融合策略可以显著节省 17%-32% 的 LayerNorm 内核延迟。然而,当 K/N 增长到超过 128 时,增益开始消失甚至变为巨大的倒退。这是因为随着 N 的增长,强制 tile_n = N 与未融合 GEMM 内核的最优数据块大小偏差越来越大,节省内存 IO 的好处逐渐被扭曲基础 GEMM 算法的弊端所掩盖。 ## 2. 延迟预归一化:一种将预归一化与线性层融合的新技术 概述:本节我们介绍一种新颖的前序融合技术,用于将预 RMSNorm 融合到 GEMM 内核中。我们讨论了此类前序融合的动机和挑战,并提出了一种名为延迟预归一化的新颖算法,通过巧妙延迟部分预归一化计算到完成 GEMM 之后(利用数学技巧),从而应对这些挑战,并获得了良好的性能加速。 我们提出的第一个想法通过将预归一化与后续 GEMM 融合(作为前序融合)来避免上述问题。尽管通常应避免前序融合,但仍有几个原因值得探讨: 1. 前序融合绕过了后序融合中遇到的块划分问题——在 GEMM 内核中,每个 CTA 无法访问输出张量的整行。相比之下,每个 CTA 在算法上本就会扫描输入张量 A 的整行! 2. 实际上,预归一化比后归一化更为普遍,尤其是在大语言模型中。 为了使前序融合有效,我们针对预归一化的一个特例——不带逐元素仿射的 RMSNorm——设计了一种优化技术,称为延迟预归一化。具体来说,我们要融合: C = rmsnorm(A) @ B 其中 rmsnorm(A) = A * rstd(A)[:, None] 且 rstd(A) = rsqrt((A ** 2).sum(dim=-1) / A.shape[-1] + 1e-5) 延迟预归一化旨在解决以下关键困难:在典型的分块 GEMM 中,虽然我们最终能访问整行数据(这允许我们计算规约结果 rstd),但我们是通过逐块方式访问,并且我们实际上需要在处理每个数据块时就用到 rstd!这就产生了循环依赖:我们需要等到 k 循环结束才能计算 rstd,但我们需要 rstd 才能开始循环中的工作!为了解决这个问题,第一个关键观察是,这两个相互依赖的部分本质上属于不同类型的计算:规约和逐元素应用。 1. rstd 计算部分是规约,需要扫描整个行。 2. 使用 rstd 应用归一化是对 A 中每个单独元素的逐元素计算。 让我们分别处理这两个部分。对于规约部分,首先要注意到它本身不会阻塞任何东西,这是一个很好的特性,因为这意味着我们可以将其与 TensorCore 计算并行执行。由于每个 CTA 自然沿 A 的内部维度扫描,我们可以与矩阵乘法并行地累计 A 的平方和。逐元素部分问题更大,因为它依赖于规约结果,导致循环依赖。解决方法来自一个数学技巧,其关键观察是,在无仿射的 RMSNorm 中,逐元素乘法实际上是逐行乘法,即 A 同一行中的所有元素都乘以相同的 rstd。这意味着以下关键性质: (A * rstd[:, None]) @ B = (A @ B) * rstd[:, None] 证明:逐行乘法等价于 M @ A,其中 M 是对角矩阵。因此 (A * rstd) @ B = (M @ A) @ B = M @ (A @ B) = (A @ B) * rstd 这非常棒,因为它意味着逐元素计算可以“延迟计算”,等到整个 k 循环结束后再进行,从而实际上变成了一个后序操作!综上所述,以下是延迟预归一化算法的内核伪代码: def GEMM_norm_fusion_kernel(A, B, C): compute the m_tile and n_tile of this CTA square_sum = zeros(m_tile) acc = zeros(m_tile, n_tile) for each k_tile: tile_A = A[m_tile][k_tile] tile_B = B[k_tile][n_tile] acc += tile_A @ tile_B square_sum += (tile_A * tile_A).sum(-1) # 与 GEMM 并行计算! rstd = rsqrt(square_sum / A.shape[-1] + 1e-5) acc *= rstd[:, None] C[m_tile][n_tile] = acc 注意,尽管每次 k 迭代中仍会产生一些额外计算,但它们可以与矩阵乘法重叠。通过 warp 专业化,内核的 warp 划分和执行如下所示: 请注意,这种算法仍然具有前序融合的一个关键缺点:RMSNorm 计算在许多 CTA 间是冗余的(想想所有 CTA 计算输出张量的相同行但不同列;更多内容请见第 3 节)。然而,由于延迟预归一化确保这部分计算大部分与 TensorCore 完全重叠,这种冗余是可接受的,并且仍然能带来良好的性能提升。 延迟预归一化算法的一些限制: 1. 它不能轻松支持逐元素仿射,因为逐元素仿射是列方向的乘法。这会破坏我们的先决条件——逐元素操作必须是逐行乘法。 2. 它不适用于 LayerNorm,因为 LayerNorm 的逐元素部分涉及减法,不是简单的逐行乘法。 3. 这种融合的反向实现会很棘手,因为在前向中我们从未具体化 rmsnorm(A)。因此,在计算 dA 和 dB 时,我们需要即时从 A 和 rstd 重新构建 rmsnorm(A)。 ## 3. 多 CTA 归一化:将后归一化与线性层融合作为后序操作 概述:尽管延迟预归一化能获得良好的加速,但它仍有局限性,无法通用化到大多数归一化用例。在本节中,我们讨论一种更通用的技术,用于将后归一化与 GEMM 融合,并回到后序融合的领域,直接解决第 1 节中提到的块划分不匹配问题,这要归功于CTA 集群分布式共享内存。我们借鉴了 Quack (https://github.com/Dao-AILab/quack/blob/main/media/2025-07-10-membound-sol.md) 的思想,并将其从独立的归一化内核扩展到融合内核。Quack 归一化内核利用 CTA 集群 (https://docs.nvidia.com/cuda/parallel-thread-execution/#cluster-of-cooperative-thread-arrays) 将大的 N 在同一集群的不同 CTA 间分区,并让它们通过分布式共享内存彼此协作,进行单一的跨 N 规约。这使我们能够有多个 CTA 协作划分并处理相同行的数据,并根据需要在归一化过程中相互通信,而无需承担全局内存 IO 的开销。 如上所述,大多数归一化操作可以分解为规约部分(例如 RMSNorm 的 rstd,LayerNorm 的均值和方差)和随后使用规约结果的逐元素部分。只有规约部分需要扫描整个 N 维度,我们通过 CTA 集群进行分而治之。由于规约结果通常很小(因为它是一个规约),发送/接收它到其他 CTA 所需的 DSMEM 通信开销非常小。 注意,这个想法正好解决了我们面临的正规化融合问题——仅仅是 N 太大!(尽管 N 太大的原因和阈值不同)。这意味着我们可以简单地将这个多 CTA 算法放入 GEMM 的后序中,融合就完成了! def GEMM_norm_fusion_kernel(A, B, C): compute the m_tile and n_tile of this CTA acc = zeros(m_tile, n_tile) for each k_tile: tile_A = A[m_tile][k_tile] tile_B = B[k_tile][n_tile] acc += tile_A @ tile_B acc = multi_cta_norm(acc) # 在此处进行 DSMEM 通信 C[m_tile][n_tile] = acc 请注意,这种融

相似文章

面向Tensix架构的大语言模型推理中的算子融合

arXiv cs.LG

本文提出了一种针对Tenstorrent Tensix架构上大语言模型推理的算子融合策略,将RMSNorm与矩阵乘法融合,以提高数据局部性并减少DRAM访问。在Wormhole平台上,使用Qwen2.5-0.5B、Qwen3-0.6B和Qwen3-4B进行的实验显示,注意力模块延迟降低高达37.44%,MLP延迟降低15.89%。