基于 Strassen 与 LCMA 低复杂度矩阵乘,腾讯 FalconGEMM 探索超越硬件峰值的矩阵乘优化

8 月 1 日,HyperAI 主办的 Meet AI Compiler 技术沙龙第 9 期在北京举办。本期活动聚焦 AI 编译技术的最新进展,多位来自产业界和科研机构的专家围绕编程语言、算子开发、编译优化与推理执行展开分享,呈现 AI 编译器从上层语言表达到硬件执行的协同演进。

其中,腾讯高性能计算工程师朱泓霖以「FalconGEMM: Surpassing Hardware Peaks with Lower-Complexity Matrix Multiplication」为题,分享了团队围绕低复杂度矩阵乘法开展的算法与算子优化实践。

面对 cuBLAS 等成熟算子库已经将矩阵乘性能推至接近硬件峰值、传统 Kernel 级优化空间日益收窄的问题,团队重新从算法复杂度入手,以 Strassen、AlphaTensor 等低复杂度矩阵乘算法为基础,构建统一的 LCMA(Low-Complexity Matrix Algorithms)框架,并结合 QDSL、算子融合、Persistent Kernel、细粒度调度与 Cost Model 等技术,将「减少乘法次数」的理论优势转化为 GPU 上的实际性能收益。

在 NVIDIA H20 的 FP16、BF16 测试中,FalconGEMM 在大量矩阵 Shape 上超过 cuBLAS,峰值性能提升约 10%---16%,同时在语言模型 Benchmark 中保持了与标准矩阵乘基本一致的数值精度。

朱泓霖老师与观众深度交流

HyperAI 在不违原意的前提下,对分享内容进行了整理汇总。

关注微信公众号「HyperAI超神经」,后台回复关键字「 0801 AI 编译器 」,即可获取确认授权的讲师演讲 PPT。

从 Strassen 出发,重新寻找矩阵乘的优化空间

**矩阵乘是深度学习中最重要的基础算子之一,也通常占据模型主要的计算耗时。**CUDA、MKL、cuBLAS 等软件栈已经经过多年优化,在不少场景下,单个 GEMM Kernel 的性能已经非常接近硬件峰值。这意味着,如果继续只在指令、流水线和访存层面做局部优化,能够挖掘的空间已经越来越有限。团队因此将目光重新投向经典的 Strassen 算法。

Strassen 算法由 Volker Strassen 于 1969 年提出。对于最基础的 2×2 矩阵乘,传统方法需要完成 8 次乘法,而 Strassen 通过重新组合输入矩阵,只需要执行 7 次乘法,再通过额外的加减操作恢复最终结果,相当于减少了 1/8 的乘法计算量。如果操作对象只是标量,这笔交易并不划算;但当操作单元变成子矩阵后,矩阵加法 O(N²) 与矩阵乘法 O(N³) 之间的复杂度差异,使得「少做一次矩阵乘、多做若干矩阵加法」开始具备实际价值。

如果进一步递归使用 Strassen,乘法次数还可以继续下降。例如,一个 4×4 的分块矩阵乘,传统方式需要 64 次块乘法,而两层 Strassen 只需要 49 次。不过,递归层数增加也会带来更多加法、数据组织和访存开销,因此实际系统往往只采用有限层数,在计算削减和额外开销之间寻找平衡。

2022 年,DeepMind 提出的 AlphaTensor 又进一步拓展了这一算法空间。它将矩阵乘法转化为 Tensor Decomposition 问题,并利用强化学习搜索更低 Rank 的分解方式,说明除了经典 Strassen 之外,不同 M、N、K Shape 下还可能存在大量不同的低复杂度矩阵乘算法。

