diff --git a/BMM/BMM_Theory/tests/test_branches.py b/BMM/BMM_Theory/tests/test_branches.py index 86c259f..c6156c5 100644 --- a/BMM/BMM_Theory/tests/test_branches.py +++ b/BMM/BMM_Theory/tests/test_branches.py @@ -12,7 +12,7 @@ import unittest from bmm_theory.models import BmmCase -from bmm_theory.router import BranchRouter, BRANCH_TO_MATMUL, BRANCH_SPECIAL +from bmm_theory.router import BranchRouter from bmm_theory.branches.merge_batch import MergeBatchBranch from bmm_theory.branches.iter_batch import IterBatchBranch @@ -56,11 +56,43 @@ class TestBranchEntry(unittest.TestCase): 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"], BRANCH_TO_MATMUL) + self.assertEqual(r["branch"], "转Matmul") def test_special_k1(self): r = self.router.route(mkcase(128, 256, 256, 1)) - self.assertEqual(r["branch"], BRANCH_SPECIAL) + 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):