Files
matmul-analysis/BMM分块计算数学公式.md

11 KiB
Raw Blame History

BMM 分块计算的数学公式表达

昇腾 BatchMatMulV3 算子 | 目标芯片 950PR (DAV_3510) | 4 维可切Batch × M × N × K


一、问题定义

BMM 的语义是:


C[b, m, n] = \sum_{k=0}^{K-1} A[b, m, k] \cdot B[b, k, n]

其中:


\begin{aligned}
b &= b_0 \times b_1 \times b_2 \times b_3 \quad &\text{(展平后的 batch最多 4 级)}\\
m &\in [0, M), \quad n \in [0, N), \quad k \in [0, K) &\text{(矩阵维度)}
\end{aligned}

A,B 的 batch 维支持广播(任一维值为 1即实际的 A 索引为:


b_A[i] = \begin{cases} 0 & \text{if } b_i^A = 1 \\ b_i & \text{otherwise} \end{cases}

分块计算的目标:将上述 4 重循环 (b, m, n, k) 的所有迭代点重新组织为核级任务图,每个核负责一个子任务(一次 Load → Compute → Store 流水),子任务序列按 swizzle 重排以最大化 L2 复用。


二、四维切分的通用形式

2.1 分块参数

参数 含义 典型值
B^t batch 分块粒度iterBatch 1 ~ 8
M^t M 向基本块baseM 16 ~ 256
N^t N 向基本块baseN 16 ~ 256
K^t K 向基本块baseK 16 ~ 128fp16 下)

2.2 分块数


\tilde{B} = \lceil B / B^t \rceil,\quad
\tilde{M} = \lceil M / M^t \rceil,\quad
\tilde{N} = \lceil N / N^t \rceil,\quad
\tilde{K} = \lceil K / K^t \rceil

2.3 一个基本块的完整计算

一个基本块tile是四维索引 $(\beta, \mu, \nu) \in [0, \tilde{B}) \times [0, \tilde{M}) \times [0, \tilde{N})$,它要完成的计算是:


\boxed{
\begin{aligned}
C_{(\beta, \mu, \nu)}
&= \sum_{\kappa=0}^{\tilde{K}-1}
A\big[\beta B^t : (\beta{+}1)B^t,\; \mu M^t : (\mu{+}1)M^t,\; \kappa K^t : (\kappa{+}1)K^t\big] \\
&\qquad \times \;
B\big[\beta B^t : (\beta{+}1)B^t,\; \kappa K^t : (\kappa{+}1)K^t,\; \nu N^t : (\nu{+}1)N^t\big]
\end{aligned}
}

其中 A[\cdots], B[\cdots] 表示对应子矩阵。K 循环 (\kappa) 在 L0C 上累加,不写回 GM——这是 L0C 累加的核心:每次 mmad 的 L0C 结果驻留,下一轮 K 步的 mmad 直接在 L0C 上累加(cmatrixInitVal=false + cmatrixSource),省去 CO1→GM→UB 的来回搬运。


三、Swizzle块到核的映射函数

设总核数 $C$AIC 数量,如 32。核心问题是如何将 \tilde{B} \times \tilde{M} \times \tilde{N} 个基本块分配到 C 个核上使得时间上相邻的块在空间上相邻L2 友好)?

3.1 线性展开

首先将 3 维 block 索引展开为 1 维:


\text{idx}(\beta, \mu, \nu) = \beta \cdot \tilde{M}\tilde{N} + \sigma(\mu, \nu)

其中 \sigma(\mu, \nu)swizzle 函数——它在 M \times N 平面上定义遍历顺序。

3.2 ASW 滑窗 Swizzle

ASW 定义窗口宽度 $W = \max{,d \mid d \mid C,; d \le \lfloor\sqrt{C}\rfloor ,}$(如 32 核 → $W=4$)。将 M 方向按窗口分组:


\mu_{\text{row}} = \lfloor \mu / W \rfloor,\quad
\mu_{\text{col}} = \mu \bmod W

