11 KiB
BMM 算子最优实现分析 —— 播客文字稿
对话者:主持人(M)& 技术专家(E) 时长:约 25 分钟
M:今天我们要聊一份技术文档,全称叫《BMM 算子优化分析 v0.7》,是讲昇腾 NPU 上批量矩阵乘怎么做到最优的。说实话我第一眼看到这份文档,里面全是什么 B、M、N、K 的公式,什么 MergeBatch、StreamK,什么 L0C、L2 切分,看不太懂。所以我把作者请来了,咱们从最浅的地方开始,一步步把它讲明白。先问一个最基本的问题——BMM 算子到底是干什么的?
E:BMM 就是 Batch Matrix Multiply,带 batch 的矩阵乘。你有一个左矩阵 A,形状是 [Batch, M, K],一个右矩阵 B,形状是 [Batch, K, N],对 batch 维里的每一个索引,独立做一次矩阵乘,得到一个输出 [Batch, M, N]。跟普通矩阵乘的区别就是多了一个 batch 维度——同一套操作,做很多份。比如一个 batch 里有 128 个样本,每个样本都要做同一个矩阵乘,你就得高效地并行完成这 128 次。
M:那"最优"怎么定义?是不是跑得越快越好?
E:对,但"快"的度量要精确。NPU 上计算不是一条线,是好几条流水线同时跑——Cube 单元在做乘加(MMAD),同时 MTE 引擎在把数据从全局内存搬到 L1(MTE2)、从 L1 搬到 L0(MTE1),Fixpipe 在把结果写出去。这些流水线可以并行,但总时延取决于最慢的那一条。所以公式就是 T_total = max(T_MMAD, T_MTE2, T_MTE1, T_Fixpipe)——总时延等于最慢的那一级。优化就是找到瓶颈在哪一级,然后想办法缩短它。
M:那如果搬移是瓶颈,计算不是,怎么办?
E:这就是这份文档里一个很重要的概念,叫"瓶颈交换"。搬移是瓶颈的时候,你可以故意多算一点——就是让 Cube 多做一点冗余计算,但换来的好处是数据搬移变少了或变高效了。反之,如果计算是瓶颈,你可以牺牲搬移效率换计算效率。这份文档里有两个典型例子:MergeBatch 就是"牺牲算力换搬移效率"——把多个 batch 合并成一个大矩阵算,算了很多交叉项最后扔掉,浪费了算力,但搬移效率上去了,因为瓶颈本来就在搬移,浪费的算力被掩盖了。ASW_Basic 切 M/N 是反过来——"牺牲搬移换并行度"——把矩阵按行或列分给不同核,同一个右矩阵被多个核重复读,多搬了几次,但换来了更多核并行干活。
M:好,现在我知道什么叫最优了。那文档里怎么决定一个 case 用哪种方案呢?我看到有 7 个分支。
E:这个推导过程是整份文档最核心的逻辑。BMM 的实现本质就是把数据切成小块,分给 32 个 AIC 核去算。切分有 4 个维度:B(batch)、M(行)、N(列)、K(内积维)。核间怎么分这 4 个维度,就是第一性问题。文档首先做了一件事——把 4 个维度的任意非空子集全部列出来,一共 15 种切分组合。然后问:每种切法有什么代价?
三个维度代价完全不同:切 B 是免费的——因为 batch 在数学上就是独立的,每个核领几个 batch,各干各的,零重复读、零依赖、零中间结果。切 M 或 N 是廉价的——共享矩阵会被多核重复读,但如果 128MB 的 L2 Cache 能装下共享数据,重复读的代价大部分被 L2 的 5.2TB/s 带宽吸收了,而不是去走慢得多的 GM 1.6TB/s。切 K 是最贵的——因为要把 K 拆给多个核,每个核算一段,最后要把各段的部分和汇总起来,这就引入了核间通信和归约开销。
M:所以整条分支决策树就是"按价格从低到高购买并行度"?
E:对,就是这个思路。先买免费的 B 维——如果 B 够大(≥32),切 B 就能填满核。然后看单 batch 的 M×N 大不大:大就逐个 batch 算(IterBatch),小就把多个 batch 合并成一个大块算(MergeBatch)。如果 B 不够 32,买不满核,就加价买廉价的 M/N 维——切行或列,重复读的代价靠 L2 和 swizzle 吸收(ASW_Basic)。如果 M/N 也买不满(B、M、N 都小),才买昂贵的 K 维——切 K,付归约代价(StreamK)。如果 K 也不够大,不值得切,那就认了,只用几个核,剩下的闲置(降核 ASW)。加上前面两个特殊通路——K=0 或 K=1 走 AIV 向量核、一侧 batch=1 折叠成普通 Matmul——一共 7 个路径。
M:这个逻辑很清晰。但具体到每个分支,怎么判断一个 case 能不能进?我注意到文档里的进入条件全是 B、M、N、K 和芯片参数的公式,没有凭感觉的。
E:这正是这份文档严谨的地方。我拿 MergeBatch 举例。MergeBatch 是把多个 batch 合并计算,有浪费,浪费比例是 (b-1)/b。所以第一个条件:浪费必须能被掩盖,这就要求 case 必须是访存 Bound——瓶颈在搬移,算力有余量。文档用算存比公式来表达:AI = 2MN/(M+N) < R₁₆/b₀,其中 R₁₆ 是芯片的平衡点 ≈ 607.5 FLOP/元素。这保证了合并 b₀ 个 batch 后算存比仍然低于平衡点,算力浪费不被感知。
然后是容量约束:合并后的输出块 [b₀M, b₀N] 必须放得进 L0C(256KB)。L0C 是 Cube 累加器的片上缓存,放不下就白合并了。接着是搬移效率:单核搬移总量必须 ≥ 480KB,这是 950PR 实测出来保证 GM 带宽利用率的下限;单次搬移的 tile 大小 ≥ 16KB,搬移太小启动开销盖过传输。最后是 batch 量:每核至少分到 2×2=4 个 batch,因为合并最小粒度是 2 个 batch,还要两组才能乒乓流水。
M:每一条条件都有道理。那 IterBatch 呢?我注意到文档里特别强调了一点——"IterBatch 能否进入不由算存比判定"。
E:没错,这是 v0.6 到 v0.7 的一个关键修正。IterBatch 是逐 batch 算,看起来简单,但有一个隐藏风险:如果 L1 放不下单 batch 的完整输入数据(M×K 和 K×N 两个矩阵),核内就会在 K 循环中重复读——每切一次 M 就要重新读一遍 B,每切一次 N 就要重新读一遍 A。这个重复读会引入额外的搬移,可能把一个本来计算 Bound 的 case 重新拖回访存 Bound。所以进入条件不能靠算存比判定,必须直接看 L1 能不能装下。文档为此设计了四种 L1 驻留形态:全驻留、双 batch 乒乓、一侧驻留+对侧切 K、两侧都切 K。形态之间还有 batch 间流水掩盖的分析——比如多 batch 时一侧驻留的方案,必须在 batch 边界把下一 batch 的驻留侧预取进来,否则会产生气泡。这个分析在 v0.7 里做了合并和简化,逻辑比 v0.6 更清晰。
M:StreamK 那个 θ = 109 的推导,我看了好几遍才明白。能再讲一遍吗?
E:StreamK 是最后一招——B、M、N 都切不动了,只能切 K。切 K 的代价是归约。归约的时延到底多大?v0.7 按实际实现流程拆成了四段:第一段,AIC 把部分和通过 fixpipe 写到 workspace,workspace 驻留 L2,走 5.2TB/s 的 L2 写口,这是很快的。第二段,AIV 从 L2 读回各部分和,走 5.2TB/s 的 L2 读口。第三段,AIV 做向量求和——AIV 是独立硬件,有自己的 UB 内存和向量单元,64 个 AIV 聚合做 fp32 加法的速率大约 13.5×10¹² 次每秒。第四段,结果写回。把这四段时间加起来,要求每核的计算时间至少是归约时间的 10 倍,才能保证归约不成为新的瓶颈。代入 950PR 的数值,算出来 K 必须 ≥ grid_K² × 109。grid_K 是切 K 的份数,grid_K=8 时 K 需要 ≥ 7.0K。有意思的是,源码里 StreamK 的固定门槛是 K ≥ 8192——就是说,当 grid_K 取到 8 左右,本文的解析门槛和源码的固定门槛是一致的,互相印证。之前 v0.6 的推导用的是 GM 带宽(1.6TB/s),算出来 θ=1700,比现在高 15 倍——那是因为 v0.6 假设部分和走 GM,而实际上部分和可以驻留 L2,归约的读写成本低得多。v0.7 修正了这个模型,也更符合实际硬件。
M:ASW_Basic 是占比最大的分支,35.2%。它的 swizzle 和 L2 切分,我感觉是最难理解的部分。
E:swizzle 的本质其实很简单。切 M/N 之后,同一时刻 32 个核各算一个输出块,这些块需要的 A 行和 B 列就是"活跃工作集"。如果按自然顺序分配,这 32 个块横跨很多行和列,活跃工作集很大,L2 放不下就掉到 GM 去读,重复读的代价就真实发生了。swizzle 做的事情就是重新编排输出块的执行顺序,让同一波 32 个块尽量落在一个小窗口里——窗口里只有 W 个 A 行块和 C/W 条 B 列带,活跃工作集最小。W 取多少?一波 32 个块的 L2 足迹近似是 (W + C/W)·K·M^t·dtype,均值不等式告诉我们 W 接近 √C 时足迹最小,加上 W 必须整除 C 才能保证窗口边界不碎波,所以 W 取 C 的 ≤√C 的最大因子。32 核就是 4。
蛇形只在窗口行边界做——上一个窗口的最后一列 B 带和下一个窗口的第一列 B 带是同一条,热数据不丢。窗内不用蛇形,因为窗内所有的 A 行块全程驻留 L2,顺序不影响。
L2 切分是另一层:当整个工作集超过 128MB,即使滑窗压缩了同一波内的足迹,跨波次还是会有数据被挤出。所以要把 M×N 平面切成大块,一块的工作集控制在 L2 容量内,逐块计算。同时还要考虑输出怎么处理——输出如果驻留 L2,会挤占输入的空间,让输入被挤出、增加重复读;如果直写 GM,会占用共享的 GM 总线。文档给出了三个场景的决策,核心逻辑是:输入能驻留 L2 就优先驻留(保住 r_in=1),输出只写不读、零复用收益,所以优先直写 GM 给输入腾空间,只要校验总线不爆。
M:最后问一个我猜很多人会问的问题——这份分析靠谱吗?
E:从几个维度来看。第一,根基来自源码——CAN N ops-nn 仓库的 batch_mat_mul_v3 算子,arch35 路径,分支条件一一行对得上,差异处文档都标注了并解释了为什么。第二,硬件参数全部来自 950PR 架构白皮书和 CANN 9.0 编程文档,不是拍脑袋。第三,数学推导闭环——15 种组合完备枚举 → 价格表 → 7 分支,极小性有反例表格验证,每个分支去掉都会有一个 case 失去最优方案。第四,遍历验证——取 B 从 1 到 2048、M/N/K 从 1 到 10240 的 20736 个 case,按进入条件严格分类,7 个分支全部有真实命中,无覆盖空洞。
需要注意的地方也诚实标注了:经验常数(480KB、16KB 这些)是 950PR 实测值,换芯片要重新标定;MergeBatch 和 IterBatch 在 M×N 中段有重叠区,归属要由端到端时延模型精确仲裁;StreamK 的 θ=109 基于 workspace 驻留 L2 的前提,如果工作集太大溢出到 GM,θ 会升高到 1700。但整体来说,这份分析的逻辑体系是完整的,推导是自洽的,可以指导工程实现。
M:好,感谢你耐心讲解。这份文档和 PPT 都在 Gitea 仓库里,感兴趣的朋友可以去看原文。
参考文档:BMM算子优化分析_v0.7.md(Gitea: matmul-analysis 仓库 BMM算子优化分析_Release/)