代码基线: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 文档),文中所有结论均标注源码文件/函数出处
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 的原因。
| 目录 | 内容 | 关键点 |
|---|---|---|
op_host/op_tiling/batch_mat_mul_v3_tiling.cpp | tiling 总入口、平台信息提取(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/block | 7 字段 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 相关的增量逻辑 |
ExtractMatrixBatchInfo)、batch 相关策略类与 kernel block。理解这一点是理解其分支结构的前提。
BMM 所有分支的进入条件中的魔法数字(256、48KB、64KB、aicNum/2、4×aicNum……)都能在 950PR 的硬件规格中找到根源。950PR 软件编程架构为 NPU 架构版本 351x(DAV_3510),第三代达芬奇架构,AIC/AIV 分离设计。
| 规格项 | 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:2) | StreamK/fixpipe 1V2 分支要求 aivNum == 2*aicNum;K==0/K==1 纯 AIV 分支的核数 |
| Cube 算力 BF16/FP16 | 486 / 425 TFLOPS(含 Vector) | 算存比 ≈ 486TFLOPS / 1.6TB/s ≈ 304 FLOP/B,极高——绝大多数 case 是访存受限,减少 GM 搬运 = 性能,这是 L1 全载/iterbatch 复用类分支的根本动机 |
| 片上内存 | 128GB / 1.6TB/s(降配 112GB/1.4TB/s) | |
| L1 Buffer / 核 | 512KB | L1 全载条件 alignMatSize×2 ≤ l1Size;单次 L1 搬运效率阈值 L1_SINGLE_SIZE_LIMIT=48KB;iterBatchL1 = L1 能驻留的 batch 数 |
| L0A / L0B / 核 | 各 64KB | baseM×baseK、baseN×baseK(×DB)必须 ≤ 64KB;mergebatch 把多 batch 拼进 L0 的容量上界 |
| L0C / 核 | 256KB(较上代增大,白皮书明确动机是"更灵活的 Tiling 策略") | fp32 累加块 baseM×baseN×4B×DB ≤ 256KB;iterbatch 在 L0C 攒多 batch 输出(batchOutNum) |
| UB / 核 | 512KB | TO_MUL/Vector 分支单轮驻留 batch 数 = ubSize/singleBatchSize |
| L2 Cache | 128MB(降配 112MB),512B cacheline | ASW 滑窗/对角错位 swizzle 的复用目标就是 L2 命中 |
| Cube 基本节拍 | 一拍完成 FP16 16×16×16 MAC | baseM/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 后处理)分支的硬件基础 |
| Fixpipe | L0C→GM/UB 随路量化/转置(NZ2ND 等) | L0C2OUT_MODEL = ND_FIXPIPE_1_1 / 1_2 的使能条件 |
| issue queue | 连续 8 次 mmad 会撑爆(经验值 mmadCount=8) | iterbatch 中 iterBatchL1 被 8 截断、iterBatch≤4 时 K 向 DB 减半的依据 |
TilingParse 阶段(TilingPrepareForBatchMatMulV3)从平台信息提取全部硬件参数:aicNum/aivNum、L1/L0A/L0B/L0C/L2/UB 容量、supportL0c2out(fixpipe L0C 直出)、supportL12BtBf16(L1→BT 直通)、btSize(1024/4096),存入 MatmulV3CompileInfo 供各策略使用。
| 优先级 | 策略常量 | 分支名 | 一句话定位 |
|---|---|---|---|
| 0 | BATCH_MATMUL_INPUT_K_EQUAL_ZERO | K==0 清零 | K=0 → 输出零矩阵,AIV 直写 |
| 1 | BATCH_MATMUL_TO_MUL | matmul2mul | K=1 → 退化为向量乘,AIV 计算 |
| 2 | BATCH_STREAM_K | StreamK | K 极大且 MN 并行度不足 → 切 K 填核 |
| 3 | MERGE_BATCH_BASICAPI | mergebatch | M/N 小、batch 巨大 → 多 batch 拼进 L0 大块 |
| 4 | ITER_BATCH_BROADCAST_BASICAPI | iterbatch 广播 | 单边单轴 batch 广播 → 广播算子 L1 驻留一份 |
| 5 | ITER_BATCH_BASICAPI | iterbatch(基础API) | batch 相等且 > 核数 → L1/L0 多 batch 流水 |
| 6 | ITER_BATCH | iterbatch(高阶API) | 同上,自定义搬移(IterateBatch) |
| 7 | AL1_FULL_LOAD_BASIC | A L1 全载 | A 无 batch 且 M≤256 → A 常驻 L1 |
| 8 | BL1_FULL_LOAD_BASIC | B L1 全载 | B 无 batch 且 N≤256 → B 常驻 L1 |
| 9 | ASW_BASIC | ASW 基础API通用 | cubeBound 模型寻优 + 自适应滑窗 |
| 999 | BASE | ASW 高阶API兜底 | 无条件命中,保证任何合法输入可 tiling |
DAV_RESV(s8s4 量化保留平台)只保留 5 个分支:ITER_BATCH_BASICAPI → AL1_FULL_LOAD → BL1_FULL_LOAD → ASW_BASIC → BASE。