移动后清理旧路径 BMMv3算子分支实现分析.html
This commit is contained in:
@@ -1,176 +0,0 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>BatchMatMulV3(BMM v3)算子分支实现深度分析 —— 面向昇腾950PR(DAV_3510)</title>
|
||||
<style>
|
||||
:root{
|
||||
--ink:#1a2233; --muted:#5a6478; --accent:#0b5fff; --accent2:#7c3aed;
|
||||
--bg:#f6f8fc; --card:#ffffff; --line:#e3e8f2; --code-bg:#0f172a; --code-ink:#e2e8f0;
|
||||
--ok:#0f9d58; --warn:#b45309; --bad:#b91c1c;
|
||||
}
|
||||
*{box-sizing:border-box}
|
||||
body{margin:0;font-family:"PingFang SC","Microsoft YaHei","Helvetica Neue",Arial,sans-serif;
|
||||
color:var(--ink);background:var(--bg);line-height:1.75;font-size:15.5px}
|
||||
.wrap{max-width:1180px;margin:0 auto;padding:32px 28px 96px}
|
||||
header.hero{background:linear-gradient(135deg,#0b2a6b 0%,#0b5fff 55%,#7c3aed 100%);
|
||||
color:#fff;border-radius:14px;padding:38px 40px;margin-bottom:28px;box-shadow:0 8px 30px rgba(11,95,255,.18)}
|
||||
header.hero h1{margin:0 0 10px;font-size:26px;letter-spacing:.5px}
|
||||
header.hero p{margin:4px 0;opacity:.92;font-size:14.5px}
|
||||
nav.toc{background:var(--card);border:1px solid var(--line);border-radius:12px;padding:22px 28px;margin-bottom:28px}
|
||||
nav.toc h2{margin:0 0 10px;font-size:17px}
|
||||
nav.toc ol{margin:0;padding-left:22px;columns:2;column-gap:48px}
|
||||
nav.toc li{margin:3px 0;font-size:14.5px}
|
||||
nav.toc a{color:var(--accent);text-decoration:none}
|
||||
nav.toc a:hover{text-decoration:underline}
|
||||
section{background:var(--card);border:1px solid var(--line);border-radius:12px;padding:28px 34px;margin-bottom:26px}
|
||||
h2{font-size:21px;border-left:5px solid var(--accent);padding-left:12px;margin:6px 0 18px}
|
||||
h3{font-size:17.5px;color:#0b2a6b;margin:26px 0 10px}
|
||||
h4{font-size:15.5px;color:var(--accent2);margin:18px 0 8px}
|
||||
table{border-collapse:collapse;width:100%;margin:14px 0;font-size:14px}
|
||||
th,td{border:1px solid var(--line);padding:8px 11px;text-align:left;vertical-align:top}
|
||||
th{background:#eef3ff;color:#0b2a6b;white-space:nowrap}
|
||||
tr:nth-child(even) td{background:#fafbfe}
|
||||
code,kbd{font-family:"JetBrains Mono","SF Mono",Consolas,monospace;font-size:13px;
|
||||
background:#eef1f8;border-radius:4px;padding:1px 6px;color:#0b2a6b}
|
||||
pre{background:var(--code-bg);color:var(--code-ink);border-radius:10px;padding:16px 18px;
|
||||
overflow-x:auto;font-size:13px;line-height:1.6}
|
||||
pre code{background:none;color:inherit;padding:0}
|
||||
.tag{display:inline-block;border-radius:20px;padding:1px 12px;font-size:12.5px;font-weight:600;margin-right:6px;white-space:nowrap}
|
||||
.tag.aic{background:#e0f2e9;color:var(--ok)} .tag.aiv{background:#fdeaea;color:var(--bad)}
|
||||
.tag.mix{background:#fff3e0;color:var(--warn)} .tag.br{background:#e8edff;color:var(--accent)}
|
||||
.callout{border-left:4px solid var(--accent);background:#f0f5ff;border-radius:0 8px 8px 0;padding:12px 18px;margin:14px 0}
|
||||
.callout.warn{border-color:var(--warn);background:#fff8ee}
|
||||
.callout.crit{border-color:var(--accent2);background:#f6f0ff}
|
||||
.branch{border:1px solid var(--line);border-radius:10px;padding:18px 22px;margin:16px 0;background:#fdfeff}
|
||||
.branch h4{margin-top:0}
|
||||
.src{color:var(--muted);font-size:12.5px}
|
||||
ul.tight li,ol.tight li{margin:4px 0}
|
||||
.flow{font-family:"JetBrains Mono",Consolas,monospace;font-size:13px;background:#0f172a;color:#dbeafe;
|
||||
border-radius:10px;padding:18px 20px;overflow-x:auto;line-height:1.55;white-space:pre}
|
||||
footer{color:var(--muted);font-size:13px;text-align:center;margin-top:30px}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="wrap">
|
||||
|
||||
<header class="hero">
|
||||
<h1>BatchMatMulV3 算子分支实现深度分析</h1>
|
||||
<p>代码基线:CANN ops-nn 仓 <code style="background:rgba(255,255,255,.15);color:#fff">matmul/batch_mat_mul_v3</code>(arch35 路径) · 目标芯片:昇腾 950PR(NPU 架构 DAV_3510,Atlas 350 加速卡)</p>
|
||||
<p>主题:为高性能覆盖 BMM 全部 case,当前算子划分了哪些分支、为什么是这些分支(系统视角 + 定量依据)、每个分支的 tiling/swizzle 与 kernel 具体实现</p>
|
||||
<p>资料来源:本地昇腾NPU知识库(代码仓完整源码 + 昇腾950 NPU 架构白皮书 + CANN 9.0.0 AscendC 文档),文中所有结论均标注源码文件/函数出处</p>
|
||||
</header>
|
||||
|
||||
<nav class="toc">
|
||||
<h2>目录</h2>
|
||||
<ol>
|
||||
<li><a href="#sec1">算子概览与代码地图</a></li>
|
||||
<li><a href="#sec2">硬件基础:950PR 微架构规格与 tiling 约束</a></li>
|
||||
<li><a href="#sec3">Tiling 总体框架:入口、分流与短路遍历</a></li>
|
||||
<li><a href="#sec4">分支全景:为什么是这些分支(case 空间论证)</a></li>
|
||||
<li><a href="#sec5">逐分支详解(11 个分支)</a></li>
|
||||
<li><a href="#sec6">Swizzle 专题:核间分块顺序与块到核映射</a></li>
|
||||
<li><a href="#sec7">tiling_key 编码与 kernel 分发树</a></li>
|
||||
<li><a href="#sec8">完备性论证与批判性评估</a></li>
|
||||
<li><a href="#sec9">附:非 arch35 老路径简述</a></li>
|
||||
<li><a href="#sec10">参考来源清单</a></li>
|
||||
</ol>
|
||||
</nav>
|
||||
|
||||
<section id="sec1">
|
||||
<h2>1. 算子概览与代码地图</h2>
|
||||
<p>BatchMatMulV3 是 CANN ops-nn 仓中批量矩阵乘的第三代实现(对应 aclnnBatchMatMul / aclnnBaddbmm / aclnnAddbmm / aclnnEinsum 等接口),语义为 <code>C[batch, M, N] = A[batch, M, K] × B[batch, K, N] (+bias)</code>,batch 维最多 4 级且支持广播。与单矩阵乘 MatMulV3 相比,BMM 多出的核心复杂度全部来自 <b>batch 维的处理策略</b>——这正是其分支数量远多于 MatMulV3 的原因。</p>
|
||||
|
||||
<h3>1.1 目录结构(与分支的对应关系)</h3>
|
||||
<table>
|
||||
<tr><th>目录</th><th>内容</th><th>关键点</th></tr>
|
||||
<tr><td><code>op_host/op_tiling/batch_mat_mul_v3_tiling.cpp</code></td><td>tiling 总入口、平台信息提取(TilingParse)</td><td>arch35 / 老架构分流</td></tr>
|
||||
<tr><td><code>op_host/op_tiling/arch35/</code></td><td>DAV_3510(昇腾950)高级 tiling:1 个策略表 + 11 个策略类</td><td><b>本文分析主体</b></td></tr>
|
||||
<tr><td><code>op_host/op_tiling/batch_mat_mul_v3_base_tiling.cpp</code></td><td>老架构(910B 等)基类 tiling,67KB</td><td>顺序执行 + flag 覆盖模式</td></tr>
|
||||
<tr><td><code>op_kernel/arch35/</code></td><td>arch35 kernel 入口 + 各策略 kernel/block</td><td>7 字段 tilingKey 编译期分发</td></tr>
|
||||
<tr><td><code>op_kernel/batch_mat_mul_v3*.h</code></td><td>老架构通用 kernel(Common/UnAligned/MultiBatch/Vector)</td><td>5 字段 tilingKey</td></tr>
|
||||
<tr><td><code>matmul/mat_mul_v3/</code>、<code>matmul/common/cmct/</code>、<code>blaze/</code></td><td>被大量复用的公共基类:MatmulImpl 高阶 API、Cmct GEMM 框架、Blaze GEMM 框架</td><td>BMM 只写 batch 相关的增量逻辑</td></tr>
|
||||
</table>
|
||||
|
||||
<div class="callout">
|
||||
<b>设计观察:</b>BMM v3 的 arch35 实现本质上是「MatMulV3 的 tiling/kernel 骨架」+「batch 维策略层」。策略类注册表(MMTilingRegistry)、基本块寻优(GetRebalanceBlock)、L1 tiling 计算(CalL1Tiling)、MatmulImpl 主流水全部复用 MatMulV3;BMM 自己只新增 batch 信息提取(<code>ExtractMatrixBatchInfo</code>)、batch 相关策略类与 kernel block。理解这一点是理解其分支结构的前提。
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section id="sec2">
|
||||
<h2>2. 硬件基础:950PR 微架构规格与 tiling 约束</h2>
|
||||
<p>BMM 所有分支的进入条件中的魔法数字(256、48KB、64KB、aicNum/2、4×aicNum……)都能在 950PR 的硬件规格中找到根源。950PR 软件编程架构为 <b>NPU 架构版本 351x(DAV_3510)</b>,第三代达芬奇架构,AIC/AIV 分离设计。</p>
|
||||
|
||||
<h3>2.1 关键规格表(昇腾950 NPU 架构白皮书 表3-1/表4-2)</h3>
|
||||
<table>
|
||||
<tr><th>规格项</th><th>950PR 数值</th><th>对 BMM tiling 的直接约束</th></tr>
|
||||
<tr><td>AIC(Cube Core)数</td><td>32(满配)/ 28(降配)</td><td>并行度基准:batch×mCnt×nCnt 需 ≥ aicNum 才能填满核;aicNum×2、aicNum/2、4×aicNum 等阈值由此来</td></tr>
|
||||
<tr><td>AIV(Vector Core)数</td><td>64 / 56(<b>AIC:AIV = 1:2</b>)</td><td>StreamK/fixpipe 1V2 分支要求 <code>aivNum == 2*aicNum</code>;K==0/K==1 纯 AIV 分支的核数</td></tr>
|
||||
<tr><td>Cube 算力 BF16/FP16</td><td>486 / 425 TFLOPS(含 Vector)</td><td rowspan="2">算存比 ≈ 486TFLOPS / 1.6TB/s ≈ 304 FLOP/B,<b>极高</b>——绝大多数 case 是访存受限,减少 GM 搬运 = 性能,这是 L1 全载/iterbatch 复用类分支的根本动机</td></tr>
|
||||
<tr><td>片上内存</td><td>128GB / 1.6TB/s(降配 112GB/1.4TB/s)</td></tr>
|
||||
<tr><td>L1 Buffer / 核</td><td><b>512KB</b></td><td>L1 全载条件 <code>alignMatSize×2 ≤ l1Size</code>;单次 L1 搬运效率阈值 <code>L1_SINGLE_SIZE_LIMIT=48KB</code>;iterBatchL1 = L1 能驻留的 batch 数</td></tr>
|
||||
<tr><td>L0A / L0B / 核</td><td>各 <b>64KB</b></td><td>baseM×baseK、baseN×baseK(×DB)必须 ≤ 64KB;mergebatch 把多 batch 拼进 L0 的容量上界</td></tr>
|
||||
<tr><td>L0C / 核</td><td><b>256KB</b>(较上代增大,白皮书明确动机是"更灵活的 Tiling 策略")</td><td>fp32 累加块 baseM×baseN×4B×DB ≤ 256KB;iterbatch 在 L0C 攒多 batch 输出(batchOutNum)</td></tr>
|
||||
<tr><td>UB / 核</td><td><b>512KB</b></td><td>TO_MUL/Vector 分支单轮驻留 batch 数 = ubSize/singleBatchSize</td></tr>
|
||||
<tr><td>L2 Cache</td><td>128MB(降配 112MB),512B cacheline</td><td>ASW 滑窗/对角错位 swizzle 的复用目标就是 L2 命中</td></tr>
|
||||
<tr><td>Cube 基本节拍</td><td>一拍完成 FP16 16×16×16 MAC</td><td>baseM/baseN/baseK 以 16 对齐;M/N 过小时 mmad 粒度浪费 → mergebatch 分支动机</td></tr>
|
||||
<tr><td>分形格式</td><td>L0A=FRACTAL_NZ,L0B=FRACTAL_ZN,L0C=FRACTAL_NZ;L0A/B 需 512B 对齐,K 内轴 128B/256B 对齐</td><td>各分支对齐常数(16、128B/dtype、256B/dtype)的来源</td></tr>
|
||||
<tr><td>MTE 通路变化(351x)</td><td>删除 GM→L0 直通;新增 L0C→UB、UB↔L1(CV 硬通道)、SSBuffer 核间通信</td><td>ND_FIXPIPE_1_2(1 AIC + 2 AIV fixpipe 后处理)分支的硬件基础</td></tr>
|
||||
<tr><td>Fixpipe</td><td>L0C→GM/UB 随路量化/转置(NZ2ND 等)</td><td>L0C2OUT_MODEL = ND_FIXPIPE_1_1 / 1_2 的使能条件</td></tr>
|
||||
<tr><td>issue queue</td><td>连续 8 次 mmad 会撑爆(经验值 mmadCount=8)</td><td>iterbatch 中 iterBatchL1 被 8 截断、iterBatch≤4 时 K 向 DB 减半的依据</td></tr>
|
||||
</table>
|
||||
|
||||
<div class="callout warn">
|
||||
<b>950PR 的产品定位强化了这些分支的价值:</b>950PR 面向 LLM Prefill/推荐等计算受限场景,带宽(1.6TB/s)显著低于 950DT(4TB/s)而算力接近(486 vs 547 TFLOPS),算存比更高。这意味着在 PR 上,<b>任何能减少 HBM/L2 搬运的分支(L1 全载、batch 复用、滑窗 L2 复用)收益都被放大</b>。LLM 推理中的典型 BMM(attention 的 QK^T/AV:大 batch、中小 M/N、K 为 head_dim)恰好落在 iterbatch / mergebatch / 全载分支的甜蜜区。
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section id="sec3">
|
||||
<h2>3. Tiling 总体框架:入口、分流与短路遍历</h2>
|
||||
|
||||
<h3>3.1 主调用链</h3>
|
||||
<div class="flow">BatchMatMulV3TilingFunc <span style="color:#7dd3fc">// op_tiling/batch_mat_mul_v3_tiling.cpp:39</span>
|
||||
├─ IsAdvancedSocVersion(context)? <span style="color:#7dd3fc">// NpuArch ∈ {DAV_3510, DAV_RESV} → arch35 高级路径</span>
|
||||
│ └─ batch_matmul_v3_advanced::BatchMatMulV3Tiling(context).DoTiling()
|
||||
│ ├─ GetShapeAttrsInfo / CheckArgs / GetArgs <span style="color:#7dd3fc">// 格式、dtype、transpose、M/K/N 提取校验(复用 MatMulV3)</span>
|
||||
│ ├─ ExtractMatrixBatchInfo() <span style="color:#7dd3fc">// BMM 特有:提取 batchA0~A3/B0~B3/C0~C3(≤4 级 batch,总维数≤6)</span>
|
||||
│ ├─ ValidateMatrixBatchInfo() <span style="color:#7dd3fc">// 广播合法性:对应位相等或其一为 1</span>
|
||||
│ │ └─ MergeBatchAndMAxis() <span style="color:#7dd3fc">// 关键优化:batchB==1 且 A 不转置 → batch 折叠进 M 轴,退化为 MatMul</span>
|
||||
│ └─ MMTilingRegistry::DoTilingImpl(priorities) <span style="color:#7dd3fc">// 按优先级表逐分支尝试</span>
|
||||
│ └─ for priority in priorities:
|
||||
│ ├─ 构造策略类 → DoTiling() <span style="color:#7dd3fc">// 模板方法:GetShapeAttrsInfo → IsCapable → DoOpTiling → Adjust → Post</span>
|
||||
│ ├─ GRAPH_SUCCESS → 直接 return(短路) <span style="color:#7dd3fc">// 第一个命中的分支生效</span>
|
||||
│ └─ GRAPH_PARAM_INVALID(IsCapable==false)→ 继续下一个
|
||||
└─ 否则走老路径 TilingRegistry → BatchMatmulV3BaseTiling(见 §9)</div>
|
||||
|
||||
<p>TilingParse 阶段(<code>TilingPrepareForBatchMatMulV3</code>)从平台信息提取全部硬件参数:aicNum/aivNum、L1/L0A/L0B/L0C/L2/UB 容量、<code>supportL0c2out</code>(fixpipe L0C 直出)、<code>supportL12BtBf16</code>(L1→BT 直通)、btSize(1024/4096),存入 <code>MatmulV3CompileInfo</code> 供各策略使用。</p>
|
||||
|
||||
<h3>3.2 策略优先级表(batch_matmul_v3_tiling_strategy.h)</h3>
|
||||
<table>
|
||||
<tr><th>优先级</th><th>策略常量</th><th>分支名</th><th>一句话定位</th></tr>
|
||||
<tr><td>0</td><td><code>BATCH_MATMUL_INPUT_K_EQUAL_ZERO</code></td><td>K==0 清零</td><td>K=0 → 输出零矩阵,AIV 直写</td></tr>
|
||||
<tr><td>1</td><td><code>BATCH_MATMUL_TO_MUL</code></td><td>matmul2mul</td><td>K=1 → 退化为向量乘,AIV 计算</td></tr>
|
||||
<tr><td>2</td><td><code>BATCH_STREAM_K</code></td><td>StreamK</td><td>K 极大且 MN 并行度不足 → 切 K 填核</td></tr>
|
||||
<tr><td>3</td><td><code>MERGE_BATCH_BASICAPI</code></td><td>mergebatch</td><td>M/N 小、batch 巨大 → 多 batch 拼进 L0 大块</td></tr>
|
||||
<tr><td>4</td><td><code>ITER_BATCH_BROADCAST_BASICAPI</code></td><td>iterbatch 广播</td><td>单边单轴 batch 广播 → 广播算子 L1 驻留一份</td></tr>
|
||||
<tr><td>5</td><td><code>ITER_BATCH_BASICAPI</code></td><td>iterbatch(基础API)</td><td>batch 相等且 > 核数 → L1/L0 多 batch 流水</td></tr>
|
||||
<tr><td>6</td><td><code>ITER_BATCH</code></td><td>iterbatch(高阶API)</td><td>同上,自定义搬移(IterateBatch)</td></tr>
|
||||
<tr><td>7</td><td><code>AL1_FULL_LOAD_BASIC</code></td><td>A L1 全载</td><td>A 无 batch 且 M≤256 → A 常驻 L1</td></tr>
|
||||
<tr><td>8</td><td><code>BL1_FULL_LOAD_BASIC</code></td><td>B L1 全载</td><td>B 无 batch 且 N≤256 → B 常驻 L1</td></tr>
|
||||
<tr><td>9</td><td><code>ASW_BASIC</code></td><td>ASW 基础API通用</td><td>cubeBound 模型寻优 + 自适应滑窗</td></tr>
|
||||
<tr><td>999</td><td><code>BASE</code></td><td>ASW 高阶API兜底</td><td>无条件命中,保证任何合法输入可 tiling</td></tr>
|
||||
</table>
|
||||
<p>DAV_RESV(s8s4 量化保留平台)只保留 5 个分支:<code>ITER_BATCH_BASICAPI → AL1_FULL_LOAD → BL1_FULL_LOAD → ASW_BASIC → BASE</code>。</p>
|
||||
|
||||
<div class="callout">
|
||||
<b>为什么按这个顺序短路?</b>顺序 = "对计算模式的改变程度"从大到小:
|
||||
<ol class="tight">
|
||||
<li><b>K 退化特判最靠前</b>(K==0、K==1):计算模式彻底改变(不需要 Cube),必须先拦截,否则会被后面的 cube 模板接住造成数量级浪费;</li>
|
||||
<li><b>StreamK 次之</b>:它改变的是全局并行结构(K 维跨核拆分 + workspace 归约),要在 batch 维优化之前决定;</li>
|
||||
<li><b>batch 维优化居中</b>(mergebatch → iterbatch 广播 → iterbatch×2):只改变 batch 维的调度与驻留方式,不改变单 batch 的计算模式;</li>
|
||||
<li><b>L1 全载靠后</b>:只改变单侧操作数的搬移次数,是局部优化;</li>
|
||||
<li><b>通用 ASW 兜底</b>:BASE 用 999 保证永远最后尝试、必然命中——这是分支体系"完备性"的工程保证。</li>
|
||||
</ol>
|
||||
</div>
|
||||
</section>
|
||||
Reference in New Issue
Block a user