1.9 KiB
1.9 KiB
转Matmul 分支理论
整理自《BMM算子优化分析 v0.98》§四. 对应软件实现
bmm_theory/branches/to_matmul.py.
1. 一句话本质
单边 batch=1 的 BMM 与普通 Matmul 只差一个维度标签——把 batch 维折叠掉即可复用 Matmul 的成熟优化体系,没必要在 BMM 框架内重新造一套。
2. 进入条件
BatchA = 1 \;\lor\; BatchB = 1
无 broadcast 歧义时(另一边 batch 任意)即可折叠。
3. 两种折叠形态(代价不对称)
| 情形 | 折叠方式 | 代价 |
|---|---|---|
| BatchB = 1(右矩阵单边) | 左矩阵 [B, M, K] 的 batch 维与 M 维在 ND 下内存相邻紧排,直接视图为 [B·M, K],输出布局逐元素一致 |
零重排、零 split,免费转换 |
| BatchA = 1(左矩阵单边) | 右矩阵需折叠为 [K, B·N],要一次真实转置重排 O(B·K·N),且输出存在置换需 scatter |
有代价 |
4. BatchA=1 的细分决策
A 矩阵较小(M·K·dtype ≤ L1,可常驻 L1)时:
- 优先留在 BMM 分支内做广播友好形态(A 常驻 L1、逐 batch 复用),避免转置重排开销;
- A 较大(放不进 L1)时,广播扩展后分别预估"BMM 分支"与"重排 + 转Matmul"两条路的时延,择优——不是无条件转。
5. 实现方案要点
- BatchB=1:直接视图折叠,零成本。折叠后按 Matmul 的标准分支(切 M/N,必要时切 K)走;
- BatchA=1 且 A 小:BMM 广播形态,A 常驻 L1,对 B 的每 batch 做
[M,K]@[K,N]; - BatchA=1 且 A 大:广播扩展为
[B,M,K]@[B,K,N],比较 (a) BMM 框架内 ASW_Basic 切分 vs (b) 重排转置 + 转 Matmul,取时延小者。
软件中本分支当前只做路由标注(
router.py前置归约),详细方案生成待与 Matmul 理论体系打通后补齐——折叠后的 Matmul 复用 MM 的切分逻辑,不在 BMM 范围内重复建设。