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

@@ -1,8 +1,10 @@
"""ASW_Basic 分支: 核间切 M/N (或混合切) 的兜底分支, 含尾轮处理.
理论依据:
- 《BMM算子优化分析 v0.98》§八 + docs/02_分支理论/06_ASW_Basic分支.md
- 《BMM尾轮处理策略对比分析 v1.5》+ docs/02_分支理论/07_尾轮处理策略.md
- 《ASW_Basic分支分析 v1.91》(L0C 单缓冲方形 Base tile §5.1 + SingleCoreM/N 有界
枚举 §5.2) + docs/02_分支理论/06_ASW_Basic分支.md
- 《BMM尾轮处理策略对比分析 v1.5》(单块搬入稳态口径 K(sM+sN)·dt/BW, k_L1 约掉)
+ docs/02_分支理论/07_尾轮处理策略.md
核心: 兜底分支, 核间切 M/N, 重复读交给 L2 + swizzle. 尾轮处理是必要组成环节:
默认方案 B (工程简洁, 主流场景与 A1b 严格打平), 周长型且 rho>=rho_dv 时 A1b.
@@ -13,7 +15,7 @@ from __future__ import annotations
import math
from ..hardware import NpuSpec, ASCEND950PR
from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div, align_down
from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div, align_down, align_up
from ..timing import assemble_timing
from .base import Branch, ConditionCheck
@@ -60,22 +62,31 @@ class AswBasicBranch(Branch):
return self._plan_reduced_core(case, p)
# ---- 正常模式: Step 0~6 ----
# Step 0: BaseM/BaseN 用满 L0C (32768 元素双缓冲)
base_mn = int(math.sqrt(s.l0c_bytes // 8)) # L0C/(2*4B)
base_m = align_down(base_mn, s.fractal)
base_n = align_down(base_mn, s.fractal)
base_k = align_down(int(min(
s.l0a_bytes / (2 * base_m * dt),
s.l0b_bytes / (2 * base_n * dt),
)), s.fractal)
# Step 0: BaseM/BaseN (v1.91 §5.1): UnitFlag 单缓冲 (tile 内 16x16x16 细粒度
# 流水替代 tile 间粗粒度双缓冲), BaseM*BaseN = L0C/4B = 65536 元素, 方形
# 优先 (256x256, L0A/L0B 同时装满、baseK 加倍); 仅当 M/N < 方形边长时被迫
# 跟随 M/N (另一侧按 L0C 面积余量放大); BaseK 由 L0A/L0B (双缓冲) 反推.
from ..constraints import clamp_base_k
cap_l0c = s.l0c_bytes // 4 # 65536 元素 (单缓冲)
sq = align_down(int(math.sqrt(cap_l0c)), s.fractal) # 256
# L0A/L0B 在 base_k>=16 (fractal) 下的单边上限 (v1.91 §5.7 核内约束;
# 防止例外路径 (如 M=16 -> BaseN=4096) 超 L0B)
side_cap = s.l0a_bytes // (2 * s.fractal * dt) # L0A=L0B=64KB
if m >= sq and n >= sq:
base_m = base_n = sq
elif m < sq:
base_m = max(align_down(m, s.fractal), s.fractal)
base_n = max(align_down(min(cap_l0c // base_m, n, side_cap),
s.fractal), s.fractal)
else: # n < sq
base_n = max(align_down(n, s.fractal), s.fractal)
base_m = max(align_down(min(cap_l0c // base_n, m, side_cap),
s.fractal), s.fractal)
base_k = clamp_base_k(base_m, base_n, dt, case.k, s) # v1.91 §5.7
# Step 1: SingleCoreM/N 尽量大 (满足并行度下限)
min_blocks = ceil_div(s.aic_num, b)
single_m, single_n = self._pick_single_core(m, n, min_blocks, base_m, base_n, dt)
# Step 2: mCnt/nCnt
m_cnt = ceil_div(m, single_m)
n_cnt = ceil_div(n, single_n)
# Step 1+2+6: SingleCoreM/N + k_L1 + mCnt/nCnt 有界枚举 (v1.91 §5.2):
(single_m, single_n, k_l1, m_cnt, n_cnt, pick_note) = \
self._pick_single_core(case, base_m, base_n)
# Step 4: swizzle 窗口 W = max{d | d|C, d <= floor(sqrt(C))}
swizzle_w = self._swizzle_w()
@@ -90,12 +101,6 @@ class AswBasicBranch(Branch):
if scene["to_l2"] else
"direct_gm(输出直写GM, 输入优先驻留L2)")
# Step 6: k_l1 (GM->L1 K 向粒度, 须 >= dValue 下限; K 小于下限时整 K 一次搬入不切)
dv_min_elems = max(s.dvalue_hw_min // dt, 1)
k_l1 = min(k, s.dvalue_recommend // dt)
if k_l1 < dv_min_elems:
k_l1 = k # K 本身小于 dValue 下限: 不切 K, 整段搬入 (K 非连续维, dValue 由 M/N 保证)
# ---- 尾轮决策 (必要组成环节) ----
n_blk = b * m_cnt * n_cnt
tail = self._decide_tail(case, single_m, single_n, n_blk, k_l1)
@@ -116,9 +121,11 @@ class AswBasicBranch(Branch):
tail_m_cnt=tail["tail_m_cnt"], tail_n_cnt=tail["tail_n_cnt"],
tail_k_cnt=1, tail_m_main=tail["tail_m_main"], tail_n_main=tail["tail_n_main"],
tail_block_cnt=tail["r"], tail_wave_num=tail["n_wave"],
fixpipe_unitflag=True,
fixpipe_unitflag=True, # UnitFlag 单缓冲: tile 内 16x16x16 细粒度流水
out_dtype_bytes=case.dtype_out_bytes,
note=(f"L2场景: {scene['label']} (V_in={case.input_bytes/1048576:.1f}MB, "
note=(f"tile枚举: {pick_note}; Base tile {base_m}x{base_n} "
f"(L0C 单缓冲方形用满); "
f"L2场景: {scene['label']} (V_in={case.input_bytes/1048576:.1f}MB, "
f"V_out={case.output_bytes/1048576:.1f}MB, "
f"L2={s.l2_bytes/1048576:.0f}MB); "
+ (scene["note_extra"] + "; " if scene["note_extra"] else "")
@@ -163,24 +170,116 @@ class AswBasicBranch(Branch):
)
# ------------------------------------------------------------------
def _pick_single_core(self, m, n, min_blocks, base_m, base_n, dt):
"""Step 1: 满足并行度下限前提下 SingleCoreM/N 尽量大, 长宽比跟随 M/N."""
def _pick_single_core(self, case: BmmCase, base_m: int, base_n: int) -> tuple:
"""SingleCoreM/N + k_L1 + mCnt/nCnt 有界枚举 (v1.91 §5.2 + v1.5 §2.1 修正口径).
- P = ⌈C/B⌉; P=1 (B>=C): 先试不切分 (mCnt=nCnt=1, tile 跟随 M/N, 情形1),
约束不满足则进入枚举强制切分;
- 枚举空间: mCnt ∈ [1, ⌈M/BaseM⌉], nCnt ∈ [1, ⌈N/BaseN⌉] (约束4 上界:
SingleCore 为 Base 整数倍), 且 B·mCnt·nCnt >= C (约束1 并行度);
- 约束 2 (L1 双缓冲): k_L1 = min(K, ⌊L1/(2(sM+sN)·dt)⌋16), 须 ≥ 一个
fractal 且满足 dValue 下限;
- 约束 3 (搬移效率, 转置感知): A 非转置 dValue = k_L1·dt / A 转置 = sM·dt;
B 非转置 dValue = sN·dt / B 转置 = k_L1·dt; 均须 ≥ 256B (dValue 硬件
突发下限); 且 sM·k_L1·dt 与 k_L1·sN·dt ≥ min_TileSize (16KB);
- 目标 (v1.5 §2.1 修正: 稳态流水下单块搬入 = K·(sM+sN)·dt, 与分次粒度
k_L1 无关, k_L1 只进约束): 每 batch 搬入量 K·dt·(nCnt·M + mCnt·N) 最小
(共享块重复读最少 = tile 尽可能大/少); 并列取 r = B·mCnt·nCnt mod C 最大
(尾轮块越多, 重切收益越大, v1.91 §5.2.d);
返回 (single_m, single_n, k_l1, m_cnt, n_cnt, note).
"""
s = self.spec
# 从大到小枚举 (m_cnt*n_cnt >= min_blocks), 取最大 tile
best = (base_m, base_n)
for m_cnt in range(1, int(math.sqrt(min_blocks)) + 2):
n_cnt = ceil_div(min_blocks, m_cnt)
if m_cnt * n_cnt < min_blocks:
continue
sm = align_down(ceil_div(m, m_cnt), base_m) or base_m
sn = align_down(ceil_div(n, n_cnt), base_n) or base_n
# 约束 2: L1 容量 2(sM+sN)*k_l1*dt <= L1, k_l1 取 256B/dt
k_l1_min = s.dvalue_hw_min // dt
if 2 * (sm + sn) * k_l1_min * dt > s.l1_bytes:
continue
if sm * sn > best[0] * best[1]:
best = (sm, sn)
return best
b = case.batch_c
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
c = s.aic_num
p_min = ceil_div(c, b) # 最少切分块数 (约束1)
dv_min = s.dvalue_hw_min # 256B
min_tile = s.min_tile_size # 16KB
def eval_tile(sm: int, sn: int):
"""约束 2/3 + 目标值; 返回 (ok, k_l1, traffic_per_batch)."""
k_cap = int(s.l1_bytes / (2 * (sm + sn) * dt))
k_l1 = align_down(min(k, k_cap), s.fractal)
if k_l1 < s.fractal:
return False, 0, None
# dValue 下限 (转置感知连续维, 同 issue#19 口径)
dv_a = sm * dt if case.trans_a else k_l1 * dt
dv_b = k_l1 * dt if case.trans_b else sn * dt
if dv_a < dv_min or dv_b < dv_min:
return False, 0, None
# 单次搬移量下限
if sm * k_l1 * dt < min_tile or k_l1 * sn * dt < min_tile:
return False, 0, None
mc = ceil_div(m, sm)
nc = ceil_div(n, sn)
traffic = k * dt * (nc * m + mc * n) # 每 batch 总搬入字节
return True, k_l1, traffic
# 情形 1: P=1 (B >= C) 先试不切分, tile 跟随 M/N
if p_min == 1:
sm = max(align_up(m, s.fractal), s.fractal)
sn = max(align_up(n, s.fractal), s.fractal)
ok, k_l1, traffic = eval_tile(sm, sn)
if ok:
return (sm, sn, k_l1, 1, 1,
"P=1(B>=C) 不切分, tile 跟随 M/N (v1.91 §5.2 情形1)")
# 情形 2: 有界枚举
best = None
m_max = ceil_div(m, base_m)
n_max = ceil_div(n, base_n)
for mc in range(1, m_max + 1):
sm = align_up(ceil_div(m, mc), base_m)
for nc in range(1, n_max + 1):
sn = align_up(ceil_div(n, nc), base_n)
ok, k_l1, traffic = eval_tile(sm, sn)
if not ok:
continue
mc_r = ceil_div(m, sm)
nc_r = ceil_div(n, sn)
# 约束 1: 并行度 (真实块数, 考虑 align 上取后的收缩)
if b * mc_r * nc_r < c:
continue
r = (b * mc_r * nc_r) % c
key = (traffic, -r) # 主键搬入量最小; 并列 r 最大
if best is None or key < best[0]:
best = (key, sm, sn, k_l1, mc_r, nc_r, traffic, r)
if best is None:
# 兜底: 放开约束 4 (Base 整数倍), 在 16 对齐网格上按硬下限 (dValue>=128B
# 硬底 / L1 容量) 找最优可行 —— 极端形状 (如 M=2 与大 N/K 组合) 下
# Base 粒度可能无可行解; 若仍无解则退回 Base tile 由自检标注违规
dv_hard = s.dvalue_min # 128B 硬下限
for sm in range(s.fractal, align_up(min(m, 1024), s.fractal) + 1, s.fractal):
for sn in range(s.fractal, align_up(min(n, 1024), s.fractal) + 1, s.fractal):
k_cap = int(s.l1_bytes / (2 * (sm + sn) * dt))
k_l1 = align_down(min(k, k_cap), s.fractal)
if k_l1 < s.fractal:
continue
dv_a = sm * dt if case.trans_a else k_l1 * dt
dv_b = k_l1 * dt if case.trans_b else sn * dt
if dv_a < dv_hard or dv_b < dv_hard:
continue
mc_r = ceil_div(m, sm)
nc_r = ceil_div(n, sn)
if b * mc_r * nc_r < c:
continue
traffic = k * dt * (nc_r * m + mc_r * n)
r = (b * mc_r * nc_r) % c
key = (traffic, -r)
if best is None or key < best[0]:
best = (key, sm, sn, k_l1, mc_r, nc_r, traffic, r)
if best is None:
# 极端兜底: 退回 Base tile (约束校验会标注违规, 方案不可行可人工处置)
k_l1 = max(align_down(min(k, int(s.l1_bytes /
(2 * (base_m + base_n) * dt))),
s.fractal), s.fractal)
return (base_m, base_n, k_l1, ceil_div(m, base_m), ceil_div(n, base_n),
"枚举无可行候选, 退回 Base tile (自检会标注)")
_, sm, sn, k_l1, mc_r, nc_r, traffic, r = best
return (sm, sn, k_l1, mc_r, nc_r,
f"P={p_min}, 有界枚举最优 mCnt={mc_r} x nCnt={nc_r} "
f"(tile {sm}x{sn}, 每batch搬入{traffic/1048576:.1f}MB, r={r})")
def _swizzle_w(self) -> int:
s = self.spec