Files
matmul-analysis/BMM/BMM从分块计算到四大分支的逻辑推导.md

20 KiB
Raw Blame History

从分块计算到四大分支BMM 最优实现的逻辑推导

昇腾 BatchMatMulV3 算子 | 目标芯片 950PRDAV_351032 AIC / 1.65GHz 上游文档《BMM分块计算数学公式》分块计算"是什么" · 下游文档《BMM最优软件实现方案设计》第六章四大分支"怎么用" 本文定位:补上中间缺失的一层——为什么从"分块计算 + 最优实现"出发,恰好推导出 MergeBatch / IterBatch / ASW_Basic / StreamK 这四大分支,不多也不少。


0. 问题陈述:缺失的那层逻辑是什么

《BMM分块计算数学公式》告诉我们BMM 的 4 重循环 (b, m, n, k) 可以按硬件容量做嵌套分块,并用 swizzle 函数 \sigma 重排块到核的映射。《BMM最优软件实现方案设计》第六章直接给出了四大分支。但中间有一跳没有论证

分块参数 $(B^t, M^t, N^t, K^t)$、核间切分方式、swizzle 函数 \sigma 有无穷多种取法。凭什么最优实现的搜索空间恰好收敛到 4 个分支

本文的推导链条如下,每一节对应链条上的一环:

分块计算公式(自由度)
  → 硬件强加的两条不变式(规则)
  → 四个维度核间切分的代价表(代价不对称)
  → "最优"的三个必要条件(目标函数展开)
  → 决策树:按代价从低到高购买并行度(推导)
  → 15 种切分组合坍缩为 4 个等价类(完备性 + 极小性证明)
  → 与分块公式的参数特化对应(闭环)

1. 起点:分块计算公式给出的自由度

分块计算把 BMM 重组为两层结构记号同《BMM分块计算数学公式》

核间\tilde{B} \times \tilde{M} \times \tilde{N} 个基本块,经 swizzle \sigma 映射到 C 个核:


(\beta, \mu, \nu) = \sigma^{-1}\big((c + rC) \bmod \tilde{B}\tilde{M}\tilde{N}\big), \qquad c \in [0, C)

核内K 循环在 L0C 上原地累加,经 L1/L0 两级缓冲流水:


C_{(\beta,\mu,\nu)}^{\text{L0C}} = \sum_{\kappa=0}^{\tilde{K}-1} \mathrm{mmad}\big(A[\beta,\mu,\kappa],\; B[\beta,\kappa,\nu]\big)

这个数学结构里,自由决策只有三类

自由度 数学对象 物理含义
F1核间切分 \tilde B, \tilde M, \tilde N 如何对 C 分解;\kappa 是否跨核 哪些维度的块被分到不同核
F2核内组织 B^t, M^t, N^t, K^t 及 L1/L0 驻留深度 d_A, d_B 单核内如何折叠/驻留/流水
F3执行顺序 swizzle 函数 \sigma 时间上相邻的块在空间上的排布L2 复用)

"四大分支"问题,本质上就是 F1 和 F2 的选择问题F3 是每个分支内部的二级优化)。所以推导的主线是:F1 有哪些本质上不同的选择,各自的代价是什么。


2. 规则:硬件强加的两条不变式

分块公式里有句话写得轻描淡写,却是整个分支体系的根:

"K 循环 (\kappa) 在 L0C 上累加,不写回 GM。"

把它和输出写回方式放在一起,就是硬件强加给所有实现的两条不变式:

不变式 I1L0C 累加不变式):一个输出块 C_{(\beta,\mu,\nu)}\tilde K 轮 mmad 结果驻留在 L0C256KBFP32上原地累加中间不产生任何 GM/L2 流量。

  • 推论 1核内切 K 是免费的\tilde K 只是循环次数)。
  • 推论 2核间切 K 必然打破 I1——每个核只算了一段 K 的部分和L0C 驻留不住"别人的 K",部分和必须写出到 GM/L2 workspace再由 AIV 做核间归约。分块公式第六章的 STREAM_K 行写的就是这个:

C_{(\beta,\mu,\nu)} = \sum_{c \in \text{group}} C_{(\beta,\mu,\nu)}^{(c)} \quad \text{(部分和经 workspace 归约)}

