diff --git a/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py b/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py new file mode 100644 index 0000000..d6cb910 --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py @@ -0,0 +1,203 @@ +"""IterBatch 分支: 核间切 B, 核内逐 batch 标准 Matmul 分块计算. + +理论依据: + - 《BMM算子优化分析 v0.98》§六 (进入条件四形态 a/b/c/d + 实现方案) + - 《MergeBatch_vs_IterBatch分析 v1.1》§二/§四 (执行模型 + 分界) + +核心思想: "切B"最朴素形态, 无算力浪费、无跨 batch 依赖. +进入的核心要求是单 batch 计算核内零重复读 (由 L1 驻留四形态刻画); +batch 间流水靠 L1 双 buffer / 驻留侧预取 (c 形态半预算) 掩盖. +""" + +from __future__ import annotations + +from ..hardware import NpuSpec, ASCEND950PR +from ..models import BmmCase, ImplPlan, HardwareTiming, align_down +from ..timing import assemble_timing +from .base import Branch, BranchResult, ConditionCheck + + +class IterBatchBranch(Branch): + name = "IterBatch" + + def __init__(self, spec: NpuSpec = ASCEND950PR): + super().__init__(spec) + + # ------------------------------------------------------------------ + # L1 驻留形态判定 (v0.98 §六 条件 3, 四选一) + # ------------------------------------------------------------------ + def l1_form(self, case: BmmCase) -> tuple: + """返回 (形态 'a'/'b'/'c'/'d'/None, k_l1, 说明).""" + s = self.spec + m, n, k = case.m, case.n, case.k + dt = case.dtype_in_bytes + b_core = -(-case.batch_c // s.aic_num) # ceil + single = (m * k + k * n) * dt + + # a) 单 batch 全驻留 + if b_core == 1 and single <= s.l1_bytes: + return "a", k, f"单batch全驻留: (MK+KN)*dtype={single/1024:.0f}KB <= L1" + + # b) 双 batch 乒乓 + if b_core > 1 and 2 * single <= s.l1_bytes: + return "b", k, f"双batch乒乓: 2*(MK+KN)*dtype={2*single/1024:.0f}KB <= L1" + + # c) 一侧驻留 + 对侧切 K, 预算按 b_core 分档 + l1_budget = s.l1_bytes / min(b_core, 2) + for resident, side in ((m * k * dt, "A"), (k * n * dt, "B")): + other = n * dt if side == "A" else m * dt + if resident <= l1_budget: + k_l1 = min(int((l1_budget - resident) / other / 2), k) + k_l1 = align_down(max(k_l1, s.fractal), s.fractal) + if k_l1 >= s.fractal and k_l1 * dt >= s.dvalue_min: + return "c", k_l1, ( + f"一侧驻留({side})+对侧切K: {side}驻留{resident/1024:.0f}KB, " + f"预算L1/{min(b_core,2)}, k_L1={k_l1}") + # b_core>=2 时另一半 L1 预取下一 batch 驻留侧, 边界无气泡 + + # d) 两侧都切 K (兜底) + k_l1 = align_down(int(s.l1_bytes / (2 * (m + n) * dt)), s.fractal) + if k_l1 >= s.fractal and k_l1 * dt >= s.dvalue_min: + return "d", k_l1, f"两侧都切K: k_L1={k_l1}, K段成对流水, batch边界天然无缝" + + return None, 0, "L1 四形态均不满足 (M/N 相对 L1 过大)" + + # ------------------------------------------------------------------ + # 进入条件 (v0.98 §六, 四条同时满足) + # ------------------------------------------------------------------ + def check_conditions(self, case: BmmCase) -> list: + s = self.spec + b = case.batch_c + b_core = -(-b // s.aic_num) + + checks = [] + + c1 = (case.batch_a == case.batch_b) and (b_core >= 1) + checks.append(ConditionCheck( + "1_batch关系与每核份额: BatchA==BatchB 且 b_core>=1", + c1, f"batchA={case.batch_a}, batchB={case.batch_b}, b_core={b_core}")) + + # 条件 2: 负载均衡 (整除 或 尾波活跃核数 >= minCoreNum) + rem = b % s.aic_num + c2 = (rem == 0) or (rem >= s.min_core_num) + checks.append(ConditionCheck( + "2_负载均衡: B mod C == 0 或 >= minCoreNum", + c2, f"B mod C={rem}, minCoreNum={s.min_core_num}")) + + form, k_l1, form_desc = self.l1_form(case) + checks.append(ConditionCheck( + "3_L1驻留形态(四选一, 核心要求: 单batch核内零重复读)", + form is not None, form_desc)) + + # 条件 4: 搬移效率下限 (c/d 切分后) + m, n, k = case.m, case.n, case.k + dt = case.dtype_in_bytes + if form in ("c", "d"): + tile_ok = (k_l1 * m * dt >= s.min_tile_size) or (k_l1 * n * dt >= s.min_tile_size) + dv_ok = k_l1 * dt >= s.dvalue_min + c4 = tile_ok and dv_ok + checks.append(ConditionCheck( + "4_搬移效率: 搬移分块>=min_TileSize 且 dValue>=128B", + c4, f"tile={max(k_l1*m*dt, k_l1*n*dt)/1024:.1f}KB, dValue={k_l1*dt}B")) + else: + checks.append(ConditionCheck( + "4_搬移效率(a/b形态不切分, 恒满足)", True, "")) + + return checks + + # ------------------------------------------------------------------ + # 实现方案 (v0.98 §六 实现方案 a/b/c/d) + # ------------------------------------------------------------------ + 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_core = -(-case.batch_c // s.aic_num) + form, k_l1, form_desc = self.l1_form(case) + form = form or "d" + + # L0 级 tile (v0.98 §六 (a) 伪代码): + # L0C 放得下完整输出 -> BaseM=M, BaseN=N, BaseK 由 L0A/L0B 决定 (双缓冲, 故除 2); + # 放不下 -> 按较小维切, 再定 BaseK. + # 注意 L0A/L0B 双缓冲两份, 且 base_k 不超 k_l1 (L1 tile 的 K 粒度) 与 K 本身. + if s.l0c_bytes >= m * n * 4 * 2: + base_m, base_n = m, n + else: + if m < n: + base_m = align_down(max(m, s.fractal), s.fractal) + base_n = align_down(int(s.l0c_bytes / (2 * 4 * base_m)), s.fractal) + else: + base_n = align_down(max(n, s.fractal), s.fractal) + base_m = align_down(int(s.l0c_bytes / (2 * 4 * base_n)), s.fractal) + base_m = max(base_m, s.fractal) + base_n = max(base_n, s.fractal) + base_k = align_down(int(min( + s.l0a_bytes / (2 * base_m * dt), + s.l0b_bytes / (2 * base_n * dt), + )), s.fractal) + base_k = max(min(base_k, k_l1 if k_l1 > 0 else k, k), s.fractal) + + form_name = {"a": "a_单batch全驻留", "b": "b_双batch乒乓", + "c": "c_一侧驻留+对侧切K", "d": "d_两侧都切K"}[form] + + 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=1, + single_core_m=m, single_core_n=n, single_core_k=k, + k_l1=k_l1, b_l1=2 if form == "b" else 1, l1_form=form_name, + base_m=base_m, base_n=base_n, base_k=base_k, + l2_policy_in="allocate(GM->L1随路驻留L2)", + l2_policy_out="direct_gm(输出仅写一次,直写GM不占L2)", + swizzle_w=0, workspace_bytes=0, + tail_strategy="不涉及(核内不切M/N)", + fixpipe_unitflag=True, + out_dtype_bytes=case.dtype_out_bytes, + note=form_desc, + ) + + # ------------------------------------------------------------------ + # 时延评估 (v1.1 §4 端到端模型, IterBatch 侧) + # ------------------------------------------------------------------ + 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, k_l1 = plan.b_core, plan.k_l1 + + k_truncated = k_l1 >= k + n_k = 1 if k_truncated else -(-k // k_l1) + + t_load = min(k_l1, k) * (m + n) * dt / s.bw_pc + t_comp_chunk = 2.0 * m * n * min(k_l1, k) / s.q16 + t_write = m * n * out_b / s.bw_pc + + # 搬移: 每 batch n_K 次 GM->L1, 每次含 T_cmd + dma_cmds = b_core * n_k + t_mte2_data = dma_cmds * t_load + t_dma_cmd = dma_cmds * s.t_cmd + t_mte2 = t_mte2_data + t_dma_cmd + + # Cube: 无冗余 + flops_pc = b_core * 2.0 * m * n * k + t_mmad = flops_pc / s.q16 + + # Fixpipe: b_core 个 batch 输出, 按 C dtype + fix_bytes_pc = b_core * m * n * out_b + t_fix = fix_bytes_pc / s.bw_pc + + # drain: 末 batch 排空 = T_comp + T_write + t_drain = 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, + )