主窗口区($\mu_{\text{row}} < \lfloor \tilde{M}/W \rfloor$


\sigma_{\text{main}}(\mu, \nu) =
\big(\mu_{\text{row}} \cdot W + \mu_{\text{col}}\big) \cdot \tilde{N} + \nu'

其中 \nu' 由蛇形snake决定


\nu' = \begin{cases}
\nu & \text{若 } \mu_{\text{row}} \text{ 为偶数(正向)} \\
\tilde{N} - 1 - \nu & \text{若 } \mu_{\text{row}} \text{ 为奇数(反向)}
\end{cases}

核号与轮次


c = \mathrm{idx}(\beta, \mu, \nu) \bmod C,\qquad
r = \lfloor \mathrm{idx}(\beta, \mu, \nu) / C \rfloor

即第 r 轮核 c 处理全局第 c + r \cdot C 个基本块。

为什么窗口取 \lfloor\sqrt{C}\rfloor 的最大因子:窗口越接近正方形,同窗口内 A 行块与 B 列块的 L2 足迹越小,且因子保证整窗被核数均分、窗口边界不碎。

3.3 对角线错位 Swizzle老路径


\sigma_{\text{diag}}(\mu, \nu) = \mu \cdot \tilde{N} +
\left(\nu + \left\lfloor \frac{\mu \cdot C}{\mathrm{lcm}(\tilde{M}, \tilde{N})} \right\rfloor \right) \bmod \tilde{N}

这使同一时刻各核落在 M \times N 平面的不同对角线上,避免多核同时抢同一行 A / 同一列 B 的 GM 带宽。


四、单核计算流水GM → L1 → L0 → Cube → L0C → GM

c 在第 r 轮处理的块为 $(\beta, \mu, \nu)$(由 swizzle 逆映射得到)。

4.1 L1 驻留

该核在 L1 上 驻留的数据量为:


\begin{aligned}
A_{c}^{\text{L1}} &= A\big[\beta B^t : (\beta{+}1)B^t,\; \mu M^t : (\mu{+}1)M^t,\; \kappa_0 K^t : (\kappa_0 + d_A) K^t\big] \\[4pt]
B_{c}^{\text{L1}} &= B\big[\beta B^t : (\beta{+}1)B^t,\; \kappa_0 K^t : (\kappa_0 + d_B) K^t,\; \nu N^t : (\nu{+}1)N^t\big]
\end{aligned}

其中:

  • d_A = \mathrm{stepM} \cdot \mathrm{stepKa} 是 L1 上 A 的驻留深度(基本块份数)
  • $d_B = \mathrm{stepKb} \cdot 2$(双缓冲 DB
  • 容量约束:(d_A \cdot B^t M^t K^t + d_B \cdot B^t K^t N^t) \cdot \mathrm{dtype} \le 512\text{KB}

4.2 L1 → L0 搬运(一个 K 步)


\begin{aligned}
A_{c}^{\text{L0A}} &= A_{c}^{\text{L1}}\big[\;:\;,\;:\;,\; \kappa:\kappa{+}1 \big] \quad (\text{大小 } B^t \cdot M^t \cdot K^t) \\[4pt]
B_{c}^{\text{L0B}} &= B_{c}^{\text{L1}}\big[\;:\;,\; \kappa:\kappa{+}1,\; :\;\big] \quad (\text{大小 } B^t \cdot K^t \cdot N^t)
\end{aligned}

L0A/L0B 各 64KB512B 对齐fractal 分形,一个分形恰好 512B

4.3 Cube 计算与 L0C 累加

一次 mmadfp16 下为 16×16×16 fractal一拍完成


C_{c}^{\text{L0C}} \mathrel{+}= \mathrm{mmad}\big(A_{c}^{\text{L0A}},\; B_{c}^{\text{L0B}}\big)

L0C 256KBfp32 累加。K 循环 \kappa = 0, 1, \dots, \tilde{K}-1 全部在 L0C 上累加。

4.4 写出(经 fixpipe


C\big[\beta B^t : (\beta{+}1)B^t,\; \mu M^t : (\mu{+}1)M^t,\; \nu N^t : (\nu{+}1)N^t\big] \leftarrow C_{c}^{\text{L0C}}

fixpipe 可随路完成 NZ2ND 排布转换、量化FP32→BF16/FP16/FP8等。


五、两层嵌套的完整表达

外层(核间)—— swizzle 调度:


\forall c \in [0, C),\quad \forall r \in [0, R):\quad
(\beta, \mu, \nu) = \sigma^{-1}\big((c + rC) \bmod \tilde{B}\tilde{M}\tilde{N}\big)

其中 R = \lceil \tilde{B}\tilde{M}\tilde{N} / C \rceil 是轮数。

内层(核内)—— K 循环 + 多级流水:


C_{(\beta,\mu,\nu)}^{\text{L0C}} = \sum_{\kappa=0}^{\tilde{K}-1}
\underbrace{\mathrm{mmad}\Big(
\underbrace{A[\beta,\mu,\kappa]}_{\text{L0A, } B^t M^t K^t},\;
\underbrace{B[\beta,\kappa,\nu]}_{\text{L0B, } B^t K^t N^t}
\Big)}_{\text{16×16×16 fractal, fp16 一拍}}

六、各分支的数学特化

分支的差异归结为:哪些维度被"折叠"(合并/全载/驻留)以减少循环层数或搬运次数,以及 swizzle 函数 \sigma 的形式

分支 分块参数的特殊化 数学变化
K_EQUAL_ZERO K=0 $C = 0$,无计算,纯 AIV 写零
TO_MUL K=1, K^t=1, \tilde{K}=1 $C = A \odot B$逐元素乘K 循环消失,走 UB→Vector→GM
STREAM_K $\tilde{B}\tilde{M}\tilde{N} \le C/2$\tilde{K} 跨核拆分 $C_{(\beta,\mu,\nu)} = \sum_{c \in \text{group}} C_{(\beta,\mu,\nu)}^{(c)}$,部分和经 workspace 归约
MERGE_BATCH $B^t > 1$(合并 batchM^t, N^t 按合并倍数放大 $\tilde{B} = \lceil B/(B^t \cdot C) \rceil$batch 维折叠进 M/N 块
ITER_BATCH $B^t > 1$batch 组在 L1/L0 驻留 单核的 A/B 块大小 $\times B^t$Cube 一次算 B^t 个 batch
ITER_BROADCAST 广播侧 batch 维为 1B^t 只对非广播侧 广播侧数据 L1 驻留一份,对端 B^t 个 batch 共享
AL1_FULL_LOAD $M^t = M$(整个 MB 侧无 batch A 的 L1 驻留 = 整个 A所有 (\beta, \nu) 复用GM→L1 仅 1 次
BL1_FULL_LOAD $N^t = N$(整个 NA 侧无 batch 镜像,B 的 L1 驻留 = 整个 B
ASW / BASE 通用参数,B^t=1 或由 L1 容量决定 通用公式swizzle \sigma 为滑窗蛇形

特殊化详解

MERGE_BATCH:将 B^t 个 batch 折叠进 M^t 或 $N^t$


\tilde{B}_{\text{eff}} = \lceil B / (B^t \cdot C) \rceil,\quad
M_{\text{eff}}^t = B^t \cdot M^t,\quad
N_{\text{eff}}^t = B^t \cdot N^t

最优合并 batch 数 B^t 由 L0C 容量二次方程求解CalBatchL0WithPolynomial


\text{令 } a = \frac{M}{16},\; t = \frac{15}{16},\; p = \frac{a \cdot \text{l0CSize}}{16 \cdot \text{alignN} \cdot 4 \cdot \text{DB}}

(ax)^2 + t(ax) - p = 0 \quad\Rightarrow\quad B^t_{\text{opt}} = \left\lfloor \min\!\left(\frac{p}{\lceil y \rceil \cdot a},\; \frac{\lceil y \rceil}{a}\right) \right\rfloor

其中 $y = \sqrt{p + t^2/4} - t/2$。

AL1_FULL_LOAD$\tilde{M} = 1$A 的 \mu 循环消失,且 A 的 GM→L1 搬移只发生一次:


A_{c}^{\text{L1}} = A\big[0 : M,\; 0 : K\big] \quad \text{(完整 A 常驻 L1depthA1 = stepM × stepKa = 整个 A}

搬运次数对比192篇无全载时总搬运 M_{\text{blocks}} \times (1 + N_{\text{blocks}}) = 2 \times 3 = 6 次;全载后 1 + M_{\text{blocks}} = 3 次。

StreamK\tilde{K} 被跨核拆分K 循环变成两阶段:


C_{(\beta,\mu,\nu)} = \sum_{c \in \text{group}(\beta,\mu,\nu)} \underbrace{\sum_{\kappa \in \text{slice}(c)} \mathrm{mmad}\big(A[\kappa], B[\kappa]\big)}_{\text{核 c 的部分和 } C_{(\beta,\mu,\nu)}^{(c)}}

部分和写 workspace$C \times 256 \times 256 \times 4\text{B}$),再经 AIV 归约AIC:AIV = 1:2


七、参数取值与硬件约束的对应关系

参数 约束来源 公式
M^t, N^t 16 对齐 Cube 一拍 16×16×16 M^t = 16 \cdot \lceil M^t/16 \rceil
K^t 内轴 128B/256B/512B 对齐 MTE 搬运拆分粒度056篇 K^t \cdot \text{dtype} \in \{128, 256, 512\}
L1 双缓冲 d_A \cdot B^t M^t K^t \cdot 2 + d_B \cdot B^t K^t N^t \cdot 2 \le 512\text{KB} DB 乒乓 ×2
L0A/B 双缓冲 B^t M^t K^t \cdot 2 \le 64\text{KB} L0A/L0B 各 64KB
L0C 双缓冲 B^t M^t N^t \cdot 4\text{B} \cdot 2 \le 256\text{KB} L0C 256KBfp32 累加
L0C 累加 K 循环 \tilde{K} = \lceil K / K^t \rceil 每轮 mmad 结果在 L0C 上累加
Swizzle 窗长 W = \max\{d \mid d \mid C, d \le \lfloor\sqrt{C}\rfloor\} L2 128MB 全局共享

八、总结

BMM 分块计算在数学上就是:

把 4 重循环 (b,m,n,k) 的迭代空间,按硬件容量(核数 $C$、L1 512KB、L0A/B 64KB、L0C 256KB做嵌套分块并用 swizzle 函数 \sigma 重排块到核的映射顺序使得存储层级L2→L1→L0上的数据复用最大化。

分支体系就是针对不同形状的迭代空间,选择不同的"折叠维度"和"swizzle 策略"

  • 折叠维度MERGE_BATCH 折叠 batch 进 M/NL1_FULL_LOAD 折叠 M 或 N 进驻留StreamK 展开 K 到核间
  • Swizzle 策略ASW 滑窗蛇形(窗长 \sqrt{C} 因子)、对角错位、简单轮询
  • 数据通路Cube 通路GM→L1→L0→Cube→L0C→GMvs Vector 通路GM→UB→Mul→GMK=0/1 时)