不变式 I2输出专属不变式:若输出块 C_{(\beta,\mu,\nu)} 由一个核独占负责,则该核只写出最终结果,无核间同步;反之(切 K则引入核间依赖与归约流量。

两条不变式合起来给出本文最重要的结构性事实:

K 维和其他三维在核间切分中的地位根本不对等:切 B/M/N 保持 I1、I2切 K 同时打破 I1、I2。这就是为什么"是否切 K"天然是分支的第一分界线StreamK 必然自成一支。


3. 关键一步:四个维度核间切分的代价表

对 F1核间切分的本质选择是"4 个维度中切哪些"。逐维度分析切分代价依据batch 语义、数据通路、GM/L2 带宽差):

3.1 切 B —— 零代价并行

不同 batch 的 A[b], B[b], C[b] 在内存中天然不相交。因此:

  • 读入:每个数据块被且仅被 1 个核读取,零重复;
  • 计算:每核独立产出最终结果,无核间依赖(保持 I1、I2
  • 写出:只有最终结果,无中间写出。

代价 = 0。但并行度上限 = $B$,且单核分到的计算是单个 batch 的 $[M,K]\times[K,N]$——Cube 效率完全由单 batch 的 M, N 决定(第 4 节会看到这是 MergeBatch 与 IterBatch 分家的原因)。

3.2 切 M 或切 N —— 低代价并行(代价可被 L2 吸收)

切 M同一 batch 的右矩阵 $B[b,:,:]$(大小 $K \times N$)被负责不同 M 行块的所有核重复读取;切 N 对称(左矩阵 M \times K 被重复读)。

  • 读入:重复因子 = M或 N方向的核间切分数 $g_M$$g_N$)。若被共享矩阵能驻留 L2$K N \cdot \text{dtype} \le 128\text{MB}$),则重复读取以 5.2TB/s 命中 L2 而非 1.6TB/s 的 GM——代价大部分被 L2 吸收swizzle $\sigma$(滑窗蛇形)进一步压缩同时活跃的工作集。
  • 计算:无核间依赖(保持 I1、I2
  • 写出:只有最终结果。

代价 = 共享矩阵的重复读取,受 L2 容量约束,属于有上限的低代价

3.3 切 K —— 高代价并行(结构性代价,无法吸收)

  • 读入:无重复(各核读不同 K 段);
  • 计算打破 I1、I2——每个输出块由 \text{grid}_K 个核共同产生,部分和写 workspace再归约
  • 写出:存在中间结果写出。额外时延项 $T_{\text{REDUCE}} \propto \text{grid}_K \times B_c M_c N_c \times 4\text{B} / \text{BW}$AtomicAdd 到 L2 约 5.2TB/s、到 GM 约 1.6TB/s另有同步开销

代价 = 归约流量 + 核间同步,是随 grid_K 线性增长的结构性代价L2 吸收不掉。

3.4 代价不对称性总结


\boxed{\;\text{cost}(\text{切}B) = 0 \;<\; \text{cost}(\text{切}M) \approx \text{cost}(\text{切}N) \;\ll\; \text{cost}(\text{切}K)\;}

这个排序不是经验,是三条硬件事实的推论:① batch 维在数学上独立BMM 语义);② L0C 累加机制I1使切 K 产生独一份的归约开销;③ L25.2TB/s / 128MB恰好能把 M/N 共享读取的代价吸收掉大半。整条决策树就是按这个代价排序"从便宜到贵"地购买并行度。


4. 目标:最短端到端时延展开为三个必要条件

最优实现的判据是 T_{total} = \max(T_{MMAD}, T_{MTE2}, T_{MTE1}, T_{FIXPIPE} [, T_{REDUCE}]) 最小。逐项拆开,\min\max 等价于三个必要条件:

R1 核占满(并行度条件)T_{MMAD} 与活跃核数成反比GM 带宽也要求 ≥ 3/4 核24 核)并发才能达到 90%+ 利用率。设不切 K 时的独立输出块数


P \;=\; B \times \Big\lceil \tfrac{M}{16} \Big\rceil \times \Big\lceil \tfrac{N}{16} \Big\rceil

16 是 mmad 的 fractal 粒度,即输出块的最小可分单元)。P \ge C = 32 是满核的必要条件;P < 32 时必然有核闲置,闲置部分纯浪费。

R2 搬运最快(复用条件)。MTE2 有效带宽由 L2 命中率决定:


BW_{MTE2} \approx (1 - h) \times 1.6 + h \times 5.2 \ \text{TB/s}

baseM = baseN = 256 时达成 Cube Bound 需要 BW_{MTE2} \ge 2.64 TB/s即 $h \gtrsim 28.9%$。GM 流量必须最小化(每份数据 GM 只读一次),重复读取尽量落在 L2/L1。

