用户澄清: MergeBatch vs IterBatch 的本质区别不只是 DMA 命令数 —— 合并 b0 个 batch 的左/右矩阵一起搬入 L1, 使单块 tile = nValue*dValue*dt 放大 b0 倍 (堆叠 方向视转置: A ND 非转置沿 M(nValue), B ND 非转置沿 N(dValue)), 搬移效率更高, 即便 T_cmd=0 也有效益。 - models.move_eff: 单命令搬移效率 eff = min(1, tile/min_TileSize) (16KB 饱和, 与进入条件4效率下限语义同源); gm_move_time 按 A/B 两侧字节加权 t = (V_A/eff_A + V_B/eff_B)/BW_gm; 只影响时间列, GM 字节量仍 = V_in - IterBatch: l1_form 补驻留侧返回; move_tiles 分侧口径 (a/b 双侧整K, c 驻留侧 整K+对侧k_l1, d 双侧k_l1), evaluate 接入效率加权 - MergeBatch: 合并 tile 放大 b0 倍接入效率加权; beats_iterbatch 净收益 = 命令节省(cmds差×T_cmd) + 效率节省(t_data差) − drain惩罚, K截断且效率打平且 T_cmd>0 时严格退化为 v1.1 §4.5 闭式; 退役 T_cmd<=0 策略特判 - router: 退役 "T_cmd<=0 策略优先 MergeBatch" 覆盖, 时延模型统一终审 - hardware: t_cmd_ns 50 -> 0 (未标定按 0; 合并收益不再依赖 T_cmd 估计值) - 作用域: 仅切B 两分支接入 (逐命令 tile 小、效率差显著); ASW/StreamK 单命令 tile 通常已饱和, 极端小 tile 走 issue#34 效率降级标注通道 - 用户 case 家族 B=128,M=1~16,N=128,K=512: m=1~8 -> MergeBatch (效率节省 ~0.61us > drain), m=16 -> IterBatch (iter A tile 恰达 16KB 饱和, 效率打平, drain 决定); 分界与时延全家族一致 - demo: merge_demo_k_trunc 形状 (2048,32,32,256)->(2048,16,64,128) (原形状 两侧 tile 均已 16KB 饱和, t_cmd=0 下无收益转 IterBatch; 新形状 iter A tile 4KB eff=0.25 vs 合并 16KB eff=1.0, 保持 MergeBatch 胜出演示且仍 K截断) - 测试: 74/74 (新增 TestIssue36 5 例: 效率曲线/字节不变/效率差胜出/家族; TestArbitration/TestZeroCmdHandling 按 t_cmd=0+效率语义重写; TestIssue35 家族期望更新) - 文档: 01_MergeBatch §4/§5 效率模型+泛化净收益; 02_IterBatch 口径注; 00_总纲胜出条件; 01_软件架构 T_cmd 标定说明; 05 时间列效率口径注; README 要点 - 验证: examples 重生成可复现 0 diff; 压力 10000 例 0 崩溃/0 NaN/0 违规/ 0 GM<V_in, 七分支覆盖 (MergeBatch 386 例)
936 lines
46 KiB
Python
936 lines
46 KiB
Python
"""单元测试: 用理论文档中的典型边界 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.hardware import ASCEND950PR
|
|
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 截断 + 合并 tile 效率差: MergeBatch 应胜
|
|
# (issue#36: (2048,16,64,128) IterBatch A tile=4KB eff=0.25 vs 合并
|
|
# 16KB eff=1.0, 效率节省 15.7us >> drain 0.17us; T_cmd=0 默认下仍胜)
|
|
case = mkcase(2048, 16, 64, 128)
|
|
self.assertTrue(MergeBatchBranch().analyze(case).capable)
|
|
win, detail = self.mb.beats_iterbatch(case)
|
|
self.assertTrue(win, detail)
|
|
|
|
def test_no_eff_diff_mergebatch_loses(self):
|
|
# issue#36: K 截断但 IterBatch 单命令 tile 已达 16KB 饱和 (效率打平),
|
|
# T_cmd=0 下命令节省为 0, 合并只剩 drain 惩罚 -> IterBatch 优
|
|
# (原 t_cmd=50ns 时代的 MergeBatch 胜例 (2048,32,32,256), 语义迁移)
|
|
case = mkcase(2048, 32, 32, 256)
|
|
win, detail = self.mb.beats_iterbatch(case)
|
|
self.assertFalse(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", "MMAD"))
|
|
|
|
def test_fixpipe_dtype_conversion(self):
|
|
# C 指定 fp16 输出时, 写出量按 2B 而非 L0C 的 4B;
|
|
# 字节列为整芯片口径 (issue#29): fixpipe_bytes = B*MN*outB
|
|
case = mkcase(128, 64, 64, 512, dtype_c="fp16")
|
|
ib = IterBatchBranch().analyze(case)
|
|
expect = case.batch_c * 64 * 64 * 2
|
|
self.assertAlmostEqual(ib.timing.fixpipe_bytes, expect)
|
|
self.assertAlmostEqual(ib.timing.fixpipe_bytes, case.output_bytes)
|
|
|
|
|
|
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/#17: 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):
|
|
"""第三轮复评问题 (#17-#20) + 恢复 #11-#15 回归."""
|
|
|
|
def setUp(self):
|
|
self.router = BranchRouter()
|
|
|
|
def test_issue11_streamk_fixpipe_no_double_count(self):
|
|
# issue#11/#17: 部分和写出只经 t_reduce 计账一次; 稳态 fixpipe 不得再计
|
|
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
|
|
self.assertAlmostEqual(t.t_fixpipe, 0.0) # 归约串行口径下无稳态 fixpipe 账
|
|
self.assertAlmostEqual(t.fixpipe_bytes, 0.0)
|
|
# 端到端 = 稳态(MTE2搬移链) + 归约, fixpipe 无重复账
|
|
self.assertEqual(t.bottleneck, "MTE2")
|
|
self.assertAlmostEqual(t.t_total, t.t_steady + t.t_drain)
|
|
|
|
def test_issue12_k1_pingpong_still_ok(self):
|
|
# issue#12/#17: 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
|
|
# (issue#31 后该样例 MergeBatch 时延已不再胜出, 用 mock 压低其时延以固定
|
|
# "胜者违规 -> 回退" 机制路径; MergeBatch 方案 dValue=112B<128 仍违规)
|
|
from unittest import mock
|
|
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
|
from bmm_theory.constraints import check_plan_constraints
|
|
case = mkcase(256, 1, 256, 4096, dtype_a="int8", dtype_b="int8")
|
|
mb_plan = MergeBatchBranch().analyze(case).plan
|
|
self.assertTrue(check_plan_constraints(case, mb_plan),
|
|
"该样例 MergeBatch 方案应仍自检违规 (dValue)")
|
|
real_eval = MergeBatchBranch.evaluate
|
|
|
|
def fake_win(self, c, p):
|
|
t = real_eval(self, c, p)
|
|
t.t_total = 1e-9 # 强制 MergeBatch 时延胜出 -> 仲裁选它 -> 走自检/回退
|
|
return t
|
|
with mock.patch.object(MergeBatchBranch, "evaluate", fake_win):
|
|
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")
|
|
|
|
def test_issue18_placeholder_plan_infeasible_in_evaluate(self):
|
|
# issue#18: 占位方案(used_core_num=0)在 evaluate 中必须不可行,
|
|
# 不得被当作可行方案给出正常时延
|
|
from bmm_theory.evaluator import PlanEvaluator
|
|
from bmm_theory.models import ImplPlan
|
|
# K=1 且 B<128 已恢复真实单缓冲方案 (#17), 故直接构造占位 plan 验证约束层
|
|
case = mkcase(64, 8192, 32, 1, dtype_a="int8", dtype_b="int8")
|
|
r = self.router.route(case)
|
|
self.assertGreater(r["plan"].used_core_num, 0) # 真实方案
|
|
ph = ImplPlan(case_id="ph", branch="特殊分支", used_core_num=0,
|
|
note="该区域暂无理论方案(进入条件不满足)")
|
|
er = PlanEvaluator().evaluate(case, ph)
|
|
self.assertFalse(er.feasible, "占位方案应判不可行")
|
|
self.assertIn("used_core_num", er.violations)
|
|
|
|
def test_issue19_transpose_dvalue_guard_effective(self):
|
|
# issue#19: 转置感知 dValue 判据在生成守卫/条件4/约束三处同源后真正生效.
|
|
# 判别形状: B=64 M=4096 N=64 K=4096 bf16 —— d 形态 k_l1=16,
|
|
# 均不转置时 dv_a=k_l1*2=32B <128 挡下 (B 侧 dv_b=N*2=128B 恰好达标也不放行,
|
|
# 因为两侧切 K 两侧都要高效);
|
|
# A 转置后 dv_a=M*2=8192B, 应能走 IterBatch 形态 d.
|
|
from bmm_theory.constraints import check_plan_constraints
|
|
c_not = BmmCase(case_id="x", batch_a=64, batch_b=64, m=4096, n=64, k=4096,
|
|
trans_a=False, trans_b=False)
|
|
r_not = self.router.route(c_not)
|
|
self.assertNotEqual(r_not["branch"], "IterBatch") # 非转置被 dValue 守卫挡下
|
|
c_tr = BmmCase(case_id="x", batch_a=64, batch_b=64, m=4096, n=64, k=4096,
|
|
trans_a=True, trans_b=False)
|
|
r_tr = self.router.route(c_tr)
|
|
self.assertEqual(r_tr["branch"], "IterBatch") # A 转置 M 向连续, 守卫放行
|
|
self.assertIn("d_", r_tr["plan"].l1_form)
|
|
self.assertEqual(check_plan_constraints(c_tr, r_tr["plan"]), [])
|
|
|
|
|
|
class TestFp4Support(unittest.TestCase):
|
|
"""fp4 (0.5B) dtype 支持 (对齐 bmmv3, 2026-09-03)."""
|
|
|
|
def test_fp4_dtype_bytes(self):
|
|
from bmm_theory.models import dtype_bytes
|
|
self.assertEqual(dtype_bytes("fp4"), 0.5)
|
|
self.assertEqual(dtype_bytes("fp4_e2m1"), 0.5)
|
|
|
|
def test_fp4_case_creation(self):
|
|
case = mkcase(32, 64, 64, 256, dtype_a="fp4", dtype_b="fp4", dtype_c="fp16")
|
|
self.assertEqual(case.dtype_in_bytes, 0.5)
|
|
# fp4 输入 + fp16 输出: 输入 0.5B, 输出 2B
|
|
self.assertEqual(case.dtype_out_bytes, 2)
|
|
|
|
def test_fp4_dvalue_threshold(self):
|
|
# fp4 (0.5B) 时 dValue 128B 需要 k_l1 >= 256
|
|
from bmm_theory.constraints import check_plan_constraints
|
|
case = mkcase(128, 64, 64, 128, dtype_a="fp4", dtype_b="fp4", dtype_c="fp16")
|
|
r = BranchRouter().route(case)
|
|
v = check_plan_constraints(case, r["plan"])
|
|
# fp4 小 K 场景应能正常路由且不报 dValue 违规 (k_l1 连续维是 M/N)
|
|
self.assertEqual(v, [], f"fp4 case 不应报违规: {v}")
|
|
|
|
|
|
class TestTransposeModeling(unittest.TestCase):
|
|
"""转置对 dValue 连续维的影响建模 (对齐 bmmv3, 2026-09-03)."""
|
|
|
|
def test_transpose_affects_dvalue_judgment(self):
|
|
# A 不转置: K 向连续, dValue 判 K*dt
|
|
# A 转置: M 向连续, dValue 判 M*dt
|
|
# B 不转置: N 向连续, dValue 判 N*dt
|
|
# B 转置: K 向连续, dValue 判 K*dt
|
|
from bmm_theory.constraints import _k_segment_is_contiguous
|
|
# 当前版本不建模转置, 默认按不转置处理
|
|
case = mkcase(128, 64, 64, 512)
|
|
case.trans_a = False
|
|
case.trans_b = False
|
|
# 验证 trans 字段存在且可读写 (为后续建模做准备)
|
|
self.assertFalse(case.trans_a)
|
|
self.assertFalse(case.trans_b)
|
|
case.trans_a = True
|
|
self.assertTrue(case.trans_a)
|
|
|
|
def test_transpose_a_large_m_small_k(self):
|
|
# A 转置 + 大 M 小 K: dValue 应判 M*dt (M 向连续), 不受 K 小影响
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
# M=1024 (M*dt=2048B >= 128B), K=8 (K*dt=16B < 128B)
|
|
case = mkcase(128, 1024, 64, 8, dtype_a="bf16", dtype_b="bf16")
|
|
case.trans_a = True
|
|
case.trans_b = False
|
|
ib = IterBatchBranch().analyze(case)
|
|
# A 转置时 dValue 判 M*dt=2048B >= 128B, 应通过
|
|
c4 = [c for c in ib.checks if "搬移效率" in c.name][0]
|
|
self.assertTrue(c4.passed, f"A 转置时应判 M 向连续: {c4.detail}")
|
|
|
|
def test_no_transpose_small_k_fails(self):
|
|
# A 不转置 + 小 K + c/d 形态: dValue 判 K*dt, K=8 时 16B < 128B 应失败
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
# 大 M/N 让 L1 放不下整 K, 走 c/d 形态
|
|
case = mkcase(128, 1024, 1024, 8, dtype_a="bf16", dtype_b="bf16")
|
|
case.trans_a = False
|
|
case.trans_b = False
|
|
ib = IterBatchBranch().analyze(case)
|
|
c4 = [c for c in ib.checks if "搬移效率" in c.name][0]
|
|
# c/d 形态下 A 不转置时 K=8 应报 dValue 违规
|
|
if ib.plan.l1_form.startswith(("c_", "d_")):
|
|
self.assertFalse(c4.passed, f"A 不转置时 K=8 应报 dValue 违规: {c4.detail}")
|
|
|
|
|
|
class TestZeroCmdHandling(unittest.TestCase):
|
|
"""950PR 默认 T_cmd=0 (未标定, issue#36): 整链路不得除零/崩溃; 合并收益由
|
|
搬移效率模型 (move_eff, 合并 tile 放大 b0 倍) 刻画, 不再设策略覆盖."""
|
|
|
|
def test_beats_iterbatch_zero_cmd(self):
|
|
# T_cmd=0: 命令节省项为 0, 由效率节省 vs drain 惩罚决定 (issue#36)
|
|
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
|
mb = MergeBatchBranch() # 默认 spec 即 t_cmd_ns=0
|
|
# (b, m, n, k, MergeBatch应胜与否=效率节省>drain惩罚)
|
|
# (2048,16,64,128): iter A tile=4KB eff=0.25 vs 合并 16KB eff=1.0 -> 大胜
|
|
# (2048,32,32,256): iter A tile=16KB 已饱和, 效率打平 -> drain 惩罚 -> 恒劣
|
|
# (128,64,64,512): issue#35 第三情形 (dValue cap 截断, 命令/效率均打平) -> 恒劣
|
|
# (256,128,128,4096): L1 绑定, tile 均 >=128KB 饱和 -> 恒劣
|
|
cases = [(2048, 16, 64, 128, True), (2048, 32, 32, 256, False),
|
|
(128, 64, 64, 512, False), (256, 128, 128, 4096, False)]
|
|
for b, m, n, k, mb_wins in cases:
|
|
win, detail = mb.beats_iterbatch(mkcase(b, m, n, k))
|
|
self.assertEqual(win, mb_wins, f"{b},{m},{n},{k}: {detail}")
|
|
self.assertIn("效率节省", detail)
|
|
|
|
def test_route_with_zero_cmd_efficiency_decides(self):
|
|
# 默认 t_cmd=0: 有效率低下的 IterBatch 小 tile case 由 MergeBatch 胜;
|
|
# tile 均饱和的 case 由 IterBatch 胜 (drain 惩罚, 无策略覆盖)
|
|
router = BranchRouter()
|
|
r = router.route(mkcase(2048, 16, 64, 128)) # 效率差显著 -> MergeBatch
|
|
self.assertEqual(r["branch"], "MergeBatch")
|
|
self.assertIsNotNone(r["timing"])
|
|
r2 = router.route(mkcase(2048, 32, 32, 256)) # tile 均饱和 -> IterBatch
|
|
self.assertEqual(r2["branch"], "IterBatch")
|
|
shapes = [(128, 64, 64, 512), (64, 64, 64, 8192), (512, 128, 128, 128),
|
|
(128, 128, 128, 1024), (32, 4096, 4096, 4096)]
|
|
for b, m, n, k in shapes:
|
|
rr = router.route(mkcase(b, m, n, k))
|
|
t = rr["timing"]
|
|
self.assertIsNotNone(t, f"{b},{m},{n},{k} 应有 timing")
|
|
self.assertGreater(t.t_total, 0.0)
|
|
self.assertEqual(t.t_total, t.t_total) # 非 NaN
|
|
|
|
def test_zero_cmd_random_smoke(self):
|
|
# 随机小样本冒烟: 无异常/无 NaN
|
|
import random
|
|
from bmm_theory.hardware import NpuSpec
|
|
router = BranchRouter(NpuSpec(t_cmd_ns=0.0))
|
|
rng = random.Random(9)
|
|
for _ in range(300):
|
|
b = rng.choice([2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048])
|
|
m = rng.choice([8, 16, 32, 64, 128, 256, 512, 1024, 4096])
|
|
n = rng.choice([8, 16, 32, 64, 128, 256, 512, 1024, 4096])
|
|
k = rng.choice([16, 32, 64, 128, 256, 512, 1024, 4096, 16384])
|
|
c = BmmCase(case_id="z", batch_a=b, batch_b=b, m=m, n=n, k=k)
|
|
rr = router.route(c)
|
|
self.assertIsNotNone(rr["plan"])
|
|
t = rr["timing"]
|
|
self.assertIsNotNone(t)
|
|
self.assertGreater(t.t_total, 0.0)
|
|
self.assertEqual(t.t_total, t.t_total)
|
|
|
|
|
|
class TestIssue27to30(unittest.TestCase):
|
|
"""第四轮评审 (issue#27-#30, 设计文档 docs/05).
|
|
|
|
#27 MergeBatch Cube 公式复核 (每步 2*(b0M)(b0N)K x b_core/b0 步);
|
|
#28 dtype 感知算力 (Cube/AIV 速率表);
|
|
#29 GM 首读下限 GM>=V_in + 字节列整芯片口径;
|
|
#30 Fixpipe 输出落点 (整case驻留 -> L2, 否则直写 GM).
|
|
"""
|
|
|
|
def setUp(self):
|
|
from bmm_theory.hardware import ASCEND950PR
|
|
self.s = ASCEND950PR
|
|
self.router = BranchRouter()
|
|
|
|
# ---------------- #27 ----------------
|
|
def test_issue27_merge_cube_flops_matches_formula(self):
|
|
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
|
case = mkcase(2048, 32, 32, 256)
|
|
mb = MergeBatchBranch().analyze(case)
|
|
self.assertTrue(mb.capable)
|
|
p, t = mb.plan, mb.timing
|
|
b0, bc = p.merge_b0, p.b_core
|
|
# 整芯片列 = B*b0*2MNK; 等价 (b_core/b0) 步 x 每步 2*(b0M)(b0N)K x C 核
|
|
step_flops = 2.0 * (b0 * 32) * (b0 * 32) * 256
|
|
self.assertAlmostEqual(t.cube_flops,
|
|
step_flops * (bc / b0) * self.s.aic_num)
|
|
self.assertAlmostEqual(t.cube_flops,
|
|
case.batch_c * b0 * 2.0 * 32 * 32 * 256)
|
|
# t_mmad = 每核 flops / 单核算力
|
|
self.assertAlmostEqual(t.t_mmad,
|
|
(bc * b0 * 2.0 * 32 * 32 * 256) / self.s.q16)
|
|
|
|
def test_issue27_merge_gm_chip_col_equals_input(self):
|
|
# K 截断合并: GM 每字节一次 -> gm 整芯片列 = V_in (issue#29 口径)
|
|
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
|
case = mkcase(2048, 32, 32, 256)
|
|
t = MergeBatchBranch().analyze(case).timing
|
|
self.assertAlmostEqual(t.gm_read_bytes, case.input_bytes)
|
|
self.assertAlmostEqual(t.l2_read_bytes, 0.0)
|
|
|
|
# ---------------- #28 ----------------
|
|
def test_issue28_cube_dtype_rate(self):
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
base = mkcase(128, 64, 64, 512) # IterBatch 形态 b, 各 dtype 均可行
|
|
t16 = IterBatchBranch().analyze(base).timing
|
|
t32 = IterBatchBranch().analyze(
|
|
mkcase(128, 64, 64, 512, dtype_a="fp32", dtype_b="fp32")).timing
|
|
t8 = IterBatchBranch().analyze(
|
|
mkcase(128, 64, 64, 512, dtype_a="fp8", dtype_b="fp8")).timing
|
|
# 同计算量, 时延反比于 dtype 算力: fp32=1/2, fp8=2x (相对 bf16)
|
|
self.assertAlmostEqual(t32.t_mmad / t16.t_mmad, 2.0, places=6)
|
|
self.assertAlmostEqual(t8.t_mmad / t16.t_mmad, 0.5, places=6)
|
|
|
|
def test_issue28_aiv_dtype_rate_k1(self):
|
|
# K=1 AIV 逐元素: bf16 通量 2x fp32 -> fp32 时延为 bf16 的 2 倍
|
|
base = mkcase(64, 8192, 512, 1)
|
|
t_bf16 = BranchRouter().route(base)["timing"]
|
|
t_fp32 = BranchRouter().route(
|
|
mkcase(64, 8192, 512, 1, dtype_a="fp32", dtype_b="fp32"))["timing"]
|
|
self.assertAlmostEqual(t_fp32.t_mmad / t_bf16.t_mmad, 2.0, places=6)
|
|
|
|
# ---------------- #29 ----------------
|
|
def test_issue29_gm_floor_and_chip_columns(self):
|
|
# 各分支代表 case: GM 整芯片列 >= V_in (R1), 且迭代/合并/转matmul/special
|
|
# 等于 V_in (无重复读结构)
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
|
from bmm_theory.branches.stream_k import StreamKBranch
|
|
from bmm_theory.branches.asw_basic import AswBasicBranch
|
|
checks = [
|
|
(IterBatchBranch().analyze(mkcase(128, 64, 64, 512)).timing,
|
|
mkcase(128, 64, 64, 512)),
|
|
(MergeBatchBranch().analyze(mkcase(2048, 32, 32, 256)).timing,
|
|
mkcase(2048, 32, 32, 256)),
|
|
(StreamKBranch().analyze(mkcase(4, 128, 128, 10240)).timing,
|
|
mkcase(4, 128, 128, 10240)),
|
|
(AswBasicBranch().analyze(mkcase(8, 4096, 4096, 1024)).timing,
|
|
mkcase(8, 4096, 4096, 1024)),
|
|
]
|
|
for t, case in checks:
|
|
self.assertGreaterEqual(t.gm_read_bytes + 1e-6, case.input_bytes,
|
|
f"{case.case_id} GM < V_in")
|
|
# 转Matmul / 特殊分支 (K=1) 亦等于 V_in
|
|
r = self.router.route(mkcase(1, 2048, 2048, 2048))
|
|
t = r["timing"]
|
|
self.assertGreaterEqual(t.gm_read_bytes + 1e-6,
|
|
BmmCase(case_id="x", batch_a=1, batch_b=1,
|
|
m=2048, n=2048, k=2048).input_bytes)
|
|
r = self.router.route(mkcase(128, 256, 256, 1))
|
|
self.assertAlmostEqual(r["timing"].gm_read_bytes, 128 * (256 + 256) * 2)
|
|
|
|
def test_issue29_gm_floor_random_smoke(self):
|
|
# 随机多 dtype 冒烟: 所有可行方案 GM 整芯片列 >= V_in, 无 NaN
|
|
import random
|
|
rng = random.Random(20260609)
|
|
dtypes = ["bf16", "fp16", "fp32", "int8", "fp8", "fp4"]
|
|
for _ in range(400):
|
|
dt = rng.choice(dtypes)
|
|
kw = dict(dtype_a=dt, dtype_b=dt, dtype_c=rng.choice(["bf16", "fp16"]))
|
|
b = rng.choice([1, 2, 4, 8, 16, 32, 64, 128, 256, 2048])
|
|
m = rng.choice([1, 8, 32, 64, 128, 256, 1024, 4096])
|
|
n = rng.choice([1, 8, 32, 64, 128, 256, 1024, 4096])
|
|
k = rng.choice([0, 1, 32, 64, 128, 256, 512, 1024, 4096])
|
|
c = BmmCase(case_id="gm", batch_a=b, batch_b=b, m=m, n=n, k=k, **kw)
|
|
r = self.router.route(c)
|
|
t = r["timing"]
|
|
self.assertIsNotNone(t, c.case_id)
|
|
self.assertEqual(t.t_total, t.t_total) # 非 NaN
|
|
self.assertGreaterEqual(t.gm_read_bytes + 1e-6, c.input_bytes,
|
|
f"{dt} b={b} m={m} n={n} k={k} gm={t.gm_read_bytes}")
|
|
|
|
# ---------------- #30 ----------------
|
|
def test_issue30_iter_output_l2_when_whole_case_fits(self):
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
# 整 case 输入+输出 16.8MB+1.0MB <= 128MB -> 输出写 L2 写口, GM 写=0
|
|
case = mkcase(128, 64, 64, 512)
|
|
self.assertLess(case.input_bytes + case.output_bytes, self.s.l2_bytes)
|
|
t = IterBatchBranch().analyze(case).timing
|
|
self.assertAlmostEqual(t.fixpipe_bytes, case.output_bytes)
|
|
self.assertAlmostEqual(t.t_fixpipe, case.output_bytes / self.s.bw_l2)
|
|
|
|
def test_issue30_iter_output_gm_when_case_overflows(self):
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
# 输出 268MB > L2 余量 -> 直写 GM (输入优先驻留)
|
|
case = mkcase(128, 1024, 1024, 64)
|
|
self.assertGreater(case.input_bytes + case.output_bytes, self.s.l2_bytes)
|
|
t = IterBatchBranch().analyze(case).timing
|
|
self.assertAlmostEqual(t.fixpipe_bytes, case.output_bytes)
|
|
self.assertAlmostEqual(t.t_fixpipe, case.output_bytes / self.s.bw_gm)
|
|
|
|
def test_issue30_asw_scene_whole_case_labels(self):
|
|
# S_A (整case全驻留): l2_policy_out=resident, t_fixpipe 走 L2 写口
|
|
case_sa = mkcase(2, 4096, 4096, 512)
|
|
self.assertLess(case_sa.input_bytes + case_sa.output_bytes, self.s.l2_bytes)
|
|
r = self.router.route(case_sa)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
self.assertTrue(r["plan"].l2_policy_out.startswith("resident"), r["plan"].l2_policy_out)
|
|
self.assertIn("A_整case全驻留", r["plan"].note)
|
|
self.assertAlmostEqual(r["timing"].t_fixpipe,
|
|
case_sa.output_bytes / self.s.bw_l2)
|
|
# S_B (整case超L2, 单batch输入可驻留): 输出直写 GM
|
|
case_sb = mkcase(64, 4096, 4096, 1024)
|
|
self.assertGreater(case_sb.input_bytes + case_sb.output_bytes, self.s.l2_bytes)
|
|
self.assertLess(4096 * 1024 * 2 * 2, self.s.l2_bytes) # 单batch输入 16.8MB
|
|
r = self.router.route(case_sb)
|
|
self.assertIn("ASW", r["branch"])
|
|
self.assertTrue(r["plan"].l2_policy_out.startswith("direct_gm"),
|
|
r["plan"].l2_policy_out)
|
|
self.assertAlmostEqual(r["timing"].t_fixpipe,
|
|
case_sb.output_bytes / self.s.bw_gm)
|
|
|
|
def test_issue30_to_matmul_output_l2_toggle(self):
|
|
# 转Matmul 粗估: 整case可驻留 -> L2; 大 case -> GM
|
|
small = mkcase(1, 2048, 2048, 2048)
|
|
self.assertLess(small.input_bytes + small.output_bytes, self.s.l2_bytes)
|
|
r = self.router.route(small)
|
|
self.assertEqual(r["branch"], "转Matmul")
|
|
self.assertAlmostEqual(r["timing"].t_fixpipe,
|
|
small.output_bytes / self.s.bw_l2)
|
|
big = mkcase(1, 8192, 8192, 4096)
|
|
self.assertGreater(big.input_bytes + big.output_bytes, self.s.l2_bytes)
|
|
r = self.router.route(big)
|
|
self.assertEqual(r["branch"], "转Matmul")
|
|
self.assertAlmostEqual(r["timing"].t_fixpipe,
|
|
big.output_bytes / self.s.bw_gm)
|
|
|
|
|
|
class TestIssue31to32(unittest.TestCase):
|
|
"""第五轮评审: #31 切B类 GM 每字节恰读一次 = V_in; #32 ASW 场景升级
|
|
(单侧全驻留 -> S_B 使 GM=V_in; 双侧超 L2 -> S_C 最小替换 2D 分组 + 窗口 L2 计账)."""
|
|
|
|
def setUp(self):
|
|
from bmm_theory.hardware import ASCEND950PR
|
|
self.s = ASCEND950PR
|
|
self.router = BranchRouter()
|
|
|
|
# ---------------- #31 ----------------
|
|
def test_issue31_iter_form_c_gm_equals_vin(self):
|
|
# b64_m16_n256_k512: 形态 c, K=512 切 n_K=3 (k_L1=240) 非整除 -> GM 仍 = V_in
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
case = mkcase(64, 16, 256, 512)
|
|
ib = IterBatchBranch().analyze(case)
|
|
self.assertIn("c_", ib.plan.l1_form)
|
|
self.assertLess(ib.plan.k_l1, case.k)
|
|
t = ib.timing
|
|
self.assertAlmostEqual(t.gm_read_bytes, case.input_bytes)
|
|
self.assertAlmostEqual(t.l2_read_bytes, 0.0)
|
|
self.assertAlmostEqual(t.t_mte2_gm, case.input_bytes / self.s.bw_gm)
|
|
|
|
def test_issue31_iter_form_d_gm_equals_vin(self):
|
|
# K=5000 非整除切段 (k_L1=1024, n_K=5, 末段剩 904): GM 仍恰一次 = V_in
|
|
from bmm_theory.branches.iter_batch import IterBatchBranch
|
|
case = mkcase(128, 64, 64, 5000)
|
|
ib = IterBatchBranch().analyze(case)
|
|
self.assertIn("d_", ib.plan.l1_form)
|
|
self.assertLess(ib.plan.k_l1, case.k)
|
|
self.assertAlmostEqual(ib.timing.gm_read_bytes, case.input_bytes)
|
|
self.assertAlmostEqual(ib.timing.t_mte2_gm,
|
|
case.input_bytes / self.s.bw_gm)
|
|
|
|
def test_issue31_merge_l1_bound_gm_equals_vin(self):
|
|
# MergeBatch L1 绑定 (k_L1<K): 切 K 段互不重叠, GM = V_in
|
|
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
|
case = mkcase(811, 33, 1, 2459, dtype_a="fp32", dtype_b="fp32")
|
|
mb = MergeBatchBranch().analyze(case)
|
|
self.assertTrue(mb.capable)
|
|
self.assertLess(mb.plan.k_l1, case.k) # L1 绑定
|
|
self.assertAlmostEqual(mb.timing.gm_read_bytes, case.input_bytes)
|
|
self.assertAlmostEqual(mb.timing.l2_read_bytes, 0.0)
|
|
|
|
# ---------------- #32 ----------------
|
|
def test_issue32_asw_single_side_resident_gm_equals_vin(self):
|
|
# b32_m8192_n4096_k7168: 单batch输入 176MB > L2, 但 B=58.7MB 可全驻留
|
|
# + A 行块滑窗 (58.7+2x2.5 <= 128) -> S_B: GM = V_in, 重复读全命中 L2
|
|
case = mkcase(32, 8192, 4096, 7168)
|
|
r = self.router.route(case)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
p, t = r["plan"], r["timing"]
|
|
a_b = case.m * case.k * case.dtype_in_bytes
|
|
bb_b = case.k * case.n * case.dtype_in_bytes
|
|
m_cnt, n_cnt = p.m_cnt, p.n_cnt
|
|
self.assertGreater(a_b + bb_b, self.s.l2_bytes) # 单 batch 确实超 L2
|
|
self.assertLessEqual(bb_b + 2 * (a_b / m_cnt),
|
|
self.s.l2_bytes) # 单侧全驻留成立
|
|
self.assertAlmostEqual(t.gm_read_bytes, case.input_bytes)
|
|
exp_l2 = case.batch_c * ((n_cnt - 1) * a_b + (m_cnt - 1) * bb_b)
|
|
self.assertAlmostEqual(t.l2_read_bytes, exp_l2)
|
|
|
|
def test_issue32_asw_scene_c_min_gm(self):
|
|
# b8_m131072_n8192_k8192: 双侧均不可全驻留 -> S_C 最小替换分组;
|
|
# gm == 容量约束最小解 (测试内复算), 且 >= V_in
|
|
case = mkcase(8, 131072, 8192, 8192)
|
|
r = self.router.route(case)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
p, t = r["plan"], r["timing"]
|
|
self.assertIn("C_", p.note)
|
|
m_cnt, n_cnt = p.m_cnt, p.n_cnt
|
|
dt = case.dtype_in_bytes
|
|
a_b = case.m * case.k * dt
|
|
bb_b = case.k * case.n * dt
|
|
blk_a = a_b / m_cnt
|
|
blk_b = bb_b / n_cnt
|
|
best = None
|
|
for mg in range(1, m_cnt + 1):
|
|
for ng in range(1, n_cnt + 1):
|
|
if mg * blk_a + ng * blk_b > self.s.l2_bytes:
|
|
continue
|
|
gm = (n_cnt + ng - 1) // ng * a_b + (m_cnt + mg - 1) // mg * bb_b
|
|
if best is None or gm < best:
|
|
best = gm
|
|
self.assertIsNotNone(best)
|
|
self.assertAlmostEqual(t.gm_read_bytes, case.batch_c * best)
|
|
self.assertGreaterEqual(t.gm_read_bytes, case.input_bytes)
|
|
# 最小替换解应严格优于"每块独立落 GM"的保守上界
|
|
self.assertLessEqual(t.gm_read_bytes,
|
|
case.batch_c * (n_cnt * a_b + m_cnt * bb_b))
|
|
|
|
|
|
class TestIssue33(unittest.TestCase):
|
|
"""ASW_Basic tile 选择 (issue#33, v1.91 §5.1/§5.2 + 尾轮 v1.5 §2.1 修正口径):
|
|
BaseM/N = 256 方形 (UnitFlag 单缓冲 L0C 用满 65536 元素), SingleCoreM/N 有界枚举."""
|
|
|
|
def setUp(self):
|
|
from bmm_theory.hardware import ASCEND950PR
|
|
self.s = ASCEND950PR
|
|
self.router = BranchRouter()
|
|
|
|
def test_base_tile_square_single_buffer(self):
|
|
# 默认 BaseM=BaseN=256 (非 176 双缓冲口径), L0C 恰用满 (256*256*4=256KB)
|
|
from bmm_theory.branches.asw_basic import AswBasicBranch
|
|
p = AswBasicBranch().make_plan(mkcase(32, 4096, 4096, 4096))
|
|
self.assertEqual((p.base_m, p.base_n), (256, 256))
|
|
self.assertEqual(p.base_m * p.base_n * 4, self.s.l0c_bytes)
|
|
self.assertEqual(p.base_k, 64) # 64KB/(2*256*2) = 64
|
|
self.assertTrue(p.fixpipe_unitflag) # UnitFlag 单缓冲
|
|
|
|
def test_base_tile_exception_small_m(self):
|
|
# M=128 < 256: BaseM=128 (被迫跟随), BaseN = min(65536/128=512, N)=512
|
|
from bmm_theory.branches.asw_basic import AswBasicBranch
|
|
p = AswBasicBranch().make_plan(mkcase(64, 128, 4096, 1024))
|
|
self.assertEqual((p.base_m, p.base_n), (128, 512))
|
|
self.assertEqual(p.base_k, 32) # min(64K/(2*128*2)=128, 64K/(2*512*2)=32)
|
|
|
|
def test_v191_example_singlecore(self):
|
|
# v1.91 §5.2 完整实例: B=8 M=N=2048 K=1024 bf16 -> (4,4) 512x512, k_l1=128, r=0
|
|
case = mkcase(8, 2048, 2048, 1024)
|
|
r = self.router.route(case)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
p = r["plan"]
|
|
self.assertEqual((p.single_core_m, p.single_core_n), (512, 512))
|
|
self.assertEqual((p.m_cnt, p.n_cnt), (4, 4))
|
|
self.assertEqual(p.k_l1, 128)
|
|
self.assertEqual(p.tail_block_cnt, 0) # 8*4*4=128 % 32 = 0 完美整除
|
|
|
|
def test_singlecore_not_constant_176(self):
|
|
# 回归原 bug: SingleCoreM/N 恒为 176 (双缓冲 Base 兜底); 现在自适应
|
|
from bmm_theory.branches.asw_basic import AswBasicBranch
|
|
p = AswBasicBranch().make_plan(mkcase(8, 4096, 4096, 1024))
|
|
self.assertNotEqual((p.single_core_m, p.single_core_n), (176, 176))
|
|
self.assertEqual((p.single_core_m, p.single_core_n), (512, 512))
|
|
# 且为 Base 的整数倍 (约束 4)
|
|
self.assertEqual(p.single_core_m % p.base_m, 0)
|
|
self.assertEqual(p.single_core_n % p.base_n, 0)
|
|
|
|
def test_l0c_single_buffer_constraint_pass(self):
|
|
# ASW 单缓冲: base 256x256 通过约束 (无违规)
|
|
from bmm_theory.constraints import check_plan_constraints
|
|
from bmm_theory.branches.asw_basic import AswBasicBranch
|
|
case = mkcase(8, 4096, 4096, 1024)
|
|
p = AswBasicBranch().make_plan(case)
|
|
self.assertEqual(check_plan_constraints(case, p), [])
|
|
|
|
def test_p1_fallback_to_enum_when_no_split_infeasible(self):
|
|
# B>=C 时不切分 k_l1 过 dValue 下限失败 -> 强制切分枚举 -> 512x512 (8,8)
|
|
case = mkcase(128, 4096, 4096, 8192)
|
|
r = self.router.route(case)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
p = r["plan"]
|
|
self.assertEqual((p.single_core_m, p.single_core_n), (512, 512))
|
|
self.assertEqual((p.m_cnt, p.n_cnt), (8, 8))
|
|
|
|
def test_issue34_k_l1_decomposition_derived(self):
|
|
# b32_m16_n8192_k7168: 分解 = Base 16x1024 (例外: M=16<256; N 侧受 L0B 单边
|
|
# 上限收敛), SingleCore 16x1024 (mCnt=1,nCnt=8), k_l1 = L1 双缓冲反推
|
|
# ⌊L1/(2·(sM+sN)·dt)⌋16 = ⌊524288/(2·1040·2)⌋16 = 112 (v1.91 §5.2)
|
|
from bmm_theory.branches.asw_basic import AswBasicBranch
|
|
case = mkcase(32, 16, 8192, 7168)
|
|
p = AswBasicBranch().make_plan(case)
|
|
self.assertEqual((p.base_m, p.base_n), (16, 1024))
|
|
self.assertEqual((p.single_core_m, p.single_core_n), (16, 1024))
|
|
self.assertEqual(p.k_l1, 112)
|
|
|
|
def test_pathological_shape_degraded_warning_always_plan(self):
|
|
# issue#34: 兜底分支恒出方案 —— 极端形状 (N=8 int8 大K, B 侧 dValue=8B
|
|
# 物理不可满足) 照常给方案 + 标注效率降级 (warning), 不判违规/不判不可行
|
|
from bmm_theory.constraints import check_plan_constraints
|
|
from bmm_theory.evaluator import PlanEvaluator
|
|
case = mkcase(4096, 2048, 8, 1024, dtype_a="int8", dtype_b="int8")
|
|
r = self.router.route(case)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
self.assertEqual(check_plan_constraints(case, r["plan"]), [])
|
|
self.assertIn("效率降级", r["plan"].note)
|
|
er = PlanEvaluator().evaluate(case, r["plan"])
|
|
self.assertTrue(er.feasible)
|
|
self.assertIn("效率降级", er.advice)
|
|
self.assertGreater(er.timing.t_total, 0.0)
|
|
|
|
def test_pathological_n8_hard_floor_fallback(self):
|
|
# 更极端: N=8 int8 连 128B 硬下限也不满足 -> 仍恒出方案, 效率降级标注
|
|
from bmm_theory.evaluator import PlanEvaluator
|
|
case = mkcase(512, 4096, 1, 8192, dtype_a="int8", dtype_b="int8")
|
|
r = self.router.route(case)
|
|
self.assertEqual(r["branch"], "ASW_Basic")
|
|
self.assertIsNotNone(r["plan"])
|
|
self.assertIn("效率降级", r["plan"].note)
|
|
er = PlanEvaluator().evaluate(case, r["plan"])
|
|
self.assertTrue(er.feasible)
|
|
|
|
|
|
class TestIssue35(unittest.TestCase):
|
|
"""issue#35: MergeBatch L1 绑定情形 DMA 命令数多计 b0 倍修复 + 分界泛化口径.
|
|
|
|
用户 case 家族: B=128, M=1~16, N=128, K=512, bf16. b_core=4, b0=4,
|
|
合并后 k_l1^m=240/224 < K=512 (L1 绑定区), 真实每核命令数 = 1x3=3 条
|
|
(< IterBatch 的 4 条), 修复前被多计为 12 条导致仲裁翻错方向.
|
|
"""
|
|
|
|
def setUp(self):
|
|
self.router = BranchRouter()
|
|
self.mb = MergeBatchBranch()
|
|
self.ib = IterBatchBranch()
|
|
|
|
def test_merged_cmd_count_formula(self):
|
|
# 每核命令数 = ceil(b_core/b0) * ceil(K/k_l1^m) (K截断时 = b_core/b0)
|
|
case = mkcase(128, 16, 128, 512)
|
|
r = self.mb.analyze(case)
|
|
self.assertTrue(r.capable)
|
|
p = r.plan
|
|
expect = -(-p.b_core // p.merge_b0) * (-(-case.k // p.k_l1))
|
|
self.assertEqual(r.timing.dma_cmd_count, expect)
|
|
# 本 case: b0=4, k_l1=224 -> 1*3 = 3 条 (修复前 12 条)
|
|
self.assertEqual((p.merge_b0, p.k_l1), (4, 224))
|
|
self.assertEqual(r.timing.dma_cmd_count, 3)
|
|
|
|
def test_cmd_count_truncated_unchanged(self):
|
|
# K 截断情形数值不变: cmds = b_core/b0 = IterBatch 的 1/b0
|
|
case = mkcase(2048, 32, 32, 256)
|
|
mb = self.mb.analyze(case)
|
|
ib = self.ib.analyze(case)
|
|
self.assertGreaterEqual(mb.plan.k_l1, case.k) # 合并后仍截断
|
|
self.assertEqual(mb.timing.dma_cmd_count,
|
|
-(-mb.plan.b_core // mb.plan.merge_b0))
|
|
self.assertAlmostEqual(
|
|
mb.timing.dma_cmd_count / ib.timing.dma_cmd_count,
|
|
1.0 / mb.plan.merge_b0, places=6)
|
|
|
|
def test_boundary_uses_merged_k_l1(self):
|
|
# 截断判定与 plan.k_l1 口径一致: k_l1^m < K 时不得声称 K截断
|
|
case = mkcase(128, 1, 128, 512)
|
|
_, detail = self.mb.beats_iterbatch(case)
|
|
p = self.mb.make_plan(case)
|
|
self.assertLess(p.k_l1, case.k)
|
|
self.assertIn("L1绑定", detail)
|
|
self.assertNotIn("K截断", detail)
|
|
|
|
def test_user_case_family_routing(self):
|
|
# B=128,M=1~16,N=128,K=512: issue#36 效率模型后, iter A tile=m*1KB 未饱和
|
|
# (m<16), 合并 tile 4 倍大 -> 效率节省 ~0.61us 恒定, drain 随 M 线性增长;
|
|
# m<=8 效率节省 > drain -> MergeBatch; m=16 iter tile 恰达 16KB 饱和 ->
|
|
# 效率打平, drain 0.66us 决定 -> IterBatch
|
|
expect = {1: "MergeBatch", 2: "MergeBatch", 4: "MergeBatch",
|
|
8: "MergeBatch", 16: "IterBatch"}
|
|
for m, branch in expect.items():
|
|
r = self.router.route(mkcase(128, m, 128, 512))
|
|
self.assertEqual(r["branch"], branch, f"m={m}: {r['arbitration']}")
|
|
self.assertEqual(r["candidates"], {"MergeBatch": True, "IterBatch": True})
|
|
self.assertEqual(r["self_check_violations"], [])
|
|
|
|
def test_arbitration_text_final_winner_consistent(self):
|
|
# 仲裁文本 [裁决] 位必须是最终胜者 (分界与时延不一致时括注说明)
|
|
import re
|
|
r = self.router.route(mkcase(128, 2, 128, 512))
|
|
m = re.search(r"\[裁决\] (\w+)", r["arbitration"])
|
|
self.assertIsNotNone(m)
|
|
self.assertEqual(m.group(1), r["branch"])
|
|
|
|
|
|
class TestIssue36(unittest.TestCase):
|
|
"""issue#36: 合并搬移效率建模 (tile=nValue*dValue*dt 放大 b0 倍) + t_cmd_ns=0.
|
|
|
|
用户澄清: MergeBatch 合并多 batch 左/右矩阵一起搬移, 单块 tile 放大 ->
|
|
搬移效率更高, 即便 T_cmd=0 也有效益; 堆叠方向视转置 (A ND 非转置沿
|
|
M(nValue), B ND 非转置沿 N(dValue)), 乘积口径不变。
|
|
"""
|
|
|
|
def test_t_cmd_default_zero(self):
|
|
self.assertEqual(ASCEND950PR.t_cmd_ns, 0.0)
|
|
self.assertEqual(ASCEND950PR.t_cmd, 0.0)
|
|
|
|
def test_move_eff_curve(self):
|
|
from bmm_theory.models import move_eff
|
|
cap = ASCEND950PR.min_tile_size # 16KB 饱和点
|
|
self.assertEqual(move_eff(cap, cap), 1.0) # 饱和
|
|
self.assertEqual(move_eff(2 * cap, cap), 1.0) # 超出仍饱和
|
|
self.assertAlmostEqual(move_eff(cap / 4, cap), 0.25) # 之下线性
|
|
self.assertEqual(move_eff(0, cap), 1.0) # 零值守卫
|
|
|
|
def test_merge_eff_gain_beats_iterbatch(self):
|
|
# (2048,16,64,128): iter A tile=4KB eff=0.25, 合并 A'=16KB eff=1.0
|
|
# -> 效率节省 ~15.7us >> drain 0.17us, T_cmd=0 下 MergeBatch 大胜
|
|
case = mkcase(2048, 16, 64, 128)
|
|
mb = MergeBatchBranch().analyze(case)
|
|
ib = IterBatchBranch().analyze(case)
|
|
self.assertTrue(mb.capable and ib.capable)
|
|
win, detail = MergeBatchBranch().beats_iterbatch(case)
|
|
self.assertTrue(win, detail)
|
|
self.assertLess(mb.timing.t_total, ib.timing.t_total)
|
|
|
|
def test_gm_bytes_unchanged_by_eff(self):
|
|
# 效率模型只影响时间列: GM 字节量仍 = V_in (issue#31 口径不被破坏)
|
|
for shp in [(128, 1, 128, 512), (2048, 16, 64, 128), (128, 64, 64, 512)]:
|
|
case = mkcase(*shp)
|
|
for br in (MergeBatchBranch(), IterBatchBranch()):
|
|
r = br.analyze(case)
|
|
if r.capable:
|
|
self.assertAlmostEqual(r.timing.gm_read_bytes,
|
|
case.input_bytes, places=3)
|
|
|
|
def test_iter_small_tile_slower_than_merge(self):
|
|
# 用户 case m=1: IterBatch A tile=1KB eff=1/16 -> t_mte2_gm 显著高于
|
|
# MergeBatch (A'=1.9KB eff=0.117), 且两者 GM 字节相同
|
|
case = mkcase(128, 1, 128, 512)
|
|
mb = MergeBatchBranch().analyze(case)
|
|
ib = IterBatchBranch().analyze(case)
|
|
self.assertGreater(ib.timing.t_mte2_gm, mb.timing.t_mte2_gm)
|
|
self.assertEqual(mb.timing.gm_read_bytes, ib.timing.gm_read_bytes)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|