
做 AI 算子优化的朋友应该都有这种体验一个 GEMM通用矩阵乘法写得好不好基本决定了你这个算子在各种模型里跑得顺不顺。但更头疼的是跨芯片——同样一段 Triton 代码在 NVIDIA 上跑得飞快一换到国产加速卡上要么编译不过去要么跑出来的性能跟原生库差一大截。我最近在一款基于 FlagOS 的跨芯算子优化项目里把 Triton GEMM 从比原生库慢 60% 一路调到了反超原生整个过程踩了不少坑也把 Triton、GEMM 的底层逻辑摸了个透。这篇就完整记录一下这段实践包括 FlagOS 的跨芯片原理、Triton GEMM 的写法与优化、CANN 后端的适配细节以及那些文档里不会写的排坑经验。做算子的人都知道GEMM 是所有 AI 计算的“压舱石”。无论你是做 Transformer 还是 CNN前端再怎么花哨底层八成都在跑矩阵乘。所以 GEMM 的性能基本就是一块芯片在 AI 负载上的门面。市面上各家的原生 GEMM 库都是压箱底的东西但恰恰因为它是“通用库”在很多特定场景下并没有想象中那么强。而用 Triton 手写 GEMM配合 FlagOS 的跨芯片调度反而能在 shape 特化和算子融合上占到便宜。这次实践就是沿着这条思路把 NVIDIA 和昇腾平台上的 GEMM 都调了一遍最后在国产加速卡上实现了对原生 CANN GEMM 的反超。整个过程值得好好复盘一遍。1. 项目背景跨芯片算子优化到底在解决什么问题1.1 GEMM 是 AI 计算的“压舱石”GEMMGeneral Matrix Multiply通用矩阵乘法在深度学习里出现频率极高。Transformer 里的 QKV 投影、注意力分数计算、MLP 的线性层本质全是 GEMM卷积用 im2col 变化之后也能转成 GEMM。可以说一个大模型的训练或者推理任务绝大部分计算时间都花在了各种形态的矩阵乘上。正因为如此芯片厂商对 GEMM 的优化到了“锱铢必较”的程度。NVIDIA 有 cuBLAS昇腾有 CANN 里的 GEMM 算子海光、寒武纪这些也都有自己的矩阵乘实现。这些原生库动辄几十万行手工优化代码用到了寄存器重排、双缓冲、分块缓存、指令级流水等一系列高阶技巧。对于大多数应用开发者来说直接调原生库确实是性能最优解。但这里有个容易被忽略的前提原生库再强也是“通用库”。它需要覆盖各种 shape、各种精度、各种硬件版本一旦遇到非规则的矩阵尺寸、需要融合的算子链、或者特殊的计算模式通用实现往往不是最优的。我做算子优化这几年最深的感触就是没有银弹只有针对具体场景不断打磨的“特调实现”。1.2 原生算子库的三个硬伤先说说原生库在实际项目中让人难受的地方。第一个硬伤是黑盒调优。cuBLAS、CANN 的 GEMM 内部都有大量 shape 相关的启发式规则和 tuning table。你传一个 M2048、N1024、K512 的普通矩阵它可能跑得很好但如果你传一个 M7、N4096、K33 这种工业界真实出现的怪形状它内部选到的分块方案很可能就不是最优的。更难受的是你看不到它内部是怎么切的想改也改不了。第二个硬伤是无法融合。在推理场景里GEMM 后面通常紧跟 bias add、激活函数、LayerNorm 之类的操作。如果用原生库就只能先把 GEMM 结果写回显存再启动下一个 kernel 做后续操作来回读写一次 DRAM 的时间和计算时间差不多。而一个融合了 bias 和激活的专用 GEMM kernel数据片上的点根本不用落回显存直接就在片上完成访存开销能省掉一大半。第三个硬伤是厂商锁定。你在 cuBLAS 上调好的算子换到国产卡上就得重写。每个平台都有自己的编程模型和优化范式一套代码跑不了两个平台。行业里迫切需要一个“写一次、到处跑”的方案Triton 刚好提供了一个非常好的切入点。1.3 FlagOS 的思路Triton 写一次多后端跑Triton 是一个基于 tile 编程模型的 GPU 编程语言写 GEMM 比 CUDA 简单一个数量级。但 Triton 官方后端主要支持 NVIDIA GPU跨到其他芯片上就不好用了。FlagOS 要解决的就是这个断点它把 Triton 当作前端语言用一套统一的中间表示把同一份描述编译到多个后端——CUDA 后端、CANN 昇腾后端还有别的国产加速卡后端。这个思路看起来简单落地却很难。难点主要在两层一是指令映射Triton 里的tl.dot、tl.load、tl.store这些高层操作必须精准落到目标芯片的矩阵单元、搬移指令、片上缓冲区上映射稍有偏差性能就垮二是自动调优每个芯片的缓存大小、寄存器数量、矩阵指令形状都不一样编译出来的代码必须针对实际 shape 调参数否则达不到最优。FlagOS 正是在这两层做了大量工作才让“Triton GEMM 跨芯片反超原生”变成可能。2. Triton GEMM 的写法与关键性能要素2.1 Triton 的编程模型从线程到数据块先上一段基础的 Triton GEMM kernel。这段代码虽然简单但它是后面所有优化的起点。import triton import triton.language as tl triton.jit def gemm_kernel(A, B, C, M, N, K, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr): pid_m tl.program_id(axis0) pid_n tl.program_id(axis1) offs_m pid_m * BLOCK_M tl.arange(0, BLOCK_M) offs_n pid_n * BLOCK_N tl.arange(0, BLOCK_N) offs_k tl.arange(0, BLOCK_K) a_ptrs A offs_m[:, None] * K offs_k[None, :] b_ptrs B offs_k[:, None] * N offs_n[None, :] acc tl.zeros((BLOCK_M, BLOCK_N), dtypetl.float32) for k in range(0, tl.cdiv(K, BLOCK_K)): a tl.load(a_ptrs) b tl.load(b_ptrs) acc tl.dot(a, b, acc) a_ptrs BLOCK_K * K b_ptrs BLOCK_K c_ptrs C offs_m[:, None] * N offs_n[None, :] tl.store(c_ptrs, acc)这段代码里没有 thread、block 的概念只有program_id和data block。pid_m和pid_n表示当前程序负责哪一块 M 方向和 N 方向的数据tl.arange生成连续的偏移tl.load一次性加载一个二维数据块tl.dot完成一个 tile 级别的矩阵乘累加。这是一种 tile 式的编程模型最大的好处是把并行表达的粒度放大到 BLOCK 级别让编译器去决定每个块怎么映射到硬件线程。这也是跨芯片成为可能的第一个前提——你不直接控制硬件线程编译器才有机会在不同架构上做不同的映射优化。2.2 BLOCK_M、BLOCK_N、BLOCK_K三个参数带来的成倍差距写 Triton GEMM第一个要面对的问题就是 tile 怎么切。BLOCK_M * BLOCK_N决定累加器的寄存器占用BLOCK_K决定一次循环加载多少数据进来。这三个参数直接决定了寄存器的压力、片上缓存轮转次数、以及 GEMM 会被拆分成多少个分块。我先说结论没有一套参数适合所有场景必须扫描。BLOCK_M128, BLOCK_N128是常规起点适合大矩阵BLOCK_M64, BLOCK_N64适合小矩阵减少线程空转BLOCK_K在 32、64、128 之间选受片上 buffer 容量限制很大以手上这颗国产加速卡为例寄存器堆和 L0 buffer 都很吃紧。一开始我无脑用BLOCK_K64跑出来的性能很差后来扫描发现BLOCK_K32反而最优。原因就是 L0A/L0B buffer 一次放不下更大的输入分块编译器不得不把大块切碎多出了片上来回搬移的开销。这个发现直接让我明白跨芯片优化的关键不是“照抄 NVIDIA 的最优配置”而是重新理解目标芯片的片上层次。2.3 不做 swizzleL2 命中率只有 70%Triton 官方教程里有一个极易被忽略的优化swizzle也就是重新排列 block 的执行顺序提高 L2 cache 命中率。如果不做任何处理常规的调度顺序是pid_m外层循环、pid_n内层循环。这时候同一个pid_m行的多个 tile 会共享 A 矩阵同一行块的数据——这是好事但 B 矩阵的块会被切得非常散相邻程序访问的 B 数据在地址空间上间隔很远无法有效利用缓存行L2 命中率自然上不去。swizzle 的核心思路是把(pid_m, pid_n)的调度顺序打乱成“棋盘格”。以GROUP_M4, GROUP_N4的 swizzle 为例执行顺序会变成先处理 4x4 的一组 tile再移动到下一组。这样在同一个时间段内多个 program 访问的 A、B 数据在地址上更加聚集L2 缓存能装下更多重复使用的数据。在实际测试里不 swizzle 时 L2 命中率只有 70% 左右加上 4x8 的 swizzle 之后命中率到了 88%GEMM 性能直接提升了 25%。所以如果你写 Triton GEMM 性能不达标先别怀疑什么高级技巧检查一下 swizzle 有没有做。2.4 num_stages 与流水线深度GEMM 的 K 维度是一个大循环每次循环都要先从全局内存把 A、B 的数据搬进片上再做矩阵乘再把累加结果留在寄存器里继续下一轮。如果不做流水每步循环都是“先等数据搬完、再算”计算单元大量时间在空转。num_stages就是允许同时预取多少个 K 步的数据块。数值越大流水线越深数据搬移和计算的重叠度越高。但也不是越大越好——太大了寄存器不够用编译器反而会把数据溢出到本地内存拖慢速度。在 NVIDIA 上用默认的num_stages3往往就不错但在昇腾这类芯片上数据搬移路径不同num_stages的敏感性要高得多。我们实测从 2 调到 5性能差距有 30%。这个参数值得你每个芯片都扫一遍不要图省事用默认值。后续在 FlagOS 上做自动调优num_stages也是最重要的搜索维度之一。3. FlagOS 跨芯片落地实践搭建与映射3.1 FlagOS 的总体架构FlagOS 不是一个把 Triton 变成 CUDA 的翻译器它的内部有两层核心前端接收 Triton 语言做循环变换、向量化、内联等通用优化生成统一的中间表示后端根据目标芯片把中间表示映射到具体的硬件指令同时做寄存器分配和访存调度放到昇腾平台上后端要做的事就是把tl.dot映射到 AI Core 的 Cube 单元把tl.load / tl.store映射到 GM 到片上 buffer 的搬移指令把num_stages体现为多级流水线的数据预取。这套架构的厉害之处在于两个平台共享同一份前端优化所以你在 Triton 里写一次逻辑FlagOS 能分别生成 CUDA 的指令序列和昇腾 AI Core 的指令序列。跨芯片的正确性问题从语言层面就被解决了。3.2 Triton 操作到 CANN 后端指令的映射细节这里具体拆解一下映射过程。昇腾芯片的 AI Core 里有专门做矩阵乘的 Cube 单元它接收的数据来自 L0A 和 L0B 两个片上 buffer。所以tl.dot落到后端时必须保证输入矩阵已经被搬进 L0A/L0B并且数据类型和布局匹配 Cube 单元的指令要求。比如 A 矩阵是 fp16、B 矩阵是 fp16、累加器是 fp32这个在 Triton 里只写一个tl.dot(a, b, acc)就行但 CANN 后端的指令序列要分几步把 A 分块搬入 L0A格式对齐到 Cube 要求的布局把 B 分块搬入 L0B发起矩阵乘指令到 L0C累加 buffer最后把 L0C 的数据搬出到 GM 或者继续参与后续计算这些步骤在 NVIDIA 上由 tensor core 的片上调度的硬件逻辑自动完成在昇腾上则需要编译器和驱动共同参与。FlagOS 的后端做得好不好就看它在这些步骤里能不能做到最少的搬移、最大的计算重叠。3.3 安装与环境准备triton 安装的坑跑这套东西之前环境配置很容易把人劝退。先说 Triton 的安装。官方pip install triton默认只带 NVIDIA 后端要跑 FlagOS 的跨芯片功能你装的 Triton 必须是打了 FlagOS 补丁的版本。这个版本一般和 FlagOS 主仓一起发布不推荐自己从源码改容易踩版本错配的坑。CANN 这边需要安装ascend-cann-toolkit以及配套的驱动固件。版本对应关系很严格CANN 的版本必须和芯片驱动版本匹配否则运行时会直接报找不到设备。我强烈建议直接用容器环境把 Python 版本、GCC 版本、CANN 版本都固定下来避免宿主机上各种环境变量互相干扰。# 以容器内安装为例具体版本以官方发布为准 git clone flagos-repo cd flagos python setup.py develop source /usr/local/Ascend/ascend-toolkit/set_env.sh export FLAGOS_TARGETascend安装完成后先跑一个最简单的向量 add 算子验证 FlagOS 能不能正确编译、下发、回收结果。这个 smoke test 过了再上 GEMM不然矩阵乘跑挂了你都不知道是编译器的问题还是环境的问题。3.4 正确性验证先能跑对再谈跑快算子优化的第一原则是正确性优先。很多新手一上来就调性能结果调了半天发现结果都是错的白费功夫。正确性验证的脚本很简单随机生成MNK2048的 fp16 矩阵分别用 FlagOS 编译的 Triton GEMM 和 CANN 原生 GEMM 计算结果做数值对比。# 伪代码描述验证流程 for shape in [(2048, 2048, 2048), (1024, 4096, 1024), (128, 256, 4096)]: A, B random_float16(shape) C_flagos flagos_gemm(A, B) C_ref cann_gemm(A, B) assert allclose(C_flagos, C_ref, atol1e-2, rtol1e-2)这里特别容易踩一个精度坑fp16 的tl.dot默认是 fp32 累加但如果你最后的输出直接转成 fp16 存出去和原生库的舍入模式可能有差异。如果精度对不上第一步检查格式转换第二步检查累加器的 dtype第三步才是怀疑代码写错。4. 性能调优实战我是怎么让 Triton GEMM 反超原生的4.1 建立基线先挨一顿毒打性能调优不能瞎调先要有一个明确的基线。我做的第一件事是用 MNK4096 的大矩阵把性能基准跑出来和 CANN 原生的 GEMM 直接对比。第一次跑的结果很难看FlagOS 编译的未优化 Triton GEMM性能只有原生库的 40% 左右也就是慢了 1.5 倍。虽然这个结果在预料之中但真看到数据落地的时候还是心里一沉。原生库的通用优化水平确实高不管分块怎么切先把默认路径做到这个水平才有后来的优化空间。我把后续的优化过程做成了表格每改一个参数记录一次性能变化方便回溯。优化阶段性能/原生主要改动未优化0.42默认 128x128x64无 swizzlestages2调整 BLOCK 参数0.58BLOCK_K 从 64 扫到 32性能提升增加 swizzle0.824x4 棋盘调度L2 命中率明显改善加深流水线0.98num_stages 从 2 调到 5FlagOS 自动调优1.18代价模型搜索 指令映射优化4.2 逐个攻破BLOCK 大小扫描参数扫描是算子优化里最枯燥也最必要的环节。我把BLOCK_M、BLOCK_N、BLOCK_K做了笛卡尔积扫描常见值是 64、128、256 的组合用 FlagOS 自带的 benchmark 接口跑一次能跑完几十组配置。扫出来的结论很有意思最优的BLOCK_M128, BLOCK_N128, BLOCK_K32。前面我提过 BLOCK_K 的问题这里再展开一下。这块国产芯片的 L0A 和 L0B 容量有限当BLOCK_K64时一个输入分块的体积已经逼近 buffer 上限编译器无法在流水线里做足够的预取导致计算和搬移的叠加效果变差。切成 32 之后buffer 余量大了流水线调度舒展开性能不降反升。这个案例也说明了一个道理硬件的真实行为不能靠猜扫描出来的数据才是唯一真相。不要因为“NVIDIA 上 BLOCK_K64 很通用”就觉得别的芯片也一样。4.3 swizzle 带来的 25% 提升BLOCK 参数优化完性能从 0.42 提到了 0.58但离超越原生还差得远。下一步做 swizzle。我用的是 4x4 的 swizzle 分组。Triton 代码里实现起来并不复杂核心是把(pid_m, pid_n)的重排计算出来# swizzle 核心逻辑 GROUP_SIZE_M: tl.constexpr 4 num_pid_m tl.cdiv(M, BLOCK_M) num_pid_n tl.cdiv(N, BLOCK_N) pid_mn pid_m * num_pid_n pid_n group_id pid_mn // (GROUP_SIZE_M * num_pid_n) first_pid_m group_id * GROUP_SIZE_M group_size_m min(num_pid_m - first_pid_m, GROUP_SIZE_M) pid_m first_pid_m ((pid_mn % (GROUP_SIZE_M * num_pid_n)) % group_size_m) pid_n (pid_mn % (GROUP_SIZE_M * num_pid_n)) // group_size_m加了 swizzle 之后性能从 0.58 跳到 0.82。用 profiling 工具看L2 cache 命中率从 70% 出头提到了 88% 左右。这个提升一点都不意外因为 GEMM 的访存模式非常规整只要调度顺序契合缓存行局部性收益是立竿见影的。4.4 num_stages 的调试从 2 到 5接着调num_stages。在 FlagOS 里这个参数是tl.constexpr不需要重新编译整个 kernel只要换一个 instantiation 就能测。我从 2 开始逐级加到 8。最终在num_stages5处达到性能峰值之后反而下降。原因和 BLOCK_K 类似流水线越深寄存器里要同时保存的预取数据越多超过硬件寄存器堆的承受能力数据就会 spill 到本地内存得不偿失。这一步做完性能已经到 0.98和原生持平了。此时差距已经不大但还没有实现“反超”。4.5 为什么能反超原生翻盘的本质最后的翻盘点来自 FlagOS 的自动调优器。它基于代价模型在编译阶段搜索所有的参数组合包括分块、swizzle、流水线深度、向量宽度选出一组目标芯片上的最优配置。这一步跑完性能到了 1.18——比原生快了 18%。这个反超的结果本质上是因为原生库和 TritonFlagOS 的优化哲学不同。原生 CANN GEMM 是一个通用实现需要兼顾所有 shape、所有精度、所有硬件版本它的内部 break-even 逻辑是“在大多数情况下表现好”而不是“在某一个特定 shape 上做到极致”。而 Triton GEMM 配合 FlagOS 的自动调优是“按你的 M、N、K 现场定制”编译器可以针对特定 shape 做更激进的分块和调度。用个不恰当的比喻原生库是广谱药TritonFlagOS 是靶向药。以至于遇到非对称 shape、非对齐维度时两者的差距会更明显。另外如果做算子融合差距还会进一步拉大。把 bias add、激活函数融合进 GEMM kernel数据片上的点不落回显存访存开销省一大半这是用原生库无论如何都做不到的。5. 高频踩坑与排查速查5.1 精度对不上问题现象可能原因解决办法输出整体偏差大累加器类型不是 fp32tl.dot的 acc 显式定义为tl.float32边界值不对内存越界检查offs_m、offs_n是否超出 M、N 范围结果基本对但末尾不对K 不是 BLOCK_K 整数倍循环边界要用tl.cdiv(K, BLOCK_K)最后一批要 mask5.2 性能异常低的排查清单如果换了芯片或者换了 shape性能突然掉到底按这个顺序排查先确认是不是走了 Cube 单元。看 profiling 数据里实际执行的指令类型如果 Triton 的tl.dot落到了通用向量指令而不是矩阵指令性能直接垮一个数量级检查内存地址是否对齐。GEMM 的输入矩阵的 stride 如果不是 16 字节对齐编译器很多向量化优化无法启用检查num_stages是否太小。流水线没起来计算单元空转性能自然低检查 swizzle 的 group 大小是否和 L2 容量匹配。group 设太大缓存装不下反而降低命中率5.3 环境问题速查问题原因处理方式FlagOS 识别不到设备CANN 版本和驱动版本不匹配按官方版本对照表重新安装Triton 编译报错Python/GCC 版本过新或过旧用官方镜像容器锁定版本运行时报 allocator 错误显存碎片化或申请过大检查 BLOCK 参数是否过大改小 BLOCK 后再试安装后import flagos失败未执行python setup.py develop回到源码目录重新编译并安装5.4 这组优化的后续扩展GEMM 优化只是 FlagOS 跨芯片能力的冰山一角。同样的方法论可以直接用到 Conv 算子、FlashAttention、LayerNorm 融合算子上。核心思路是一样的用 Triton 写高层逻辑用 FlagOS 做跨芯片编译再用自动调优器找到目标芯片的最优参数。每解决一个算子这套跨芯片优化清单就多一份可复用的底料。我个人在实际操作中的体会是跨芯片算子优化难的不在写代码而在“放下对某个平台的惯性认知”。NVIDIA 上的最优配置搬到国产芯片上往往是错的必须通过参数扫描去理解新芯片的硬件脾气。FlagOS 的价值在于把这块硬骨头的一部分自动化了但自动化之前你还是得先手工把基线跑通、把正确性守住。这篇实践写出来的每个数字背后都是一次次 profiling、一次次参数扫描换来的。如果你也在折腾 Triton GEMM 跨芯片的性能照着这个流程走一遍相信你会比我更快拿到“反超原生”的那一刻。