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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user