diff --git a/BMM/BMM_Theory/bmm_theory/branches/merge_batch.py b/BMM/BMM_Theory/bmm_theory/branches/merge_batch.py new file mode 100644 index 0000000..d9a08fc --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/branches/merge_batch.py @@ -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} 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