Add BMM_Theory: bmm_theory/branches/merge_batch.py

This commit is contained in:
2026-09-03 08:09:31 +00:00
parent d93ae35729
commit 2550c0e9ac

View File

@@ -0,0 +1,230 @@
"""MergeBatch 分支: 核间切 B, 核内合并 b0 个 batch 计算.
理论依据:
- 《BMM算子优化分析 v0.98》§五 (进入条件 + 实现方案 Step1~3)
- 《MergeBatch_vs_IterBatch分析 v1.1》§三/§四 (执行模型 + 分界条件)
核心思想: 合并 b0 个 batch 的 A'[b0*M,K] @ B'[K,b0*N] 为单次 DMA 搬入,
减少 GM->L1 搬移命令数 (省 b0 倍 T_cmd); 交叉项被算出但丢弃 (冗余比例 (b0-1)/b0),
进入条件 5 保证 case 为访存 Bound, 冗余算力被搬移时延掩盖.
"""
from __future__ import annotations
import math
from ..hardware import NpuSpec, ASCEND950PR
from ..models import BmmCase, ImplPlan, HardwareTiming, align_down
from ..timing import MoveInPlan, assemble_timing
from .base import Branch, BranchResult, ConditionCheck
MIN_B0 = 2 # b0: 合并搬移有收益的最小合并数
class MergeBatchBranch(Branch):
name = "MergeBatch"
def __init__(self, spec: NpuSpec = ASCEND950PR):
super().__init__(spec)
# ------------------------------------------------------------------
# 进入条件 (v0.98 §五, 五条同时满足)
# ------------------------------------------------------------------
def check_conditions(self, case: BmmCase) -> list:
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
b = case.batch_c
b_core = b // s.aic_num
checks = []
# 条件 1: BatchA=BatchB 且 b_core >= 2*b0
c1 = (case.batch_a == case.batch_b) and (b_core >= 2 * MIN_B0)
checks.append(ConditionCheck(
"1_batch关系与每核份额: BatchA==BatchB 且 b_core=B/C>=2*b0",
c1, f"batchA={case.batch_a}, batchB={case.batch_b}, b_core={b_core}"))
# 条件 2: 2*(b0*M)*(b0*N)*4B <= L0C
l0c_need = 2 * (MIN_B0 * m) * (MIN_B0 * n) * 4
c2 = l0c_need <= s.l0c_bytes
checks.append(ConditionCheck(
"2_L0C容量: 2*(b0*M)*(b0*N)*4B <= L0C",
c2, f"need={l0c_need}B, L0C={s.l0c_bytes}B"))
# 条件 3: b_core*(MK+KN)*dtype >= min_DatamountPerCore
amt = b_core * (m * k + k * n) * dt
c3 = amt >= s.min_datamount_per_core
checks.append(ConditionCheck(
"3_单核搬移总量: b_core*(MK+KN)*dtype >= min_DatamountPerCore",
c3, f"{amt/1024:.0f}KB vs {s.min_datamount_per_core/1024:.0f}KB"))
# 条件 4: max(MK, KN)*dtype >= min_TileSize
tile = max(m * k, k * n) * dt
c4 = tile >= s.min_tile_size
checks.append(ConditionCheck(
"4_搬移tile大小: max(MK,KN)*dtype >= min_TileSize",
c4, f"{tile/1024:.1f}KB vs {s.min_tile_size/1024:.0f}KB"))
# 条件 5: 2MN/(M+N) < R16/b0 (合并后仍访存Bound)
ai = 2.0 * m * n / (m + n)
c5 = ai < s.r16 / MIN_B0
checks.append(ConditionCheck(
"5_访存Bound: 2MN/(M+N) < R16/b0",
c5, f"AI={ai:.1f} vs R16/b0={s.r16/MIN_B0:.1f}"))
return checks
# ------------------------------------------------------------------
# 实现方案 (v0.98 §五 Step1~3)
# ------------------------------------------------------------------
def make_plan(self, case: BmmCase) -> ImplPlan:
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
b = case.batch_c
b_core = b // s.aic_num
# Step 1: 合并数 b0 (L0C + 算存比双上限, 尽量取 b_core 的因子)
b0_l0c = math.sqrt(s.l0c_bytes / (2 * m * n * 4))
b0_ai = s.r16 * (m + n) / (2 * m * n)
b0_max = int(min(b0_l0c, b0_ai, b_core))
b0 = max(MIN_B0, self._factor_floor(b_core, b0_max))
# Step 2: L0 级 K 粒度 k_L0
k_l0 = align_down(int(min(
s.l0a_bytes / (2 * b0 * m * dt),
s.l0b_bytes / (2 * b0 * n * dt),
)), s.fractal)
k_l0 = max(k_l0, s.fractal)
# Step 3: L1 级 k_L1 (先反推再截断: 不超过 K, 不超过 dValue 推荐 512B)
k_l1_star = s.l1_bytes / (2 * b0 * (m + n) * dt)
dvalue_cap = s.dvalue_recommend // dt
k_l1 = align_down(int(min(k_l1_star, k, dvalue_cap)), s.fractal)
k_l1 = max(k_l1, s.fractal)
k_truncated = k_l1 >= k # K 截断: L1 一次装下整个 K 维
# Step 4: b_L1 最大化 (提升 batch 间流水深度)
b_l1 = int(min(
s.l1_bytes / (2 * k_l1 * (m + n) * dt),
b_core,
))
b_l1 = max(b_l1, b0)
note = (f"b0={b0} (L0C上限{b0_l0c:.1f}/算存比上限{b0_ai:.1f}/b_core={b_core}); "
f"{'K截断' if k_truncated else 'L1绑定'}; "
f"合并后单次DMA搬入 A'[{b0*m},{k_l1}]+B'[{k_l1},{b0*n}]")
return ImplPlan(
case_id=case.case_id, branch=self.name,
used_core_num=s.aic_num,
split_b=s.aic_num, m_cnt=1, n_cnt=1, grid_k=1,
core_map="切B均分(核间零重复读零依赖)",
b_core=b_core, merge_b0=b0,
single_core_m=b0 * m, single_core_n=b0 * n, single_core_k=k,
k_l1=k_l1, b_l1=b_l1, l1_form="合并驻留",
# L0C 双缓冲约束已由进入条件 2 保证 (2*(b0*M)*(b0*N)*4B <= L0C)
base_m=b0 * m, base_n=b0 * n, base_k=max(min(k_l0, k_l1, k), s.fractal),
l2_policy_in="allocate(GM->L1随路驻留L2)",
l2_policy_out="direct_gm(输出仅写一次,直写GM不占L2)" if not case.out_nd else "direct_gm",
swizzle_w=0, workspace_bytes=0,
tail_strategy="不涉及(核内不切M/N)",
fixpipe_unitflag=True,
out_dtype_bytes=case.dtype_out_bytes,
note=note,
)
# ------------------------------------------------------------------
# 时延评估 (v1.1 §4 端到端模型)
# ------------------------------------------------------------------
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
out_b = case.dtype_out_bytes
b_core, b0, k_l1 = plan.b_core, plan.merge_b0, plan.k_l1
k_truncated = k_l1 >= k
# 每 K 分块搬移/计算时延 (未合并基准, v1.1 §4.1 符号)
t_load = k_l1 * (m + n) * dt / s.bw_pc
t_comp_chunk = 2.0 * m * n * k_l1 / s.q16
t_write = m * n * out_b / s.bw_pc # 单 batch 输出写回 (单核带宽份额)
if k_truncated:
# K 截断: n_K=1, 合并后单次搬移量 b0 倍
n_move = b_core / b0
t_mte2_data = n_move * b0 * t_load
dma_cmds = n_move
else:
# L1 绑定: k_L1^m = k_L1/b0, n_K^m = b0*n_K, 搬移次数与 IterBatch 相同
n_k = -(-k // k_l1)
n_move = b_core * n_k
t_mte2_data = n_move * t_load
dma_cmds = n_move
t_dma_cmd = dma_cmds * s.t_cmd
t_mte2 = t_mte2_data + t_dma_cmd
# Cube: 合并计算含冗余 (b0^2 输出, 有效 b0) -> 计算量 = b_core*b0*2MNK
flops_pc = b_core * b0 * 2.0 * m * n * k
t_mmad = flops_pc / s.q16
# Fixpipe: 只写对角块, 写出量 = b_core*MN*outB (C 矩阵 dtype, 随路转换)
fix_bytes_pc = b_core * m * n * out_b
t_fix = fix_bytes_pc / s.bw_pc
# drain: 末合并 batch 排空 = b0*(T_comp + T_write)
t_drain = b0 * (t_comp_chunk + t_write)
gm_bytes_pc = b_core * (m * k + k * n) * dt
return assemble_timing(
t_mte2_gm=t_mte2_data, t_mte2_l2=0.0, t_dma_cmd=t_dma_cmd,
t_mmad=t_mmad, t_fixpipe=t_fix, t_reduce=0.0, t_drain=t_drain,
gm_read_bytes=gm_bytes_pc, l2_read_bytes=0.0,
dma_cmd_count=dma_cmds, cube_flops=flops_pc,
fixpipe_bytes=fix_bytes_pc,
)
# ------------------------------------------------------------------
# MergeBatch vs IterBatch 分界 (v1.1 §4.5 统一分界条件)
# ------------------------------------------------------------------
def beats_iterbatch(self, case: BmmCase) -> tuple:
"""返回 (MergeBatch是否更优, 说明).
MergeBatch 最优 ⟺ K截断 (k_L1=K) 且 b_core > b0*(T_comp+T_write)/T_cmd
L1 绑定情形 MergeBatch 恒劣于 IterBatch (搬移次数相同, 只放大 drain).
"""
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
out_b = case.dtype_out_bytes
b_core = case.batch_c // s.aic_num
# IterBatch 基准的 k_L1 (未合并): L1 双缓冲单 batch
k_l1_iter = min(k, s.l1_bytes / (2 * (m + n) * dt))
k_truncated = k_l1_iter >= k
plan = self.make_plan(case)
b0 = plan.merge_b0
t_comp = 2.0 * m * n * min(k_l1_iter, k) / s.q16
t_write = m * n * out_b / s.bw_pc
threshold = b0 * (t_comp + t_write) / s.t_cmd
win = k_truncated and (b_core > threshold)
detail = (f"k_L1={'K(截断)' if k_truncated else f'{k_l1_iter:.0f}<K(L1绑定)'}; "
f"b_core={b_core} vs 阈值 b0*(T_comp+T_write)/T_cmd={threshold:.1f}; "
f"drain惩罚=(b0-1)*(T_comp+T_write)={((b0-1)*(t_comp+t_write))*1e6:.2f}us, "
f"搬移节省=b_core*(1-1/b0)*T_cmd={b_core*(1-1/b0)*s.t_cmd*1e6:.2f}us")
return win, detail
# ------------------------------------------------------------------
@staticmethod
def _factor_floor(b_core: int, b0_max: int) -> int:
"""取不超过 b0_max 的 b_core 的最大因子 (>=MIN_B0), 无则返回 MIN_B0."""
best = MIN_B0
for d in range(MIN_B0, b0_max + 1):
if b_core % d == 0:
best = d
return best