Files
matmul-analysis/BMM/BMM_Theory/bmm_theory/router.py
admin b0b48b9073 Fix #36: MergeBatch 合并搬移效率收益建模 (move_eff) + t_cmd_ns 置 0
用户澄清: 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 例)
2026-09-07 21:09:49 +08:00

189 lines
9.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""分支决策路由: 按决策树推导分支, 重叠区由端到端时延模型仲裁.
决策树 (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": []}