Add BMM_Theory: bmm_theory/branches/merge_batch.py
This commit is contained in:
230
BMM/BMM_Theory/bmm_theory/branches/merge_batch.py
Normal file
230
BMM/BMM_Theory/bmm_theory/branches/merge_batch.py
Normal 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
|
||||
Reference in New Issue
Block a user