但从算法发现走向工程应用,还有一个现实问题:**如果每一种低复杂度算法都需要单独手写 GPU Kernel,开发和维护成本显然过高。**为此,团队将这类算法统一抽象为 LCMA(Low-Complexity Matrix Algorithms),统一描述输入矩阵中哪些子块需要预先组合、实际执行多少次矩阵乘,以及中间结果最终如何组合成输出矩阵,再通过 Codegen 自动生成对应实现。

由此,问题从「如何实现一个 Strassen Kernel」,转变为「如何构建一个能够承载多种低复杂度矩阵算法、同时保持高性能的统一框架」。

与此同时,低复杂度算法还必须面对数值精度问题。Strassen 在代数上与标准矩阵乘等价,但浮点运算并不严格满足结合律,计算顺序变化可能带来额外的舍入误差。因此,LCMA 在追求性能的同时,也需要控制低精度计算中的误差传播。

有了统一的算法描述,下一步就是寻找合适的 GPU 实现方式。团队先后尝试了 CUDA、Triton、TiLang 和 QDSL。CUDA 的硬件控制能力最强,但面对大量不同 LCMA 算法时,寄存器、Shared Memory 和中间求和结构都需要针对性调整,扩展和维护成本较高。

Triton 在基础 Strassen 场景中几乎可以达到 CUDA 的性能,但当算法扩展到更大的分块结构后,需要在多个中间计算之间精确复用寄存器 Buffer,Triton 容易产生额外 Spill。TiLang 在寄存器和 Shared Memory 控制上更加灵活,但团队测试中性能仍比 Triton 低约 5%---10%。对于理论收益只有 12.5% 的基础 Strassen 而言,这一损失已经足以明显侵蚀算法收益。

最终,**团队选择 QDSL 作为 FalconGEMM 的主要实现后端。**QDSL 的开发粒度接近 CUDA,同时具备 Codegen 能力,并支持嵌入 PTX,既方便迁移已有高性能实现,也适合根据不同 LCMA 描述批量生成代码,为后续融合与定制优化提供了更大的空间。

从 LCMA 到 FalconGEMM,把算法收益落到 GPU 上

最直接的 Strassen GPU 实现可以分为几个环节:分别组合 A 和 B 的子矩阵,生成 7 对新的输入;执行 7 个 Batched GEMM;最后再将 7 组中间结果组合成最终矩阵 C。相比常规 GEMM,其中真正计算密集的矩阵乘部分只有原来的 7/8,因此只要前后处理的额外耗时低于省下来的 1/8 计算量,整体就有机会获得收益。

团队首先在 NVIDIA H20 上进行测试。H20 具有较高的显存带宽和相对较低的计算峰值,比较适合这种「增加部分数据处理、换取计算量下降」的方案。在约 2048³ 及以上的矩阵规模上,基础实现已经能够观察到稳定收益。但在更小的 Shape 上,输入组合、中间结果写回和输出组合的占比会迅速上升,很容易吃掉节省下来的计算量。

因此,后续优化的重点从 GEMM 本身转向了中间访存。最直接的思路是算子融合,**让中间结果尽量留在片上,而不是反复写回 Global Memory。**不过,输入端 Combine A/B 并不适合直接融合进 GEMM,因为同一个子块可能被多个 SM 使用,容易产生重复加载和重复求和。相比之下,Batched GEMM 与 Combine H 的后处理融合更具可行性。

真正的难点在于,Strassen 的 7 个中间结果会以不同方式贡献给最终 4 个输出子矩阵。如果以 H 为并行单元,多个 SM 可能同时向同一个 C 写回,带来严重的 Atomic 冲突;如果以 C 为并行单元,又会导致部分 H 被不同 SM 重复计算。两种方式都会抵消减少乘法带来的收益。

团队最终不再按照 Strassen 的中间结果组织任务,而是按照矩阵的空间坐标进行分组:将 7 组 Batched GEMM 中处于相同位置的 7 个乘法 Tile 组成一个 Group,并放在同一个 SM 上执行。这样,一个 Group 完成相关计算后,可以直接在片上将结果累加至最终 C,既省去了中间结果写回 Global Memory,也避免了明显的跨 SM 写冲突。

