BMM 算子最优实现分析

昇腾 950PR — BatchMatMulV3

v0.7|基于 CANN 源码 + 950PR 白皮书 + 架构规格的系统推导

这份分析靠谱吗?→ 最后一页有验证

1. 什么是 BMM 算子?

Batch MatMul:对 batch 维的每个索引独立做一次矩阵乘

C[b, m, n] = Σₖ A[b, m, k] · B[b, k, n] + bias

输入:左矩阵 [BatchA, M, K]、右矩阵 [BatchB, K, N]

输出:[BatchC, M, N],batch 维支持广播

核心挑战:B、M、N、K 四个维度都可能很大或很小,必须在32 个 AIC 核上高效并行完成计算

一个 case 由 (B, M, N, K, dtype, 转置, 广播形态) 完全决定

2. 硬件基础:昇腾 950PR

规格数值对 tiling 的意义
AIC / AIV 核数32 / 64(1:2)核间并行度上限 C=32
Cube 算力 BF16486 TFLOPS算存比分子
GM 带宽1.6 TB/s(读写共享)访存 Bound 分母
L2 Cache128MB / 5.2 TB/s重复读取的吸收层
L1 / L0A / L0B / L0C512KB / 64KB / 64KB / 256KB核内分块的容量约束

关键关系:L2(5.2TB/s)→ L1(512KB)→ L0A/B(64KB)→ Cube(16×16×16 一拍)→ L0C(256KB)→ Fixpipe → GM

GM 带宽 1.6TB/s 相对 L2 的 5.2TB/s 是瓶颈——数据复用就是一切

3. 什么叫"最优"?

T_total = max(T_MMAD, T_MTE2, T_MTE1, T_Fixpipe [, T_Reduce])

总时延 = 流水线最慢的一级

核内多级流水并行:Cube 计算(MMAD)、GM→L1 搬移(MTE2)、L1→L0 搬移(MTE1)、L0C 写出(Fixpipe)

瓶颈交换:搬移是瓶颈时,可牺牲算力换搬移效率;计算是瓶颈时,可牺牲搬移换计算效率

算存比:AI = 2MN/(M+N) FLOP/元素

16bit 平衡点:R₁₆ = 486 TFLOPS / (1.6TB/s ÷ 2B) ≈ 607.5 FLOP/元素

AI < R₁₆ → 访存 Bound(瓶颈在搬移)|AI > R₁₆ → 计算 Bound(瓶颈在 Cube)

4. 分块计算的本质:4 维可切

BMM 的实现 = 把数据按 B、M、N、K 切块,由 32 个 AIC 核并行 + 串行完成

核间怎么分这 4 个维度,就是分支划分的第一性问题

切分维度读入特征计算特征写出特征
切 B核间零重复读无核间依赖无中间结果
切 M / 切 N共享矩阵被多核重复读无核间依赖无中间结果
切 K零重复读多核共同完成,有依赖需核间归约

4 维任意非空子集共 15 种切分组合 → 完备枚举

为什么切 K 最贵?L0C 累加机制:核内切 K 时多轮 mmad 在 L0C 原地累加不出核;切到核间,部分和必须写出 workspace 再归约

5. 切分维度的"价格表"

cost(切 B) = 0 < cost(切 M/N) ≪ cost(切 K)

切 B 免费:零重复读、零依赖、零中间写出——BMM 语义就是逐 batch 独立

切 M/N 廉价:共享矩阵被重复读,但 128MB L2(5.2TB/s)吸收大部分代价

切 K 昂贵:归约流量 ∝ 切 K 份数 × 输出量,且引入核间同步——结构性代价

整条分支决策树就是一句话:

按价格从低到高购买并行度,买不够才加价

6. 从价格表到 7 大分支

第0层:K=0/1 → AIV 向量通路(Cube 无用) BatchA=1 或 BatchB=1 → 折叠转普通 Matmul 第1层:B ≥ 32(切 B 可满核) M×N 大 → IterBatch(逐 batch 算) M×N 小 → MergeBatch(多 batch 合并成大 tile 算) 第2层:B < 32,P = B·MN·4B/L0C ≥ 32 → ASW_Basic(切 M/N 补并行,重复读交 L2+swizzle 吸收) 第3层:P < 32(B/M/N 都填不满核) K 大 → StreamK(切 K,付归约代价) K 小 → 降核 ASW(只用 ⌈P⌉ 个核,其余闲置)

