Fix #36: MergeBatch 合并搬移效率收益建模 (move_eff) + t_cmd_ns 置 0

用户澄清: 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 例)
This commit is contained in:
2026-09-07 21:09:49 +08:00
parent fdd3c8883c
commit b0b48b9073
16 changed files with 277 additions and 139 deletions

View File

@@ -12,6 +12,7 @@
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
@@ -108,12 +109,22 @@ class TestArbitration(unittest.TestCase):
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)
# 大 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):
"""时延模型自洽性."""
@@ -385,34 +396,34 @@ class TestTransposeModeling(unittest.TestCase):
class TestZeroCmdHandling(unittest.TestCase):
"""T_cmd=0 (无命令时延/未标定) 时整链路不得除零/崩溃, 按策略优先 MergeBatch."""
"""950PR 默认 T_cmd=0 (未标定, issue#36): 整链路不得除零/崩溃; 合并收益由
搬移效率模型 (move_eff, 合并 tile 放大 b0 倍) 刻画, 不再设策略覆盖."""
def test_beats_iterbatch_policy(self):
# T_cmd<=0: 阈值 +inf 不可除零; 合并后每核命令数更少即按策略判 MergeBatch 胜
# (指令级收益未建模), 命令数打平/更多则恒劣 (issue#35 泛化口径)
from bmm_theory.hardware import NpuSpec
def test_beats_iterbatch_zero_cmd(self):
# T_cmd=0: 命令节省项为 0, 由效率节省 vs drain 惩罚决定 (issue#36)
from bmm_theory.branches.merge_batch import MergeBatchBranch
spec0 = NpuSpec(t_cmd_ns=0.0)
mb = MergeBatchBranch(spec0)
# (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)]
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("T_cmd=0", detail)
self.assertIn("效率节省", detail)
def test_route_with_zero_cmd_prefers_merge(self):
# T_cmd=0 且 K截断时路由应优先 MergeBatch (命令/指令次数少 b0 倍, 结构性收益)
from bmm_theory.hardware import NpuSpec
router = BranchRouter(NpuSpec(t_cmd_ns=0.0))
r = router.route(mkcase(2048, 32, 32, 256)) # 大 B 小 MN 典型合并场景
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.assertIn("策略", r["arbitration"])
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:
@@ -848,10 +859,12 @@ class TestIssue35(unittest.TestCase):
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 增长, 节省固定)
# 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: "IterBatch", 16: "IterBatch"}
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']}")
@@ -867,5 +880,56 @@ class TestIssue35(unittest.TestCase):
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()