Fix #33: ASW_Basic tile 选择重写 (v1.91 §5.1/§5.2 + 尾轮 v1.5 §2.1)

- BaseM/BaseN: UnitFlag 单缓冲方形 256x256 (L0C/4B=65536 元素用满, 替代双缓冲 176x176);
  M/N < 256 被迫跟随 M/N (另一侧按 L0C 余量放大, 且收敛 L0A/L0B baseK>=16 单边上限);
  baseK = min(L0A/(2*BaseM*dt), L0B/(2*BaseN*dt)) 向下16对齐 (64KB 两侧双缓冲)
- SingleCoreM/N: 有界枚举取代旧"仅 sqrt(P)+2 范围按面积取最大"(恒退回176兜底):
  P=1 (B>=C) 先试不切分 tile 跟随 M/N; 否则枚举 mCnt<=ceil(M/BaseM) x nCnt<=ceil(N/BaseN)
  且 B*mCnt*nCnt>=C, sM/sN 为 Base 整数倍, 约束2 L1 双缓冲反推 k_L1,
  约束3 dValue>=256B 转置感知 + minTile 16KB; 目标 = 每batch搬入 K*dt*(nCnt*M+mCnt*N) 最小
  (v1.5 修正: 稳态下 k_L1 约掉, 只进约束); 并列取 r 最大
- 兜底: 约束4无解时放开到 16 对齐网格按硬下限(dValue>=128B)再搜; 极端形状
  (如 N=8 int8 大K)仍无解时退回 Base tile 并自检标注违规 (不静默产伪方案)
- constraints: ASW 双分支 L0C 口径改 UnitFlag 单缓冲 (factor 1)
- docs/06 Step0/Step1/Step6 重写为 v1.91 口径, 头部标注 2026-09 更新与 issue#33
- 回归: v1.91 §5.2 完整实例 (B=8 M=N=2048 K=1024 -> (4,4) 512x512, k_l1=128, r=0);
  方形例外 (M=128 -> 128x512); 单缓冲约束通过; 极端形状违规标注; 60->61 测试全过;
  压力 seed7/6000 + seed2024/4000: 0 崩溃/0 NaN/0 GM<V_in, 违规仅剩极端形状如实标注;
  examples 重生成可复现 0 diff (L2 重复读降 2-3x, 如 b128_m8192_n8192_k7168 1.38TB->451GB)
This commit is contained in:
2026-09-07 15:49:24 +08:00
parent 18f59599e7
commit 05ca91e291
7 changed files with 321 additions and 123 deletions

View File

@@ -697,5 +697,83 @@ class TestIssue31to32(unittest.TestCase):
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_pathological_shape_flagged_infeasible(self):
# 极端形状 (N=8 int8, 大 K): B 侧 dValue = sN*dt = 8B 物理上不可能满足
# 256B/128B 下限 -> 方案必须被自检标注违规 (不再静默产出伪可行方案)
from bmm_theory.constraints import check_plan_constraints
from bmm_theory.branches.asw_basic import AswBasicBranch
case = mkcase(4096, 2048, 8, 1024, dtype_a="int8", dtype_b="int8")
r = self.router.route(case)
self.assertEqual(r["branch"], "ASW_Basic")
v = check_plan_constraints(case, r["plan"])
self.assertTrue(v, "极端形状应被自检标注违规 (搬移效率崩塌)")
self.assertIn("自检违规", r["arbitration"])
# 时延仍应给出 (供人工评估), 但不为 0 / NaN
self.assertGreater(r["timing"].t_total, 0.0)
if __name__ == "__main__":
unittest.main()