Files
matmul-analysis/BMM/BMM_Theory/tests/test_branches.py
admin 8d42d958e4 Fix #38: 补全昇腾950系列硬件规格 (新增 Ascend950DT 等 4 个 SKU), 默认加载不变
- 数据源: 昇腾950_NPU架构白皮书 表3-1 (系列 SKU) / 表4-2 (Memory 层次)
- 新增 hardware/ascend950dt.py: ASCEND950DT (36AIC/72AIV, HBM 4TB/s 144GB,
  L2 128MB 主bin) + ASCEND950DT_C32 + ASCEND950DT_C28
- ascend950pr.py: 新增 ASCEND950PR_C28 (28核, 1.4TB/s, L2 112MB) 与
  gm_capacity_gb 信息字段; 注明 cube_peak_tflops 口径 (= 白皮书
  "Cube+Vector 总算力" 行, 与 Cube 单行 432 的 12.5% 偏差待用户裁决)
- __init__.py: SPECS 注册表 + get_spec(name), 默认仍为 ASCEND950PR,
  现有 import 与调用点零改动
- 共架构交叉验证: 单核 Cube 13.5T (432/32=486/36), Vector fp32
  27T/64核≈30T/72核 -> 128lane@1.65GHz; 表4-2 L1/L0/UB 各档一致
- docs/01_软件架构 + README 硬件层说明更新
- tests: TestIssue38 六例 (默认不变/注册表/DT派生量/AIV白皮书验证/DT路由
  冒烟/未知KeyError); 87/87 通过; examples 44例 0 diff;
  压力回归 10000 例干净, 分支分布与上轮逐数一致
2026-09-09 10:33:22 +08:00

1107 lines
54 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.

