From ced23432d2e0a6db1a03e1eb0c2a5a0802cd44a9 Mon Sep 17 00:00:00 2001 From: admin Date: Thu, 20 Aug 2026 14:40:48 +0000 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=20BMMv3=E7=AE=97=E5=AD=90?= =?UTF-8?q?=E5=88=86=E6=94=AF=E5=AE=9E=E7=8E=B0=E5=88=86=E6=9E=90.html?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- BMMv3算子分支实现分析.html | 176 +++++++++++++++++++++++++++++++++++++ 1 file changed, 176 insertions(+) create mode 100644 BMMv3算子分支实现分析.html diff --git a/BMMv3算子分支实现分析.html b/BMMv3算子分支实现分析.html new file mode 100644 index 0000000..a55c5f2 --- /dev/null +++ b/BMMv3算子分支实现分析.html @@ -0,0 +1,176 @@ + + + + + +BatchMatMulV3(BMM v3)算子分支实现深度分析 —— 面向昇腾950PR(DAV_3510) + + + +
+ +
+

BatchMatMulV3 算子分支实现深度分析

+

代码基线:CANN ops-nn 仓 matmul/batch_mat_mul_v3(arch35 路径) · 目标芯片:昇腾 950PR(NPU 架构 DAV_3510,Atlas 350 加速卡)

+

主题:为高性能覆盖 BMM 全部 case,当前算子划分了哪些分支、为什么是这些分支(系统视角 + 定量依据)、每个分支的 tiling/swizzle 与 kernel 具体实现

+

资料来源:本地昇腾NPU知识库(代码仓完整源码 + 昇腾950 NPU 架构白皮书 + CANN 9.0.0 AscendC 文档),文中所有结论均标注源码文件/函数出处

+
+ + + +
+

1. 算子概览与代码地图

+

BatchMatMulV3 是 CANN ops-nn 仓中批量矩阵乘的第三代实现(对应 aclnnBatchMatMul / aclnnBaddbmm / aclnnAddbmm / aclnnEinsum 等接口),语义为 C[batch, M, N] = A[batch, M, K] × B[batch, K, N] (+bias),batch 维最多 4 级且支持广播。与单矩阵乘 MatMulV3 相比,BMM 多出的核心复杂度全部来自 batch 维的处理策略——这正是其分支数量远多于 MatMulV3 的原因。

+ +

1.1 目录结构(与分支的对应关系)

+ + + + + + + + +
目录内容关键点
op_host/op_tiling/batch_mat_mul_v3_tiling.cpptiling 总入口、平台信息提取(TilingParse)arch35 / 老架构分流
op_host/op_tiling/arch35/DAV_3510(昇腾950)高级 tiling:1 个策略表 + 11 个策略类本文分析主体
op_host/op_tiling/batch_mat_mul_v3_base_tiling.cpp老架构(910B 等)基类 tiling,67KB顺序执行 + flag 覆盖模式
op_kernel/arch35/arch35 kernel 入口 + 各策略 kernel/block7 字段 tilingKey 编译期分发
op_kernel/batch_mat_mul_v3*.h老架构通用 kernel(Common/UnAligned/MultiBatch/Vector)5 字段 tilingKey
matmul/mat_mul_v3/matmul/common/cmct/blaze/被大量复用的公共基类:MatmulImpl 高阶 API、Cmct GEMM 框架、Blaze GEMM 框架BMM 只写 batch 相关的增量逻辑
+ +
+设计观察:BMM v3 的 arch35 实现本质上是「MatMulV3 的 tiling/kernel 骨架」+「batch 维策略层」。策略类注册表(MMTilingRegistry)、基本块寻优(GetRebalanceBlock)、L1 tiling 计算(CalL1Tiling)、MatmulImpl 主流水全部复用 MatMulV3;BMM 自己只新增 batch 信息提取(ExtractMatrixBatchInfo)、batch 相关策略类与 kernel block。理解这一点是理解其分支结构的前提。 +
+
+ +
+

2. 硬件基础:950PR 微架构规格与 tiling 约束

+

BMM 所有分支的进入条件中的魔法数字(256、48KB、64KB、aicNum/2、4×aicNum……)都能在 950PR 的硬件规格中找到根源。950PR 软件编程架构为 NPU 架构版本 351x(DAV_3510),第三代达芬奇架构,AIC/AIV 分离设计。

