Add BMM_Theory: bmm_theory/branches/asw_basic.py

This commit is contained in:
2026-09-03 09:25:40 +00:00
parent 113769b657
commit 6a24b5d945

View File

@@ -0,0 +1,289 @@
"""ASW_Basic 分支: 核间切 M/N (或混合切) 的兜底分支, 含尾轮处理.
理论依据:
- 《BMM算子优化分析 v0.98》§八 + docs/02_分支理论/06_ASW_Basic分支.md
- 《BMM尾轮处理策略对比分析 v1.5》+ docs/02_分支理论/07_尾轮处理策略.md
核心: 兜底分支, 核间切 M/N, 重复读交给 L2 + swizzle. 尾轮处理是必要组成环节:
默认方案 B (工程简洁, 主流场景与 A1b 严格打平), 周长型且 rho>=rho_dv 时 A1b.
"""
from __future__ import annotations
import math
from ..hardware import NpuSpec, ASCEND950PR
from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div, align_down
from ..timing import assemble_timing
from .base import Branch, ConditionCheck
class AswBasicBranch(Branch):
name = "ASW_Basic"
def __init__(self, spec: NpuSpec = ASCEND950PR):
super().__init__(spec)
# ------------------------------------------------------------------
def _p_value(self, case: BmmCase) -> float:
return case.batch_c * case.m * case.n * 4 / self.spec.l0c_bytes
# ------------------------------------------------------------------
def check_conditions(self, case: BmmCase) -> list:
s = self.spec
p = self._p_value(case)
checks = []
# 条件 1: P >= C (并行度补齐); P < C 时看是否走降核模式
if p >= s.aic_num:
checks.append(ConditionCheck(
"1_并行度补齐: P=B*MN*4B/L0C >= C", True, f"P={p:.2f} >= C={s.aic_num}"))
else:
# 降核模式: P < C 且不满足 StreamK —— 此处由 router 保证 (StreamK 优先)
checks.append(ConditionCheck(
"1_降核模式: P<C, 只用 ceil(P) 核", True,
f"P={p:.2f} < C={s.aic_num}, 降核用 {math.ceil(p)}"))
# 条件 2: 无 batch 结构限制 (恒满足)
checks.append(ConditionCheck("2_无batch结构限制", True, "兜底分支承接各种 batch 结构"))
return checks
# ------------------------------------------------------------------
def make_plan(self, case: BmmCase) -> ImplPlan:
s = self.spec
b = case.batch_c
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
p = self._p_value(case)
# 降核模式
if p < s.aic_num:
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 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 4: swizzle 窗口 W = max{d | d|C, d <= floor(sqrt(C))}
swizzle_w = self._swizzle_w()
# Step 5: L2 分组判断
s_in = b * (m * k + k * n) * dt
s_out = b * m * n * case.dtype_out_bytes
if s_in + s_out <= s.l2_bytes:
l2_scene = "A_全驻留"
l2_out = "resident(输出驻留L2异步回写)"
r_in = 1.0
elif s_in <= s.l2_bytes:
l2_scene = "B_输入驻留输出直写GM"
l2_out = "direct_gm(输出直写GM不占L2)"
r_in = 1.0
else:
l2_scene = "C_输入超L2分组执行"
l2_out = "direct_gm(输出直写GM不占L2)"
r_in = self._l2_group_r_in(case, single_m, single_n)
# Step 6: k_l1 (GM->L1 K 向粒度, 须 >= 256B/dt)
k_l1 = min(k, s.dvalue_recommend // dt)
k_l1 = max(k_l1, s.dvalue_hw_min // dt)
# ---- 尾轮决策 (必要组成环节) ----
n_blk = b * m_cnt * n_cnt
tail = self._decide_tail(case, single_m, single_n, n_blk, k_l1)
return ImplPlan(
case_id=case.case_id, branch=self.name,
used_core_num=s.aic_num,
split_b=1, m_cnt=m_cnt, n_cnt=n_cnt, grid_k=1,
core_map=f"B->M->N线性映射+ASW滑窗蛇形(W={swizzle_w})",
b_core=0, merge_b0=1,
single_core_m=single_m, single_core_n=single_n, single_core_k=k,
k_l1=k_l1, b_l1=1, l1_form="双缓冲驻留当前tile输入",
base_m=base_m, base_n=base_n, base_k=base_k,
l2_policy_in="allocate(输入驻留L2吸收重复读)",
l2_policy_out=l2_out,
swizzle_w=swizzle_w, workspace_bytes=0,
tail_strategy=tail["strategy"],
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,
out_dtype_bytes=case.dtype_out_bytes,
note=f"L2场景{l2_scene}, r_in={r_in:.2f}; 尾轮: {tail['reason']}",
)
# ------------------------------------------------------------------
def _plan_reduced_core(self, case: BmmCase, p: float) -> ImplPlan:
"""降核模式: 只用 ceil(P) 核, 每核一个 L0C 满载输出块."""
s = self.spec
used = max(1, math.ceil(p))
# SingleCoreM/N 在 L0C 容量内取最大
single_mn = int(math.sqrt(s.l0c_bytes // 4)) # L0C/4B 单份
single_m = min(case.m, align_down(single_mn, s.fractal))
single_n = min(case.n, align_down(single_mn, s.fractal))
return ImplPlan(
case_id=case.case_id, branch=self.name + "_降核",
used_core_num=used,
split_b=1, m_cnt=ceil_div(case.m, single_m), n_cnt=ceil_div(case.n, single_n),
grid_k=1,
core_map=f"降核: 只用{used}核, 每核一个L0C满载输出块, 其余核闲置",
b_core=0, merge_b0=1,
single_core_m=single_m, single_core_n=single_n, single_core_k=case.k,
k_l1=min(case.k, s.dvalue_recommend // case.dtype_in_bytes), b_l1=1,
l1_form="标准核内流水",
base_m=single_m, base_n=single_n,
base_k=min(case.k, 64),
l2_policy_in="allocate", l2_policy_out="direct_gm",
swizzle_w=0, workspace_bytes=0,
tail_strategy="不涉及(每核一块无尾轮)",
fixpipe_unitflag=True, out_dtype_bytes=case.dtype_out_bytes,
note=f"P={p:.2f}<C, 降核是理性选择 (强切则 tile 跌破搬移效率下限反而更慢)",
)
# ------------------------------------------------------------------
def _pick_single_core(self, m, n, min_blocks, base_m, base_n, dt):
"""Step 1: 满足并行度下限前提下 SingleCoreM/N 尽量大, 长宽比跟随 M/N."""
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
def _swizzle_w(self) -> int:
s = self.spec
w = 1
for d in range(1, int(math.sqrt(s.aic_num)) + 1):
if s.aic_num % d == 0:
w = d
return w
def _l2_group_r_in(self, case, sm, sn) -> float:
"""场景 C: L2 分组的重复读倍率 r_in = (n_grp*M + m_grp*N)/(M+N)."""
s = self.spec
m, n, k, b = case.m, case.n, case.k, case.batch_c
dt = case.dtype_in_bytes
d = s.l2_bytes / (b * k * dt)
m_grp = max(1, int(d / (2 * sm)))
n_grp = max(1, int(d / (2 * sn)))
m_cnt = ceil_div(m, sm)
n_cnt = ceil_div(n, sn)
return (ceil_div(n_cnt, n_grp) * m + ceil_div(m_cnt, m_grp) * n) / (m + n)
# ------------------------------------------------------------------
def _decide_tail(self, case, sm, sn, n_blk, k_l1) -> dict:
"""尾轮策略决策 (v1.5 闭式流程). 默认方案 B, 周长型且 rho>=rho_dv 时 A1b."""
s = self.spec
c = s.aic_num
n_wave = ceil_div(n_blk, c)
r = n_blk % c
rho = r / c
dt = case.dtype_in_bytes
out_b = case.dtype_out_bytes
base = dict(r=r, n_wave=n_wave, tail_m_cnt=1, tail_n_cnt=1,
tail_m_main=0, tail_n_main=0)
if r == 0:
return {**base, "strategy": "A0", "reason": "r=0 无尾轮"}
# 主导项判定
bw_eff = s.bw_l2_pc # L2 命中
t_mmad = 2 * sm * sn * case.k / s.q16
t_mte2 = case.k * (sm + sn) * dt / bw_eff
t_fix = sm * sn * out_b / s.bw_pc
t_block = max(t_mmad, t_mte2, t_fix)
area_dominated = t_block != t_mte2 # 面积型 = MMAD 或 FIX 主导
# 方案 B: 全局均匀重切, g = n_wave*C/N_blk
g = n_wave * c / n_blk
s_b = int(math.sqrt(sm * sn / g))
s_b = align_down(max(s_b, s.fractal), s.fractal)
plan_b_m_cnt = ceil_div(case.m, s_b) if case.m >= s_b else 1
plan_b_n_cnt = ceil_div(case.n, s_b) if case.n >= s_b else 1
if area_dominated:
# 面积型: 方案 B 与 A1b 严格打平, 默认方案 B (工程简洁); r 小翻出时方案 B 微优
return {**base, "strategy": "方案B",
"tail_m_cnt": plan_b_m_cnt, "tail_n_cnt": plan_b_n_cnt,
"tail_m_main": plan_b_m_cnt, "tail_n_main": plan_b_n_cnt,
"reason": f"面积型主导, A1b与方案B严格打平(T={t_block*(n_wave-1+rho)*1e6:.2f}us), "
f"选方案B工程简洁 (一套tile)"}
else:
# 周长型: rho >= rho_dv 时 A1b 恒优
rho_dv = (s.dvalue_hw_min / (sn * dt)) ** 2
if rho >= rho_dv:
# A1b: 尾轮 tile 缩 sqrt(rho)
s_t = align_down(int(sn * math.sqrt(rho)), s.fractal) or s.fractal
t_m = ceil_div(sm, s_t) if sm >= s_t else 1
t_n = ceil_div(sn, s_t) if sn >= s_t else 1
return {**base, "strategy": "A1b",
"tail_m_cnt": t_m, "tail_n_cnt": t_n,
"tail_m_main": t_m, "tail_n_main": t_n,
"reason": f"周长型主导, rho={rho:.2f}>=rho_dv={rho_dv:.2f}, A1b恒优 "
f"(尾轮tile缩√rho={math.sqrt(rho):.2f}倍凑满核)"}
else:
return {**base, "strategy": "方案B",
"tail_m_cnt": plan_b_m_cnt, "tail_n_cnt": plan_b_n_cnt,
"tail_m_main": plan_b_m_cnt, "tail_n_main": plan_b_n_cnt,
"reason": f"周长型主导, rho={rho:.2f}<rho_dv={rho_dv:.2f}, A1b被dValue卡死, 方案B反超"}
# ------------------------------------------------------------------
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
s = self.spec
b = case.batch_c
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
out_b = case.dtype_out_bytes
flops = 2.0 * b * m * n * k
t_mmad = flops / (plan.used_core_num * s.q16)
in_bytes = b * (m * k + k * n) * dt
out_bytes = b * m * n * out_b
# r_in 从 note 里解析困难, 重新计算
s_in = in_bytes
s_out = out_bytes
if s_in + s_out <= s.l2_bytes or s_in <= s.l2_bytes:
r_in = 1.0
else:
r_in = self._l2_group_r_in(case, plan.single_core_m, plan.single_core_n)
t_mte2 = r_in * in_bytes / (plan.used_core_num * s.bw_pc)
to_l2 = "resident" in plan.l2_policy_out
t_fix = out_bytes / (plan.used_core_num * (s.bw_l2_pc if to_l2 else s.bw_pc))
# drain: 尾轮暴露 (方案 B 已均匀重切, drain 小; A1b 尾轮凑满, drain 小; A0 尾轮 r 核空转)
t_drain = 0.0
if plan.tail_strategy == "A0" and plan.tail_block_cnt > 0:
t_drain = max(t_mmad, t_mte2, t_fix) # 尾轮空转一个整块
return assemble_timing(
t_mte2_gm=t_mte2, t_mte2_l2=0.0, t_dma_cmd=0.0,
t_mmad=t_mmad, t_fixpipe=t_fix, t_reduce=0.0, t_drain=t_drain,
gm_read_bytes=r_in * in_bytes, l2_read_bytes=0.0, dma_cmd_count=0.0,
cube_flops=flops, fixpipe_bytes=out_bytes,
)