From 6a24b5d945594b00dde8c8660f0f37a89f3a7183 Mon Sep 17 00:00:00 2001 From: admin Date: Thu, 3 Sep 2026 09:25:40 +0000 Subject: [PATCH] Add BMM_Theory: bmm_theory/branches/asw_basic.py --- .../bmm_theory/branches/asw_basic.py | 289 ++++++++++++++++++ 1 file changed, 289 insertions(+) create mode 100644 BMM/BMM_Theory/bmm_theory/branches/asw_basic.py diff --git a/BMM/BMM_Theory/bmm_theory/branches/asw_basic.py b/BMM/BMM_Theory/bmm_theory/branches/asw_basic.py new file mode 100644 index 0000000..c996e33 --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/branches/asw_basic.py @@ -0,0 +1,289 @@ +"""ASW_Basic 分支: 核间切 M/N (或混合切) 的兜底分支, 含尾轮处理. + +理论依据: + - 《BMM算子优化分析 v0.98》§八 + docs/02_分支理论/06_ASW_Basic分支.md + - 《BMM尾轮处理策略对比分析 v1.5》+ docs/02_分支理论/07_尾轮处理策略.md + +核心: 兜底分支, 核间切 M/N, 重复读交给 L2 + swizzle. 尾轮处理是必要组成环节: +默认方案 B (工程简洁, 主流场景与 A1b 严格打平), 周长型且 rho>=rho_dv 时 A1b. +""" + +from __future__ import annotations + +import math + +from ..hardware import NpuSpec, ASCEND950PR +from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div, align_down +from ..timing import assemble_timing +from .base import Branch, ConditionCheck + + +class AswBasicBranch(Branch): + name = "ASW_Basic" + + def __init__(self, spec: NpuSpec = ASCEND950PR): + super().__init__(spec) + + # ------------------------------------------------------------------ + def _p_value(self, case: BmmCase) -> float: + return case.batch_c * case.m * case.n * 4 / self.spec.l0c_bytes + + # ------------------------------------------------------------------ + def check_conditions(self, case: BmmCase) -> list: + s = self.spec + p = self._p_value(case) + + checks = [] + # 条件 1: P >= C (并行度补齐); P < C 时看是否走降核模式 + if p >= s.aic_num: + checks.append(ConditionCheck( + "1_并行度补齐: P=B*MN*4B/L0C >= C", True, f"P={p:.2f} >= C={s.aic_num}")) + else: + # 降核模式: P < C 且不满足 StreamK —— 此处由 router 保证 (StreamK 优先) + checks.append(ConditionCheck( + "1_降核模式: P ImplPlan: + s = self.spec + b = case.batch_c + m, n, k = case.m, case.n, case.k + dt = case.dtype_in_bytes + p = self._p_value(case) + + # 降核模式 + if p < s.aic_num: + return self._plan_reduced_core(case, p) + + # ---- 正常模式: Step 0~6 ---- + # Step 0: BaseM/BaseN 用满 L0C (32768 元素双缓冲) + base_mn = int(math.sqrt(s.l0c_bytes // 8)) # L0C/(2*4B) + base_m = align_down(base_mn, s.fractal) + base_n = align_down(base_mn, s.fractal) + base_k = align_down(int(min( + s.l0a_bytes / (2 * base_m * dt), + s.l0b_bytes / (2 * base_n * dt), + )), s.fractal) + + # Step 1: SingleCoreM/N 尽量大 (满足并行度下限) + min_blocks = ceil_div(s.aic_num, b) + single_m, single_n = self._pick_single_core(m, n, min_blocks, base_m, base_n, dt) + + # Step 2: mCnt/nCnt + m_cnt = ceil_div(m, single_m) + n_cnt = ceil_div(n, single_n) + + # Step 4: swizzle 窗口 W = max{d | d|C, d <= floor(sqrt(C))} + swizzle_w = self._swizzle_w() + + # Step 5: L2 分组判断 + s_in = b * (m * k + k * n) * dt + s_out = b * m * n * case.dtype_out_bytes + if s_in + s_out <= s.l2_bytes: + l2_scene = "A_全驻留" + l2_out = "resident(输出驻留L2异步回写)" + r_in = 1.0 + elif s_in <= s.l2_bytes: + l2_scene = "B_输入驻留输出直写GM" + l2_out = "direct_gm(输出直写GM不占L2)" + r_in = 1.0 + else: + l2_scene = "C_输入超L2分组执行" + l2_out = "direct_gm(输出直写GM不占L2)" + r_in = self._l2_group_r_in(case, single_m, single_n) + + # Step 6: k_l1 (GM->L1 K 向粒度, 须 >= 256B/dt) + k_l1 = min(k, s.dvalue_recommend // dt) + k_l1 = max(k_l1, s.dvalue_hw_min // dt) + + # ---- 尾轮决策 (必要组成环节) ---- + n_blk = b * m_cnt * n_cnt + tail = self._decide_tail(case, single_m, single_n, n_blk, k_l1) + + return ImplPlan( + case_id=case.case_id, branch=self.name, + used_core_num=s.aic_num, + split_b=1, m_cnt=m_cnt, n_cnt=n_cnt, grid_k=1, + core_map=f"B->M->N线性映射+ASW滑窗蛇形(W={swizzle_w})", + b_core=0, merge_b0=1, + single_core_m=single_m, single_core_n=single_n, single_core_k=k, + k_l1=k_l1, b_l1=1, l1_form="双缓冲驻留当前tile输入", + base_m=base_m, base_n=base_n, base_k=base_k, + l2_policy_in="allocate(输入驻留L2吸收重复读)", + l2_policy_out=l2_out, + swizzle_w=swizzle_w, workspace_bytes=0, + tail_strategy=tail["strategy"], + tail_m_cnt=tail["tail_m_cnt"], tail_n_cnt=tail["tail_n_cnt"], + tail_k_cnt=1, tail_m_main=tail["tail_m_main"], tail_n_main=tail["tail_n_main"], + tail_block_cnt=tail["r"], tail_wave_num=tail["n_wave"], + fixpipe_unitflag=True, + out_dtype_bytes=case.dtype_out_bytes, + note=f"L2场景{l2_scene}, r_in={r_in:.2f}; 尾轮: {tail['reason']}", + ) + + # ------------------------------------------------------------------ + def _plan_reduced_core(self, case: BmmCase, p: float) -> ImplPlan: + """降核模式: 只用 ceil(P) 核, 每核一个 L0C 满载输出块.""" + s = self.spec + used = max(1, math.ceil(p)) + # SingleCoreM/N 在 L0C 容量内取最大 + single_mn = int(math.sqrt(s.l0c_bytes // 4)) # L0C/4B 单份 + single_m = min(case.m, align_down(single_mn, s.fractal)) + single_n = min(case.n, align_down(single_mn, s.fractal)) + return ImplPlan( + case_id=case.case_id, branch=self.name + "_降核", + used_core_num=used, + split_b=1, m_cnt=ceil_div(case.m, single_m), n_cnt=ceil_div(case.n, single_n), + grid_k=1, + core_map=f"降核: 只用{used}核, 每核一个L0C满载输出块, 其余核闲置", + b_core=0, merge_b0=1, + single_core_m=single_m, single_core_n=single_n, single_core_k=case.k, + k_l1=min(case.k, s.dvalue_recommend // case.dtype_in_bytes), b_l1=1, + l1_form="标准核内流水", + base_m=single_m, base_n=single_n, + base_k=min(case.k, 64), + l2_policy_in="allocate", l2_policy_out="direct_gm", + swizzle_w=0, workspace_bytes=0, + tail_strategy="不涉及(每核一块无尾轮)", + fixpipe_unitflag=True, out_dtype_bytes=case.dtype_out_bytes, + note=f"P={p:.2f}= min_blocks), 取最大 tile + best = (base_m, base_n) + for m_cnt in range(1, int(math.sqrt(min_blocks)) + 2): + n_cnt = ceil_div(min_blocks, m_cnt) + if m_cnt * n_cnt < min_blocks: + continue + sm = align_down(ceil_div(m, m_cnt), base_m) or base_m + sn = align_down(ceil_div(n, n_cnt), base_n) or base_n + # 约束 2: L1 容量 2(sM+sN)*k_l1*dt <= L1, k_l1 取 256B/dt + k_l1_min = s.dvalue_hw_min // dt + if 2 * (sm + sn) * k_l1_min * dt > s.l1_bytes: + continue + if sm * sn > best[0] * best[1]: + best = (sm, sn) + return best + + def _swizzle_w(self) -> int: + s = self.spec + w = 1 + for d in range(1, int(math.sqrt(s.aic_num)) + 1): + if s.aic_num % d == 0: + w = d + return w + + def _l2_group_r_in(self, case, sm, sn) -> float: + """场景 C: L2 分组的重复读倍率 r_in = (n_grp*M + m_grp*N)/(M+N).""" + s = self.spec + m, n, k, b = case.m, case.n, case.k, case.batch_c + dt = case.dtype_in_bytes + d = s.l2_bytes / (b * k * dt) + m_grp = max(1, int(d / (2 * sm))) + n_grp = max(1, int(d / (2 * sn))) + m_cnt = ceil_div(m, sm) + n_cnt = ceil_div(n, sn) + return (ceil_div(n_cnt, n_grp) * m + ceil_div(m_cnt, m_grp) * n) / (m + n) + + # ------------------------------------------------------------------ + def _decide_tail(self, case, sm, sn, n_blk, k_l1) -> dict: + """尾轮策略决策 (v1.5 闭式流程). 默认方案 B, 周长型且 rho>=rho_dv 时 A1b.""" + s = self.spec + c = s.aic_num + n_wave = ceil_div(n_blk, c) + r = n_blk % c + rho = r / c + dt = case.dtype_in_bytes + out_b = case.dtype_out_bytes + + base = dict(r=r, n_wave=n_wave, tail_m_cnt=1, tail_n_cnt=1, + tail_m_main=0, tail_n_main=0) + + if r == 0: + return {**base, "strategy": "A0", "reason": "r=0 无尾轮"} + + # 主导项判定 + bw_eff = s.bw_l2_pc # L2 命中 + t_mmad = 2 * sm * sn * case.k / s.q16 + t_mte2 = case.k * (sm + sn) * dt / bw_eff + t_fix = sm * sn * out_b / s.bw_pc + t_block = max(t_mmad, t_mte2, t_fix) + area_dominated = t_block != t_mte2 # 面积型 = MMAD 或 FIX 主导 + + # 方案 B: 全局均匀重切, g = n_wave*C/N_blk + g = n_wave * c / n_blk + s_b = int(math.sqrt(sm * sn / g)) + s_b = align_down(max(s_b, s.fractal), s.fractal) + plan_b_m_cnt = ceil_div(case.m, s_b) if case.m >= s_b else 1 + plan_b_n_cnt = ceil_div(case.n, s_b) if case.n >= s_b else 1 + + if area_dominated: + # 面积型: 方案 B 与 A1b 严格打平, 默认方案 B (工程简洁); r 小翻出时方案 B 微优 + return {**base, "strategy": "方案B", + "tail_m_cnt": plan_b_m_cnt, "tail_n_cnt": plan_b_n_cnt, + "tail_m_main": plan_b_m_cnt, "tail_n_main": plan_b_n_cnt, + "reason": f"面积型主导, A1b与方案B严格打平(T={t_block*(n_wave-1+rho)*1e6:.2f}us), " + f"选方案B工程简洁 (一套tile)"} + else: + # 周长型: rho >= rho_dv 时 A1b 恒优 + rho_dv = (s.dvalue_hw_min / (sn * dt)) ** 2 + if rho >= rho_dv: + # A1b: 尾轮 tile 缩 sqrt(rho) + s_t = align_down(int(sn * math.sqrt(rho)), s.fractal) or s.fractal + t_m = ceil_div(sm, s_t) if sm >= s_t else 1 + t_n = ceil_div(sn, s_t) if sn >= s_t else 1 + return {**base, "strategy": "A1b", + "tail_m_cnt": t_m, "tail_n_cnt": t_n, + "tail_m_main": t_m, "tail_n_main": t_n, + "reason": f"周长型主导, rho={rho:.2f}>=rho_dv={rho_dv:.2f}, A1b恒优 " + f"(尾轮tile缩√rho={math.sqrt(rho):.2f}倍凑满核)"} + else: + return {**base, "strategy": "方案B", + "tail_m_cnt": plan_b_m_cnt, "tail_n_cnt": plan_b_n_cnt, + "tail_m_main": plan_b_m_cnt, "tail_n_main": plan_b_n_cnt, + "reason": f"周长型主导, rho={rho:.2f} HardwareTiming: + s = self.spec + b = case.batch_c + m, n, k = case.m, case.n, case.k + dt = case.dtype_in_bytes + out_b = case.dtype_out_bytes + + flops = 2.0 * b * m * n * k + t_mmad = flops / (plan.used_core_num * s.q16) + + in_bytes = b * (m * k + k * n) * dt + out_bytes = b * m * n * out_b + # r_in 从 note 里解析困难, 重新计算 + s_in = in_bytes + s_out = out_bytes + if s_in + s_out <= s.l2_bytes or s_in <= s.l2_bytes: + r_in = 1.0 + else: + r_in = self._l2_group_r_in(case, plan.single_core_m, plan.single_core_n) + + t_mte2 = r_in * in_bytes / (plan.used_core_num * s.bw_pc) + to_l2 = "resident" in plan.l2_policy_out + t_fix = out_bytes / (plan.used_core_num * (s.bw_l2_pc if to_l2 else s.bw_pc)) + + # drain: 尾轮暴露 (方案 B 已均匀重切, drain 小; A1b 尾轮凑满, drain 小; A0 尾轮 r 核空转) + t_drain = 0.0 + if plan.tail_strategy == "A0" and plan.tail_block_cnt > 0: + t_drain = max(t_mmad, t_mte2, t_fix) # 尾轮空转一个整块 + + return assemble_timing( + t_mte2_gm=t_mte2, t_mte2_l2=0.0, t_dma_cmd=0.0, + t_mmad=t_mmad, t_fixpipe=t_fix, t_reduce=0.0, t_drain=t_drain, + gm_read_bytes=r_in * in_bytes, l2_read_bytes=0.0, dma_cmd_count=0.0, + cube_flops=flops, fixpipe_bytes=out_bytes, + )