"""单元测试: 用理论文档中的典型边界 case 固化分支判定与仲裁逻辑.
覆盖:
- v0.98 §十 典型边界 case (MergeBatch/IterBatch 分界)
- v1.1 §4.5 统一分界条件 (K 截断 + b_core 阈值)
- MergeBatch 五条进入条件逐条触发
- 时延模型自洽性 (MergeBatch 胜时确实更优)
运行: python -m unittest discover -s tests -v
"""
import unittest
from bmm_theory.models import BmmCase
from bmm_theory.hardware import ASCEND950PR
from bmm_theory.router import BranchRouter
from bmm_theory.branches.merge_batch import MergeBatchBranch
from bmm_theory.branches.iter_batch import IterBatchBranch
def mkcase(b, m, n, k, **kw):
return BmmCase(case_id=f"B{b}_M{m}_N{n}_K{k}", batch_a=b, batch_b=b,
m=m, n=n, k=k, **kw)
class TestBranchEntry(unittest.TestCase):
"""v0.98 §十 典型边界 case 的分支归属."""
def setUp(self):
self.router = BranchRouter()
def test_mergebatch_boundary_case(self):
# 文档: B=128 M=N=64 K=512 -> MergeBatch 五条全过
r = self.router.route(mkcase(128, 64, 64, 512))
self.assertEqual(r["candidates"].get("MergeBatch"), True)
def test_iterbatch_when_datamount_insufficient(self):
# 文档: 同上但 K=256 -> 条件 3 不满足 -> IterBatch
r = self.router.route(mkcase(128, 64, 64, 256))
self.assertEqual(r["branch"], "IterBatch")
mb = MergeBatchBranch().analyze(mkcase(128, 64, 64, 256))
self.assertFalse(mb.capable)
def test_iterbatch_when_l0c_too_small(self):
# 文档: B=512 M=N=128 K=128 -> 条件 2 不满足 (MN=16384>8192) -> IterBatch
mb = MergeBatchBranch().analyze(mkcase(512, 128, 128, 128))
self.assertFalse(mb.capable)
r = self.router.route(mkcase(512, 128, 128, 128))
self.assertEqual(r["branch"], "IterBatch")
def test_iterbatch_form_d(self):
# 文档: B=64 M=N=64 K=8192 -> IterBatch 形态 d
r = self.router.route(mkcase(64, 64, 64, 8192))
self.assertEqual(r["branch"], "IterBatch")
self.assertIn("d_", r["plan"].l1_form)
def test_to_matmul(self):
r = self.router.route(BmmCase(case_id="t", batch_a=1, batch_b=1,
m=2048, n=2048, k=2048))
self.assertEqual(r["branch"], "转Matmul")
def test_special_k1(self):
r = self.router.route(mkcase(128, 256, 256, 1))
self.assertEqual(r["branch"], "特殊分支")
def test_special_k0(self):
r = self.router.route(mkcase(128, 256, 256, 0))
self.assertEqual(r["branch"], "特殊分支")
def test_streamk_entry(self):
# 文档: B=4 M=N=128 K=10240 -> StreamK (P=1 < C/2, K>=8192)
r = self.router.route(mkcase(4, 128, 128, 10240))
self.assertEqual(r["branch"], "StreamK")
self.assertGreaterEqual(r["plan"].grid_k, 2)
def test_streamk_partial_sum_4byte(self):
# StreamK 中间部分和按 4B (L0C dtype), 不随 C 的 fp16 转换
r = self.router.route(mkcase(4, 128, 128, 10240, dtype_c="fp16"))
self.assertEqual(r["plan"].out_dtype_bytes, 4)
def test_asw_basic_fallback(self):
# 文档: B=32 M=N=4096 K=4096 -> ASW_Basic (L1 四形态不满足)
r = self.router.route(mkcase(32, 4096, 4096, 4096))
self.assertEqual(r["branch"], "ASW_Basic")
def test_asw_reduced_core(self):
# 文档: B=16 M=N=256 K=128 -> 降核 ASW (P=4 < 32 且 K 不满足 StreamK)
r = self.router.route(mkcase(16, 256, 256, 128))
self.assertEqual(r["branch"], "ASW_Basic_降核")
self.assertLess(r["plan"].used_core_num, 32)
def test_asw_tail_strategy_embedded(self):
# ASW_Basic 的尾轮策略是必要组成, r>0 时不应为 A0
r = self.router.route(mkcase(8, 512, 512, 512)) # N_blk 不整除
if r["branch"] == "ASW_Basic" and r["plan"].tail_block_cnt > 0:
self.assertIn(r["plan"].tail_strategy, ("A1b", "方案B"))
class TestArbitration(unittest.TestCase):
"""v1.1 §4.5: MergeBatch 仅 K 截断且 b_core 足够大时胜."""
def setUp(self):
self.mb = MergeBatchBranch()
def test_l1_bound_mergebatch_loses(self):
# L1 绑定 (k_L1 < K): MergeBatch 恒劣
win, detail = self.mb.beats_iterbatch(mkcase(256, 128, 128, 4096))
self.assertFalse(win)
self.assertIn("L1绑定", detail)
def test_large_batch_mergebatch_wins(self):
# 大 B + 小 MN + K 截断 + 合并 tile 效率差: MergeBatch 应胜
# (issue#36: (2048,16,64,128) IterBatch A tile=4KB eff=0.25 vs 合并
# 16KB eff=1.0, 效率节省 15.7us >> drain 0.17us; T_cmd=0 默认下仍胜)
case = mkcase(2048, 16, 64, 128)
self.assertTrue(MergeBatchBranch().analyze(case).capable)
win, detail = self.mb.beats_iterbatch(case)
self.assertTrue(win, detail)
def test_no_eff_diff_mergebatch_loses(self):
# issue#36: K 截断但 IterBatch 单命令 tile 已达 16KB 饱和 (效率打平),
# T_cmd=0 下命令节省为 0, 合并只剩 drain 惩罚 -> IterBatch 优
# (原 t_cmd=50ns 时代的 MergeBatch 胜例 (2048,32,32,256), 语义迁移)
case = mkcase(2048, 32, 32, 256)
win, detail = self.mb.beats_iterbatch(case)
self.assertFalse(win, detail)
class TestTimingSanity(unittest.TestCase):
"""时延模型自洽性."""
def test_mergebatch_dma_cmd_saved(self):
# K 截断时 MergeBatch 的 DMA 命令数 = IterBatch 的 1/b0
# B=2048 M=N=32 K=256 是 K 截断 case (b0=4, k_l1=256=K)
case = mkcase(2048, 32, 32, 256)
mb = MergeBatchBranch().analyze(case)
ib = IterBatchBranch().analyze(case)
self.assertTrue(mb.capable)
self.assertTrue(ib.capable)
self.assertGreaterEqual(mb.plan.k_l1, case.k) # 确认 K 截断前提
ratio = mb.timing.dma_cmd_count / ib.timing.dma_cmd_count
self.assertAlmostEqual(ratio, 1.0 / mb.plan.merge_b0, places=1)
def test_bottleneck_memory_bound_for_small_mn(self):
# MergeBatch case 必为访存 Bound (进入条件 5)
case = mkcase(128, 64, 64, 512)
mb = MergeBatchBranch().analyze(case)
self.assertIn(mb.timing.bottleneck, ("MTE2", "MMAD"))
def test_fixpipe_dtype_conversion(self):
# C 指定 fp16 输出时, 写出量按 2B 而非 L0C 的 4B;
# 字节列为整芯片口径 (issue#29): fixpipe_bytes = B*MN*outB
case = mkcase(128, 64, 64, 512, dtype_c="fp16")
ib = IterBatchBranch().analyze(case)
expect = case.batch_c * 64 * 64 * 2
self.assertAlmostEqual(ib.timing.fixpipe_bytes, expect)
self.assertAlmostEqual(ib.timing.fixpipe_bytes, case.output_bytes)
class TestIssueRegression(unittest.TestCase):
"""issue 复现 case 固化 (测评报告 §5)."""
def setUp(self):
self.router = BranchRouter()
def test_issue4_k1_small_batch_real_plan(self):
# issue#4/#12/#17: K=1 且 B<128 不崩溃, 且给出 AIV 单缓冲真实方案 (不再是无方案占位)
case = mkcase(64, 8192, 32, 1, dtype_a="int8", dtype_b="int8")
r = self.router.route(case)
self.assertIsNotNone(r["plan"])
self.assertEqual(r["branch"], "特殊分支")
self.assertEqual(r["plan"].used_core_num, 64) # AIV 核
self.assertNotIn("暂无理论方案", r["arbitration"])
self.assertIn("单缓冲", r["plan"].note)
from bmm_theory.constraints import check_plan_constraints
self.assertEqual(check_plan_constraints(case, r["plan"]), [])
def test_issue5_asw_reduced_core_base_k_dtype_aware(self):
# issue#5: ASW 降核 base_k 按 dtype 反推, fp32 不再 L0A 溢出
from bmm_theory.constraints import check_plan_constraints
case = mkcase(51, 255, 42, 682, dtype_a="fp32", dtype_b="fp32", dtype_c="fp32")
r = self.router.route(case)
self.assertEqual(r["plan"].branch, "StreamK") # 该 case 满足 StreamK
v = check_plan_constraints(case, r["plan"])
self.assertEqual(v, [], f"应无违规: {v}")
def test_issue5_asw_reduced_core_fp16_l0c(self):
# issue#5: ASW 降核 fp32 场景 base tile 不越 L0C
from bmm_theory.constraints import check_plan_constraints
case = mkcase(2, 64, 16, 32, dtype_a="fp16", dtype_b="fp16", dtype_c="fp16")
r = self.router.route(case)
self.assertIn("降核", r["plan"].branch)
# base_k 应由 L0A/L0B 反推, 不再硬编码 64
self.assertLessEqual(
r["plan"].base_m * r["plan"].base_k * 2 * 2, 64 * 1024) # L0A 上限
def test_issue6_iterbatch_small_k_whole_resident(self):
# issue#6: IterBatch a/b 形态 K 整驻留, dValue 下限对 K 向豁免
from bmm_theory.constraints import check_plan_constraints
case = mkcase(256, 128, 1024, 8) # K=8 bf16, 整驻留
r = self.router.route(case)
self.assertEqual(r["branch"], "IterBatch")
v = check_plan_constraints(case, r["plan"])
self.assertEqual(v, [], f"整驻留形态不应报 dValue 违规: {v}")
def test_issue9_streamk_reduce_not_double_counted(self):
# issue#9: StreamK 归约串行追加, 不进稳态 max (否则双倍计账)
from bmm_theory.branches.stream_k import StreamKBranch
case = mkcase(4, 128, 128, 10240)
sk = StreamKBranch().analyze(case)
self.assertTrue(sk.capable)
t = sk.timing
# t_total = max(稳态) + drain(=reduce), 不应是 max(...,reduce)+reduce
self.assertAlmostEqual(t.t_total, t.t_steady + t.t_drain)
self.assertAlmostEqual(t.t_drain, t.t_reduce)
def test_out_nd_disables_streamk(self):
# issue#7: out_nd=False 禁用 StreamK (经 CLI 层解析)
case = mkcase(4, 128, 128, 10240)
case.out_nd = False
r = self.router.route(case)
self.assertNotEqual(r["branch"], "StreamK")
class TestIssueRegression2(unittest.TestCase):
"""第三轮复评问题 (#17-#20) + 恢复 #11-#15 回归."""
def setUp(self):
self.router = BranchRouter()
def test_issue11_streamk_fixpipe_no_double_count(self):
# issue#11/#17: 部分和写出只经 t_reduce 计账一次; 稳态 fixpipe 不得再计
from bmm_theory.branches.stream_k import StreamKBranch
case = mkcase(4, 128, 128, 10240)
sk = StreamKBranch().analyze(case)
self.assertTrue(sk.capable)
t = sk.timing
self.assertAlmostEqual(t.t_fixpipe, 0.0) # 归约串行口径下无稳态 fixpipe 账
self.assertAlmostEqual(t.fixpipe_bytes, 0.0)
# 端到端 = 稳态(MTE2搬移链) + 归约, fixpipe 无重复账
self.assertEqual(t.bottleneck, "MTE2")
self.assertAlmostEqual(t.t_total, t.t_steady + t.t_drain)
def test_issue12_k1_pingpong_still_ok(self):
# issue#12/#17: K=1 且 B>=128 仍走 UB 乒乓 (原行为不变)
r = self.router.route(mkcase(128, 256, 256, 1))
self.assertEqual(r["branch"], "特殊分支")
self.assertIn("乒乓", r["plan"].l1_form)
self.assertIsNotNone(r["timing"])
def test_issue13_merge_b0_l0ab_capped(self):
# issue#13: MergeBatch 瘦长 case 的 b0 受 L0A/L0B 容量约束 (B=811 M=33 N=1 fp32)
from bmm_theory.constraints import check_plan_constraints
case = mkcase(811, 33, 1, 2459, dtype_a="fp32", dtype_b="fp32", dtype_c="fp32")
r = self.router.route(case)
self.assertEqual(r["branch"], "MergeBatch")
p = r["plan"]
self.assertLessEqual(p.base_m * p.base_k * 4 * 2, 64 * 1024) # L0A 容量内
self.assertLessEqual(p.base_n * p.base_k * 4 * 2, 64 * 1024) # L0B 容量内
self.assertEqual(check_plan_constraints(case, p), [])
def test_issue13_router_fallback_when_winner_infeasible(self):
# issue#13: 仲裁胜出的 MergeBatch 自检违规时, 回退到可行候选 IterBatch
# (issue#31 后该样例 MergeBatch 时延已不再胜出, 用 mock 压低其时延以固定
# "胜者违规 -> 回退" 机制路径; MergeBatch 方案 dValue=112B<128 仍违规)
from unittest import mock
from bmm_theory.branches.merge_batch import MergeBatchBranch
from bmm_theory.constraints import check_plan_constraints
case = mkcase(256, 1, 256, 4096, dtype_a="int8", dtype_b="int8")
mb_plan = MergeBatchBranch().analyze(case).plan
self.assertTrue(check_plan_constraints(case, mb_plan),
"该样例 MergeBatch 方案应仍自检违规 (dValue)")
real_eval = MergeBatchBranch.evaluate
def fake_win(self, c, p):
t = real_eval(self, c, p)
t.t_total = 1e-9 # 强制 MergeBatch 时延胜出 -> 仲裁选它 -> 走自检/回退
return t
with mock.patch.object(MergeBatchBranch, "evaluate", fake_win):
r = self.router.route(case)
self.assertEqual(r["branch"], "IterBatch") # 回退
self.assertIn("自检违规", r["arbitration"])
self.assertIn("回退", r["arbitration"])
self.assertEqual(check_plan_constraints(case, r["plan"]), [])
def test_issue15_input_validation(self):
# issue#15: 非法维度/负值必须抛错, 不再静默产出伪方案
for kw in (dict(m=0), dict(m=-5), dict(n=0), dict(k=-1),
dict(batch_a=0), dict(batch_b=-3)):
with self.assertRaises(ValueError, msg=str(kw)):
BmmCase(case_id="bad", **kw)
with self.assertRaises(ValueError):
BmmCase(case_id="bad", m=64, n=64, k=1, dtype_a="xxx")
def test_issue18_placeholder_plan_infeasible_in_evaluate(self):
# issue#18: 占位方案(used_core_num=0)在 evaluate 中必须不可行,
# 不得被当作可行方案给出正常时延
from bmm_theory.evaluator import PlanEvaluator
from bmm_theory.models import ImplPlan
# K=1 且 B<128 已恢复真实单缓冲方案 (#17), 故直接构造占位 plan 验证约束层
case = mkcase(64, 8192, 32, 1, dtype_a="int8", dtype_b="int8")
r = self.router.route(case)
self.assertGreater(r["plan"].used_core_num, 0) # 真实方案
ph = ImplPlan(case_id="ph", branch="特殊分支", used_core_num=0,
note="该区域暂无理论方案(进入条件不满足)")
er = PlanEvaluator().evaluate(case, ph)
self.assertFalse(er.feasible, "占位方案应判不可行")
self.assertIn("used_core_num", er.violations)
def test_issue19_transpose_dvalue_guard_effective(self):
# issue#19: 转置感知 dValue 判据在生成守卫/条件4/约束三处同源后真正生效.
# 判别形状: B=64 M=4096 N=64 K=4096 bf16 —— d 形态 k_l1=16,
# 均不转置时 dv_a=k_l1*2=32B <128 挡下 (B 侧 dv_b=N*2=128B 恰好达标也不放行,
# 因为两侧切 K 两侧都要高效);
# A 转置后 dv_a=M*2=8192B, 应能走 IterBatch 形态 d.
from bmm_theory.constraints import check_plan_constraints
c_not = BmmCase(case_id="x", batch_a=64, batch_b=64, m=4096, n=64, k=4096,
trans_a=False, trans_b=False)
r_not = self.router.route(c_not)
self.assertNotEqual(r_not["branch"], "IterBatch") # 非转置被 dValue 守卫挡下
c_tr = BmmCase(case_id="x", batch_a=64, batch_b=64, m=4096, n=64, k=4096,
trans_a=True, trans_b=False)
r_tr = self.router.route(c_tr)
self.assertEqual(r_tr["branch"], "IterBatch") # A 转置 M 向连续, 守卫放行
self.assertIn("d_", r_tr["plan"].l1_form)
self.assertEqual(check_plan_constraints(c_tr, r_tr["plan"]), [])
class TestFp4Support(unittest.TestCase):
"""fp4 (0.5B) dtype 支持 (对齐 bmmv3, 2026-09-03)."""
def test_fp4_dtype_bytes(self):
from bmm_theory.models import dtype_bytes
self.assertEqual(dtype_bytes("fp4"), 0.5)
self.assertEqual(dtype_bytes("fp4_e2m1"), 0.5)
def test_fp4_case_creation(self):
case = mkcase(32, 64, 64, 256, dtype_a="fp4", dtype_b="fp4", dtype_c="fp16")
self.assertEqual(case.dtype_in_bytes, 0.5)
# fp4 输入 + fp16 输出: 输入 0.5B, 输出 2B
self.assertEqual(case.dtype_out_bytes, 2)
def test_fp4_dvalue_threshold(self):
# fp4 (0.5B) 时 dValue 128B 需要 k_l1 >= 256
from bmm_theory.constraints import check_plan_constraints
case = mkcase(128, 64, 64, 128, dtype_a="fp4", dtype_b="fp4", dtype_c="fp16")
r = BranchRouter().route(case)
v = check_plan_constraints(case, r["plan"])
# fp4 小 K 场景应能正常路由且不报 dValue 违规 (k_l1 连续维是 M/N)
self.assertEqual(v, [], f"fp4 case 不应报违规: {v}")
class TestTransposeModeling(unittest.TestCase):
"""转置对 dValue 连续维的影响建模 (对齐 bmmv3, 2026-09-03)."""
def test_transpose_affects_dvalue_judgment(self):
# A 不转置: K 向连续, dValue 判 K*dt
# A 转置: M 向连续, dValue 判 M*dt
# B 不转置: N 向连续, dValue 判 N*dt
# B 转置: K 向连续, dValue 判 K*dt
from bmm_theory.constraints import _k_segment_is_contiguous
# 当前版本不建模转置, 默认按不转置处理
case = mkcase(128, 64, 64, 512)
case.trans_a = False
case.trans_b = False
# 验证 trans 字段存在且可读写 (为后续建模做准备)
self.assertFalse(case.trans_a)
self.assertFalse(case.trans_b)
case.trans_a = True
self.assertTrue(case.trans_a)
def test_transpose_a_large_m_small_k(self):
# A 转置 + 大 M 小 K: dValue 应判 M*dt (M 向连续), 不受 K 小影响
from bmm_theory.branches.iter_batch import IterBatchBranch
# M=1024 (M*dt=2048B >= 128B), K=8 (K*dt=16B < 128B)
case = mkcase(128, 1024, 64, 8, dtype_a="bf16", dtype_b="bf16")
case.trans_a = True
case.trans_b = False
ib = IterBatchBranch().analyze(case)
# A 转置时 dValue 判 M*dt=2048B >= 128B, 应通过
c4 = [c for c in ib.checks if "搬移效率" in c.name][0]
self.assertTrue(c4.passed, f"A 转置时应判 M 向连续: {c4.detail}")
def test_no_transpose_small_k_fails(self):
# A 不转置 + 小 K + c/d 形态: dValue 判 K*dt, K=8 时 16B < 128B 应失败
from bmm_theory.branches.iter_batch import IterBatchBranch
# 大 M/N 让 L1 放不下整 K, 走 c/d 形态
case = mkcase(128, 1024, 1024, 8, dtype_a="bf16", dtype_b="bf16")
case.trans_a = False
case.trans_b = False
ib = IterBatchBranch().analyze(case)
c4 = [c for c in ib.checks if "搬移效率" in c.name][0]
# c/d 形态下 A 不转置时 K=8 应报 dValue 违规
if ib.plan.l1_form.startswith(("c_", "d_")):
self.assertFalse(c4.passed, f"A 不转置时 K=8 应报 dValue 违规: {c4.detail}")
class TestZeroCmdHandling(unittest.TestCase):
"""950PR 默认 T_cmd=0 (未标定, issue#36): 整链路不得除零/崩溃; 合并收益由
搬移效率模型 (move_eff, 合并 tile 放大 b0 倍) 刻画, 不再设策略覆盖."""
def test_beats_iterbatch_zero_cmd(self):
# T_cmd=0: 命令节省项为 0, 由效率节省 vs drain 惩罚决定 (issue#36)
from bmm_theory.branches.merge_batch import MergeBatchBranch
mb = MergeBatchBranch() # 默认 spec 即 t_cmd_ns=0
# (b, m, n, k, MergeBatch应胜与否=效率节省>drain惩罚)
# (2048,16,64,128): iter A tile=4KB eff=0.25 vs 合并 16KB eff=1.0 -> 大胜
# (2048,32,32,256): iter A tile=16KB 已饱和, 效率打平 -> drain 惩罚 -> 恒劣
# (128,64,64,512): issue#35 第三情形 (dValue cap 截断, 命令/效率均打平) -> 恒劣
# (256,128,128,4096): L1 绑定, tile 均 >=128KB 饱和 -> 恒劣
cases = [(2048, 16, 64, 128, True), (2048, 32, 32, 256, False),
(128, 64, 64, 512, False), (256, 128, 128, 4096, False)]
for b, m, n, k, mb_wins in cases:
win, detail = mb.beats_iterbatch(mkcase(b, m, n, k))
self.assertEqual(win, mb_wins, f"{b},{m},{n},{k}: {detail}")
self.assertIn("效率节省", detail)
def test_route_with_zero_cmd_efficiency_decides(self):
# 默认 t_cmd=0: 有效率低下的 IterBatch 小 tile case 由 MergeBatch 胜;
# tile 均饱和的 case 由 IterBatch 胜 (drain 惩罚, 无策略覆盖)
router = BranchRouter()
r = router.route(mkcase(2048, 16, 64, 128)) # 效率差显著 -> MergeBatch
self.assertEqual(r["branch"], "MergeBatch")
self.assertIsNotNone(r["timing"])
r2 = router.route(mkcase(2048, 32, 32, 256)) # tile 均饱和 -> IterBatch
self.assertEqual(r2["branch"], "IterBatch")
shapes = [(128, 64, 64, 512), (64, 64, 64, 8192), (512, 128, 128, 128),
(128, 128, 128, 1024), (32, 4096, 4096, 4096)]
for b, m, n, k in shapes:
rr = router.route(mkcase(b, m, n, k))
t = rr["timing"]
self.assertIsNotNone(t, f"{b},{m},{n},{k} 应有 timing")
self.assertGreater(t.t_total, 0.0)
self.assertEqual(t.t_total, t.t_total) # 非 NaN
def test_zero_cmd_random_smoke(self):
# 随机小样本冒烟: 无异常/无 NaN
import random
from bmm_theory.hardware import NpuSpec
router = BranchRouter(NpuSpec(t_cmd_ns=0.0))
rng = random.Random(9)
for _ in range(300):
b = rng.choice([2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048])
m = rng.choice([8, 16, 32, 64, 128, 256, 512, 1024, 4096])
n = rng.choice([8, 16, 32, 64, 128, 256, 512, 1024, 4096])
k = rng.choice([16, 32, 64, 128, 256, 512, 1024, 4096, 16384])
c = BmmCase(case_id="z", batch_a=b, batch_b=b, m=m, n=n, k=k)
rr = router.route(c)
self.assertIsNotNone(rr["plan"])
t = rr["timing"]
self.assertIsNotNone(t)
self.assertGreater(t.t_total, 0.0)
self.assertEqual(t.t_total, t.t_total)
class TestIssue27to30(unittest.TestCase):
"""第四轮评审 (issue#27-#30, 设计文档 docs/05).
#27 MergeBatch Cube 公式复核 (每步 2*(b0M)(b0N)K x b_core/b0 步);
#28 dtype 感知算力 (Cube/AIV 速率表);
#29 GM 首读下限 GM>=V_in + 字节列整芯片口径;
#30 Fixpipe 输出落点 (整case驻留 -> L2, 否则直写 GM).
"""
def setUp(self):
from bmm_theory.hardware import ASCEND950PR
self.s = ASCEND950PR
self.router = BranchRouter()
# ---------------- #27 ----------------
def test_issue27_merge_cube_flops_matches_formula(self):
from bmm_theory.branches.merge_batch import MergeBatchBranch
case = mkcase(2048, 32, 32, 256)
mb = MergeBatchBranch().analyze(case)
self.assertTrue(mb.capable)
p, t = mb.plan, mb.timing
b0, bc = p.merge_b0, p.b_core
# 整芯片列 = B*b0*2MNK; 等价 (b_core/b0) 步 x 每步 2*(b0M)(b0N)K x C 核
step_flops = 2.0 * (b0 * 32) * (b0 * 32) * 256
self.assertAlmostEqual(t.cube_flops,
step_flops * (bc / b0) * self.s.aic_num)
self.assertAlmostEqual(t.cube_flops,
case.batch_c * b0 * 2.0 * 32 * 32 * 256)
# t_mmad = 每核 flops / 单核算力
self.assertAlmostEqual(t.t_mmad,
(bc * b0 * 2.0 * 32 * 32 * 256) / self.s.q16)
def test_issue27_merge_gm_chip_col_equals_input(self):
# K 截断合并: GM 每字节一次 -> gm 整芯片列 = V_in (issue#29 口径)
from bmm_theory.branches.merge_batch import MergeBatchBranch
case = mkcase(2048, 32, 32, 256)
t = MergeBatchBranch().analyze(case).timing
self.assertAlmostEqual(t.gm_read_bytes, case.input_bytes)
self.assertAlmostEqual(t.l2_read_bytes, 0.0)
# ---------------- #28 ----------------
def test_issue28_cube_dtype_rate(self):
from bmm_theory.branches.iter_batch import IterBatchBranch
base = mkcase(128, 64, 64, 512) # IterBatch 形态 b, 各 dtype 均可行
t16 = IterBatchBranch().analyze(base).timing
t32 = IterBatchBranch().analyze(
mkcase(128, 64, 64, 512, dtype_a="fp32", dtype_b="fp32")).timing
t8 = IterBatchBranch().analyze(
mkcase(128, 64, 64, 512, dtype_a="fp8", dtype_b="fp8")).timing
# 同计算量, 时延反比于 dtype 算力: fp32=1/2, fp8=2x (相对 bf16)
self.assertAlmostEqual(t32.t_mmad / t16.t_mmad, 2.0, places=6)
self.assertAlmostEqual(t8.t_mmad / t16.t_mmad, 0.5, places=6)
def test_issue28_aiv_dtype_rate_k1(self):
# K=1 AIV 逐元素: bf16 通量 2x fp32 -> fp32 时延为 bf16 的 2 倍
base = mkcase(64, 8192, 512, 1)
t_bf16 = BranchRouter().route(base)["timing"]
t_fp32 = BranchRouter().route(
mkcase(64, 8192, 512, 1, dtype_a="fp32", dtype_b="fp32"))["timing"]
self.assertAlmostEqual(t_fp32.t_mmad / t_bf16.t_mmad, 2.0, places=6)
# ---------------- #29 ----------------
def test_issue29_gm_floor_and_chip_columns(self):
# 各分支代表 case: GM 整芯片列 >= V_in (R1), 且迭代/合并/转matmul/special
# 等于 V_in (无重复读结构)
from bmm_theory.branches.iter_batch import IterBatchBranch
from bmm_theory.branches.merge_batch import MergeBatchBranch
from bmm_theory.branches.stream_k import StreamKBranch
from bmm_theory.branches.asw_basic import AswBasicBranch
checks = [
(IterBatchBranch().analyze(mkcase(128, 64, 64, 512)).timing,
mkcase(128, 64, 64, 512)),
(MergeBatchBranch().analyze(mkcase(2048, 32, 32, 256)).timing,
mkcase(2048, 32, 32, 256)),
(StreamKBranch().analyze(mkcase(4, 128, 128, 10240)).timing,
mkcase(4, 128, 128, 10240)),
(AswBasicBranch().analyze(mkcase(8, 4096, 4096, 1024)).timing,
mkcase(8, 4096, 4096, 1024)),
]
for t, case in checks:
self.assertGreaterEqual(t.gm_read_bytes + 1e-6, case.input_bytes,
f"{case.case_id} GM < V_in")
# 转Matmul / 特殊分支 (K=1) 亦等于 V_in
r = self.router.route(mkcase(1, 2048, 2048, 2048))
t = r["timing"]
self.assertGreaterEqual(t.gm_read_bytes + 1e-6,
BmmCase(case_id="x", batch_a=1, batch_b=1,
m=2048, n=2048, k=2048).input_bytes)
r = self.router.route(mkcase(128, 256, 256, 1))
self.assertAlmostEqual(r["timing"].gm_read_bytes, 128 * (256 + 256) * 2)
def test_issue29_gm_floor_random_smoke(self):
# 随机多 dtype 冒烟: 所有可行方案 GM 整芯片列 >= V_in, 无 NaN
import random
rng = random.Random(20260609)
dtypes = ["bf16", "fp16", "fp32", "int8", "fp8", "fp4"]
for _ in range(400):
dt = rng.choice(dtypes)
kw = dict(dtype_a=dt, dtype_b=dt, dtype_c=rng.choice(["bf16", "fp16"]))
b = rng.choice([1, 2, 4, 8, 16, 32, 64, 128, 256, 2048])
m = rng.choice([1, 8, 32, 64, 128, 256, 1024, 4096])
n = rng.choice([1, 8, 32, 64, 128, 256, 1024, 4096])
k = rng.choice([0, 1, 32, 64, 128, 256, 512, 1024, 4096])
c = BmmCase(case_id="gm", batch_a=b, batch_b=b, m=m, n=n, k=k, **kw)
r = self.router.route(c)
t = r["timing"]
self.assertIsNotNone(t, c.case_id)
self.assertEqual(t.t_total, t.t_total) # 非 NaN
self.assertGreaterEqual(t.gm_read_bytes + 1e-6, c.input_bytes,
f"{dt} b={b} m={m} n={n} k={k} gm={t.gm_read_bytes}")
# ---------------- #30 ----------------
def test_issue30_iter_output_l2_when_whole_case_fits(self):
from bmm_theory.branches.iter_batch import IterBatchBranch
# 整 case 输入+输出 16.8MB+1.0MB <= 128MB -> 输出写 L2 写口, GM 写=0
case = mkcase(128, 64, 64, 512)
self.assertLess(case.input_bytes + case.output_bytes, self.s.l2_bytes)
t = IterBatchBranch().analyze(case).timing
self.assertAlmostEqual(t.fixpipe_bytes, case.output_bytes)
self.assertAlmostEqual(t.t_fixpipe, case.output_bytes / self.s.bw_l2)
def test_issue30_iter_output_gm_when_case_overflows(self):
from bmm_theory.branches.iter_batch import IterBatchBranch
# 输出 268MB > L2 余量 -> 直写 GM (输入优先驻留)
case = mkcase(128, 1024, 1024, 64)
self.assertGreater(case.input_bytes + case.output_bytes, self.s.l2_bytes)
t = IterBatchBranch().analyze(case).timing
self.assertAlmostEqual(t.fixpipe_bytes, case.output_bytes)
self.assertAlmostEqual(t.t_fixpipe, case.output_bytes / self.s.bw_gm)
def test_issue30_asw_scene_whole_case_labels(self):
# S_A (整case全驻留): l2_policy_out=resident, t_fixpipe 走 L2 写口
case_sa = mkcase(2, 4096, 4096, 512)
self.assertLess(case_sa.input_bytes + case_sa.output_bytes, self.s.l2_bytes)
r = self.router.route(case_sa)
self.assertEqual(r["branch"], "ASW_Basic")
self.assertTrue(r["plan"].l2_policy_out.startswith("resident"), r["plan"].l2_policy_out)
self.assertIn("A_整case全驻留", r["plan"].note)
self.assertAlmostEqual(r["timing"].t_fixpipe,
case_sa.output_bytes / self.s.bw_l2)
# S_B (整case超L2, 单batch输入可驻留): 输出直写 GM
case_sb = mkcase(64, 4096, 4096, 1024)
self.assertGreater(case_sb.input_bytes + case_sb.output_bytes, self.s.l2_bytes)
self.assertLess(4096 * 1024 * 2 * 2, self.s.l2_bytes) # 单batch输入 16.8MB
r = self.router.route(case_sb)
self.assertIn("ASW", r["branch"])
self.assertTrue(r["plan"].l2_policy_out.startswith("direct_gm"),
r["plan"].l2_policy_out)
self.assertAlmostEqual(r["timing"].t_fixpipe,
case_sb.output_bytes / self.s.bw_gm)
def test_issue30_to_matmul_output_l2_toggle(self):
# 转Matmul 粗估: 整case可驻留 -> L2; 大 case -> GM
small = mkcase(1, 2048, 2048, 2048)
self.assertLess(small.input_bytes + small.output_bytes, self.s.l2_bytes)
r = self.router.route(small)
self.assertEqual(r["branch"], "转Matmul")
self.assertAlmostEqual(r["timing"].t_fixpipe,
small.output_bytes / self.s.bw_l2)
big = mkcase(1, 8192, 8192, 4096)
self.assertGreater(big.input_bytes + big.output_bytes, self.s.l2_bytes)
r = self.router.route(big)
self.assertEqual(r["branch"], "转Matmul")
self.assertAlmostEqual(r["timing"].t_fixpipe,
big.output_bytes / self.s.bw_gm)
class TestIssue31to32(unittest.TestCase):
"""第五轮评审: #31 切B类 GM 每字节恰读一次 = V_in; #32 ASW 场景升级
(单侧全驻留 -> S_B 使 GM=V_in; 双侧超 L2 -> S_C 最小替换 2D 分组 + 窗口 L2 计账)."""
def setUp(self):
from bmm_theory.hardware import ASCEND950PR
self.s = ASCEND950PR
self.router = BranchRouter()
# ---------------- #31 ----------------
def test_issue31_iter_form_c_gm_equals_vin(self):
# b64_m16_n256_k512: 形态 c, K=512 切 n_K=3 (k_L1=240) 非整除 -> GM 仍 = V_in
from bmm_theory.branches.iter_batch import IterBatchBranch
case = mkcase(64, 16, 256, 512)
ib = IterBatchBranch().analyze(case)
self.assertIn("c_", ib.plan.l1_form)
self.assertLess(ib.plan.k_l1, case.k)
t = ib.timing
self.assertAlmostEqual(t.gm_read_bytes, case.input_bytes)
self.assertAlmostEqual(t.l2_read_bytes, 0.0)
self.assertAlmostEqual(t.t_mte2_gm, case.input_bytes / self.s.bw_gm)
def test_issue31_iter_form_d_gm_equals_vin(self):
# K=5000 非整除切段 (k_L1=1024, n_K=5, 末段剩 904): GM 仍恰一次 = V_in
from bmm_theory.branches.iter_batch import IterBatchBranch
case = mkcase(128, 64, 64, 5000)
ib = IterBatchBranch().analyze(case)
self.assertIn("d_", ib.plan.l1_form)
self.assertLess(ib.plan.k_l1, case.k)
self.assertAlmostEqual(ib.timing.gm_read_bytes, case.input_bytes)
self.assertAlmostEqual(ib.timing.t_mte2_gm,
case.input_bytes / self.s.bw_gm)
def test_issue31_merge_l1_bound_gm_equals_vin(self):
# MergeBatch L1 绑定 (k_L1<K): 切 K 段互不重叠, GM = V_in
from bmm_theory.branches.merge_batch import MergeBatchBranch
case = mkcase(811, 33, 1, 2459, dtype_a="fp32", dtype_b="fp32")
mb = MergeBatchBranch().analyze(case)
self.assertTrue(mb.capable)
self.assertLess(mb.plan.k_l1, case.k) # L1 绑定
self.assertAlmostEqual(mb.timing.gm_read_bytes, case.input_bytes)
self.assertAlmostEqual(mb.timing.l2_read_bytes, 0.0)
# ---------------- #32 ----------------
def test_issue32_asw_single_side_resident_gm_equals_vin(self):
# b32_m8192_n4096_k7168: 单batch输入 176MB > L2, 但 B=58.7MB 可全驻留
# + A 行块滑窗 (58.7+2x2.5 <= 128) -> S_B: GM = V_in, 重复读全命中 L2
case = mkcase(32, 8192, 4096, 7168)
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
p, t = r["plan"], r["timing"]
a_b = case.m * case.k * case.dtype_in_bytes
bb_b = case.k * case.n * case.dtype_in_bytes
m_cnt, n_cnt = p.m_cnt, p.n_cnt
self.assertGreater(a_b + bb_b, self.s.l2_bytes) # 单 batch 确实超 L2
self.assertLessEqual(bb_b + 2 * (a_b / m_cnt),
self.s.l2_bytes) # 单侧全驻留成立
self.assertAlmostEqual(t.gm_read_bytes, case.input_bytes)
exp_l2 = case.batch_c * ((n_cnt - 1) * a_b + (m_cnt - 1) * bb_b)
self.assertAlmostEqual(t.l2_read_bytes, exp_l2)
def test_issue32_asw_scene_c_min_gm(self):
# b8_m131072_n8192_k8192: 双侧均不可全驻留 -> S_C 最小替换分组;
# gm == 容量约束最小解 (测试内复算), 且 >= V_in
case = mkcase(8, 131072, 8192, 8192)
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
p, t = r["plan"], r["timing"]
self.assertIn("C_", p.note)
m_cnt, n_cnt = p.m_cnt, p.n_cnt
dt = case.dtype_in_bytes
a_b = case.m * case.k * dt
bb_b = case.k * case.n * dt
blk_a = a_b / m_cnt
blk_b = bb_b / n_cnt
best = None
for mg in range(1, m_cnt + 1):
for ng in range(1, n_cnt + 1):
if mg * blk_a + ng * blk_b > self.s.l2_bytes:
continue
gm = (n_cnt + ng - 1) // ng * a_b + (m_cnt + mg - 1) // mg * bb_b
if best is None or gm < best:
best = gm
self.assertIsNotNone(best)
self.assertAlmostEqual(t.gm_read_bytes, case.batch_c * best)
self.assertGreaterEqual(t.gm_read_bytes, case.input_bytes)
# 最小替换解应严格优于"每块独立落 GM"的保守上界
self.assertLessEqual(t.gm_read_bytes,
case.batch_c * (n_cnt * a_b + m_cnt * bb_b))
class TestIssue33(unittest.TestCase):
"""ASW_Basic tile 选择 (issue#33, v1.91 §5.1/§5.2 + 尾轮 v1.5 §2.1 修正口径):
BaseM/N = 256 方形 (UnitFlag 单缓冲 L0C 用满 65536 元素), SingleCoreM/N 有界枚举."""
def setUp(self):
from bmm_theory.hardware import ASCEND950PR
self.s = ASCEND950PR
self.router = BranchRouter()
def test_base_tile_square_single_buffer(self):
# 默认 BaseM=BaseN=256 (非 176 双缓冲口径), L0C 恰用满 (256*256*4=256KB)
from bmm_theory.branches.asw_basic import AswBasicBranch
p = AswBasicBranch().make_plan(mkcase(32, 4096, 4096, 4096))
self.assertEqual((p.base_m, p.base_n), (256, 256))
self.assertEqual(p.base_m * p.base_n * 4, self.s.l0c_bytes)
self.assertEqual(p.base_k, 64) # 64KB/(2*256*2) = 64
self.assertTrue(p.fixpipe_unitflag) # UnitFlag 单缓冲
def test_base_tile_exception_small_m(self):
# M=128 < 256: BaseM=128 (被迫跟随), BaseN = min(65536/128=512, N)=512
from bmm_theory.branches.asw_basic import AswBasicBranch
p = AswBasicBranch().make_plan(mkcase(64, 128, 4096, 1024))
self.assertEqual((p.base_m, p.base_n), (128, 512))
self.assertEqual(p.base_k, 32) # min(64K/(2*128*2)=128, 64K/(2*512*2)=32)
def test_v191_example_singlecore(self):
# v1.91 §5.2 完整实例: B=8 M=N=2048 K=1024 bf16 -> (4,4) 512x512, k_l1=128, r=0
case = mkcase(8, 2048, 2048, 1024)
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
p = r["plan"]
self.assertEqual((p.single_core_m, p.single_core_n), (512, 512))
self.assertEqual((p.m_cnt, p.n_cnt), (4, 4))
self.assertEqual(p.k_l1, 128)
self.assertEqual(p.tail_block_cnt, 0) # 8*4*4=128 % 32 = 0 完美整除
def test_singlecore_not_constant_176(self):
# 回归原 bug: SingleCoreM/N 恒为 176 (双缓冲 Base 兜底); 现在自适应
from bmm_theory.branches.asw_basic import AswBasicBranch
p = AswBasicBranch().make_plan(mkcase(8, 4096, 4096, 1024))
self.assertNotEqual((p.single_core_m, p.single_core_n), (176, 176))
self.assertEqual((p.single_core_m, p.single_core_n), (512, 512))
# 且为 Base 的整数倍 (约束 4)
self.assertEqual(p.single_core_m % p.base_m, 0)
self.assertEqual(p.single_core_n % p.base_n, 0)
def test_l0c_single_buffer_constraint_pass(self):
# ASW 单缓冲: base 256x256 通过约束 (无违规)
from bmm_theory.constraints import check_plan_constraints
from bmm_theory.branches.asw_basic import AswBasicBranch
case = mkcase(8, 4096, 4096, 1024)
p = AswBasicBranch().make_plan(case)
self.assertEqual(check_plan_constraints(case, p), [])
def test_p1_fallback_to_enum_when_no_split_infeasible(self):
# B>=C 时不切分 k_l1 过 dValue 下限失败 -> 强制切分枚举 -> 512x512 (8,8)
case = mkcase(128, 4096, 4096, 8192)
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
p = r["plan"]
self.assertEqual((p.single_core_m, p.single_core_n), (512, 512))
self.assertEqual((p.m_cnt, p.n_cnt), (8, 8))
def test_issue34_k_l1_decomposition_derived(self):
# b32_m16_n8192_k7168: 分解 = Base 16x1024 (例外: M=16<256; N 侧受 L0B 单边
# 上限收敛), SingleCore 16x1024 (mCnt=1,nCnt=8), k_l1 = L1 双缓冲反推
# ⌊L1/(2·(sM+sN)·dt)⌋16 = ⌊524288/(2·1040·2)⌋16 = 112 (v1.91 §5.2)
from bmm_theory.branches.asw_basic import AswBasicBranch
case = mkcase(32, 16, 8192, 7168)
p = AswBasicBranch().make_plan(case)
self.assertEqual((p.base_m, p.base_n), (16, 1024))
self.assertEqual((p.single_core_m, p.single_core_n), (16, 1024))
self.assertEqual(p.k_l1, 112)
def test_pathological_shape_degraded_warning_always_plan(self):
# issue#34: 兜底分支恒出方案 —— 极端形状 (N=8 int8 大K, B 侧 dValue=8B
# 物理不可满足) 照常给方案 + 标注效率降级 (warning), 不判违规/不判不可行
from bmm_theory.constraints import check_plan_constraints
from bmm_theory.evaluator import PlanEvaluator
case = mkcase(4096, 2048, 8, 1024, dtype_a="int8", dtype_b="int8")
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
self.assertEqual(check_plan_constraints(case, r["plan"]), [])
self.assertIn("效率降级", r["plan"].note)
er = PlanEvaluator().evaluate(case, r["plan"])
self.assertTrue(er.feasible)
self.assertIn("效率降级", er.advice)
self.assertGreater(er.timing.t_total, 0.0)
def test_pathological_n8_hard_floor_fallback(self):
# 更极端: N=8 int8 连 128B 硬下限也不满足 -> 仍恒出方案, 效率降级标注
from bmm_theory.evaluator import PlanEvaluator
case = mkcase(512, 4096, 1, 8192, dtype_a="int8", dtype_b="int8")
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
self.assertIsNotNone(r["plan"])
self.assertIn("效率降级", r["plan"].note)
er = PlanEvaluator().evaluate(case, r["plan"])
self.assertTrue(er.feasible)
class TestIssue35(unittest.TestCase):
"""issue#35: MergeBatch L1 绑定情形 DMA 命令数多计 b0 倍修复 + 分界泛化口径.
用户 case 家族: B=128, M=1~16, N=128, K=512, bf16. b_core=4, b0=4,
合并后 k_l1^m=240/224 < K=512 (L1 绑定区), 真实每核命令数 = 1x3=3 条
(< IterBatch 的 4 条), 修复前被多计为 12 条导致仲裁翻错方向.
"""
def setUp(self):
self.router = BranchRouter()
self.mb = MergeBatchBranch()
self.ib = IterBatchBranch()
def test_merged_cmd_count_formula(self):
# 每核命令数 = ceil(b_core/b0) * ceil(K/k_l1^m) (K截断时 = b_core/b0)
case = mkcase(128, 16, 128, 512)
r = self.mb.analyze(case)
self.assertTrue(r.capable)
p = r.plan
expect = -(-p.b_core // p.merge_b0) * (-(-case.k // p.k_l1))
self.assertEqual(r.timing.dma_cmd_count, expect)
# 本 case: b0=4, k_l1=224 -> 1*3 = 3 条 (修复前 12 条)
self.assertEqual((p.merge_b0, p.k_l1), (4, 224))
self.assertEqual(r.timing.dma_cmd_count, 3)
def test_cmd_count_truncated_unchanged(self):
# K 截断情形数值不变: cmds = b_core/b0 = IterBatch 的 1/b0
case = mkcase(2048, 32, 32, 256)
mb = self.mb.analyze(case)
ib = self.ib.analyze(case)
self.assertGreaterEqual(mb.plan.k_l1, case.k) # 合并后仍截断
self.assertEqual(mb.timing.dma_cmd_count,
-(-mb.plan.b_core // mb.plan.merge_b0))
self.assertAlmostEqual(
mb.timing.dma_cmd_count / ib.timing.dma_cmd_count,
1.0 / mb.plan.merge_b0, places=6)
def test_boundary_uses_merged_k_l1(self):
# 截断判定与 plan.k_l1 口径一致: k_l1^m < K 时不得声称 K截断
case = mkcase(128, 1, 128, 512)
_, detail = self.mb.beats_iterbatch(case)
p = self.mb.make_plan(case)
self.assertLess(p.k_l1, case.k)
self.assertIn("L1绑定", detail)
self.assertNotIn("K截断", detail)
def test_user_case_family_routing(self):
# B=128,M=1~16,N=128,K=512: issue#36 效率模型后, iter A tile=m*1KB 未饱和
# (m<16), 合并 tile 4 倍大 -> 效率节省 ~0.61us 恒定, drain 随 M 线性增长;
# m<=8 效率节省 > drain -> MergeBatch; m=16 iter tile 恰达 16KB 饱和 ->
# 效率打平, drain 0.66us 决定 -> IterBatch
expect = {1: "MergeBatch", 2: "MergeBatch", 4: "MergeBatch",
8: "MergeBatch", 16: "IterBatch"}
for m, branch in expect.items():
r = self.router.route(mkcase(128, m, 128, 512))
self.assertEqual(r["branch"], branch, f"m={m}: {r['arbitration']}")
self.assertEqual(r["candidates"], {"MergeBatch": True, "IterBatch": True})
self.assertEqual(r["self_check_violations"], [])
def test_arbitration_text_final_winner_consistent(self):
# 仲裁文本 [裁决] 位必须是最终胜者 (分界与时延不一致时括注说明)
import re
r = self.router.route(mkcase(128, 2, 128, 512))
m = re.search(r"\[裁决\] (\w+)", r["arbitration"])
self.assertIsNotNone(m)
self.assertEqual(m.group(1), r["branch"])
class TestIssue36(unittest.TestCase):
"""issue#36: 合并搬移效率建模 (tile=nValue*dValue*dt 放大 b0 倍) + t_cmd_ns=0.
用户澄清: MergeBatch 合并多 batch 左/右矩阵一起搬移, 单块 tile 放大 ->
搬移效率更高, 即便 T_cmd=0 也有效益; 堆叠方向视转置 (A ND 非转置沿
M(nValue), B ND 非转置沿 N(dValue)), 乘积口径不变。
"""
def test_t_cmd_default_zero(self):
self.assertEqual(ASCEND950PR.t_cmd_ns, 0.0)
self.assertEqual(ASCEND950PR.t_cmd, 0.0)
def test_move_eff_curve(self):
from bmm_theory.models import move_eff
cap = ASCEND950PR.min_tile_size # 16KB 饱和点
self.assertEqual(move_eff(cap, cap), 1.0) # 饱和
self.assertEqual(move_eff(2 * cap, cap), 1.0) # 超出仍饱和
self.assertAlmostEqual(move_eff(cap / 4, cap), 0.25) # 之下线性
self.assertEqual(move_eff(0, cap), 1.0) # 零值守卫
def test_merge_eff_gain_beats_iterbatch(self):
# (2048,16,64,128): iter A tile=4KB eff=0.25, 合并 A'=16KB eff=1.0
# -> 效率节省 ~15.7us >> drain 0.17us, T_cmd=0 下 MergeBatch 大胜
case = mkcase(2048, 16, 64, 128)
mb = MergeBatchBranch().analyze(case)
ib = IterBatchBranch().analyze(case)
self.assertTrue(mb.capable and ib.capable)
win, detail = MergeBatchBranch().beats_iterbatch(case)
self.assertTrue(win, detail)
self.assertLess(mb.timing.t_total, ib.timing.t_total)
def test_gm_bytes_unchanged_by_eff(self):
# 效率模型只影响时间列: GM 字节量仍 = V_in (issue#31 口径不被破坏)
for shp in [(128, 1, 128, 512), (2048, 16, 64, 128), (128, 64, 64, 512)]:
case = mkcase(*shp)
for br in (MergeBatchBranch(), IterBatchBranch()):
r = br.analyze(case)
if r.capable:
self.assertAlmostEqual(r.timing.gm_read_bytes,
case.input_bytes, places=3)
def test_iter_small_tile_slower_than_merge(self):
# 用户 case m=1: IterBatch A tile=1KB eff=1/16 -> t_mte2_gm 显著高于
# MergeBatch (A'=1.9KB eff=0.117), 且两者 GM 字节相同
case = mkcase(128, 1, 128, 512)
mb = MergeBatchBranch().analyze(case)
ib = IterBatchBranch().analyze(case)
self.assertGreater(ib.timing.t_mte2_gm, mb.timing.t_mte2_gm)
self.assertEqual(mb.timing.gm_read_bytes, ib.timing.gm_read_bytes)
class TestIssue37(unittest.TestCase):
"""issue#37: ASW_Basic evaluate 尾轮残余 drain 闭式化 (方案甲).
- P1: A0 + r>0 的 drain 由"全 case 三级 max"(≈2x 稳态) 修为 (1-ρ)·T_block;
- P2: 周长型 A1b/方案B 补残余 (√ρ−ρ) / (√(n_wave(n_wave1+ρ))(n_wave1+ρ))·T_load;
- 面积型 A1b/方案B 与 r=0 残余恒 0 (v1.5 §4.3 严格相等);
- 块级三段时延与 _decide_tail 同源 (_block_times, L2 命中口径).
"""
def setUp(self):
from bmm_theory.branches.asw_basic import AswBasicBranch
self.br = AswBasicBranch()
self.s = ASCEND950PR
@staticmethod
def _plan(case, sm, sn, m_cnt, n_cnt, strategy, r, n_wave):
from bmm_theory.models import ImplPlan
return ImplPlan(case_id=case.case_id, branch="ASW_Basic",
used_core_num=32, m_cnt=m_cnt, n_cnt=n_cnt,
single_core_m=sm, single_core_n=sn,
single_core_k=case.k, k_l1=128,
base_m=256, base_n=256, base_k=64,
tail_strategy=strategy, tail_block_cnt=r,
tail_wave_num=n_wave, fixpipe_unitflag=True,
out_dtype_bytes=case.dtype_out_bytes)
def test_block_times_formula(self):
# 块级三段 = v1.5 §2.1 口径 (k_L1 约掉, L2 命中带宽)
case = mkcase(3, 1024, 1024, 1024)
t_mm, t_mv, t_fx = self.br._block_times(case, 256, 256)
self.assertAlmostEqual(t_mm, 2 * 256 * 256 * 1024
/ self.s.q_cube("bf16", "bf16"))
self.assertAlmostEqual(t_mv, 1024 * (256 + 256) * 2 / self.s.bw_l2_pc)
self.assertAlmostEqual(t_fx, 256 * 256 * 2 / self.s.bw_pc)
def test_a0_drain_is_block_level_residual(self):
# P1 修复: A0 + r>0 的 drain = (1-ρ)·T_block (修复前误用全 case 三级
# max -> t_total ≈ 2x 稳态)
case = mkcase(3, 1024, 1024, 1024)
p = self._plan(case, 512, 128, 2, 8, "A0", r=16, n_wave=2)
t = self.br.evaluate(case, p)
t_block = max(self.br._block_times(case, 512, 128))
self.assertAlmostEqual(t.t_drain, (1 - 16 / 32) * t_block)
self.assertAlmostEqual(t.t_total, t.t_steady + t.t_drain)
self.assertLess(t.t_drain, t.t_steady) # 不再 ~2x 稳态
def test_area_dominated_tail_residual_zero(self):
# 面积型 (MMAD/FIX 主导): A1b/方案B 残余恒 0 (v1.5 §4.3 严格相等)
case = mkcase(3, 1024, 1024, 1024)
for strat in ("A1b", "方案B"):
p = self._plan(case, 512, 128, 2, 8, strat, r=16, n_wave=2)
t = self.br.evaluate(case, p)
self.assertEqual(t.t_drain, 0.0, strat)
self.assertAlmostEqual(t.t_total, t.t_steady)
def test_perimeter_a1b_residual(self):
# P2: 周长型 (块级 MTE2 主导) + A1b: drain = (√ρ−ρ)·T_load
case = mkcase(1, 640, 1408, 4096)
t_mm, t_mv, t_fx = self.br._block_times(case, 64, 64)
self.assertGreater(t_mv, max(t_mm, t_fx)) # 确认为周长型前提
p = self._plan(case, 64, 64, 10, 22, "A1b", r=28, n_wave=7)
t = self.br.evaluate(case, p)
rho = 28 / 32
self.assertAlmostEqual(t.t_drain, (rho ** 0.5 - rho) * t_mv)
self.assertAlmostEqual(t.t_total, t.t_steady + t.t_drain)
def test_perimeter_planb_residual(self):
# P2: 周长型 + 方案B: drain = (√(n_wave(n_wave1+ρ))(n_wave1+ρ))·T_load
case = mkcase(1, 640, 1408, 4096)
p = self._plan(case, 64, 64, 10, 22, "方案B", r=28, n_wave=7)
t = self.br.evaluate(case, p)
t_mv = self.br._block_times(case, 64, 64)[1]
rho, x = 28 / 32, 7 - 1 + 28 / 32
self.assertAlmostEqual(t.t_drain, ((7 * x) ** 0.5 - x) * t_mv)
self.assertGreater(t.t_drain, 0.0)
def test_r0_drain_zero(self):
case = mkcase(3, 1024, 1024, 1024)
p = self._plan(case, 512, 128, 2, 8, "A0", r=0, n_wave=2)
self.assertEqual(self.br.evaluate(case, p).t_drain, 0.0)
def test_make_plan_perimeter_a1b_end_to_end(self):
# 端到端 (make_plan 自产方案): 瘦长 case 周长型 + ρρ_dv -> A1b,
# drain 与闭式一致 (_decide_tail 与 evaluate 同源)
case = mkcase(33, 16, 8192, 7168)
p = self.br.make_plan(case)
self.assertGreater(p.tail_block_cnt, 0)
self.assertEqual(p.tail_strategy, "A1b", p.note)
t_mm, t_mv, t_fx = self.br._block_times(
case, p.single_core_m, p.single_core_n)
self.assertGreater(t_mv, max(t_mm, t_fx)) # 周长型前提
t = self.br.evaluate(case, p)
rho = p.tail_block_cnt / 32
self.assertAlmostEqual(t.t_drain, (rho ** 0.5 - rho) * t_mv)
self.assertAlmostEqual(t.t_total, t.t_steady + t.t_drain)
class TestIssue38(unittest.TestCase):
"""issue#38: 昇腾950 系列多 SKU 硬件规格 (950PR 32/28核 + 950DT 36/32/28核,
白皮书表3-1/表4-2), 默认加载调用不变."""
def test_default_spec_unchanged(self):
# 默认加载调用不受影响: ASCEND950PR 各项与 NpuSpec() 关键字默认构造等价
from bmm_theory.hardware import NpuSpec, get_spec
s = ASCEND950PR
self.assertEqual((s.aic_num, s.aiv_num), (32, 64))
self.assertEqual(s.bw_gm, 1.6e12)
self.assertEqual(s.l2_bytes, 128 * 1024 * 1024)
self.assertAlmostEqual(s.q16, 486e12 / 32)
self.assertEqual(s.gm_capacity_gb, 128.0)
self.assertIs(get_spec(), ASCEND950PR) # 默认 = 950PR 主bin
self.assertIs(get_spec("Ascend950PR"), ASCEND950PR)
self.assertEqual(NpuSpec(), ASCEND950PR) # 无参构造 == 主bin
def test_spec_registry_five_skus(self):
from bmm_theory.hardware import SPECS
self.assertEqual(sorted(SPECS), [
"Ascend950DT", "Ascend950DT_C28", "Ascend950DT_C32",
"Ascend950PR", "Ascend950PR_C28"])
for s in SPECS.values():
# 共架构: 单核 Cube 口径全系列一致 (白皮书总算力行取整, 28核档
# 425/28=15.179T 与 15.1875T 差 0.06%, 容差 0.1%); AIV 恒为 AIC 两倍
self.assertLess(abs(s.q16 / (486e12 / 32) - 1), 1e-3)
self.assertEqual(s.aiv_num, 2 * s.aic_num)
self.assertEqual(s.l1_bytes, 512 * 1024) # 表4-2 各档一致
self.assertEqual(s.l0c_bytes, 256 * 1024)
def test_950dt_specs(self):
from bmm_theory.hardware import (ASCEND950DT, ASCEND950DT_C28,
ASCEND950DT_C32)
self.assertEqual((ASCEND950DT.aic_num, ASCEND950DT.aiv_num), (36, 72))
self.assertEqual(ASCEND950DT.bw_gm, 4.0e12) # 4TB/s HBM
self.assertEqual(ASCEND950DT.l2_bytes, 128 * 1024 * 1024)
self.assertEqual(ASCEND950DT.gm_capacity_gb, 144.0)
self.assertEqual(ASCEND950DT.cube_peak_tflops, 547.0)
self.assertEqual((ASCEND950DT_C32.aic_num,
ASCEND950DT_C32.cube_peak_tflops), (32, 486.0))
self.assertEqual((ASCEND950DT_C28.aic_num,
ASCEND950DT_C28.cube_peak_tflops), (28, 425.0))
self.assertEqual(ASCEND950DT_C28.gm_capacity_gb, 96.0)
# 派生量: 单核 GM 份额 4TB/36; r16 平衡点随带宽升至 4TB/s 而减半
self.assertAlmostEqual(ASCEND950DT.bw_pc, 4.0e12 / 36)
self.assertAlmostEqual(ASCEND950DT.r16, 547e12 / (4.0e12 / 2))
self.assertAlmostEqual(ASCEND950DT_C32.r16, 486e12 / (4.0e12 / 2))
# AIV 白皮书交叉验证: fp32 通量 72 核 x 128 lane x 1.65GHz x 2 ≈ 30T
self.assertAlmostEqual(ASCEND950DT.aiv_elem_rate_fp32 * 2 / 1e12,
30.0, places=0)
def test_950pr_c28(self):
from bmm_theory.hardware import ASCEND950PR_C28
self.assertEqual((ASCEND950PR_C28.aic_num,
ASCEND950PR_C28.aiv_num), (28, 56))
self.assertEqual(ASCEND950PR_C28.bw_gm, 1.4e12)
self.assertEqual(ASCEND950PR_C28.l2_bytes, 112 * 1024 * 1024)
self.assertEqual(ASCEND950PR_C28.gm_capacity_gb, 112.0)
def test_get_spec_unknown_raises(self):
from bmm_theory.hardware import get_spec
with self.assertRaises(KeyError):
get_spec("Ascend910")
def test_dt_router_smoke(self):
# 换芯片只换 spec: DT 路由/时延正常; 同一 case GM 段 4TB/s 快于 1.6TB/s
from bmm_theory.hardware import ASCEND950DT
case = mkcase(8, 4096, 4096, 1024)
r_dt = BranchRouter(ASCEND950DT).route(case)
r_pr = BranchRouter().route(case)
self.assertGreater(r_dt["timing"].t_total, 0.0)
self.assertLess(r_dt["timing"].t_mte2_gm, r_pr["timing"].t_mte2_gm)
if __name__ == "__main__":
unittest.main()