7 大路径:特殊(AIV)转MatmulIterBatchMergeBatchASW_BasicStreamK降核ASW

7. MergeBatch — 多 batch 合并计算

思路:单 batch M×N 太小时,把 b 个 batch 拼成 [bM,K]@[K,bN] 大矩阵,

算完取块对角线得各 batch 结果。交叉项被丢弃,浪费比例 (b−1)/b

为什么允许浪费? 进这个分支的 case 必然是访存 Bound,瓶颈在搬移

不是在计算——浪费的算力被搬移时延掩盖,用闲置算力换搬移效率

进入条件精要

典型:B=128, M=N=64, K=512 → MergeBatch;K=256 → 单核搬移不够 480KB → IterBatch

8. IterBatch — 逐 batch 计算

思路:核间切 B,核内逐个 batch 做标准 Matmul。无浪费、无跨 batch 依赖

进入条件精要:B ≥ 32;负载均衡;L1 四形态之一满足

关键:能否进 IterBatch 不由算存比判定

即使计算 Bound,L1 放不下完整 M、N 维输入 → 单 batch 内部重复读 → 额外搬移可能把算子重新拖回访存 Bound

L1 四形态

9. StreamK — K 维核间切分

思路:B/M/N 都填不满核时,把 K 切给多个核各算一段,再归约

进入条件:P = B·MN·4B/L0C < 16(=C/2);K/grid_K ≥ 128 元素(dValue 256B)

阈值取 C/2:grid_K≥2 时每块需 2 核,P×2 ≤ C 才放得下

归约代价(按实现流程推导):

第1步:grid_K 个 AIC 各算 K/grid_K 段,L0C 原地累加 第2步:fixpipe 写部分和到 workspace(驻留 L2,走 5.2TB/s 写口) 第3步:AIV 从 L2 读回各段部分和到 UB,向量求和,写回 (AIV 独立硬件,数据流 GM/L2→UB→AIV→UB→L2/GM)

归约时延 = 写部分和 + AIV 读回 + AIV 求和 + 写回,四段之和

要求计算时延 ≥ 归约时延 × 10 → K ≥ grid_K² × 109

grid_K=2→K≥0.5K;4→1.8K;8→7.0K;32→112K

源码固定门槛 8192 ≈ 本模型 grid_K≤8 的要求,互为印证

10. ASW_Basic — 通用框架(swizzle)

思路:不切 K,切 B/M/N 任意组合。最常命中的分支(35.2%)

swizzle 滑窗蛇形:编排输出块的执行顺序,压缩每一波核的活跃工作集

窗口怎么取:M 向每 W 块划窗,W = C 的 ≤ √C 的最大因子(32 核 → W=4)

一波核的 L2 足迹 ≈ (W·M^t·K + C/W·K·N^t)·dtype,W=√C 时最小

蛇形只在窗口行边界:奇数窗口行 N 方向反向,使上一窗口末尾的 B 列带延续到下一窗口(LRU 热线),跨窗切换几乎零增量。窗内不蛇形——窗内 W 个 A 行块全程驻留 L2,扫序无关。

对比:朴素行优先(Ñ=16 时)一波足迹 18 单位;滑窗 12 单位 → 缩小 1/3

11. ASW_Basic — L2 切分 & 写出策略

切的是输出平面:把 M×N 切成大矩形块,每块的输入工作集 ≤ L2 可用空间

L2 是读写共用的:输出驻留 L2 会压缩读入空间 → 写与读要联合决策

写出 Bound 先判定:BW_out = C·Q₁₆·outB/(2K),K=256 时已超 GM 总线

三场景

r_in = 1 含义:每个输入从 GM 只读一遍,后续复用全在 L2 命中

12. case 遍历验证:20736 个 case 全覆盖

B∈[1,2048]、M,N,K∈[1,10240] 对数采样,按进入条件严格分类

分支占比B 范围
ASW_Basic35.2%2~2048
IterBatch21.9%32~2048
降核 ASW15.4%2~128
特殊(AIV)8.3%任意
MergeBatch8.1%128~2048
转Matmul7.6%B=1
StreamK3.5%2~128

7 个分支全部有真实 case 命中,无覆盖空洞

降核 ASW 是 P<C 且 K 小区域的理性归宿:并行度凑不满,碎切反而不如少用核

这份分析靠谱吗?

✅ 根基扎实

⚠️ 需要注意

整体评价:逻辑体系完整,推导自洽,可指导工程实现