Add BMM_Theory: bmm_theory/branches/asw_basic.py
This commit is contained in:
289
BMM/BMM_Theory/bmm_theory/branches/asw_basic.py
Normal file
289
BMM/BMM_Theory/bmm_theory/branches/asw_basic.py
Normal 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,
|
||||
)
|
||||
Reference in New Issue
Block a user