Add BMM_Theory: bmm_theory/router.py
This commit is contained in:
112
BMM/BMM_Theory/bmm_theory/router.py
Normal file
112
BMM/BMM_Theory/bmm_theory/router.py
Normal file
@@ -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<C 或切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": {}}
|
||||
Reference in New Issue
Block a user