这一融合方式大幅削减了 Strassen 带来的额外访存,但更大的 Group 粒度又带来了负载不均衡。以 4096³ 矩阵乘为例,粗粒度调度可能造成约 21% 的额外 Wave 浪费,甚至超过 Strassen 本身 12.5% 的计算削减。

为此,**团队借鉴 Stream-K 的思路,将一个 Group 在必要时拆分到两个 SM 上执行。**调度层面仍以 Group 为基本单位,但实际执行可以进一步细化到 Tile,从而减少尾部 SM 空转,在保留 Group 级数据复用优势的同时提升硬件利用率。

不过,解决负载均衡后,新的问题随之出现:L2 Cache 抖动。Group 被拆分后,同一个 Wave 中可能混入不同类型的中间乘法,访问的数据彼此独立,导致 L2 命中率明显下降。与此同时,GEMM 已经高度占用 Tensor Core,当访存压力也接近满载时,H20 会触碰功耗墙。实测中,核心频率从约 1.8 GHz 降至 1.6 GHz,计算性能随之下降,部分融合收益再次被抵消。

针对 L2 Cache 抖动,团队进一步调整拆分后的 Group 顺序,尽可能让同一个 Wave 内处理相同类型的中间结果,只在少量尾部 Wave 中出现混排。这样既保留了细粒度调度带来的负载均衡,也恢复了较好的 L2 数据局部性,最终消除了明显的降频问题。

值得一提的是,**这些调度优化的重要基础是 Persistent Kernel。**与普通 Kernel 中 CTA 完成一个 Block 后退出不同,Persistent Kernel 让 CTA 长时间驻留在 SM 上,持续领取后续任务,使开发者能够更灵活地控制 Group 和 Tile 的执行顺序,并进行片上资源复用。前面的任务拆分、调度重排和缓存优化,也因此能够在同一个 Kernel 内完成。

从缓存重排到 Cost Model,峰值性能提升10%-16%

经过融合、负载均衡和缓存重排后,FalconGEMM 已经能够在更多 Shape 上释放 LCMA 的计算优势。但 LCMA 并不会在所有情况下都优于普通 GEMM:**它的本质仍然是以额外的数据处理换取更少的乘法计算。**如果原始 GEMM 本身已经受到访存限制,继续减少计算并不能带来足够收益;只有在计算密度较高时,低复杂度算法才更具优势。

因此,**团队进一步设计了一个类似 Roofline 的 Cost Model,用于决定什么时候采用 LCMA,以及在多种 LCMA 中选择哪一种。**由于目标是「选对算法」而非精确预测执行时间,模型主要分析不同方案的计算量、访存量,并结合目标 GPU 的算力与带宽,估算其所处的计算/访存瓶颈区间。

在这一模型中,低复杂度算法减少的乘法量对应计算侧收益,额外的数据组合和重复访存则构成新的内存开销;而前述融合优化,又进一步压缩了这部分访存成本。由此,FalconGEMM 可以根据不同 M、N、K Shape,判断常规 GEMM 与不同 LCMA 方案的收益边界,并自动选择更合适的实现。

结合 QDSL 的 Codegen 能力,整个框架最终形成了一套较完整的执行流程:首先基于 LCMA 描述生成相应的融合 Persistent Kernel,通过融合减少中间访存;随后由 Cost Model 针对具体 Shape 选择合适的矩阵乘算法;最后由 QDSL 自动生成并编译目标代码。这样一来,LCMA 不再只是某一种固定的 Strassen 实现,而是形成了可以根据工作负载动态选择的算法空间。