R3 Cube 喂饱(效率条件)。单核每次 mmad 序列的 tile 要足够大L0C 利用率


U = \frac{M^t \times N^t \times 4\text{B}}{256\text{KB}} = \frac{M^t N^t}{65536}

越接近 1Cube 空转越少;M^t, N^t 需 16 对齐,且 MTE1 不拖后腿要求 baseM, baseN ≥ 80 左右。

分支推导 = 在 R1/R2/R3 约束下,按 §3 的代价排序购买并行度与 Cube 效率。 下面逐分支看它们各自是哪一组约束矛盾的"最优解"。


5. 推导主链:四种典型矛盾 → 四个最优响应

5.1 情形一:B \ge C 且单 batch M \times N 大 → IterBatch

  • R1零代价并行度 B \ge 32 已足够,无需切 M/N省掉共享读取更无需切 K。
  • R3M \times N 大(如 ≥ 128×128单 batch 就能开出 baseM×baseN ≈ 256×256 的 tileU 接近 1Cube 喂得饱。
  • R2核内走标准三级流水L1→L0→Cube→L0CM/N 大 tile 保证复用。

此时任何额外的切分都只会引入代价而无收益,最优解就是"每核分若干 batch核内逐个 batch 做完整 Matmul"。这正是分块公式中 $B^t = 1$、batch 作为核间最外层维度的特化。

判据B \ge C 且 $M N \gtrsim 128^2$L0C 利用率 ≥ 25% 量级)。

5.2 情形二:B \ge C 但单 batch M \times N 小 → MergeBatch

  • R1满足$B \ge 32$)。
  • R3不满足M = N = 64 时 $U = 64^2/256^2 = 6.25%$Cube 大片闲置;M = N = 16 时 $U = 0.4%$mmad 几乎全在空转。并行度够,但每个核"吃不饱"。

矛盾在 R3。唯一的解法把"折叠维度"用在 batch 上——核内把 b 个 batch 的 A 沿 M 拼接、B 沿 N 拼接,等效大矩阵 $[bM, K] \times [K, bN] \to [bM, bN]$L0C 利用率从 MN/65536 提升到 $b^2MN/65536$(受 L0C 约束 $b^2MN \le 65536$,如 M{=}N{=}64 时 $b \le 4$),最后用 BlockTrace 取 b 个对角线 $M \times N 块作为有效输出。

代价是交叉项算力浪费 $(b-1)/b$。为什么这个代价值得付?因为小 M, N 意味着算术强度低:


AI = \frac{b M N K \cdot 2}{(b M K + K b N + b^2 M N)\,\text{dtype}} < 270 \ \text{FLOP/B} \;\Rightarrow\; \text{访存 Bound}

单核分摊的 GM 带宽约 50GB/sCube 要满转需要 AI ≈ 270 FLOP/BM,N 的 case 远低于此,瓶颈本来就在搬运Cube 浪费的拍数被 MTE2 时延掩盖——浪费是免费的。这就是"为什么是 MergeBatch"的定量理由:它不是牺牲算力换效率,而是在算力本就用不完时回收闲置。