+ +

2.1 关键规格表(昇腾950 NPU 架构白皮书 表3-1/表4-2)

+ + + + + + + + + + + + + + + + +
规格项950PR 数值对 BMM tiling 的直接约束
AIC(Cube Core)数32(满配)/ 28(降配)并行度基准:batch×mCnt×nCnt 需 ≥ aicNum 才能填满核;aicNum×2、aicNum/2、4×aicNum 等阈值由此来
AIV(Vector Core)数64 / 56(AIC:AIV = 1:2StreamK/fixpipe 1V2 分支要求 aivNum == 2*aicNum;K==0/K==1 纯 AIV 分支的核数
Cube 算力 BF16/FP16486 / 425 TFLOPS(含 Vector)算存比 ≈ 486TFLOPS / 1.6TB/s ≈ 304 FLOP/B,极高——绝大多数 case 是访存受限,减少 GM 搬运 = 性能,这是 L1 全载/iterbatch 复用类分支的根本动机
片上内存128GB / 1.6TB/s(降配 112GB/1.4TB/s)
L1 Buffer / 核512KBL1 全载条件 alignMatSize×2 ≤ l1Size;单次 L1 搬运效率阈值 L1_SINGLE_SIZE_LIMIT=48KB;iterBatchL1 = L1 能驻留的 batch 数
L0A / L0B / 核64KBbaseM×baseK、baseN×baseK(×DB)必须 ≤ 64KB;mergebatch 把多 batch 拼进 L0 的容量上界
L0C / 核256KB(较上代增大,白皮书明确动机是"更灵活的 Tiling 策略")fp32 累加块 baseM×baseN×4B×DB ≤ 256KB;iterbatch 在 L0C 攒多 batch 输出(batchOutNum)
UB / 核512KBTO_MUL/Vector 分支单轮驻留 batch 数 = ubSize/singleBatchSize
L2 Cache128MB(降配 112MB),512B cachelineASW 滑窗/对角错位 swizzle 的复用目标就是 L2 命中
Cube 基本节拍一拍完成 FP16 16×16×16 MACbaseM/baseN/baseK 以 16 对齐;M/N 过小时 mmad 粒度浪费 → mergebatch 分支动机
分形格式L0A=FRACTAL_NZ,L0B=FRACTAL_ZN,L0C=FRACTAL_NZ;L0A/B 需 512B 对齐,K 内轴 128B/256B 对齐各分支对齐常数(16、128B/dtype、256B/dtype)的来源
MTE 通路变化(351x)删除 GM→L0 直通;新增 L0C→UB、UB↔L1(CV 硬通道)、SSBuffer 核间通信ND_FIXPIPE_1_2(1 AIC + 2 AIV fixpipe 后处理)分支的硬件基础
FixpipeL0C→GM/UB 随路量化/转置(NZ2ND 等)L0C2OUT_MODEL = ND_FIXPIPE_1_1 / 1_2 的使能条件
issue queue连续 8 次 mmad 会撑爆(经验值 mmadCount=8)iterbatch 中 iterBatchL1 被 8 截断、iterBatch≤4 时 K 向 DB 减半的依据
+ +
+950PR 的产品定位强化了这些分支的价值:950PR 面向 LLM Prefill/推荐等计算受限场景,带宽(1.6TB/s)显著低于 950DT(4TB/s)而算力接近(486 vs 547 TFLOPS),算存比更高。这意味着在 PR 上,任何能减少 HBM/L2 搬运的分支(L1 全载、batch 复用、滑窗 L2 复用)收益都被放大。LLM 推理中的典型 BMM(attention 的 QK^T/AV:大 batch、中小 M/N、K 为 head_dim)恰好落在 iterbatch / mergebatch / 全载分支的甜蜜区。 +
+
+ +
+

3. Tiling 总体框架:入口、分流与短路遍历

+ +

3.1 主调用链

+
BatchMatMulV3TilingFunc // op_tiling/batch_mat_mul_v3_tiling.cpp:39 +├─ IsAdvancedSocVersion(context)? // NpuArch ∈ {DAV_3510, DAV_RESV} → arch35 高级路径 +│ └─ batch_matmul_v3_advanced::BatchMatMulV3Tiling(context).DoTiling() +│ ├─ GetShapeAttrsInfo / CheckArgs / GetArgs // 格式、dtype、transpose、M/K/N 提取校验(复用 MatMulV3) +│ ├─ ExtractMatrixBatchInfo() // BMM 特有:提取 batchA0~A3/B0~B3/C0~C3(≤4 级 batch,总维数≤6) +│ ├─ ValidateMatrixBatchInfo() // 广播合法性:对应位相等或其一为 1 +│ │ └─ MergeBatchAndMAxis() // 关键优化:batchB==1 且 A 不转置 → batch 折叠进 M 轴,退化为 MatMul +│ └─ MMTilingRegistry::DoTilingImpl(priorities) // 按优先级表逐分支尝试 +│ └─ for priority in priorities: +│ ├─ 构造策略类 → DoTiling() // 模板方法:GetShapeAttrsInfo → IsCapable → DoOpTiling → Adjust → Post +│ ├─ GRAPH_SUCCESS → 直接 return(短路) // 第一个命中的分支生效 +│ └─ GRAPH_PARAM_INVALID(IsCapable==false)→ 继续下一个 +└─ 否则走老路径 TilingRegistry → BatchMatmulV3BaseTiling(见 §9)
+ +

TilingParse 阶段(TilingPrepareForBatchMatMulV3)从平台信息提取全部硬件参数:aicNum/aivNum、L1/L0A/L0B/L0C/L2/UB 容量、supportL0c2out(fixpipe L0C 直出)、supportL12BtBf16(L1→BT 直通)、btSize(1024/4096),存入 MatmulV3CompileInfo 供各策略使用。

+ +

3.2 策略优先级表(batch_matmul_v3_tiling_strategy.h)

+ + + + + + + + + + + + + +
优先级策略常量分支名一句话定位
0BATCH_MATMUL_INPUT_K_EQUAL_ZEROK==0 清零K=0 → 输出零矩阵,AIV 直写
1BATCH_MATMUL_TO_MULmatmul2mulK=1 → 退化为向量乘,AIV 计算
2BATCH_STREAM_KStreamKK 极大且 MN 并行度不足 → 切 K 填核
3MERGE_BATCH_BASICAPImergebatchM/N 小、batch 巨大 → 多 batch 拼进 L0 大块
4ITER_BATCH_BROADCAST_BASICAPIiterbatch 广播单边单轴 batch 广播 → 广播算子 L1 驻留一份
5ITER_BATCH_BASICAPIiterbatch(基础API)batch 相等且 > 核数 → L1/L0 多 batch 流水
6ITER_BATCHiterbatch(高阶API)同上,自定义搬移(IterateBatch)
7AL1_FULL_LOAD_BASICA L1 全载A 无 batch 且 M≤256 → A 常驻 L1
8BL1_FULL_LOAD_BASICB L1 全载B 无 batch 且 N≤256 → B 常驻 L1
9ASW_BASICASW 基础API通用cubeBound 模型寻优 + 自适应滑窗
999BASEASW 高阶API兜底无条件命中,保证任何合法输入可 tiling
+

DAV_RESV(s8s4 量化保留平台)只保留 5 个分支:ITER_BATCH_BASICAPI → AL1_FULL_LOAD → BL1_FULL_LOAD → ASW_BASIC → BASE

+ +
+为什么按这个顺序短路?顺序 = "对计算模式的改变程度"从大到小: +
    +
  1. K 退化特判最靠前(K==0、K==1):计算模式彻底改变(不需要 Cube),必须先拦截,否则会被后面的 cube 模板接住造成数量级浪费;
  2. +
  3. StreamK 次之:它改变的是全局并行结构(K 维跨核拆分 + workspace 归约),要在 batch 维优化之前决定;
  4. +
  5. batch 维优化居中(mergebatch → iterbatch 广播 → iterbatch×2):只改变 batch 维的调度与驻留方式,不改变单 batch 的计算模式;
  6. +
  7. L1 全载靠后:只改变单侧操作数的搬移次数,是局部优化;
  8. +
  9. 通用 ASW 兜底:BASE 用 999 保证永远最后尝试、必然命中——这是分支体系"完备性"的工程保证。
  10. +
+
+