Update BMM_Theory: bmm_theory/branches/iter_batch.py

This commit is contained in:
2026-09-03 12:43:52 +00:00
parent 517441de6e
commit af0a90bc25

View File

@@ -90,15 +90,28 @@ class IterBatchBranch(Branch):
form is not None, form_desc))
# 条件 4: 搬移效率下限 (c/d 切分后)
# 转置影响 (参考 bmmv3): A 不转置时 K 向连续, dValue 判 K*dt;
# A 转置时 M 向连续, dValue 判 M*dt;
# B 不转置时 N 向连续, dValue 判 N*dt;
# B 转置时 K 向连续, dValue 判 K*dt.
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
if form in ("c", "d"):
tile_ok = (k_l1 * m * dt >= s.min_tile_size) or (k_l1 * n * dt >= s.min_tile_size)
dv_ok = k_l1 * dt >= s.dvalue_min
# dValue 判定按转置调整连续维
if case.trans_a:
dv_a = m * dt # A 转置: M 向连续
else:
dv_a = k_l1 * dt # A 不转置: K 向连续 (切分后)
if case.trans_b:
dv_b = k_l1 * dt # B 转置: K 向连续 (切分后)
else:
dv_b = n * dt # B 不转置: N 向连续
dv_ok = dv_a >= s.dvalue_min or dv_b >= s.dvalue_min # 至少一侧满足
c4 = tile_ok and dv_ok
checks.append(ConditionCheck(
"4_搬移效率: 搬移分块>=min_TileSize 且 dValue>=128B",
c4, f"tile={max(k_l1*m*dt, k_l1*n*dt)/1024:.1f}KB, dValue={k_l1*dt}B"))
"4_搬移效率: 搬移分块>=min_TileSize 且 dValue>=128B (转置调整连续维)",
c4, f"tile={max(k_l1*m*dt, k_l1*n*dt)/1024:.1f}KB, dValueA={dv_a}B, dValueB={dv_b}B"))
else:
checks.append(ConditionCheck(
"4_搬移效率(a/b形态不切分, 恒满足)", True, ""))