diff --git a/BMM/BMM_Theory/bmm_theory/router.py b/BMM/BMM_Theory/bmm_theory/router.py new file mode 100644 index 0000000..07e4d44 --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/router.py @@ -0,0 +1,112 @@ +"""分支决策路由: 按决策树推导分支, 重叠区由端到端时延模型仲裁. + +决策树 (v0.98 §3.3): + 1. 前置归约: BatchA=1 或 BatchB=1 -> 转Matmul (本期仅标注, 详实现后续迭代) + K=0 / K=1 -> 特殊分支 (本期仅标注) + 2. B >= C: 切B -> IterBatch 与 MergeBatch 仲裁 + 仲裁规则 (v1.1 §4.5 统一分界): + MergeBatch 最优 <=> K截断(k_L1=K) 且 b_core > b0*(T_comp+T_write)/T_cmd + L1 绑定时 MergeBatch 恒劣于 IterBatch; + 两分支同时合法时用端到端时延模型 T_total 仲裁. + 3. B < C: 切 M/N -> ASW_Basic (后续迭代) + 4. B/M/N 都买不满: 切 K -> StreamK (后续迭代) +""" + +from __future__ import annotations + +from .hardware import NpuSpec, ASCEND950PR +from .models import BmmCase, ImplPlan, HardwareTiming +from .branches.base import BranchResult +from .branches.merge_batch import MergeBatchBranch +from .branches.iter_batch import IterBatchBranch + +BRANCH_TO_MATMUL = "转Matmul" +BRANCH_SPECIAL = "特殊分支" +BRANCH_ASW = "ASW_Basic" +BRANCH_STREAMK = "StreamK" + + +class BranchRouter: + """case -> 理论最优分支 + 方案 + 时延评估.""" + + def __init__(self, spec: NpuSpec = ASCEND950PR): + self.spec = spec + self.merge_batch = MergeBatchBranch(spec) + self.iter_batch = IterBatchBranch(spec) + + # ------------------------------------------------------------------ + def route(self, case: BmmCase) -> dict: + """返回 {branch, plan, timing, arbitration, candidates}.""" + s = self.spec + + # 1) 前置归约 + if case.batch_a == 1 or case.batch_b == 1: + return self._stub(case, BRANCH_TO_MATMUL, + "BatchA=1或BatchB=1, 折叠转普通Matmul (该分支详实现待后续迭代)") + if case.k <= 1: + return self._stub(case, BRANCH_SPECIAL, + f"K={case.k}, Cube 无用, 走 AIV 向量通路 (该分支详实现待后续迭代)") + + # 2) B >= C: IterBatch / MergeBatch + if case.batch_c >= s.aic_num and case.batch_a == case.batch_b: + return self._route_split_b(case) + + # 3/4) 兜底标注 + return self._stub(case, BRANCH_ASW, + "B ASW_Basic/StreamK (该分支详实现待后续迭代)") + + # ------------------------------------------------------------------ + def _route_split_b(self, case: BmmCase) -> dict: + mb = self.merge_batch.analyze(case) + ib = self.iter_batch.analyze(case) + + candidates = [] + if mb.capable: + candidates.append((self.merge_batch.name, mb)) + if ib.capable: + candidates.append((self.iter_batch.name, ib)) + + arbitration = "" + if mb.capable and ib.capable: + # 统一分界条件 + 端到端时延仲裁双保险 + mb_win, detail = self.merge_batch.beats_iterbatch(case) + t_mb = mb.timing.t_total + t_ib = ib.timing.t_total + lat_win = self.merge_batch.name if t_mb <= t_ib else self.iter_batch.name + win = self.merge_batch.name if mb_win else self.iter_batch.name + arbitration = ( + f"两分支均合法, 仲裁:\n" + f" [分界条件] MergeBatch最优={mb_win} ({detail})\n" + f" [时延模型] T_MergeBatch={t_mb*1e6:.2f}us vs T_IterBatch={t_ib*1e6:.2f}us " + f"-> {lat_win}更优\n" + f" [裁决] {win}" + ("" if win == lat_win else f" (分界条件与时延模型不一致, 以时延模型为准: {lat_win})") + ) + if win != lat_win: + win = lat_win # 时延模型为最终裁决 + elif candidates: + win = candidates[0][0] + arbitration = f"仅 {win} 条件满足" + else: + return self._stub( + case, BRANCH_ASW, + "IterBatch/MergeBatch 进入条件均不满足 (如负载均衡/搬移效率不达标), " + "回落 ASW_Basic (该分支详实现待后续迭代); " + f"IterBatch未过: {ib.failed_conditions()}; " + f"MergeBatch未过: {mb.failed_conditions()}") + + chosen = mb if win == self.merge_batch.name else ib + return { + "branch": win, + "plan": chosen.plan, + "timing": chosen.timing, + "arbitration": arbitration, + "candidates": {n: r.capable for n, r in + [(self.merge_batch.name, mb), (self.iter_batch.name, ib)]}, + } + + # ------------------------------------------------------------------ + def _stub(self, case: BmmCase, branch: str, note: str) -> dict: + plan = ImplPlan(case_id=case.case_id, branch=branch, + used_core_num=self.spec.aic_num, note=note) + return {"branch": branch, "plan": plan, "timing": None, + "arbitration": note, "candidates": {}}