Fix #31/#32: 切B类GM每字节恰一次=V_in(去K切分整段上取) / ASW场景升级(单侧全驻留+对侧滑窗->S_B, S_C最小替换2D分组+窗口L2计账) / 06文档§3+Step5与docs/05同步
#31 IterBatch/MergeBatch: K切分各(batch,K段)互不重叠+驻留侧每batch一次+末段按实际剩余 -> GM读取量=V_in(与L2容量无关), GM数据时延=V_in/W_GM; n_K仅决定DMA命令数(T_cmd) b64_m16_n256_k512 形态c: GM 25.07MB->17.83MB=V_in; 回归: 形态c/d非整除+L1绑定三类断言 #32 ASW: (1)S_B扩展单侧全驻留+对侧滑窗(a_b/b_b+2*对侧单块<=L2) -> GM=V_in, 6个场景C行回落S_B; (2)S_C在整L2容量约束下搜索最小GM=ceil(n_cnt/n_grp)a_b+ceil(m_cnt/m_grp)b_b(取代L2/2对半预算), 并计组内窗口L2流量((n_cnt-ceil)a_b+(m_cnt-ceil)b_b), 与S_B'驻留命中走L2'口径一致; b8_m131072_n8192_k8192 GM倍率4.76x->3.88x(物理下界~3.9x, 双侧均超L2) 大方形K行(如b128_m8192_n8192_k7168)由MMAD 253ms->MTE2(L2口)295ms: 共享块重复读1.38TB 经L2读口5.2TB/s, 如实计账(原C窗口流量零计低估) - docs/06 §3与Step5重写为S_A/S_B/S_C+两段链口径(旧r_in单段/除B/对半预算口径废弃) - docs/05 R5/R6/§3.1/§4.1/§4.3/§5与docs/01、02(01_MergeBatch/02_IterBatch GM口径附注)同步 - tests 54/54; 压力seed7/6000+seed2024/4000: 0违规/0占位/0NaN/0GM<V_in; examples重生成0diff
This commit is contained in:
@@ -250,9 +250,23 @@ class TestIssueRegression2(unittest.TestCase):
|
||||
|
||||
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") # 原 0.2% 违规样例
|
||||
r = self.router.route(case)
|
||||
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"])
|
||||
@@ -592,5 +606,96 @@ class TestIssue27to30(unittest.TestCase):
|
||||
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))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user