Files
matmul-analysis/BMM/BMM算子优化分析_Release/BMM理论vs源码对比分析_v1.0.md

323 lines
16 KiB
Markdown
Raw Permalink Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# BMM 理论最优 vs 源码实现对比分析
> 基于《BMM 算子优化分析 v0.93》与 `ops-nn/matmul/batch_mat_mul_v3` 源码DAV_3510
---
## 一、分支选择机制对比
### 理论v0.93):按切分代价从低到高推导
```
转Matmul → 特殊分支 → MergeBatch/IterBatch → ASW_Basic → StreamK → 降核ASW
```
核心逻辑:**按价格从低到高购买并行度**——切 B 免费 → 切 M/N 有重复读代价 → 切 K 有归约代价。每个分支有独立的进入条件B/M/N/K 闭式表达式case 落入哪个分支由条件判定。
### 源码:优先级有序匹配
```cpp
// batch_matmul_v3_tiling_strategy.h, DAV_3510 优先级
{K_EQUAL_ZERO(0), TO_MUL(1), STREAM_K(2), MERGE_BATCH(3),
ITER_BATCH_BROADCAST(4), ITER_BATCH_BASICAPI(5), ITER_BATCH(6),
AL1_FULL_LOAD(7), BL1_FULL_LOAD(8), ASW_BASIC(9), BASE(999)}
```
**优先级从高到低逐个尝试 IsCapable(),第一个返回 true 的分支胜出。**
### 关键差异
| 维度 | 理论 | 源码 |
|---|---|---|
| 选择方式 | 各分支独立条件判定 | 优先级有序匹配first-match |
| StreamK 位置 | 最后检查P<C/2 才考虑 | 优先级 3K=0/K=1 之后立即检查 |
| 转Matmul | 独立分支BatchA=1BatchB=1 | 无独立分支由广播/ASW 承接 |
| 降核 ASW | 独立分支P<C 且不满足 StreamK | 无独立分支ASW_Basic 内部处理 |
**源码优先级顺序的风险**StreamK 优先级高于 MergeBatch/IterBatch若一个 case 同时满足 StreamK IterBatch 的条件B>C 且 K≥8192源码会先进入 StreamK——但理论上 B>C 时切 B 免费,不应优先切 K。不过 StreamK 的 IsCapable 条件(`batchC×mCnt×nCnt ≤ aicNum/2`)实际上排除了 B>C 的 case所以优先级顺序在实际中不会造成误判。
---
## 二、逐分支对比
### 2.1 特殊分支K=0 / K=1
**理论条件**v0.93 §九):
- K=0无计算C = bias 或 0纯 AIV 写值
- K=1退化为逐元素乘走 AIV 向量通路GM→UB→Mul→GM触发需 B ≥ 2×64 且单 batch 输入输出驻留 UB
**源码实现**
`BATCH_MATMUL_INPUT_K_EQUAL_ZERO`(优先级 0
```cpp
// batch_matmul_v3_k_equal_zero_tiling.cpp IsCapable()
bool IsCapable() {
if (aFormat == NZ || bFormat == NZ) return false; // 不支持 NZ
if (hasBias) return false; // 不支持 bias
if (kValue != 0) return false; // K=0
return true;
}
// DoOpTiling: usedCoreNum = aivNum用 AIV 核)
```
`BATCH_MATMUL_TO_MUL`(优先级 1K=1 时退化为逐元素乘,用 AIV 核。
**对比结论**理论与源码一致。源码实现正确——K=0 和 K=1 都用 AIV 通路,不用 Cube。理论中的"B ≥ 2×64"条件在源码中未显式检查,而是通过 UB 容量隐式限制(`ubLimitBatchNum = ubSize / singleBatchSize`)。
**差异**:理论要求"单 batch 输入输出驻留 UB"作为进入条件;源码通过 `batchNum = min(singleCoreBatch, ubLimitBatchNum)` 自动处理 UB 放不下的情况(分批处理),无需在进入条件中排除。**源码更通用,理论更保守。**
---
### 2.2 MergeBatch
**理论条件**v0.93 §五):
1. BatchA = BatchB ∧ b_core = B/C ≥ 2b₀b₀=2 → b_core ≥ 4
2. 2(b₀M)(b₀N)·4B ≤ L0C
3. b_core(MK+KN)·dtype ≥ min_DatamountPerCore (480KB)
4. max(MK, KN)·dtype ≥ min_TileSize (16KB)
5. 2MN/(M+N) < R₁₆/b
**源码条件**`batch_matmul_v3_mergebatch_basicapi_tiling.cpp` IsCapable
```cpp
bool IsCapable() {
if (非连续转置) return false;
if (NZ格式) return false;
if (hasBias || (FP32 && !HF32)) return false;
if (BatchA != BatchB) return false; // 对应理论条件1
if (batchC < MIN_BATCH_L0 * aicNum) return false; // B < 128即 b_core < 4
if (alignK < 64 || M > N) return false; // K≥64M≤N
// L0 buffer check with b0=MIN_BATCH_L0=4
if (al0Size > L0A || bl0Size > L0B ||
tempAlignM * tempAlignN * 4B * 2 > L0C) return false; // 对应理论条件2
return true;
}
```
**逐条对比**
| 理论条件 | 源码对应 | 差异分析 |
|---|---|---|
| 1. b_core 2b | batchC 128b_core 4 | 理论 b₀=2 b_core 4源码 MIN_BATCH_L0=4 batchC 128 = b_core 4。**一致**b₀=2 b_core 4 与源码等价 |
| 2. L0C 容量 | L0 buffer check | 理论2(bM)(bN4B L0C源码tempAlignM × tempAlignN × 4B × 2 L0C。**一致**tempAlignM = b₀×M 对齐 |
| 3. 单核搬移量 480KB | **无** | 源码缺少此条件——可能导致小数据量 case 进入 MergeBatch GM 带宽利用率不足 |
| 4. 搬移 tile 16KB | **无** | 源码缺少此条件—— tile 时搬移效率无保障 |
| 5. AI < R₁₆/b | **无** | 源码缺少算存比约束——计算 Bound case 进入 MergeBatch 后冗余计算成为瓶颈 |
**源码额外条件**alignK 64K 64 M N理论无此限制
**Tiling 参数对比**
| 参数 | 理论 | 源码 |
|---|---|---|
| 合并数 b | b = min(L0C上限, R₁₆上限, b_core) | mergeBatchL0 = min(多项式求解, maxBatchL0, batchNumPerCore) |
| L0 K 粒度 | kL0 = min(L0A/2bM, L0B/2bN) | baseK = min(maxBasek, **64**) |
| L1 K 粒度 | kL1 max(kL0, 256B/dtype) | stepKa × baseK L1/4 容量决定 |
**关键差异**源码 `baseK ≤ 64`硬编码上限理论无此限制64 元素 = 128B BF16 = dValue 下限这是保守的安全下限
**MergeBatch vs IterBatch 分界**
- 理论b_core > b₀(T_comp+T_write)/T_cmdGM→L1 搬移命令节省 vs drain 惩罚)
- 源码无显式分界——MergeBatch 优先级高于 IterBatch先匹配先进
**结论**:源码的 MergeBatch 进入条件**偏宽**——缺少单核搬移量、tile 大小、算存比三条约束。这导致一些本应走 IterBatch 的 case如单核搬移量不足、计算 Bound被错误分配到 MergeBatch。建议补充条件 3/4/5。
---
### 2.3 IterBatch
**理论条件**v0.93 §六)——四种 L1 驻留形态:
- a) b_core=1 ∧ (MK+KN)·dtype ≤ L1单 batch 全驻留)
- b) b_core>1 ∧ 2(MK+KN)·dtype ≤ L1双 batch 乒乓)
- c) 一侧驻留 + 对侧切 K含 b_core≥2 的半预算预取档)
- d) 两侧都切 K
**源码条件**`batch_matmul_v3_iterbatch_basicapi_tiling.cpp` IsCapable
```cpp
bool IsCapable() {
if (非连续转置A) return false;
if (NZ格式) return false;
if (batchBias > 1) return false;
if (BatchA != BatchB) return false;
if (batchC <= aicNum) return false; // B ≤ 32 不进 IterBatch
// L1 双缓冲容量检查(形态 b 等价)
if ((alignM*alignK + alignK*alignN)*dtype*2 > L1) return false;
// L0 能放至少 1 个 batch
if (!l0CanLoadBatch_) {
// 负载均衡检查balanceRate ≥ 0.8
if (balanceRateOfBatch < 0.8) return false;
}
return true;
}
```
**逐条对比**
| 理论形态 | 源码对应 | 差异分析 |
|---|---|---|
| a) 单 batch 全驻留 | **无独立分支** | 源码不区分 a/b——只要 (MK+KN)×dtype×2 ≤ L1 就进 IterBatch。b_core=1 的 case 被 B>32 条件排除B≤32 时 batchC≤aicNum不进 IterBatch |
| b) 双 batch 乒乓 | ✓ 直接对应 | 源码的 L1 检查就是形态 b 的条件 |
| c) 一侧驻留+对侧切 K | **无** | 源码不实现形态 c——当 (MK+KN)×dtype×2 > L1 时直接退出 IterBatch |
| d) 两侧都切 K | **无** | 源码不实现形态 d |
**关键差异**
1. **源码只实现形态 b**,形态 a/c/d 不覆盖。形态 ab_core=1被 B>32 排除——这些 case 由 ASW 或 AL1/BL1_FULL_LOAD 承接。形态 c/dL1 放不下双 batch的 case 落到 ASW_Basic。
2. **B ≤ 32 不进 IterBatch**:源码用 `batchC <= aicNum` 排除了 B ≤ C 的 case。理论上 b_core=1 时形态 a/c 仍可行,但源码将这些 case 交给 ASW/AL1/BL1 分支。
3. **负载均衡**:源码有 `balanceRateOfBatch < 0.8` 检查(对应理论条件 2 的 B mod C ≥ minCoreNum但实现方式不同——源码用实际最大 batch 数与平均值的比率,理论用尾波核数。
**Tiling 参数对比**
| 参数 | 理论 | 源码 |
|---|---|---|
| iterBatchL1L1 级 batch 数) | 由 L1 容量决定 | min(L1容量/batch, mmadCount=8, ceil(B/C)) |
| iterBatchL0L0 级 batch 数) | 由 L0 容量决定 | min(L0A/L0B/L0C容量, iterBatchL1) |
| baseM/baseN/baseK | 由 L0 容量和 dValue 决定 | 由 MatMulV3TilingHelper::ResetBase 统一计算 |
**结论**:源码的 IterBatch **只覆盖形态 b**,且要求 B > C。形态 a/c/d 的缺失意味着:
- B ≤ C 且 L1 放得下的 case形态 a由 AL1/BL1_FULL_LOAD 或 ASW 承接——**合理**,因为这些 case 的并行度不足,需要切 M/N
- L1 放不下双 batch 但放得下单侧驻留的 case形态 c落到 ASW——**可能不是最优**,因为 ASW 引入了 M/N 维的核间重复读,而形态 c 可以避免
---
### 2.4 StreamK
**理论条件**v0.93 §七):
1. P = B·MN·4B/L0C ≤ C/2
2. K/grid_K ≥ 256B/dtypegrid_K = ⌊C/⌈P⌉⌋
3. K > grid_K²/(grid_K-1)·θ_cθ_c ≈ 12
**源码条件**`batch_matmul_v3_basic_streamk_tiling.cpp`
```cpp
bool IsCapable() {
if (deterministicLevel > 1) return false; // 确定性
if (BatchA != BatchB) return false; // batch 一致
if (FP32 && K > 2000000) return false; // FP32 精度
// K 阈值CeilAlign(K,256) ≥ max(8192, aicNum×256B/dtype)
if (CeilAlign(K,256) < max(8192, 32×256/2)) return false;
// 即 K ≥ 8192BF16
// 并行缺口batchC × mCnt × nCnt ≤ aicNum/2
mCnt = ceil(M/256); nCnt = ceil(N/256); // 256B 粒度
if (batchC * mCnt * nCnt > aicNum/2) return false; // B·⌈M/256⌉·⌈N/256⌉ ≤ 16
return true;
}
```
**逐条对比**
| 理论条件 | 源码对应 | 差异分析 |
|---|---|---|
| 1. P ≤ C/2 | B·⌈M/256⌉·⌈N/256⌉ ≤ C/2 | **粒度不同**:理论用 L0C 满载粒度M^t·N^t = L0C/4B = 65536 元素),源码用固定 256 元素粒度。**理论粒度更粗**65536 vs 65536=256²实际上一致——⌈M/256⌉·⌈N/256⌉ 就是 L0C 满载 tile 的个数 |
| 2. K/grid_K ≥ 128 元素 | K ≥ 8192固定阈值 | 源码用固定 8192 = C×256 元素 = dValue 推荐值在最大 grid_K=C 下的保障。**理论更精确**(动态 grid_K**源码更保守**(固定最大 grid_K |
| 3. K > grid_K²/(grid_K-1)·θ_c | **无** | 源码不归约代价建模,直接用固定 K 阈值覆盖。θ_c≈12 时 grid_K=2 只需 K>49远低于 8192——**源码的条件 3 被条件 2 覆盖,过于保守** |
| 4. 确定性 ≤ 1 | deterministicLevel > 1 排除 | 一致 |
**关键差异**
1. **grid_K 取值**:理论 grid_K = ⌊C/⌈P⌉⌋动态源码用固定阈值隐式假设 grid_K = C最大。这导致源码的 K 阈值8192远大于理论阈值grid_K=2 时 K>49
2. **StreamK 优先级过高**:源码中 StreamK 优先级2高于 MergeBatch3和 IterBatch5。虽然 StreamK 的 IsCapable 条件实际上排除了 B>C 的 case但如果 B≤C 且 K≥8192 且 P≤C/2源码会优先进入 StreamK 而非 ASW_Basic。这在理论上是正确的P<C/2 时切 K 是唯一剩余的并行维度)。
**结论**源码的 StreamK 实现**基本合理但过于保守**。固定 K8192 阈值排除了大量理论上可走 StreamK caseK 49~8192 之间)。建议改为动态计算 grid_K K 阈值
---
### 2.5 ASW_Basic
**理论条件**v0.93 §
1. P = B·MN·4B/L0C C
2. batch 结构限制交叉广播均可
3. 兜底B C IterBatch/MergeBatch 条件不满足
**源码条件**`batch_matmul_v3_asw_basic_tiling.cpp` IsCapable
```cpp
bool IsCapable() {
if (A/B 非连续转置不匹配) return false;
if (BatchA != BatchB) return false;
if (batchBias > 1) return false;
if (!FP16/BF16) return false; // 类型限制
return true; // 无其他条件——兜底
}
```
**对比**源码的 ASW_Basic **几乎没有进入条件**——它是优先级最低的实质分支priority 9任何不被前面分支捕获的 case 都落入 ASW这与理论的"兜底"定位一致
**Tiling 实现对比**
| 维度 | 理论 | 源码 |
|---|---|---|
| 核间切分 | B M N 混合切 | MatMulV3TilingHelper::ResetBase 统一计算 |
| swizzle | ASW 滑窗蛇形W=4 for C=32 | GetRebalanceBlock + CalL1Tiling |
| L2 切分 | 场景 A/B/C 分场景决策 | MatMulV3TilingHelper 内部处理 |
| L1 buffer | 双缓冲/4-buffer 按容量决定 | l1BufferNum = 4 if fits else 2 |
**结论**ASW_Basic 作为兜底分支源码实现合理理论和源码在"兜底"定位上一致差异主要在 tiling 参数的具体计算方法——源码用统一的 Helper 函数理论给出了更细粒度的 swizzle/L2 切分策略
---
### 2.6 ASW L1 全载特化
**理论**v0.93 §.5单边无 batch 且该侧矩阵小M 256MK·dtype·2 L1对侧每核循环 4 时小侧整个常驻 L1
**源码**`batch_matmul_v3_asw_al1_full_load_basic_tiling.cpp` IsCapable
```cpp
bool IsCapable() {
if (非连续转置) return false;
if (!FP16/BF16) return false;
if (batchA > 1) return false; // A 无 batch
if (M > 256) return false; // M ≤ 256
if (batchBias > 1) return false;
if (A数据量×2 > L1) return false; // A 放得下 L1/2
// B 数据量足够:每核至少 4 轮
if (B总数据量 < L1×aicNum && B轮数 < 4×aicNum) return false;
return true;
}
```
**对比**理论与源码**高度一致**——M 256batchA=1、A 驻留 L1B 每核 4 这是理论和源码最吻合的分支
---
### 2.7 降核 ASW
**理论**v0.93 §.3/6P < C 且不满足 StreamK 只用 P 其余核闲置
**源码****无独立的降核分支**。ASW_Basic `GetNumBlocks()` 返回 `aicNum`固定 32 不降核
**差异**理论上 P < C 时应该降核只用 P 但源码始终用全部 32 对于 P < C case部分核分到的输出块为空idle但仍参与调度这可能引入不必要的调度开销
**建议**源码应在 ASW_Basic 中加入降核逻辑—— P < C `usedCoreNum = ⌈P⌉`
---
## 三、总体结论
### 源码实现评价
| 维度 | 评价 |
|---|---|
| 分支覆盖完备性 | 7 个分支覆盖所有 case无空洞 |
| 分支选择正确性 | 优先级有序匹配基本合理 StreamK 优先级偏高实际不影响 |
| MergeBatch 条件 | 偏宽——缺少搬移量/tile/算存比约束 |
| IterBatch 形态覆盖 | 只覆盖形态 b缺少 a/c/d |
| StreamK K 阈值 | 固定 8192 过于保守 |
| 降核 ASW | 未实现 |
| L1 全载特化 | 与理论高度一致 |
| 特殊分支 | 与理论一致 |
### 理论最优方案评价
理论分析的 7 大分支体系在逻辑上完备15 种切分组合 按代价排序 6+1 分支进入条件全部为 B/M/N/K 闭式表达式与源码对比后确认
1. **理论分支体系是最优的候选集**——源码的 10 tiling 模板可以映射到理论的 7 大分支上无遗漏
2. **理论条件更精确**——MergeBatch 5 条件StreamK 3 条件都比源码更严格/精确
3. **理论缺的工程细节**——源码中的负载均衡检查balanceRate 0.8)、格式限制NZ/ND)、bias 限制等在理论中未建模
### 源码改进建议(按优先级)
1. **MergeBatch 补条件 3/4/5**搬移量/tile/算存比避免不应进入的 case 被误捕获
2. **IterBatch 补形态 c/d**一侧驻留/两侧切 K覆盖 L1 放不下双 batch case
3. **StreamK 改动态 grid_K**grid_K = ⌊C/⌈P⌉⌋K 阈值随 grid_K 变化而非固定 8192
4. **ASW_Basic 加降核逻辑**P < C usedCoreNum = ⌈P⌉
5. **补转Matmul 分支**BatchA=1 BatchB=1 时折叠为 Matmul避免走 BMM 的冗余路径