Update BMM_Theory: tests/test_branches.py

This commit is contained in:
2026-09-03 09:25:58 +00:00
parent 1ddedb5ebe
commit 5ab04bab11

View File

@@ -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):