判据$B \ge 2C$(每核至少 2 batchMN < 128^2 且 $AI < 270$。

5.3 情形三:B < CB \times \lceil M/16 \rceil \times \lceil N/16 \rceil \ge CASW_Basic

  • R1零代价维度 B 买不够 32 核,但低代价维度 M/N 可以补齐($P \ge 32$)。
  • R2切 M/N 引入共享矩阵的重复读取——必须靠 L2 驻留 + swizzle 把代价压到最低,这正是 swizzle 函数 \sigma 的主战场(滑窗 W = \max\{d \mid d \mid C, d \le \lfloor\sqrt C\rfloor\} 使同窗口 A 行块与 B 列块的 L2 足迹最小)。

矛盾是"R1 缺并行、只能靠低代价维度补"。最优解是一个通用框架:允许切 B/M/N 的任意组合(不切 K核间切分维度与 swizzle 由 shape 决定——切 B 优先零共享B 不够切 M 或 N共享一侧矩阵L2 吸收),再不够混合切。这就是 ASW_Basic它不是一个具体切法而是"所有不切 K、含 M/N 共享切法"的总框架,\sigma 取 ASW 滑窗蛇形。

判据P \ge C 且不满足情形一/二B 不够,或 M/N 大到不需要折叠 batch

与 IterBatch 的竞争边界:B \ge 32M \times N 中等时两者都可行。分野在 L2IterBatch 每核做完整 $M \times N \times K$,若单 batch 工作集 MK + KN 超 L2 则 M/N 外循环反复挤兑 L2ASW 切 M 时右矩阵 $KN \le 128$MB 可驻留 L2 供 32 核共享。谁的工作集能驻留 L2谁优——这是时延模型比较不是硬阈值见 §7

5.4 情形四:P = B \lceil M/16 \rceil \lceil N/16 \rceil < CStreamK

  • R1无法满足。B、M、N 三个便宜维度全部用尽(切到 16 粒度)仍凑不满 32 核。典型:B = 1, M = N = 64 时 $P = 16 < 32$,一半核闲置。
  • 此时唯一剩余的并行维度是 K。切 K 虽然代价高(打破 I1/I2付 $T_{REDUCE}$),但不切的代价是核闲置——两害相权,当

T_{MMAD/\text{core}} \gg T_{REDUCE} \;\Longleftrightarrow\; K \;\gtrsim\; \text{grid}_K^2 \times 1690

(推导:要求计算时延 ≥ 10× 归约时延,代入 8192 FLOP/拍、1.65GHz、GM 1.6TB/s切 K 净收益为正。K 越大,可承担的 grid_K 越大K 不够大时,减少 grid_K配合切 B/M/N即 grid_K × grid_B × grid_M × grid_N ≤ 32 的组合)来降低归约组大小。

矛盾是"便宜维度用尽仍缺并行"。最优解是付归约代价买 K 维并行,且归约组能小则小。这就是 StreamK——分块公式中 \tilde K 跨核拆分、部分和经 workspace 归约的特化。

判据$P < C$(或 P 虽够但核内 M/N 范围被压得太碎)且存在 grid_K 使 K/\text{grid}_K \ge 256 且满足上式。

5.5 一条链总结

R1 缺并行?────────────────────────────────────────────
   │ 不缺                            │ 缺
   ▼                                ▼
R3 缺 Cube 效率?              便宜维度(B/M/N)用尽?
   │ 不缺        │ 缺               │ 用尽
   ▼            ▼                  ▼
B 够切?     MergeBatch        StreamK切K付归约代价
 │B≥C  │B<C   折叠batch
 ▼      ▼     换Cube效率
Iter  ASW_Basic 浪费被访存Bound掩盖
Batch 切M/N共享
       L2+swizzle吸收代价

四大分支各自是唯一最优解的"矛盾区域"互不相同IterBatch 解"零代价并行已够"的区域MergeBatch 解"并行够但 Cube 饿"的区域ASW_Basic 解"并行缺、低代价维度可补"的区域StreamK 解"便宜维度用尽"的区域。区域不同,最优响应不同——这就是四个分支的存在性证明。


6. 完备性与极小性:为什么恰好四个,不多不少

6.1 完备性15 种切分组合按代价特征坍缩为 4 个等价类

核间切分的所有可能 = 四维 {B, M, N, K} 的非空子集,共 2^4 - 1 = 15《方案设计》§6.1 已枚举 C1~C15。关键观察一个切分组合的代价结构只由两个布尔特征决定——

  • 是否含 K(决定要不要付归约代价,打破 I1/I2
  • 是否含 M 或 N(决定有没有共享矩阵的重复读取)。
含 K 含 M/N 组合 代价结构 归入分支
任意 C4, C7, C9, C10, C12~C158 个) 必有归约grid_K×grid_B×grid_M×grid_N 只是参数差异 StreamK
C2, C3, C5, C6, C8, C116 个) 必有共享读取;切哪几维只是 grid 参数 ASW_Basic
否(纯 {B} C11 个) 零共享零归约 核内只有两种组织方式,见下

纯 {B} 的核内组织方式只有两种——把多个 batch 合并成一个大 tile 算MergeBatch或逐个 batch 算IterBatch——不存在第三种(要么利用 batch 间的 tile 级合并,要么不利用)。于是:


15 \;\xrightarrow{\text{按 (含K, 含M/N) 归并}}\; 1 + 1 + 2 \;=\; \boxed{4}

任何合法 case 的任意切法都落在这 4 个等价类之一 ⇒ 完备

