"""单元测试: 用理论文档中的典型边界 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.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 截断: MergeBatch 应胜 (v1.1 §4.5) case = mkcase(2048, 32, 32, 256) self.assertTrue(MergeBatchBranch().analyze(case).capable) win, detail = self.mb.beats_iterbatch(case) self.assertTrue(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_GM", "MTE2_L2")) def test_fixpipe_dtype_conversion(self): # C 指定 fp16 输出时, 写出量按 2B 而非 L0C 的 4B case = mkcase(128, 64, 64, 512, dtype_c="fp16") ib = IterBatchBranch().analyze(case) expect = ib.plan.b_core * 64 * 64 * 2 self.assertAlmostEqual(ib.timing.fixpipe_bytes, expect) 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: 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): """第二轮复评问题 (#11-#16) 回归.""" def setUp(self): self.router = BranchRouter() def test_issue11_streamk_fixpipe_no_double_count(self): # issue#11: 部分和写出只经 t_reduce 计账一次; 稳态 fixpipe 不得再计 from bmm_theory.branches.stream_k import StreamKBranch case = mkcase(4, 128, 128, 10240) sk = StreamKBranch().analyze(case) t = sk.timing self.assertAlmostEqual(t.t_fixpipe, 0.0) # 归约串行口径下无稳态 fixpipe 账 self.assertAlmostEqual(t.fixpipe_bytes, 0.0) # 端到端 = max(MTE2, MMAD) + 归约, 不再虚高到 55us/FIXPIPE expect = max(t.t_mte2, t.t_mmad) + t.t_reduce self.assertAlmostEqual(t.t_total, expect) self.assertEqual(t.bottleneck, "MTE2_GM") def test_issue12_k1_pingpong_still_ok(self): # issue#12: 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 from bmm_theory.constraints import check_plan_constraints case = mkcase(256, 1, 256, 4096, dtype_a="int8", dtype_b="int8") # 原 0.2% 违规样例 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") if __name__ == "__main__": unittest.main()