diff --git a/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py b/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py index 1800c12..47b6432 100644 --- a/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py +++ b/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py @@ -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, ""))