Files
matmul-analysis/BMM/BMM算子优化分析_Release/BMM算子优化分析_v0.7_播客文字稿.md

62 lines
11 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters

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 算子最优实现分析 —— 播客文字稿
> 对话者主持人M& 技术专家E
> 时长:约 25 分钟
---
**M**今天我们要聊一份技术文档全称叫《BMM 算子优化分析 v0.7》,是讲昇腾 NPU 上批量矩阵乘怎么做到最优的。说实话我第一眼看到这份文档,里面全是什么 B、M、N、K 的公式,什么 MergeBatch、StreamK什么 L0C、L2 切分看不太懂。所以我把作者请来了咱们从最浅的地方开始一步步把它讲明白。先问一个最基本的问题——BMM 算子到底是干什么的?
**E**BMM 就是 Batch Matrix Multiply带 batch 的矩阵乘。你有一个左矩阵 A形状是 `[Batch, M, K]`,一个右矩阵 B形状是 `[Batch, K, N]`,对 batch 维里的每一个索引,独立做一次矩阵乘,得到一个输出 `[Batch, M, N]`。跟普通矩阵乘的区别就是多了一个 batch 维度——同一套操作,做很多份。比如一个 batch 里有 128 个样本,每个样本都要做同一个矩阵乘,你就得高效地并行完成这 128 次。
**M**:那"最优"怎么定义?是不是跑得越快越好?
**E**:对,但"快"的度量要精确。NPU 上计算不是一条线是好几条流水线同时跑——Cube 单元在做乘加MMAD同时 MTE 引擎在把数据从全局内存搬到 L1MTE2、从 L1 搬到 L0MTE1Fixpipe 在把结果写出去。这些流水线可以并行,但总时延取决于最慢的那一条。所以公式就是 T_total = max(T_MMAD, T_MTE2, T_MTE1, T_Fixpipe)——总时延等于最慢的那一级。优化就是找到瓶颈在哪一级,然后想办法缩短它。
**M**:那如果搬移是瓶颈,计算不是,怎么办?
**E**:这就是这份文档里一个很重要的概念,叫"瓶颈交换"。搬移是瓶颈的时候,你可以故意多算一点——就是让 Cube 多做一点冗余计算但换来的好处是数据搬移变少了或变高效了。反之如果计算是瓶颈你可以牺牲搬移效率换计算效率。这份文档里有两个典型例子MergeBatch 就是"牺牲算力换搬移效率"——把多个 batch 合并成一个大矩阵算算了很多交叉项最后扔掉浪费了算力但搬移效率上去了因为瓶颈本来就在搬移浪费的算力被掩盖了。ASW_Basic 切 M/N 是反过来——"牺牲搬移换并行度"——把矩阵按行或列分给不同核,同一个右矩阵被多个核重复读,多搬了几次,但换来了更多核并行干活。
**M**:好,现在我知道什么叫最优了。那文档里怎么决定一个 case 用哪种方案呢?我看到有 7 个分支。
**E**这个推导过程是整份文档最核心的逻辑。BMM 的实现本质就是把数据切成小块,分给 32 个 AIC 核去算。切分有 4 个维度Bbatch、M、N、K内积维。核间怎么分这 4 个维度,就是第一性问题。文档首先做了一件事——把 4 个维度的任意非空子集全部列出来,一共 15 种切分组合。然后问:每种切法有什么代价?
三个维度代价完全不同:切 B 是免费的——因为 batch 在数学上就是独立的,每个核领几个 batch各干各的零重复读、零依赖、零中间结果。切 M 或 N 是廉价的——共享矩阵会被多核重复读,但如果 128MB 的 L2 Cache 能装下共享数据,重复读的代价大部分被 L2 的 5.2TB/s 带宽吸收了,而不是去走慢得多的 GM 1.6TB/s。切 K 是最贵的——因为要把 K 拆给多个核,每个核算一段,最后要把各段的部分和汇总起来,这就引入了核间通信和归约开销。
**M**:所以整条分支决策树就是"按价格从低到高购买并行度"
**E**:对,就是这个思路。先买免费的 B 维——如果 B 够大≥32切 B 就能填满核。然后看单 batch 的 M×N 大不大:大就逐个 batch 算IterBatch小就把多个 batch 合并成一个大块算MergeBatch。如果 B 不够 32买不满核就加价买廉价的 M/N 维——切行或列,重复读的代价靠 L2 和 swizzle 吸收ASW_Basic。如果 M/N 也买不满B、M、N 都小),才买昂贵的 K 维——切 K付归约代价StreamK。如果 K 也不够大,不值得切,那就认了,只用几个核,剩下的闲置(降核 ASW。加上前面两个特殊通路——K=0 或 K=1 走 AIV 向量核、一侧 batch=1 折叠成普通 Matmul——一共 7 个路径。
**M**:这个逻辑很清晰。但具体到每个分支,怎么判断一个 case 能不能进?我注意到文档里的进入条件全是 B、M、N、K 和芯片参数的公式,没有凭感觉的。
**E**:这正是这份文档严谨的地方。我拿 MergeBatch 举例。MergeBatch 是把多个 batch 合并计算,有浪费,浪费比例是 (b-1)/b。所以第一个条件浪费必须能被掩盖这就要求 case 必须是访存 Bound——瓶颈在搬移算力有余量。文档用算存比公式来表达AI = 2MN/(M+N) < R₁₆/b₀,其中 R₁₆ 是芯片的平衡点 607.5 FLOP/元素这保证了合并 b batch 后算存比仍然低于平衡点算力浪费不被感知
然后是容量约束合并后的输出块 [bM, bN] 必须放得进 L0C256KB)。L0C Cube 累加器的片上缓存放不下就白合并了接着是搬移效率单核搬移总量必须 480KB这是 950PR 实测出来保证 GM 带宽利用率的下限单次搬移的 tile 大小 16KB搬移太小启动开销盖过传输最后是 batch 每核至少分到 2×2=4 batch因为合并最小粒度是 2 batch还要两组才能乒乓流水
**M**每一条条件都有道理 IterBatch 我注意到文档里特别强调了一点——"IterBatch 能否进入不由算存比判定"。
**E**没错这是 v0.6 v0.7 的一个关键修正IterBatch 是逐 batch 看起来简单但有一个隐藏风险如果 L1 放不下单 batch 的完整输入数据M×K K×N 两个矩阵核内就会在 K 循环中重复读——每切一次 M 就要重新读一遍 B每切一次 N 就要重新读一遍 A这个重复读会引入额外的搬移可能把一个本来计算 Bound case 重新拖回访存 Bound所以进入条件不能靠算存比判定必须直接看 L1 能不能装下文档为此设计了四种 L1 驻留形态全驻留 batch 乒乓一侧驻留+对侧切 K两侧都切 K形态之间还有 batch 间流水掩盖的分析——比如多 batch 时一侧驻留的方案必须在 batch 边界把下一 batch 的驻留侧预取进来否则会产生气泡这个分析在 v0.7 里做了合并和简化逻辑比 v0.6 更清晰
**M**StreamK 那个 θ = 109 的推导我看了好几遍才明白能再讲一遍吗
**E**StreamK 是最后一招——BMN 都切不动了只能切 K K 的代价是归约归约的时延到底多大v0.7 按实际实现流程拆成了四段第一段AIC 把部分和通过 fixpipe 写到 workspaceworkspace 驻留 L2 5.2TB/s L2 写口这是很快的第二段AIV L2 读回各部分和 5.2TB/s L2 读口第三段AIV 做向量求和——AIV 是独立硬件有自己的 UB 内存和向量单元64 AIV 聚合做 fp32 加法的速率大约 13.5×10¹² 次每秒第四段结果写回把这四段时间加起来要求每核的计算时间至少是归约时间的 10 才能保证归约不成为新的瓶颈代入 950PR 的数值算出来 K 必须 grid_K² × 109grid_K 是切 K 的份数grid_K=8 K 需要 7.0K有意思的是源码里 StreamK 的固定门槛是 K 8192——就是说 grid_K 取到 8 左右本文的解析门槛和源码的固定门槛是一致的互相印证之前 v0.6 的推导用的是 GM 带宽1.6TB/s算出来 θ=1700比现在高 15 ——那是因为 v0.6 假设部分和走 GM而实际上部分和可以驻留 L2归约的读写成本低得多v0.7 修正了这个模型也更符合实际硬件
**M**ASW_Basic 是占比最大的分支35.2%。它的 swizzle L2 切分我感觉是最难理解的部分
**E**swizzle 的本质其实很简单 M/N 之后同一时刻 32 个核各算一个输出块这些块需要的 A 行和 B 列就是"活跃工作集"。如果按自然顺序分配 32 个块横跨很多行和列活跃工作集很大L2 放不下就掉到 GM 去读重复读的代价就真实发生了swizzle 做的事情就是重新编排输出块的执行顺序让同一波 32 个块尽量落在一个小窗口里——窗口里只有 W A 行块和 C/W B 列带活跃工作集最小W 取多少一波 32 个块的 L2 足迹近似是 (W + C/WK·M^t·dtype均值不等式告诉我们 W 接近 C 时足迹最小加上 W 必须整除 C 才能保证窗口边界不碎波所以 W C ≤√C 的最大因子32 核就是 4
蛇形只在窗口行边界做——上一个窗口的最后一列 B 带和下一个窗口的第一列 B 带是同一条热数据不丢窗内不用蛇形因为窗内所有的 A 行块全程驻留 L2顺序不影响
L2 切分是另一层当整个工作集超过 128MB即使滑窗压缩了同一波内的足迹跨波次还是会有数据被挤出所以要把 M×N 平面切成大块一块的工作集控制在 L2 容量内逐块计算同时还要考虑输出怎么处理——输出如果驻留 L2会挤占输入的空间让输入被挤出增加重复读如果直写 GM会占用共享的 GM 总线文档给出了三个场景的决策核心逻辑是输入能驻留 L2 就优先驻留保住 r_in=1输出只写不读、零复用收益所以优先直写 GM 给输入腾空间只要校验总线不爆
**M**最后问一个我猜很多人会问的问题——这份分析靠谱吗
**E**从几个维度来看第一根基来自源码——CAN N ops-nn 仓库的 batch_mat_mul_v3 算子arch35 路径分支条件一一行对得上差异处文档都标注了并解释了为什么第二硬件参数全部来自 950PR 架构白皮书和 CANN 9.0 编程文档不是拍脑袋第三数学推导闭环——15 种组合完备枚举 价格表 7 分支极小性有反例表格验证每个分支去掉都会有一个 case 失去最优方案第四遍历验证—— B 1 2048M/N/K 1 10240 20736 case按进入条件严格分类7 个分支全部有真实命中无覆盖空洞
需要注意的地方也诚实标注了经验常数480KB16KB 这些 950PR 实测值换芯片要重新标定MergeBatch IterBatch M×N 中段有重叠区归属要由端到端时延模型精确仲裁StreamK θ=109 基于 workspace 驻留 L2 的前提如果工作集太大溢出到 GMθ 会升高到 1700但整体来说这份分析的逻辑体系是完整的推导是自洽的可以指导工程实现
**M**感谢你耐心讲解这份文档和 PPT 都在 Gitea 仓库里感兴趣的朋友可以去看原文
---
*参考文档`BMM算子优化分析_v0.7.md`Gitea: matmul-analysis 仓库 BMM算子优化分析_Release/*