用户澄清: MergeBatch vs IterBatch 的本质区别不只是 DMA 命令数 —— 合并 b0 个 batch 的左/右矩阵一起搬入 L1, 使单块 tile = nValue*dValue*dt 放大 b0 倍 (堆叠 方向视转置: A ND 非转置沿 M(nValue), B ND 非转置沿 N(dValue)), 搬移效率更高, 即便 T_cmd=0 也有效益。 - models.move_eff: 单命令搬移效率 eff = min(1, tile/min_TileSize) (16KB 饱和, 与进入条件4效率下限语义同源); gm_move_time 按 A/B 两侧字节加权 t = (V_A/eff_A + V_B/eff_B)/BW_gm; 只影响时间列, GM 字节量仍 = V_in - IterBatch: l1_form 补驻留侧返回; move_tiles 分侧口径 (a/b 双侧整K, c 驻留侧 整K+对侧k_l1, d 双侧k_l1), evaluate 接入效率加权 - MergeBatch: 合并 tile 放大 b0 倍接入效率加权; beats_iterbatch 净收益 = 命令节省(cmds差×T_cmd) + 效率节省(t_data差) − drain惩罚, K截断且效率打平且 T_cmd>0 时严格退化为 v1.1 §4.5 闭式; 退役 T_cmd<=0 策略特判 - router: 退役 "T_cmd<=0 策略优先 MergeBatch" 覆盖, 时延模型统一终审 - hardware: t_cmd_ns 50 -> 0 (未标定按 0; 合并收益不再依赖 T_cmd 估计值) - 作用域: 仅切B 两分支接入 (逐命令 tile 小、效率差显著); ASW/StreamK 单命令 tile 通常已饱和, 极端小 tile 走 issue#34 效率降级标注通道 - 用户 case 家族 B=128,M=1~16,N=128,K=512: m=1~8 -> MergeBatch (效率节省 ~0.61us > drain), m=16 -> IterBatch (iter A tile 恰达 16KB 饱和, 效率打平, drain 决定); 分界与时延全家族一致 - demo: merge_demo_k_trunc 形状 (2048,32,32,256)->(2048,16,64,128) (原形状 两侧 tile 均已 16KB 饱和, t_cmd=0 下无收益转 IterBatch; 新形状 iter A tile 4KB eff=0.25 vs 合并 16KB eff=1.0, 保持 MergeBatch 胜出演示且仍 K截断) - 测试: 74/74 (新增 TestIssue36 5 例: 效率曲线/字节不变/效率差胜出/家族; TestArbitration/TestZeroCmdHandling 按 t_cmd=0+效率语义重写; TestIssue35 家族期望更新) - 文档: 01_MergeBatch §4/§5 效率模型+泛化净收益; 02_IterBatch 口径注; 00_总纲胜出条件; 01_软件架构 T_cmd 标定说明; 05 时间列效率口径注; README 要点 - 验证: examples 重生成可复现 0 diff; 压力 10000 例 0 崩溃/0 NaN/0 违规/ 0 GM<V_in, 七分支覆盖 (MergeBatch 386 例)
189 lines
9.8 KiB
Python
189 lines
9.8 KiB
Python
"""分支决策路由: 按决策树推导分支, 重叠区由端到端时延模型仲裁.
|
||
|
||
决策树 (v0.98 §3.3 + 各分支理论文档):
|
||
1. 前置归约: BatchA=1 或 BatchB=1 -> 转Matmul
|
||
K=0 / K=1 -> 特殊分支 (AIV 向量通路)
|
||
2. B >= C 且 BatchA==BatchB: 切B -> IterBatch 与 MergeBatch 仲裁
|
||
仲裁规则 (v1.1 §4.5 统一分界 + issue#35/#36 泛化):
|
||
净收益 = 命令节省(cmds 差 × T_cmd) + 搬移效率节省(合并 tile 放大 b0 倍,
|
||
move_eff 模型) − drain 惩罚; K截断且效率打平时退化为文档闭式
|
||
b_core > b0*(T_comp+T_write)/T_cmd;
|
||
两分支同时合法时用端到端时延模型 T_total 终审 (T_cmd 默认 0, 未标定;
|
||
合并效率收益已建模, 不再需要策略覆盖).
|
||
3. StreamK 检查: P <= C/2 且满足切K条件 -> StreamK (B/M/N 买不满时买 K)
|
||
4. 兜底: ASW_Basic 切 M/N (含降核模式)
|
||
"""
|
||
|
||
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
|
||
from .branches.to_matmul import ToMatmulBranch
|
||
from .branches.special import SpecialBranch
|
||
from .branches.stream_k import StreamKBranch
|
||
from .branches.asw_basic import AswBasicBranch
|
||
|
||
|
||
class BranchRouter:
|
||
"""case -> 理论最优分支 + 方案 + 时延评估."""
|
||
|
||
def __init__(self, spec: NpuSpec = ASCEND950PR):
|
||
self.spec = spec
|
||
self.merge_batch = MergeBatchBranch(spec)
|
||
self.iter_batch = IterBatchBranch(spec)
|
||
self.to_matmul = ToMatmulBranch(spec)
|
||
self.special = SpecialBranch(spec)
|
||
self.stream_k = StreamKBranch(spec)
|
||
self.asw_basic = AswBasicBranch(spec)
|
||
|
||
# ------------------------------------------------------------------
|
||
def route(self, case: BmmCase) -> dict:
|
||
"""返回 {branch, plan, timing, arbitration, candidates}."""
|
||
s = self.spec
|
||
|
||
# 1) 前置归约: 转Matmul / 特殊分支
|
||
if case.batch_a == 1 or case.batch_b == 1:
|
||
r = self.to_matmul.analyze(case)
|
||
return self._wrap_checked(case, r, "BatchA=1或BatchB=1, 折叠转普通Matmul")
|
||
if case.k <= 1:
|
||
r = self.special.analyze(case)
|
||
note = "K=0纯写值" if case.k == 0 else "K=1逐元素乘, 走AIV向量通路"
|
||
# issue#4 P0: 特殊分支 capable=False (如 K=1 但 B<128 不满足 UB 乒乓) 时
|
||
# r.plan=None, 必须兜底而不能把 None 传给下游 —— 显式标注"该区域暂无理论方案"
|
||
if not r.capable or r.plan is None:
|
||
return self._no_plan(case, "特殊分支",
|
||
note + f"; 但进入条件不满足 ({r.failed_conditions()}), "
|
||
f"该区域暂无理论方案, 建议参考 Cube 兜底或 AIV 单缓冲")
|
||
return self._wrap_checked(case, r, note)
|
||
|
||
# 2) B >= C 且 BatchA==BatchB: IterBatch / MergeBatch
|
||
if case.batch_c >= s.aic_num and case.batch_a == case.batch_b:
|
||
return self._route_split_b(case)
|
||
|
||
# 3) StreamK: P <= C/2 且切K条件满足
|
||
sk = self.stream_k.analyze(case)
|
||
if sk.capable:
|
||
return self._wrap_checked(case, sk, f"P<=C/2, B/M/N并行度买不满, 切K (grid_K={sk.plan.grid_k})")
|
||
|
||
# 4) 兜底: ASW_Basic (含降核模式)
|
||
asw = self.asw_basic.analyze(case)
|
||
note = "ASW_Basic兜底"
|
||
if sk.checks and not sk.capable:
|
||
note += f" (StreamK未过: {sk.failed_conditions()})"
|
||
return self._wrap_checked(case, asw, note)
|
||
|
||
# ------------------------------------------------------------------
|
||
def _route_split_b(self, case: BmmCase) -> dict:
|
||
mb = self.merge_batch.analyze(case)
|
||
ib = self.iter_batch.analyze(case)
|
||
|
||
# 候选表: [(分支名, BranchResult)], 顺序 = 仲裁优先级
|
||
cand_map = {self.merge_batch.name: mb, self.iter_batch.name: ib}
|
||
capable = {n: r.capable for n, r in cand_map.items()}
|
||
|
||
arbitration = ""
|
||
if mb.capable and ib.capable:
|
||
# 统一分界条件 + 端到端时延仲裁双保险; 时延模型为最终裁决.
|
||
# (issue#36: 合并的搬移效率收益已由 move_eff 模型量化进时延模型,
|
||
# T_cmd=0 默认下不再需要"T_cmd<=0 策略优先 MergeBatch"的覆盖逻辑)
|
||
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"两分支均合法, 仲裁: "
|
||
f"[分界条件] MergeBatch最优={mb_win} ({detail}); "
|
||
f"[时延模型] T_MergeBatch={t_mb*1e6:.2f}us vs T_IterBatch={t_ib*1e6:.2f}us -> {lat_win}更优; "
|
||
f"[裁决] {lat_win}" + ("" if win == lat_win else
|
||
f" (分界条件判{win}, 与时延模型不一致, 以时延模型为准)")
|
||
)
|
||
win = lat_win # 时延模型为最终裁决
|
||
elif any(capable.values()):
|
||
win = next(n for n, v in capable.items() if v)
|
||
arbitration = f"仅 {win} 条件满足"
|
||
else:
|
||
# 切B分支都不满足, 尝试 StreamK 再回落 ASW
|
||
sk = self.stream_k.analyze(case)
|
||
if sk.capable:
|
||
return self._wrap_checked(case, sk, "切B分支条件不满足, 落 StreamK")
|
||
asw = self.asw_basic.analyze(case)
|
||
return self._wrap_checked(
|
||
case, asw,
|
||
f"IterBatch/MergeBatch 进入条件均不满足, 回落 ASW_Basic; "
|
||
f"IterBatch未过: {ib.failed_conditions()}; "
|
||
f"MergeBatch未过: {mb.failed_conditions()}")
|
||
|
||
# 可行性保障 (issue#13): 仲裁胜出方案必须通过约束自检, 否则按
|
||
# (另一切B候选 -> StreamK -> ASW_Basic) 顺序回退到首个可行方案.
|
||
from .constraints import check_plan_constraints
|
||
|
||
def _feasible(n):
|
||
r = cand_map[n]
|
||
return r.plan is not None and not check_plan_constraints(case, r.plan, self.spec)
|
||
|
||
if _feasible(win):
|
||
chosen = cand_map[win]
|
||
else:
|
||
loser = self.merge_batch.name if win == self.iter_batch.name else self.iter_batch.name
|
||
fallback_note = (f"; 但 {win} 方案自检违规: "
|
||
f"{'; '.join(check_plan_constraints(case, cand_map[win].plan, self.spec))}")
|
||
if capable.get(loser) and _feasible(loser):
|
||
chosen, win = cand_map[loser], loser
|
||
fallback_note += f", 回退可行候选 {loser}"
|
||
else:
|
||
sk = self.stream_k.analyze(case)
|
||
if sk.capable and sk.plan is not None and \
|
||
not check_plan_constraints(case, sk.plan, self.spec):
|
||
return self._wrap_checked(case, sk, arbitration + fallback_note + ", 落 StreamK")
|
||
asw = self.asw_basic.analyze(case)
|
||
if asw.plan is not None and not check_plan_constraints(case, asw.plan, self.spec):
|
||
return self._wrap_checked(case, asw, arbitration + fallback_note + ", 回落 ASW_Basic")
|
||
chosen, win = cand_map[win], win # 无可行方案: 保留原裁决, 由自检标注
|
||
arbitration += fallback_note
|
||
|
||
result = BranchResult(capable=True, plan=chosen.plan, timing=chosen.timing)
|
||
return self._wrap_checked(case, result, arbitration,
|
||
candidates=capable)
|
||
|
||
# ------------------------------------------------------------------
|
||
def _wrap_checked(self, case: BmmCase, result: BranchResult, note: str,
|
||
candidates: dict | None = None) -> dict:
|
||
"""生成后自检 (issue#5): 推荐方案必须通过统一约束源校验, 不可行则标注违规.
|
||
|
||
约束源与 evaluate 共用 constraints.check_plan_constraints, 保证
|
||
"推荐方案 vs 自带评估器" 口径一致, 不再出现生成说可行、校验说不可行的矛盾.
|
||
"""
|
||
from .constraints import check_plan_constraints
|
||
violations = check_plan_constraints(case, result.plan, self.spec) if result.plan else []
|
||
if violations:
|
||
note = (note + " [自检违规: " + "; ".join(violations) +
|
||
"] —— 方案生成存在缺陷, 需人工复核")
|
||
# issue#34: 兜底分支 (ASW) 效率下限不满足时降级标注 (warning), 不判违规
|
||
if result.plan is not None and "效率降级" in result.plan.note:
|
||
note += " [效率降级标注: 搬移效率低于模型假设, 时延可能低估, 见 plan.note]"
|
||
return {
|
||
"branch": result.plan.branch if result.plan else "未知",
|
||
"plan": result.plan,
|
||
"timing": result.timing,
|
||
"arbitration": note + (f" | {result.note}" if result.note else ""),
|
||
"candidates": candidates or {},
|
||
"self_check_violations": violations,
|
||
}
|
||
|
||
@staticmethod
|
||
def _no_plan(case: BmmCase, branch: str, note: str) -> dict:
|
||
"""兜底: 分支 capable=False 时给出最小占位方案, 保证下游不崩溃 (issue#4).
|
||
|
||
方案标注 used_core_num=0 + 分支名, arbitration 说明"该区域暂无理论方案",
|
||
不产生时延 (timing=None), 供上层跳过或人工处理.
|
||
"""
|
||
plan = ImplPlan(case_id=case.case_id, branch=branch,
|
||
used_core_num=0, note="该区域暂无理论方案(进入条件不满足)")
|
||
return {"branch": branch, "plan": plan, "timing": None,
|
||
"arbitration": "[无方案] " + note, "candidates": {},
|
||
"self_check_violations": []}
|