Fix #35: MergeBatch L1绑定情形 DMA 命令数多计 b0 倍修复 + 分界泛化口径
- evaluate: 每核命令数 = ceil(b_core/b0) * ceil(K/k_l1^m) (K截断退化为 b_core/b0, 数值不变; L1绑定消除 b0 倍多计 —— v1.1 §4.4 恒劣恒等式的 n_K 是未合并粒度, 误代入合并后段数会多计 b0 倍, 可把仲裁方向翻错) - beats_iterbatch: 泛化为实际命令数比较 (cmds_iter=b_core*ceil(K/k_l1_iter) vs cmds_mb=ceil(b_core/b0)*ceil(K/k_l1^m), 节省>T_cmd vs drain 惩罚); K截断时严格 退化为文档闭式 b_core > b0*(T_comp+T_write)/T_cmd; 截断判定改用合并口径 plan.k_l1>=K (未合并截断不代表合并后截断); 覆盖 dValue 512B cap 第三情形; T_cmd<=0 策略路径改为 cmds_mb<cmds_iter 判 MergeBatch 优先 - router: 仲裁文案 [裁决] 位打印最终胜者 (修复分界/时延不一致时的自相矛盾表述) - 用户 case 家族 B=128,M=1~16,N=128,K=512 修复后: m=1/2/4 -> MergeBatch, m=8/16 -> IterBatch (修复前全判 IterBatch; 交叉点 m≈4~8, 物理合理) - 测试: TestIssue35 回归 5 例 (命令数公式/K截断不变/口径一致/路由家族/裁决文案); test_beats_iterbatch_policy 的 (128,64,64,512) 期望 True->False (第三情形: 合并侧 dValue cap 截断, 命令数 4=4 打平, 恒劣 —— 原期望基于误分类) - docs/01_MergeBatch分支.md: 分界小节补第三情形行 + 命令数口径警示 + 泛化净收益式 - 验证: 68/68 unittest; examples 重生成可复现 0 diff (仅仲裁文案 + 16.0->16 格式, plans.csv 不变); 压力 10000 例 (seed7/6000+seed2024/4000): 0 崩溃/0 NaN/0 违规/ 0 不可行/0 GM<V_in, 七分支全覆盖
This commit is contained in:
@@ -388,18 +388,21 @@ class TestZeroCmdHandling(unittest.TestCase):
|
||||
"""T_cmd=0 (无命令时延/未标定) 时整链路不得除零/崩溃, 按策略优先 MergeBatch."""
|
||||
|
||||
def test_beats_iterbatch_policy(self):
|
||||
# T_cmd<=0: 阈值 +inf 不可除零; K截断按策略判 MergeBatch 胜 (指令级收益未量化),
|
||||
# L1 绑定仍恒劣
|
||||
# T_cmd<=0: 阈值 +inf 不可除零; 合并后每核命令数更少即按策略判 MergeBatch 胜
|
||||
# (指令级收益未建模), 命令数打平/更多则恒劣 (issue#35 泛化口径)
|
||||
from bmm_theory.hardware import NpuSpec
|
||||
from bmm_theory.branches.merge_batch import MergeBatchBranch
|
||||
spec0 = NpuSpec(t_cmd_ns=0.0)
|
||||
mb = MergeBatchBranch(spec0)
|
||||
# (b, m, n, k, K截断与否)
|
||||
cases = [(2048, 32, 32, 256, True), (128, 64, 64, 512, True),
|
||||
# (b, m, n, k, MergeBatch应胜与否=合并后命令数更少)
|
||||
# (128,64,64,512): issue#35 —— 合并侧被 dValue 512B cap 截断 (k_l1^m=256,
|
||||
# n_K^m=2), 而 IterBatch 走 b 形态 k_l1=K=512, 每核命令数 4=4 打平, 合并只
|
||||
# 放大 drain -> 恒劣; 原期望 True 建立在未合并口径的误分类上 (第三情形)
|
||||
cases = [(2048, 32, 32, 256, True), (128, 64, 64, 512, False),
|
||||
(256, 128, 128, 4096, False)]
|
||||
for b, m, n, k, truncated in cases:
|
||||
for b, m, n, k, mb_wins in cases:
|
||||
win, detail = mb.beats_iterbatch(mkcase(b, m, n, k))
|
||||
self.assertEqual(win, truncated, f"{b},{m},{n},{k}: {detail}")
|
||||
self.assertEqual(win, mb_wins, f"{b},{m},{n},{k}: {detail}")
|
||||
self.assertIn("T_cmd=0", detail)
|
||||
|
||||
def test_route_with_zero_cmd_prefers_merge(self):
|
||||
@@ -798,5 +801,71 @@ class TestIssue33(unittest.TestCase):
|
||||
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: 修复后小 M 由 MergeBatch 胜 (命令节省 >
|
||||
# drain 惩罚), 大 M 由 IterBatch 胜 (drain 随 M 增长, 节省固定)
|
||||
expect = {1: "MergeBatch", 2: "MergeBatch", 4: "MergeBatch",
|
||||
8: "IterBatch", 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"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user