mat_mul_v3 源码分析:分支决策与 tiling/swizzle 实现
This commit is contained in:
271
Matmul/3_源码对比/3.1_mat_mul_v3源码解析.html
Normal file
271
Matmul/3_源码对比/3.1_mat_mul_v3源码解析.html
Normal file
@@ -0,0 +1,271 @@
|
||||
<!DOCTYPE html><html lang="zh-CN"><head><meta charset="utf-8"><meta name="viewport" content="width=device-width, initial-scale=1"><title>mat_mul_v3 算子源码分析 —— 分支决策与 tiling/swizzle 实现</title>
|
||||
<style>
|
||||
body { font-family: -apple-system, "PingFang SC", "Microsoft YaHei", sans-serif; max-width: 980px; margin: 40px auto; padding: 0 24px; line-height: 1.75; color: #1f2328; background:#fff; }
|
||||
h1 { font-size: 28px; border-bottom: 2px solid #d0d7de; padding-bottom: 12px; }
|
||||
h2 { font-size: 22px; border-bottom: 1px solid #d0d7de; padding-bottom: 8px; margin-top: 36px; }
|
||||
h3 { font-size: 18px; margin-top: 28px; }
|
||||
code { background:#f6f8fa; padding: 2px 6px; border-radius: 4px; font-family: "SF Mono", Consolas, monospace; font-size: 0.9em; color:#c7254e; }
|
||||
pre { background:#f6f8fa; padding: 16px; border-radius: 8px; overflow-x:auto; }
|
||||
pre code { background:none; padding:0; color:#24292e; }
|
||||
table { border-collapse: collapse; width: 100%; margin: 16px 0; font-size: 14px; }
|
||||
th, td { border: 1px solid #d0d7de; padding: 8px 12px; text-align: left; vertical-align: top; }
|
||||
th { background:#f6f8fa; font-weight: 600; }
|
||||
blockquote { border-left: 4px solid #d0d7de; margin: 16px 0; padding: 4px 16px; color:#57606a; background:#f6f8fa; }
|
||||
hr { border:none; border-top:1px solid #d0d7de; margin: 24px 0; }
|
||||
li { margin: 4px 0; }
|
||||
strong { color:#0a3069; }
|
||||
</style>
|
||||
</head><body><h1>mat_mul_v3 算子源码分析 —— 分支决策与 tiling/swizzle 实现</h1>
|
||||
</blockquote>
|
||||
<hr/>
|
||||
<h2>0. 源码结构总览</h2>
|
||||
<p>mat_mul_v3 采用 <strong>Host Tiling(CPU)+ Device Kernel(NPU)</strong> 两层结构,中间通过 <strong>tiling key(位域编码)</strong> 传递决策结果:</p>
|
||||
<pre><code>┌─────────────────────────── Host(CPU 侧,算子启动前)───────────────────────────┐
|
||||
│ MatMulV3Tiling::DoTiling() │
|
||||
│ └─ MMTilingRegistry 按 priority 顺序尝试各 tiling 类(策略模式) │
|
||||
│ ├─ MatMulV3KEqZeroTiling (priority 0) │
|
||||
│ ├─ MatMulV3ToMulTiling (priority 1) │
|
||||
│ ├─ MatMulV3ToVectorTiling (priority 2) │
|
||||
│ ├─ MatMulV3BasicStreamKTiling (priority 3) │
|
||||
│ ├─ MatMulV3BasicAswtTiling (priority 4) │
|
||||
│ └─ MatMulV3AswTiling (priority 999, 兜底) │
|
||||
│ 每个类的 IsCapable() 判定是否满足条件,满足则 DoOpTiling() 计算 tiling 参数,│
|
||||
│ 编码成 tiling key + tiling data 下发给 device。 │
|
||||
└─────────────────────────────────────────────────────────────────────────────────┘
|
||||
│ tiling key + tiling data
|
||||
▼
|
||||
┌─────────────────────────── Device(NPU 侧,每核执行)───────────────────────────┐
|
||||
│ mat_mul_v3() 入口:if constexpr 按 (ApiLevel, FullLoad, Model, L0C2Out) 分发 │
|
||||
│ ├─ MatMulInputKEqZeroClearOutput K==0 清零 │
|
||||
│ ├─ MatMulToMulActKernel / MatMulToVectorActKernel 退化向量乘 │
|
||||
│ ├─ MatMulStreamKKernel / MatMulStreamKSplitKKernel StreamK 切 K │
|
||||
│ ├─ MatMulBasicKernel / MatMulBasicSplitKKernel ASWT 基础模板 │
|
||||
│ ├─ MatMulAL1FullLoadKernel / MatMulBL1FullLoadKernel A/B 全载 │
|
||||
│ └─ MatMulFixpipeOptiTensorKernel Fixpipe 随路搬出优化 │
|
||||
└─────────────────────────────────────────────────────────────────────────────────┘</code></pre>
|
||||
<p><strong>关键枚举</strong>(<code>mat_mul_v3_tiling_key_public.h</code>):</p>
|
||||
<table><thead><tr><th>维度</th><th>枚举值</th><th>含义</th></tr></thead><tbody>
|
||||
<tr><td>MatMulV3Model</td><td>BASIC=0 / STREAM_K=1 / K_EQUAL_ZERO=2 / TO_MUL=3 / TO_MULTI_MUL=4 / SLICE=5 / BASIC_SPLIT_K=6 / SK_SPLIT_K=7</td><td>主计算模型</td></tr>
|
||||
<tr><td>MatMulV3FullLoad</td><td>NONE=0 / A_FULL_LOAD=1 / B_FULL_LOAD=2 / AB_FULL_LOAD=3</td><td>L1 全载模式</td></tr>
|
||||
<tr><td>MatMulV3L0C2Out</td><td>ON_THE_FLY=0 / ND_FIXPIPE_1_1=1 / ND_FIXPIPE_1_2=2</td><td>L0C 输出方式</td></tr>
|
||||
<tr><td>MatMulV3ApiLevel</td><td>HIGH=0 / BASIC=1 / TENSOR=2</td><td>编程接口层级(kernel 用 BASIC/TENSOR)</td></tr>
|
||||
</tbody></table>
|
||||
<hr/>
|
||||
<h2>1. 分支决策总表(DAV_3510 = 950PR)</h2>
|
||||
<p>策略优先级定义在 <code>matmul_v3_tiling_strategy.h</code>:</p>
|
||||
<pre><code>{ NpuArch::DAV_3510, {K_EQUAL_ZERO, TO_MUL, TO_MULTI_MUL, BASIC_STREAM_K, BASIC_ASWT} }</code></pre>
|
||||
<p>各 tiling 类按 priority 注册(<code>MM_REGISTER_TILING_TEMPLATE</code>),<code>MMTilingRegistry::DoTilingImpl</code> 按优先级逐个 <code>IsCapable()</code> + <code>DoTiling()</code>,<strong>第一个成功者胜出</strong>。</p>
|
||||
<table><thead><tr><th>Priority</th><th>分支</th><th>触发条件(IsCapable)</th><th>Device Kernel</th></tr></thead><tbody>
|
||||
<tr><td>0</td><td><strong>K_EQUAL_ZERO</strong></td><td>无 bias 且 <code>K == 0</code></td><td>MatMulInputKEqZeroClearOutput</td></tr>
|
||||
<tr><td>1</td><td><strong>TO_MUL</strong></td><td>高精度FP32(<code>isForceGrpAccForFp32</code>) ∧ 非slice ∧ (M==1 ∨ N==1) ∧ A/B 均 FP32</td><td>MatMulToMulActKernel</td></tr>
|
||||
<tr><td>2</td><td><strong>TO_MULTI_MUL</strong></td><td>高精度FP32 ∧ 非slice ∧ <code>!ATrans ∧ BTrans</code> ∧ 无bias ∧ A/B 均 FP32</td><td>MatMulToVectorActKernel</td></tr>
|
||||
<tr><td>3</td><td><strong>BASIC_STREAM_K</strong></td><td>deterministic≤1 ∧ A为ND ∧ 非slice ∧ aivNum==2·aicNum ∧ (SK 或 DPSK 条件)</td><td>MatMulStreamKKernel / StreamKSplitK / StreamKActKernel</td></tr>
|
||||
<tr><td>4</td><td><strong>BASIC_ASWT</strong></td><td>无条件(主分支)</td><td>MatMulBasicKernel 等(含 fullLoad/fixpipe 子分支)</td></tr>
|
||||
<tr><td>999</td><td><strong>BASE(ASW)</strong></td><td>无条件(最终兜底,老 Cmct 接口)</td><td>MatMulActKernel</td></tr>
|
||||
</tbody></table>
|
||||
</blockquote>
|
||||
<hr/>
|
||||
<h2>2. 各分支详细分析</h2>
|
||||
<h3>2.1 K_EQUAL_ZERO(priority 0)</h3>
|
||||
<p><strong>条件</strong>:<code>!hasBias && kValue == 0</code>。</p>
|
||||
<p><strong>实现</strong>(<code>matmul_v3_k_equal_zero_tiling.cpp</code> + <code>mat_mul_input_k_eq_zero_clear_output.h</code>):</p>
|
||||
<ul>
|
||||
<li><code>totalDataAmount = M * N</code>,<code>usedCoreNum = aivNum</code>(64 个 AIV)。</li>
|
||||
<li>Device 侧 <code>MatMulInputKEqZeroClearOutput</code> 直接用 AIV 对输出 <code>C</code> 清零,完全不启动 Cube。</li>
|
||||
</ul>
|
||||
<p><strong>评价</strong>:K=0 是退化空矩阵乘,输出恒为 0。用 AIV 批量清零、不浪费 Cube 计算单元,实现正确且高效。适用面窄(仅 K==0 且无 bias)。</p>
|
||||
<h3>2.2 TO_MUL(priority 1)</h3>
|
||||
<p><strong>条件</strong>:<code>isForceGrpAccForFp32</code>(op_impl_mode_enum == 0x4 高精度)∧ 非 slice ∧ (M==1 ∨ N==1) ∧ A/B 均 FP32。</p>
|
||||
<p><strong>实现</strong>(<code>matmul_v3_to_mul_tiling.cpp</code> + <code>mat_mul_to_mul_cmct.h</code>):</p>
|
||||
<ul>
|
||||
<li>矩阵乘退化为<strong>向量乘</strong>(M==1 或 N==1 时本质是点积/向量乘),全部由 AIV 完成,不占 Cube。</li>
|
||||
<li>tiling:<code>baseMN</code>、<code>baseK</code> 由 UB 容量(<code>ubSize/sizeof(float)</code>)反推;<code>loopK = ceil(K/baseK)</code>;区分 <code>dataCopyMode</code>(判断内外轴是否连续搬运)。</li>
|
||||
<li><code>usedCoreNum = min(tileNum, aivNum)</code> 分核。</li>
|
||||
</ul>
|
||||
<p><strong>评价</strong>:M==1 或 N==1 时 Cube 的 16×16×16 分形利用率极低,改用 AIV 向量乘是正确决策。局限:仅覆盖高精度 FP32 模式(<code>isForceGrpAccForFp32</code>),FP16/BF16 的 M==1/N==1 场景走不到这里。</p>
|
||||
<h3>2.3 TO_MULTI_MUL(priority 2)</h3>
|
||||
<p><strong>条件</strong>:高精度FP32 ∧ 非slice ∧ <code>!ATrans && BTrans</code> ∧ 无 bias ∧ A/B 均 FP32。</p>
|
||||
<p><strong>实现</strong>(<code>matmul_v3_to_multi_mul_tiling.cpp</code> + <code>mat_mul_to_multi_mul_cmct.h</code>):</p>
|
||||
<ul>
|
||||
<li>同样退化为 AIV 向量乘,但针对 <code>!ATrans && BTrans</code> 的特定转置组合(A 行主序、B 列主序,恰好内积方向连续)。</li>
|
||||
<li><code>CalcBasicBlock()</code>:按 mCore/nCore 的核数分配调整 baseM/baseN,使 M/N 方向的核分配均衡(<code>while baseN >= 2*baseM ...</code> 启发式均衡)。</li>
|
||||
</ul>
|
||||
<p><strong>评价</strong>:与 TO_MUL 同类,但更窄(限定转置组合、不支持 bias)。两个分支共同说明:<strong>当矩阵乘退化到 M==1 或 N==1 时,源码选择绕开 Cube 走 AIV</strong>,规避 Cube 16 对齐分形粒度的浪费。</p>
|
||||
<h3>2.4 BASIC_STREAM_K(priority 3)—— StreamK / DPSK</h3>
|
||||
<p><strong>IsCapable 前置条件</strong>:<code>deterministicLevel ≤ 1</code>(切 K 核间累加顺序不确定,强一致性场景禁用)∧ A 为 ND ∧ 非 slice ∧ <code>aivNum == 2·aicNum</code>。</p>
|
||||
<p><strong>SK 模式条件</strong>(<code>CheckStreamKSKTilingDav3510</code>):</p>
|
||||
<pre><code>align(K, 256) >= max(8192, aicNum·256B) / dtypeSize
|
||||
mCnt · nCnt <= aicNum / 2 // MN 用 base 块切的份数不超过核数一半</code></pre>
|
||||
<ul>
|
||||
<li>语义:<strong>M/N 太小、块数不够分满核,但 K 足够大</strong> → 把 K 切成多份分给多个核,每核算一段 K 后跨核累加。</li>
|
||||
<li>FP32 且非 hf32 时 base 块对齐单位退化为 32(<code>BLOCK_BYTE_SIZE</code>)。</li>
|
||||
</ul>
|
||||
<p><strong>DPSK 模式条件</strong>(<code>CheckStreamKDPSKTilingDav3510</code>):</p>
|
||||
<pre><code>M % 256 == 0 且 N % 256 == 0
|
||||
K >= max(8192, aicNum·128B) / dtypeSize
|
||||
totalMNCnt >= aicNum 且 totalMNCnt % aicNum != 0 且 余数 <= aicNum/2</code></pre>
|
||||
<ul>
|
||||
<li>语义:<strong>主轮能分满核(DP),但尾轮有 M/N 剩余块无法分满</strong> → 尾轮用 SK 方式(切 K)提前执行,让 AIV 累加时 AIC 继续算下一轮(数据并行)。</li>
|
||||
</ul>
|
||||
<p><strong>tiling 核心</strong>(<code>DoOpTiling</code>):</p>
|
||||
<ul>
|
||||
<li><code>singleCoreK = ceil(K / kCnt)</code>,<code>kCnt</code> 由核数与 MN 块数的关系决定;</li>
|
||||
<li><code>baseK = min(singleCoreK, L0A容量约束)</code>;</li>
|
||||
<li><code>workspace = aicNum · 256·256·4B + RPC·MB</code>(<strong>核间累加工作区</strong>);</li>
|
||||
<li><code>GetL0C2Out</code>:N 不对齐且大 N 时选 <code>ND_FIXPIPE_1_2</code>;</li>
|
||||
<li>model:FP32 且 <code>singleCoreK >= FP32_SPLIT_K_THRESHOLD</code> → <code>SK_SPLIT_K</code>(单核内再切 K 保精度),否则 <code>STREAM_K</code>。</li>
|
||||
</ul>
|
||||
<p><strong>Device swizzle</strong>(<code>block_scheduler_streamk.h</code>):</p>
|
||||
<ul>
|
||||
<li><code>mL1 = baseM</code>、<code>nL1 = baseN</code>(StreamK 中 L1 只装一个 base 块);</li>
|
||||
<li><code>tileNum = DP部分MN块数 + 尾轮SK部分MN块数·skKTileNum</code>;</li>
|
||||
<li><code>CheckIsSkScene(tileIdx)</code> 判定当前块属于 DP 主轮(K 不分片)还是 SK 尾轮(K 分片);</li>
|
||||
<li>同样应用 SWAT 窗口 + 蛇形扫描(见 §4)。</li>
|
||||
</ul>
|
||||
<p><strong>评价</strong>:</p>
|
||||
<ul>
|
||||
<li><strong>优点</strong>:解决"MN 块数不足分满核、但 K 大"这一 ASWT 无法高效处理的场景,大幅提升瘦高形状(大 K)的算力利用率;DPSK 让尾轮 SK 与 AIC 计算重叠,减少核空闲。</li>
|
||||
<li><strong>缺点</strong>:① 需要 workspace 做跨核累加,<strong>额外 GM 写+读流量</strong>,在 950PR 带宽紧张下成本不低;② 累加顺序不确定 → <code>deterministicLevel>1</code> 时整分支禁用;③ 仅支持 A 为 ND;④ DPSK 要求 M/N 严格 256 对齐,形状不齐时退化为纯 SK 或走 ASWT。</li>
|
||||
</ul>
|
||||
<h3>2.5 BASIC_ASWT(priority 4)—— 主分支</h3>
|
||||
<p><strong>IsCapable</strong>:无条件返回 true。其 <code>DoOpTiling</code> 先调用父类 <code>MatMulV3AswTiling::DoOpTiling()</code> 完成基础 tiling(baseM/baseN/baseK/singleCore),再按顺序做子决策:</p>
|
||||
<pre><code>DoOpTiling():
|
||||
isSlice_ = IsSelfNonContiguous() // 非连续 3D slice
|
||||
l0C2Out_ = GetL0C2Out() // 是否走 fixpipe 优化
|
||||
if (!isSlice_ && CheckAL1FullLoad()) → DoAL1FullLoad() // A 全载 L1
|
||||
elif (!isSlice_ && CheckBL1FullLoad()) → DoBL1FullLoad() // B 全载 L1
|
||||
elif (l0C2Out_ == ON_THE_FLY) → 普通场景,按 L1 剩余容量均分 stepK
|
||||
else → fixpipe 优化场景
|
||||
CheckFp32SplitK() // FP32 大 K → BASIC_SPLIT_K
|
||||
CheckApiLevelAndModel() // tensor/basic api</code></pre>
|
||||
<p><strong>A 全载条件</strong>(<code>CheckAL1FullLoad</code>):</p>
|
||||
<ul>
|
||||
<li><code>l0C2Out == ON_THE_FLY</code>(不叠加 fixpipe);</li>
|
||||
<li>非 CubeBound(<code>cubeBoundParam > cubeBoundEdge</code>,即 MTE2 是瓶颈、重复读代价高);</li>
|
||||
<li><code>nCnt > aicNum</code>(N 方向块数多于核数,存在跨核重复读 A);</li>
|
||||
<li>排除"Fixp Bound 多轮"(<code>K<=128 && mCnt!=1</code>);</li>
|
||||
<li>整个 A(M×K)+ bias ≤ <strong>3/4 L1</strong>。</li>
|
||||
</ul>
|
||||
<p><strong>B 全载条件</strong>(<code>CheckBL1FullLoad</code>):对称(非 CubeBound ∧ <code>mCnt > aicNum</code> ∧ B+bias ≤ 3/4 L1)。</p>
|
||||
<p><strong>A 全载实现</strong>(<code>DoAL1FullLoad</code>):整个 A 常驻 L1,B 按 <code>baseN</code> 分块流式搬入;<code>singleCoreM = M</code>(不再分 M)、<code>singleCoreN = baseN</code>;<code>baseN</code> 取 min(原值, L1剩余容量上限, L0C双缓冲上限, 负载均衡值);<code>stepKb</code> 由 B 的 L1 搬移量 + 256B 对齐约束确定;<code>l1BufferNum</code> 判断能否 4-buffer。</p>
|
||||
<p><strong>B 全载实现</strong>(<code>DoBL1FullLoad</code>):对称(整个 B 常驻 L1,A 流式)。</p>
|
||||
<p><strong>L0C2Out(fixpipe 优化)条件</strong>(<code>GetL0C2OutDav3510</code>):</p>
|
||||
<pre><code>isValidMKN = K<=256 && M>=256
|
||||
isMultiRound = mCnt·nCnt > aicNum
|
||||
isUnalignedN = (N·cDtypeSize % 128 != 0) && (N·cDtypeSize > 256)
|
||||
fixpipeBound = isValidMKN && isMultiRound && isUnalignedN</code></pre>
|
||||
<p>满足且 <code>aivNum == 2·aicNum</code> → FP16/BF16 选 <code>ND_FIXPIPE_1_1</code>,FP32 选 <code>ND_FIXPIPE_1_2</code>。</p>
|
||||
<p><strong>Device kernel 分发</strong>(<code>mat_mul_v3.cpp</code> 的 <code>if constexpr</code>):按 (ApiLevel, FullLoad, Model, L0C2Out) 组合映射到 <code>MatMulBasicKernel</code> / <code>MatMulAL1FullLoadKernel</code> / <code>MatMulBL1FullLoadKernel</code> / <code>MatMulFixpipeOptiTensorKernel</code> / <code>MatMulBasicSplitKKernel</code>(均基于 Blaze::Gemm 模板库)。</p>
|
||||
<p><strong>评价</strong>:</p>
|
||||
<ul>
|
||||
<li><strong>优点</strong>:① 通用性强,覆盖绝大多数 shape;② SWAT 提升 L2 命中率;③ A/B 全载在"一侧重复读多、另一侧能装进 L1"时显著减少 GM→L1 流量;④ fixpipe 优化在特定小 K 大 M 场景让搬出与计算并行;⑤ 尾轮负载均衡减少尾轮算力浪费。</li>
|
||||
<li><strong>缺点</strong>:① baseM/baseN 的选择是<strong>启发式搜索</strong>(§3 的 <code>cubeBoundParam/balanceRate</code> 权衡),不是严格最优解;② A/B 全载只支持"整个矩阵常驻 L1"一种粒度,且限定非 CubeBound、ND 场景;③ fixpipe 条件苛刻(K≤256 且 M≥256 且 N 不对齐),覆盖面窄;④ 未区分 GM 读写带宽,<code>GetHbmBW</code> 用统一换算值。</li>
|
||||
</ul>
|
||||
<h3>2.6 BASE(priority 999)—— 老 Cmct 接口兜底</h3>
|
||||
<p><strong>实现</strong>(<code>matmul_v3_asw_tiling.cpp</code> + <code>mat_mul_asw_kernel.h</code> / <code>mat_mul_asw_block.h</code>):</p>
|
||||
<ul>
|
||||
<li><code>DoOpTiling</code>: <code>ResetBase</code> → <code>GetRebalanceBlock</code> → <code>CalcTailBasicBlock</code> → <code>CalL1Tiling</code>;</li>
|
||||
<li>Device 侧走 <code>MatMulActKernel</code>(老 Cmct 接口),核内 <code>mm_.Iterate()</code> 执行标准 <code>GM→L1→L0A/L0B→Cube→L0C→Fixpipe</code> 流水;<code>SetMMLayoutTransform(true)</code> 让 Fixpipe 用 N 搬出实现 Cube 与 Fixpipe 并行;</li>
|
||||
<li>块索引计算(<code>MatmulAswBlock::UpdateBasicIndex</code>)同样实现 SWAT 窗口 + 蛇形扫描 + 尾轮重切。</li>
|
||||
</ul>
|
||||
<p><strong>评价</strong>:作为最终兜底保证正确性,逻辑与 BASIC_ASWT 基础 tiling 一致,但<strong>没有</strong>全载 / fixpipe / StreamK 等新优化,性能上限低于主分支。</p>
|
||||
<hr/>
|
||||
<h2>3. 核心 tiling 算法(决定 baseM/baseN/baseK 与单核形状)</h2>
|
||||
<h3>3.1 ResetBase(初始值,<code>matmul_v3_tiling_helper.cpp</code>)</h3>
|
||||
<pre><code>// ResetBaseDefault(DAV_3510 在此基础上改 baseM)
|
||||
usedCoreNum = aicNum; // 32
|
||||
baseM = 256; baseN = 256; // 950PR 的 base 块
|
||||
baseK = 128B / dtypeSize; // FP16=64, FP32=32
|
||||
iterateOrder = ITER_COL_FIRST; // 列优先
|
||||
singleCoreK = K; singleCoreM/N = baseM/N;</code></pre>
|
||||
<h3>3.2 GetRebalanceBlock(baseM/baseN 最优搜索,核心)</h3>
|
||||
<p>这是整个 tiling 最关键的函数,分两步:</p>
|
||||
<p><strong>① Roofline 判 CubeBound</strong>:</p>
|
||||
<pre><code>hbmBW = freq · 32核 · 31B/拍 / 1024 // ≈ 1.6TB/s
|
||||
l2BW = freq · 32核 · 100B/拍 / 1024 // ≈ 5.2TB/s
|
||||
singleCoreComputePower = freq · 8 // ≈ 13.2 TFLOPS(单核 BF16)
|
||||
computePower = singleCoreComputePower · aicNum
|
||||
|
||||
cmr = (M+N)/(M·N) // 临界算术强度相关量
|
||||
cubeBoundEdge = (l2BW/computePower) + l2CacheUsage·(1 - l2BW/hbmBW)·cmr
|
||||
- (1 + l2BW/hbmBW)/K
|
||||
cubeBoundParam = 1/baseM + 1/baseN
|
||||
// Cube Bound 条件:cubeBoundParam <= cubeBoundEdge</code></pre>
|
||||
<p><strong>② 搜索最优 (baseM, baseN)</strong>:在 <code>maxBaseM × maxBaseN</code> 解空间内双重循环,对每个候选算:</p>
|
||||
<ul>
|
||||
<li><code>curCubeBoundParam = 1/curBaseM + 1/curBaseN</code></li>
|
||||
<li><code>curBalanceRate</code>(尾轮负载均衡率,<code>GetBalanceRateWithTail</code>)</li>
|
||||
<li>目标:优先满足 CubeBound 且 balanceRate 更高;否则综合 <code>cubeBoundParam/balanceRate</code> 评选。</li>
|
||||
</ul>
|
||||
<p><code>maxBaseM/maxBaseN</code> 由 <code>GetMaxBaseWithLimit</code> 计算,受 L0A/L0C/L1/bias table/K 对齐多重约束。</p>
|
||||
<h3>3.3 GetBaseK</h3>
|
||||
<pre><code>maxBaseK = L0A_SIZE / DB_SIZE / dtypeSize / max(baseM, baseN)
|
||||
// 优先 K 全载进 L0A;否则按 256B 对齐;再退 128/64/32/16</code></pre>
|
||||
<h3>3.4 CalL1Tiling(K 方向 L1 分片,<code>CalL1TilingDefault</code>)</h3>
|
||||
<pre><code>isKInner = !ATrans || BTrans
|
||||
maxStepK = min(ceil(K/baseK), 8) // K 方向 L1 分片数上限 8
|
||||
// 遍历 stepK:满足 (aL1+bL1)·DB <= L1 且 K 256B 对齐 且 单次搬移量约束
|
||||
stepKa = stepKb = resKL1 / baseK
|
||||
depthA1 = stepKa · DB; depthB1 = stepKb · DB</code></pre>
|
||||
<h3>3.5 CalcTailBasicBlock(尾轮重切)</h3>
|
||||
<pre><code>tailCnt = (mCnt·nCnt > aicNum) ? (mCnt·nCnt % aicNum) : 0
|
||||
// 尾轮把 base 块在 M/N 方向重切 mTailCnt×nTailCnt 份,
|
||||
// 使尾轮也尽量填满核,且保持搬移效率(128B 对齐约束)</code></pre>
|
||||
<hr/>
|
||||
<h2>4. Swizzle 编排(SWAT 窗口 + 蛇形扫描)</h2>
|
||||
<p>ASWT 与 StreamK 共用的 swizzle 核心(<code>block_scheduler_aswt.h</code> 的 <code>UpdateMNTileIdx</code>、<code>mat_mul_asw_block.h</code> 的 <code>UpdateBasicIndex</code>):</p>
|
||||
<pre><code>mainWindow = min(4, mTileNum) // 固定窗口 4 行(WINDOW_LEN=4)
|
||||
mainRow = mTileNum / mainWindow - 1
|
||||
tailWindow = mTileNum - mainRow · mainWindow
|
||||
|
||||
rowIdx = tileIdx / nTileNum / mainWindow
|
||||
if (rowIdx < mainRow):
|
||||
mTileIdx = rowIdx·mainWindow + tileIdx % mainWindow // 窗口内 M 小步滑动
|
||||
nTileIdx = (tileIdx / mainWindow) % nTileNum // N 方向连续滑动
|
||||
else:
|
||||
// 尾窗口特殊处理
|
||||
if (rowIdx % 2 != 0):
|
||||
nTileIdx = nTileNum - 1 - nTileIdx // 蛇形:奇数行 N 反向</code></pre>
|
||||
<p><strong>SWAT 语义</strong>:把 M 轴按窗口(默认 4 个 base 块)分组,窗口内沿 N 连续滑动、M 小步滑动,使相邻核访问的数据在空间上邻近 → 最大化 L2 命中;奇数行 N 反向扫描(蛇形)让相邻轮的首尾块空间相邻,进一步提升 L2 复用。这是官方 <code>matmul_performance.md</code> 中 SWAT(Slide Window Adaptive Template)的落地实现。</p>
|
||||
<p><strong>尾轮重切</strong>(<code>GetBlockShape</code>):最后一轮把单个 base 块在 M/N 方向再切 <code>mTailCnt×nTailCnt</code> 份分给更多核,<code>blockIdx % tailCnt</code> 决定每核拿哪个子块,消除尾轮算力浪费。</p>
|
||||
<p><strong>StreamK 的 swizzle</strong>(<code>block_scheduler_streamk.h</code>):在 SWAT 窗口基础上叠加 <strong>DP 主轮(K 不分片)+ SK 尾轮(K 分片)</strong> 判定(<code>CheckIsSkScene</code>),主轮每核一个 MN 块算完整 K,尾轮把剩余 MN 块切 K 分多核、由 AIV 在 workspace 上确定性累加。</p>
|
||||
<hr/>
|
||||
<h2>5. 各分支优缺点对比</h2>
|
||||
<table><thead><tr><th>分支</th><th>适用场景</th><th>优点</th><th>缺点 / 局限</th></tr></thead><tbody>
|
||||
<tr><td>K_EQUAL_ZERO</td><td>K==0 且无 bias</td><td>极简,AIV 清零不浪费 Cube</td><td>仅空矩阵乘</td></tr>
|
||||
<tr><td>TO_MUL</td><td>高精度 FP32 且 M==1 或 N==1</td><td>避开 Cube 分形浪费,AIV 向量乘</td><td>仅 FP32 高精度模式</td></tr>
|
||||
<tr><td>TO_MULTI_MUL</td><td>高精度 FP32 且 !ATrans∧BTrans 且 M/N 任意</td><td>同上 + 转置组合下的内积连续</td><td>更窄(限转置组合、无 bias)</td></tr>
|
||||
<tr><td>BASIC_STREAM_K</td><td>M/N 小、K 大(K≥8192)</td><td>切 K 用满核;DPSK 尾轮与计算重叠</td><td>workspace 额外带宽;累加顺序不确定;仅 A 为 ND;DPSK 要求 256 对齐</td></tr>
|
||||
<tr><td>BASIC_ASWT</td><td>通用主分支</td><td>SWAT 提 L2 命中;A/B 全载减重复读;fixpipe 并行搬出;尾轮均衡</td><td>baseM/N 启发式搜索非严格最优;全载/fixpipe 场景窄</td></tr>
|
||||
<tr><td>BASE(ASW)</td><td>最终兜底</td><td>保证正确性</td><td>老接口,无全载/fixpipe/StreamK 优化</td></tr>
|
||||
</tbody></table>
|
||||
<hr/>
|
||||
<h2>6. 初步观察到的可改进点(衔接任务 3.2)</h2>
|
||||
<p>在通读源码过程中,已浮现若干值得深挖的改进线索,留待后续对照性能模型严格论证:</p>
|
||||
<ol>
|
||||
<li><strong>baseM/baseN 搜索目标是启发式的</strong>:<code>GetRebalanceBlock</code> 用 <code>cubeBoundParam/balanceRate</code> 复合指标剪枝,而非直接代入 §3.2 的 <code>T_total = max(T_MMAD, T_MTE2, T_MTE1, T_FIXPIPE)</code> 精确评估。理论上可用性能模型对候选解做精确打分。</li>
|
||||
<li><strong>SWAT 窗口固定为 4</strong>(<code>WINDOW_LEN=4</code>):未根据 L2 容量、shape、核数自适应调窗。窗口大小直接影响 L2 命中率与重复读量的权衡。</li>
|
||||
<li><strong>A/B 全载只有"整个矩阵常驻"一种粒度</strong>:没有"部分驻留"(多个 base 块驻留 L1 的中间态),在 A/B 稍大于 3/4 L1 时直接放弃全载,存在优化断档。</li>
|
||||
<li><strong>fixpipe 优化覆盖窄</strong>:仅 <code>K≤256 ∧ M≥256 ∧ N 不对齐</code> 场景触发,其它 Fixpipe Bound 场景(如更小 K)未覆盖。</li>
|
||||
<li><strong>GM 带宽未区分读写</strong>:<code>GetHbmBW</code> 统一按 <code>32核·31B/拍</code> 换算,未区分读/写共享 1.6TB/s 的竞争,可能高估有效带宽。</li>
|
||||
<li><strong>FP32 高精度(isForceGrpAccForFp32)的退化分支覆盖不全</strong>:M==1/N==1 的 FP16/BF16 场景无对应 AIV 退化分支,仍走 Cube。</li>
|
||||
</ol>
|
||||
<hr/>
|
||||
<h2>附录:关键源码文件索引</h2>
|
||||
<table><thead><tr><th>层</th><th>文件</th><th>内容</th></tr></thead><tbody>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_v3_tiling_strategy.h</code></td><td>分支优先级定义</td></tr>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_tiling_registry.h</code></td><td>策略注册与 DoTilingImpl 调度</td></tr>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_v3_tiling_advanced.cpp</code></td><td>主入口 + 各 Phase</td></tr>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_v3_tiling_helper.cpp</code></td><td>ResetBase/GetRebalanceBlock/CalL1Tiling/GetL0C2Out</td></tr>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_v3_basic_streamk_tiling.cpp</code></td><td>StreamK/DPSK 条件与 tiling</td></tr>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_v3_basic_aswt_tiling.cpp</code></td><td>ASWT 子分支(全载/fixpipe)</td></tr>
|
||||
<tr><td>Host</td><td><code>op_host/op_tiling/arch35/matmul_v3_{k_equal_zero,to_mul,to_multi_mul,asw}_tiling.cpp</code></td><td>其余分支</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/mat_mul_v3.cpp</code></td><td>kernel 入口 if constexpr 分发</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/mat_mul_v3_tiling_key_public.h</code></td><td>枚举定义</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/mat_mul_tiling_data.h</code></td><td>tiling data 结构</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/block_scheduler_aswt.h</code></td><td>ASWT swizzle(SWAT+蛇形+尾轮)</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/block_scheduler_streamk.h</code></td><td>StreamK swizzle(DP/SK)</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/mat_mul_asw_block.h</code> / <code>mat_mul_asw_kernel.h</code></td><td>老接口 ASW 块调度与主循环</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/mat_mul_{al1,bl1}_full_load.h</code></td><td>A/B 全载 kernel 模板</td></tr>
|
||||
<tr><td>Device</td><td><code>op_kernel/arch35/mat_mul_streamk.h</code> / <code>mat_mul_fixpipe.h</code> / <code>mat_mul_basic_split_k.h</code></td><td>StreamK/Fixpipe/SplitK 模板</td></tr>
|
||||
</tbody></table></body></html>
|
||||
Reference in New Issue
Block a user