Add BMM_Theory: bmm_theory/__main__.py

This commit is contained in:
2026-09-03 08:09:23 +00:00
parent a5fa3d56a3
commit 13022270cf

View File

@@ -0,0 +1,117 @@
"""BMM_Theory 命令行入口.
用法:
# 模式 1: 方案推荐 —— 输入 case csv, 输出理论最优方案 + 时延评估
python -m bmm_theory recommend cases.csv -o result.csv [--plans plans.csv] [-v]
# 模式 2: 方案评估 —— 输入 case csv + 方案 csv, 评估各硬件时延/瓶颈
python -m bmm_theory evaluate cases.csv plans.csv -o eval_result.csv [-v]
"""
from __future__ import annotations
import argparse
import sys
from .models import EvalResult
from .io_csv import load_cases, load_plans, save_results, save_plans
from .router import BranchRouter
from .evaluator import PlanEvaluator
from .timing import bound_type_of, HardwareTiming
def cmd_recommend(args) -> int:
cases = load_cases(args.cases)
router = BranchRouter()
results, plans = [], []
for case in cases:
r = router.route(case)
plan = r["plan"]
plans.append(plan)
res = EvalResult(case=case, plan=plan)
if r["timing"] is not None:
res.timing = r["timing"]
res.bound_type = bound_type_of(r["timing"].bottleneck)
else:
res.timing = HardwareTiming()
res.bound_type = ""
res.advice = r["arbitration"]
results.append(res)
if args.verbose:
_print_case(case, r, res)
save_results(args.output, results)
if args.plans:
save_plans(args.plans, plans)
print(f"[recommend] {len(results)} 个 case -> {args.output}"
+ (f", 方案表 -> {args.plans}" if args.plans else ""))
return 0
def cmd_evaluate(args) -> int:
cases = {c.case_id: c for c in load_cases(args.cases)}
plans = load_plans(args.plans)
ev = PlanEvaluator()
results = []
for plan in plans:
case = cases.get(plan.case_id)
if case is None:
print(f"[warn] 方案 {plan.case_id} 无对应 case, 跳过", file=sys.stderr)
continue
res = ev.evaluate(case, plan)
results.append(res)
if args.verbose:
t = res.timing
print(f"--- {case.case_id} [{plan.branch}] "
f"feasible={res.feasible} 瓶颈={res.timing.bottleneck or '-'}")
if res.violations:
print(f" 违反约束: {res.violations}")
if t.t_total:
print(f" MTE2={t.t_mte2*1e6:.2f}us (GM={t.t_mte2_gm*1e6:.2f}, "
f"L2={t.t_mte2_l2*1e6:.2f}, cmd={t.t_dma_cmd*1e6:.2f}) "
f"MMAD={t.t_mmad*1e6:.2f}us FIX={t.t_fixpipe*1e6:.2f}us "
f"drain={t.t_drain*1e6:.2f}us 总={t.t_total*1e6:.2f}us")
print(f" 建议: {res.advice}")
save_results(args.output, results)
print(f"[evaluate] {len(results)} 个方案 -> {args.output}")
return 0
def _print_case(case, r, res):
t = res.timing
print(f"=== {case.case_id}: B={case.batch_c} M={case.m} N={case.n} K={case.k} "
f"{case.dtype_a} -> [{r['branch']}]")
print(f" 仲裁: {r['arbitration']}")
p = r["plan"]
print(f" 方案: 核数={p.used_core_num} 切分=B{p.split_b}xM{p.m_cnt}xN{p.n_cnt}xK{p.grid_k} "
f"b_core={p.b_core} b0={p.merge_b0} k_L1={p.k_l1} L1形态={p.l1_form}")
if t.t_total:
print(f" 时延: 总={t.t_total*1e6:.2f}us 稳态={t.t_steady*1e6:.2f} drain={t.t_drain*1e6:.2f} "
f"| MTE2={t.t_mte2*1e6:.2f}(GM={t.t_mte2_gm*1e6:.2f}+cmd={t.t_dma_cmd*1e6:.2f}) "
f"MMAD={t.t_mmad*1e6:.2f} FIX={t.t_fixpipe*1e6:.2f} | 瓶颈={t.bottleneck}")
def main(argv=None) -> int:
ap = argparse.ArgumentParser(prog="bmm_theory",
description="BMM(batch_mat_mul_v3) Ascend950PR 理论最优实现分析")
sub = ap.add_subparsers(dest="cmd", required=True)
p1 = sub.add_parser("recommend", help="模式1: 产生理论最优实现方案")
p1.add_argument("cases", help="case 输入 csv")
p1.add_argument("-o", "--output", default="result_recommend.csv", help="结果输出 csv")
p1.add_argument("--plans", default="", help="可选: 方案表单独输出 csv (标准结构体)")
p1.add_argument("-v", "--verbose", action="store_true")
p1.set_defaults(func=cmd_recommend)
p2 = sub.add_parser("evaluate", help="模式2: 评估给定实现方案的硬件表现")
p2.add_argument("cases", help="case 输入 csv")
p2.add_argument("plans", help="实现方案 csv (标准结构体)")
p2.add_argument("-o", "--output", default="result_evaluate.csv", help="评估结果输出 csv")
p2.add_argument("-v", "--verbose", action="store_true")
p2.set_defaults(func=cmd_evaluate)
args = ap.parse_args(argv)
return args.func(args)
if __name__ == "__main__":
raise SystemExit(main())