昇腾 950PR — BatchMatMulV3
v0.7|基于 CANN 源码 + 950PR 白皮书 + 架构规格的系统推导
这份分析靠谱吗?→ 最后一页有验证
Batch MatMul:对 batch 维的每个索引独立做一次矩阵乘
输入:左矩阵 [BatchA, M, K]、右矩阵 [BatchB, K, N]
输出:[BatchC, M, N],batch 维支持广播
核心挑战:B、M、N、K 四个维度都可能很大或很小,必须在32 个 AIC 核上高效并行完成计算
一个 case 由 (B, M, N, K, dtype, 转置, 广播形态) 完全决定
| 规格 | 数值 | 对 tiling 的意义 |
|---|---|---|
| AIC / AIV 核数 | 32 / 64(1:2) | 核间并行度上限 C=32 |
| Cube 算力 BF16 | 486 TFLOPS | 算存比分子 |
| GM 带宽 | 1.6 TB/s(读写共享) | 访存 Bound 分母 |
| L2 Cache | 128MB / 5.2 TB/s | 重复读取的吸收层 |
| L1 / L0A / L0B / L0C | 512KB / 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 是瓶颈——数据复用就是一切
总时延 = 流水线最慢的一级
核内多级流水并行: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)
BMM 的实现 = 把数据按 B、M、N、K 切块,由 32 个 AIC 核并行 + 串行完成
核间怎么分这 4 个维度,就是分支划分的第一性问题
| 切分维度 | 读入特征 | 计算特征 | 写出特征 |
|---|---|---|---|
| 切 B | 核间零重复读 | 无核间依赖 | 无中间结果 |
| 切 M / 切 N | 共享矩阵被多核重复读 | 无核间依赖 | 无中间结果 |
| 切 K | 零重复读 | 多核共同完成,有依赖 | 需核间归约 |
4 维任意非空子集共 15 种切分组合 → 完备枚举
为什么切 K 最贵?L0C 累加机制:核内切 K 时多轮 mmad 在 L0C 原地累加不出核;切到核间,部分和必须写出 workspace 再归约
切 B 免费:零重复读、零依赖、零中间写出——BMM 语义就是逐 batch 独立
切 M/N 廉价:共享矩阵被重复读,但 128MB L2(5.2TB/s)吸收大部分代价
切 K 昂贵:归约流量 ∝ 切 K 份数 × 输出量,且引入核间同步——结构性代价
整条分支决策树就是一句话:
按价格从低到高购买并行度,买不够才加价
7 大路径:特殊(AIV)转MatmulIterBatchMergeBatchASW_BasicStreamK降核ASW
思路:单 batch M×N 太小时,把 b 个 batch 拼成 [bM,K]@[K,bN] 大矩阵,
算完取块对角线得各 batch 结果。交叉项被丢弃,浪费比例 (b−1)/b
为什么允许浪费? 进这个分支的 case 必然是访存 Bound,瓶颈在搬移
不是在计算——浪费的算力被搬移时延掩盖,用闲置算力换搬移效率
进入条件精要:
为何不允许 b_core=2?合并后每核仅 1 组,与 IterBatch(b_core=2) 总搬移时延相同,但多 50% 冗余计算——无收益
典型:B=128, M=N=64, K=512 → MergeBatch;K=256 → 单核搬移不够 480KB → IterBatch
思路:核间切 B,核内逐个 batch 做标准 Matmul。无浪费、无跨 batch 依赖
进入条件精要:B ≥ 32;负载均衡;L1 四形态之一满足
关键:能否进 IterBatch 不由算存比判定
即使计算 Bound,L1 放不下完整 M、N 维输入 → 单 batch 内部重复读 → 额外搬移可能把算子重新拖回访存 Bound
L1 四形态:
思路: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 才放得下
归约代价(按实现流程推导):
归约时延 = 写部分和 + AIV 读回 + AIV 求和 + 写回,四段之和
收益判据:T_Reduce < T_pipe(1−1/grid_K),T_pipe = max(T_MTE2, T_MMAD)
⟺ T_pipe > grid_K/(grid_K−1)·T_Reduce(α=grid_K/(grid_K−1),grid_K=2 时 α=2)
StreamK case 多为访存 Bound(AI < R₁₆),T_pipe = T_MTE2 是瓶颈
K > grid_K²/(grid_K−1) × 12 grid_K=2→K>49;4→K>66;8→K>112
源码 8192 = C×512B = dValue 推荐值在最大 grid_K=C 下的保障(条件2),非归约代价(条件3阈值仅 ~50~400)
思路:不切 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
切的是输出平面:把 M×N 切成大矩形块,每块的输入工作集 ≤ L2 可用空间
L2 是读写共用的:输出驻留 L2 会压缩读入空间 → 写与读要联合决策
写出 Bound 先判定:BW_out = C·Q₁₆·outB/(2K),K=256 时已超 GM 总线
三场景:
r_in = 1 含义:每个输入从 GM 只读一遍,后续复用全在 L2 命中
B∈[1,2048]、M,N,K∈[1,10240] 对数采样,按进入条件严格分类
| 分支 | 占比 | B 范围 |
|---|---|---|
| ASW_Basic | 35.2% | 2~2048 |
| IterBatch | 21.9% | 32~2048 |
| 降核 ASW | 15.4% | 2~128 |
| 特殊(AIV) | 8.3% | 任意 |
| MergeBatch | 8.1% | 128~2048 |
| 转Matmul | 7.6% | B=1 |
| StreamK | 3.5% | 2~128 |
7 个分支全部有真实 case 命中,无覆盖空洞
降核 ASW 是 P<C 且 K 小区域的理性归宿:并行度凑不满,碎切反而不如少用核
整体评价:逻辑体系完整,推导自洽,可指导工程实现