diff --git a/BMM/BMM_Theory/bmm_theory/branches/stream_k.py b/BMM/BMM_Theory/bmm_theory/branches/stream_k.py new file mode 100644 index 0000000..9a520dd --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/branches/stream_k.py @@ -0,0 +1,164 @@ +"""StreamK 分支: 核间切 K + 归约, B/M/N 并行度填不满核时的最后手段. + +理论依据: 《BMM算子优化分析 v0.98》§七 + docs/02_分支理论/05_StreamK分支.md. + +切 K 是唯一同时破坏"累加不出核"和"输出独占"的切法: +部分和必须写出 workspace (驻留 L2, 防精度丢失按 4B) 再由 AIV 归约. +""" + +from __future__ import annotations + +import math + +from ..hardware import NpuSpec, ASCEND950PR +from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div +from ..timing import assemble_timing, eval_streamk_reduce +from .base import Branch, ConditionCheck + + +class StreamKBranch(Branch): + name = "StreamK" + + def __init__(self, spec: NpuSpec = ASCEND950PR): + super().__init__(spec) + + # ------------------------------------------------------------------ + def _p_value(self, case: BmmCase) -> float: + """P = B*M*N*4B / L0C: 以 L0C 满载为基本块粒度的不切 K 最大并行度.""" + return case.batch_c * case.m * case.n * 4 / self.spec.l0c_bytes + + def _grid_k(self, case: BmmCase) -> int: + """grid_K = floor(C / ceil(P)), 最大化 K 并行度.""" + p = self._p_value(case) + return int(self.spec.aic_num // max(1, math.ceil(p))) + + def _theta_c(self) -> float: + """归约代价系数 theta_c = Q16/2 * (8B/W_L2 + 1/Q_AIV) ≈ 12.""" + s = self.spec + return s.q16 / 2 * (8 / s.bw_l2 + 1 / s.q_aiv) + + # ------------------------------------------------------------------ + def check_conditions(self, case: BmmCase) -> list: + s = self.spec + p = self._p_value(case) + grid_k = self._grid_k(case) + dt = case.dtype_in_bytes + + checks = [] + + # 条件 1: P <= C/2 (并行缺口, K 是唯一剩余并行维度) + c1 = p <= s.aic_num / 2 + checks.append(ConditionCheck( + "1_并行缺口: P=B*MN*4B/L0C <= C/2", + c1, f"P={p:.2f} vs C/2={s.aic_num/2:.0f}")) + + # 条件 2: K/grid_K >= 256B/dtype (单核 K 段下限, dValue) + seg = case.k / max(grid_k, 1) + c2 = seg >= s.dvalue_hw_min / dt and grid_k >= 2 + checks.append(ConditionCheck( + "2_单核K段下限: K/grid_K >= 256B/dtype 且 grid_K>=2", + c2, f"K/grid_K={seg:.0f} vs {s.dvalue_hw_min/dt:.0f}, grid_K={grid_k}")) + + # 条件 3: K > grid_K^2/(grid_K-1) * theta_c (归约代价可接受) + if grid_k >= 2: + theta_c = self._theta_c() + k_thresh = grid_k * grid_k / (grid_k - 1) * theta_c + c3 = case.k > k_thresh + checks.append(ConditionCheck( + "3_归约代价可接受: K > grid_K^2/(grid_K-1)*theta_c", + c3, f"K={case.k} vs 阈值={k_thresh:.1f} (theta_c={theta_c:.1f})")) + else: + checks.append(ConditionCheck("3_归约代价可接受", False, "grid_K<2 无意义")) + + # 条件 4: 确定性等级 <= 1 且 ND + c4 = case.deterministic_level <= 1 and case.out_nd + checks.append(ConditionCheck( + "4_工程约束: 确定性等级<=1 且 ND 格式", + c4, f"deterministic_level={case.deterministic_level}, out_nd={case.out_nd}")) + + return checks + + # ------------------------------------------------------------------ + def make_plan(self, case: BmmCase) -> ImplPlan: + s = self.spec + b = case.batch_c + m, n, k = case.m, case.n, case.k + out_b = case.dtype_out_bytes + grid_k = self._grid_k(case) + + blocks_per_batch = s.aic_num // b + # mCnt*nCnt 收拢为 blocksPerBatch 的因子, 剩余核预算折成 grid_K + mn_cnt = max(1, blocks_per_batch // grid_k) + m_cnt, n_cnt = self._split_mn(m, n, mn_cnt) + single_m = ceil_div(m, m_cnt) + single_n = ceil_div(n, n_cnt) + single_k = ceil_div(k, grid_k) + + # workspace: 部分和驻留 L2, 按 4B (L0C dtype, 防精度丢失) + workspace = grid_k * single_m * single_n * 4 * b + + 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=grid_k, + core_map=f"B/M/N切出{b*m_cnt*n_cnt}块, 每块{grid_k}核切K归约 (归约组内核c负责K段[c*K/{grid_k},(c+1)*K/{grid_k}))", + b_core=1, merge_b0=1, + single_core_m=single_m, single_core_n=single_n, single_core_k=single_k, + k_l1=min(single_k, s.dvalue_recommend // case.dtype_in_bytes), b_l1=1, + l1_form="K段标准分块流水", + base_m=min(single_m, 128), base_n=min(single_n, 128), base_k=min(single_k, 64), + l2_policy_in="allocate(部分和驻留L2)", + l2_policy_out="resident(部分和4B驻留L2, 防精度丢失不随C的fp16/fp8转换)", + swizzle_w=0, workspace_bytes=workspace, + tail_strategy=f"grid_K={grid_k}路切K+归约", + tail_k_cnt=grid_k, # 尾轮 K 向切分数 = grid_K + fixpipe_unitflag=True, + out_dtype_bytes=4, # 中间部分和按 4B + note=f"P={self._p_value(case):.2f}, grid_K={grid_k}, 部分和驻留L2按4B写出, AIV归约后按C dtype={out_b}B写最终", + ) + + @staticmethod + def _split_mn(m: int, n: int, mn_cnt: int) -> tuple: + """把 mn_cnt 拆成 m_cnt*n_cnt, 长宽比跟随 M/N.""" + if mn_cnt <= 1: + return 1, 1 + ratio = m / max(n, 1) + m_cnt = max(1, round(math.sqrt(mn_cnt * ratio))) + n_cnt = max(1, mn_cnt // m_cnt) + return m_cnt, n_cnt + + # ------------------------------------------------------------------ + def evaluate(self, case: BmmCase, plan: ImplPlan) -> 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 + grid_k = plan.grid_k + + # 单 tile (L0C 满载基本块) + tile_elems = s.l0c_elems # 65536 + t_mmad_tile = 2.0 * k * s.l0c_bytes / 4 / s.q16 # 2K/Q16 * L0C/4B + t_mte2_tile = k * (2 * math.sqrt(tile_elems)) * dt / s.bw_pc # 近似方形 tile + + # 切 K 后流水时延缩 grid_k 倍 + t_mmad = t_mmad_tile / grid_k + t_mte2 = t_mte2_tile / grid_k + + # 归约: 部分和 4B 驻留 L2, AIV 归约 + t_reduce = eval_streamk_reduce(tile_elems, grid_k, out_b, s) + + # 写出: 最终归约结果按 C dtype; 中间部分和按 4B + fix_bytes = tile_elems * out_b + t_fix = fix_bytes / s.bw_l2_pc + + flops_pc = 2.0 * tile_elems * k / grid_k + gm_bytes = k * (2 * math.sqrt(tile_elems)) * dt / grid_k + + 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=t_reduce, + t_drain=t_reduce, # 归约串行追加 + gm_read_bytes=gm_bytes, l2_read_bytes=0.0, dma_cmd_count=0.0, + cube_flops=flops_pc, fixpipe_bytes=fix_bytes, + )