Add BMM_Theory: tests/test_branches.py
This commit is contained in:
116
BMM/BMM_Theory/tests/test_branches.py
Normal file
116
BMM/BMM_Theory/tests/test_branches.py
Normal file
@@ -0,0 +1,116 @@
|
||||
"""单元测试: 用理论文档中的典型边界 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, BRANCH_TO_MATMUL, BRANCH_SPECIAL
|
||||
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"], BRANCH_TO_MATMUL)
|
||||
|
||||
def test_special_k1(self):
|
||||
r = self.router.route(mkcase(128, 256, 256, 1))
|
||||
self.assertEqual(r["branch"], BRANCH_SPECIAL)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user