与源码对照arch35 源码有 10 个策略,为什么这里只有 4 个?因为源码策略 = 4 个切分等价类 × 两个正交维度的笛卡尔积:① 计算通路Cube vs AIVK=0 的 K_EQUAL_ZERO、K=1 的 TO_MUL 是 K 退化时的通路切换,发生在切分决策之前,属于预处理层);② 驻留策略AL1/BL1_FULL_LOAD 是 ASW_Basic 内部 M^t = M / N^t = N 的 tiling 极限ITER_BATCH_BROADCAST 是 IterBatch 在广播输入下的数据复用特化)。它们是分支内部的参数特化,不构成新的切分等价类。源码尾部还有 BASE=999 无条件兜底,也印证"特判 ⊂ 通用"的层级结构。

6.2 极小性:每个分支都有它是唯一最优的 shape 区域

去掉任何一个分支,都存在 case 失去最优实现:

分支 独占最优的示例 caseBF16 替代方案的劣势
IterBatch B=32, M=N=K=4096 ASW 切 M/N 引入无谓共享读取MergeBatch 引入无谓浪费
MergeBatch B=64, M=N=64, K=256 IterBatch 的 L0C 利用率仅 6.25%Cube 闲置 ~94%
ASW_Basic B=2, M=N=8192, K=1024 切 B 仅 2 核活跃StreamK 付无谓归约
StreamK B=1, M=N=64, K=65536 不切 K 只用 16/32 核P=16时延差数量级文档 Case 319.5ms → 0.61μs/核量级)

⇒ 四者互为不可替代,构成极小完备集


7. 边界是软的:竞争区域由时延模型仲裁

上面的判据给出的是分支的"主场",但相邻分支的边界不是硬切换。重叠区域(如 B ≥ 32 且 M×N 中等时 IterBatch 与 ASW 切 B 等效M×N 在 64²~128² 之间时 MergeBatch 与 IterBatch 互有胜负)必须由统一点评仲裁:


\text{branch}^* = \arg\min_{\text{cand} \in \bigcup \text{4 分支候选}} \max\big(T_{MMAD}, T_{MTE2}, T_{MTE1}, T_{FIXPIPE}, T_{REDUCE}\big)

这就是《方案设计》§6.7 决策算法"全分支生成候选 → 时延模型评估 → 取 min"的合理性来源:分支体系保证候选集完备且无冗余(每个等价类只派一个代表框架),时延模型在等价类内部和边界上做精细仲裁。 两层结构缺一不可——只有时延模型没有分支体系,搜索空间是 15 种组合 × 全部 grid爆炸且重复只有分支体系没有时延模型边界 case 会被错误硬切。


8. 闭环:与分块计算公式的参数特化对应

最后把推导结果写回分块公式的语言验证链条闭合对应《BMM分块计算数学公式》第六章但现在每行都有了"理由"

分支 分块参数特化 推导出处
IterBatch $B^t = 1$$\tilde B \ge C$batch 为核间最外层;核内标准 K 累加 §5.1:零代价并行够 + Cube 饱
MergeBatch B^t = b > 1 折叠进 $M^t_{eff} = bM^t, N^t_{eff} = bN^t$L0C 二次方程求最优 $b$BlockTrace 取对角 §5.2Cube 饿 → 折叠 batch 换 $U \uparrow$,浪费被访存 Bound 掩盖
ASW_Basic 通用 (\beta, \mu, \nu) 展开;\sigma = ASW 滑窗蛇形(W \le \lfloor\sqrt C\rfloor 的最大因子) §5.3:低代价维度补并行,\sigma 压缩 L2 足迹
StreamK \tilde K 跨核拆分,C_{(\beta,\mu,\nu)} = \sum_c C^{(c)} 部分和归约grid_K×grid_B×grid_M×grid_N ≤ C §5.4:便宜维度用尽,付归约代价买 K 维并行

9. 总结:一句话逻辑链

分块计算给出自由度4 维怎么切硬件给出规则L0C 累加使切 K 独贵、batch 独立使切 B 免费、L2 使切 M/N 廉价最优目标给出约束核要满、Cube 要饱、搬运要省)。按"最便宜的维度优先购买并行度、Cube 不够大就折叠 batch、便宜维度用尽才切 K"的原则展开15 种切分组合在 (含K, 含M/N) 两个代价特征下恰好坍缩为 4 个等价类——这就是 MergeBatch / IterBatch / ASW_Basic / StreamK 四大分支,完备且极小;边界区域由端到端时延模型统一仲裁。


文档版本v1.0 | 上游《BMM分块计算数学公式》v1.1+ | 下游《BMM最优软件实现方案设计》v4.0 第六章