Update BMM_Theory: tests/test_branches.py
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user