性能测试主要在 NVIDIA H20 上进行。结果显示,在多种低精度矩阵乘场景下,**FalconGEMM 在大量 Shape 上超过 cuBLAS,峰值性能提升约 10%---16%。**完成 Group 拆分和缓存重排后,大 Shape 上能够稳定获得低复杂度算法带来的计算收益,小 Shape 的表现也得到改善,并避免了因 L2 Cache 抖动触发功耗墙而导致的降频。

从 Cost Model 的选择结果来看,在大部分跨过 LCMA 收益阈值的 Shape 上,模型都能够选择到性能更优的实现,说明「算法选择+融合 Kernel」的方式能够较好地覆盖不同计算密度的矩阵乘场景。

除了性能,数值精度也是 FalconGEMM 必须验证的问题。团队早期在低精度实验中曾出现较明显的误差,核心原因在于浮点数的累加顺序发生变化。例如,A+C+B−C 在代数上等于 A+B,但在有限精度浮点计算中并不一定严格相等。

进一步分析发现,真正显著的误差主要来自低精度 Cast,而不是 FP32 累加本身。在常见的 FP16/BF16 输入场景中,矩阵乘通常先以 FP32 完成累加,再转换回较低精度;如果中间结果频繁 Cast,FP32 尾数信息会被不断舍弃。

融合实现反而缓解了这一问题。WGMMA 的输出保持 FP32 精度,FalconGEMM 直接在片上以 FP32 对最终 C 进行组合和累加,直到计算结束后再统一 Cast 回目标精度。相比多个独立 Kernel 之间反复写回低精度中间结果,这种方式减少了一次或多次精度转换,使计算顺序变化带来的误差更多停留在 FP32 的低位。

在语言模型 Benchmark 中,使用 FalconGEMM 与标准矩阵乘得到的最终得分几乎一致,仅存在非常微弱的差异,表明当前实现并未带来明显的模型精度下降。

下一阶段,团队计划继续推进两个方向:一是进一步融合 Combine A/B,通过调整 Batch Group 和 K Loop 顺序,继续减少输入中间结果访存;二是将 LCMA 扩展至 Attention。Flash Attention 同样具有较高的计算访存比,如果低复杂度矩阵算法能够与其分块和流水线进一步结合,也可能带来新的性能空间。

从 Strassen、AlphaTensor 到 LCMA 和 FalconGEMM,这项工作的意义并不只是让已经高度优化的 GEMM 再快几个百分点。它提供了另一种思路:当 Kernel 本身已经逼近硬件上限时,性能优化不仅可以继续向下挖掘指令和流水线,也可以向上重新寻找算法复杂度上的空间,再通过编译、融合和调度,将理论上的计算削减真正转化为运行时间收益。

相关推荐
DeepIntelli3 小时前
品牌百科词条建设:从词条命名到提交的完整技术指南
人工智能
MindUp3 小时前
企业网盘选型的技术评估维度与主流产品架构简析
人工智能·安全·架构
狂师3 小时前
用AI做自动化测试,哪些是真不行,哪些是你不会用?
人工智能·面试·测试
小鹿软件办公3 小时前
微软 MAI-Image-2.6 正式发布,文字渲染与 3D 能力显著增强
人工智能·microsoft
zhangfeng11333 小时前
Windows AMD显卡跑PyTorch
人工智能·pytorch·windows
小弥儿3 小时前
Firecrawl:把整个网页变成 AI 可查询的数据库
数据库·人工智能·学习
探物 AI3 小时前
yolo检测中的激活函数19:ReLU激活函数 (Rectified Linear Unit)
网络·人工智能·深度学习·yolo
晓天衡宇•评测社区3 小时前
大语言模型 8 月榜单更新:Claude Opus 5 登顶,Gemini 3.6-Flash 与 DeepSeek-V4 Flash 展现差异化优势
大数据·人工智能
dunge20264 小时前
2026 ChatGPT Plus / Pro + Codex 实战:从 0 搭一套 AI 编程项目模板,AGENTS.md + Git + 测试一次配好
人工智能·git·chatgpt