Add BMM_Theory: bmm_theory/branches/stream_k.py

This commit is contained in:
2026-09-03 09:25:43 +00:00
parent e7da2e0a0b
commit ee2ea4d33c

View File

@@ -0,0 +1,164 @@
"""StreamK 分支: 核间切 K + 归约, B/M/N 并行度填不满核时的最后手段.
理论依据: 《BMM算子优化分析 v0.98》§七 + docs/02_分支理论/05_StreamK分支.md.
切 K 是唯一同时破坏"累加不出核""输出独占"的切法:
部分和必须写出 workspace (驻留 L2, 防精度丢失按 4B) 再由 AIV 归约.
"""
from __future__ import annotations
import math
from ..hardware import NpuSpec, ASCEND950PR
from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div
from ..timing import assemble_timing, eval_streamk_reduce
from .base import Branch, ConditionCheck
class StreamKBranch(Branch):
name = "StreamK"
def __init__(self, spec: NpuSpec = ASCEND950PR):
super().__init__(spec)
# ------------------------------------------------------------------
def _p_value(self, case: BmmCase) -> float:
"""P = B*M*N*4B / L0C: 以 L0C 满载为基本块粒度的不切 K 最大并行度."""
return case.batch_c * case.m * case.n * 4 / self.spec.l0c_bytes
def _grid_k(self, case: BmmCase) -> int:
"""grid_K = floor(C / ceil(P)), 最大化 K 并行度."""
p = self._p_value(case)
return int(self.spec.aic_num // max(1, math.ceil(p)))
def _theta_c(self) -> float:
"""归约代价系数 theta_c = Q16/2 * (8B/W_L2 + 1/Q_AIV) ≈ 12."""
s = self.spec
return s.q16 / 2 * (8 / s.bw_l2 + 1 / s.q_aiv)
# ------------------------------------------------------------------
def check_conditions(self, case: BmmCase) -> list:
s = self.spec
p = self._p_value(case)
grid_k = self._grid_k(case)
dt = case.dtype_in_bytes
checks = []
# 条件 1: P <= C/2 (并行缺口, K 是唯一剩余并行维度)
c1 = p <= s.aic_num / 2
checks.append(ConditionCheck(
"1_并行缺口: P=B*MN*4B/L0C <= C/2",
c1, f"P={p:.2f} vs C/2={s.aic_num/2:.0f}"))
# 条件 2: K/grid_K >= 256B/dtype (单核 K 段下限, dValue)
seg = case.k / max(grid_k, 1)
c2 = seg >= s.dvalue_hw_min / dt and grid_k >= 2
checks.append(ConditionCheck(
"2_单核K段下限: K/grid_K >= 256B/dtype 且 grid_K>=2",
c2, f"K/grid_K={seg:.0f} vs {s.dvalue_hw_min/dt:.0f}, grid_K={grid_k}"))
# 条件 3: K > grid_K^2/(grid_K-1) * theta_c (归约代价可接受)
if grid_k >= 2:
theta_c = self._theta_c()
k_thresh = grid_k * grid_k / (grid_k - 1) * theta_c
c3 = case.k > k_thresh
checks.append(ConditionCheck(
"3_归约代价可接受: K > grid_K^2/(grid_K-1)*theta_c",
c3, f"K={case.k} vs 阈值={k_thresh:.1f} (theta_c={theta_c:.1f})"))
else:
checks.append(ConditionCheck("3_归约代价可接受", False, "grid_K<2 无意义"))
# 条件 4: 确定性等级 <= 1 且 ND
c4 = case.deterministic_level <= 1 and case.out_nd
checks.append(ConditionCheck(
"4_工程约束: 确定性等级<=1 且 ND 格式",
c4, f"deterministic_level={case.deterministic_level}, out_nd={case.out_nd}"))
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
out_b = case.dtype_out_bytes
grid_k = self._grid_k(case)
blocks_per_batch = s.aic_num // b
# mCnt*nCnt 收拢为 blocksPerBatch 的因子, 剩余核预算折成 grid_K
mn_cnt = max(1, blocks_per_batch // grid_k)
m_cnt, n_cnt = self._split_mn(m, n, mn_cnt)
single_m = ceil_div(m, m_cnt)
single_n = ceil_div(n, n_cnt)
single_k = ceil_div(k, grid_k)
# workspace: 部分和驻留 L2, 按 4B (L0C dtype, 防精度丢失)
workspace = grid_k * single_m * single_n * 4 * b
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=grid_k,
core_map=f"B/M/N切出{b*m_cnt*n_cnt}块, 每块{grid_k}核切K归约 (归约组内核c负责K段[c*K/{grid_k},(c+1)*K/{grid_k}))",
b_core=1, merge_b0=1,
single_core_m=single_m, single_core_n=single_n, single_core_k=single_k,
k_l1=min(single_k, s.dvalue_recommend // case.dtype_in_bytes), b_l1=1,
l1_form="K段标准分块流水",
base_m=min(single_m, 128), base_n=min(single_n, 128), base_k=min(single_k, 64),
l2_policy_in="allocate(部分和驻留L2)",
l2_policy_out="resident(部分和4B驻留L2, 防精度丢失不随C的fp16/fp8转换)",
swizzle_w=0, workspace_bytes=workspace,
tail_strategy=f"grid_K={grid_k}路切K+归约",
tail_k_cnt=grid_k, # 尾轮 K 向切分数 = grid_K
fixpipe_unitflag=True,
out_dtype_bytes=4, # 中间部分和按 4B
note=f"P={self._p_value(case):.2f}, grid_K={grid_k}, 部分和驻留L2按4B写出, AIV归约后按C dtype={out_b}B写最终",
)
@staticmethod
def _split_mn(m: int, n: int, mn_cnt: int) -> tuple:
"""把 mn_cnt 拆成 m_cnt*n_cnt, 长宽比跟随 M/N."""
if mn_cnt <= 1:
return 1, 1
ratio = m / max(n, 1)
m_cnt = max(1, round(math.sqrt(mn_cnt * ratio)))
n_cnt = max(1, mn_cnt // m_cnt)
return m_cnt, n_cnt
# ------------------------------------------------------------------
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
grid_k = plan.grid_k
# 单 tile (L0C 满载基本块)
tile_elems = s.l0c_elems # 65536
t_mmad_tile = 2.0 * k * s.l0c_bytes / 4 / s.q16 # 2K/Q16 * L0C/4B
t_mte2_tile = k * (2 * math.sqrt(tile_elems)) * dt / s.bw_pc # 近似方形 tile
# 切 K 后流水时延缩 grid_k 倍
t_mmad = t_mmad_tile / grid_k
t_mte2 = t_mte2_tile / grid_k
# 归约: 部分和 4B 驻留 L2, AIV 归约
t_reduce = eval_streamk_reduce(tile_elems, grid_k, out_b, s)
# 写出: 最终归约结果按 C dtype; 中间部分和按 4B
fix_bytes = tile_elems * out_b
t_fix = fix_bytes / s.bw_l2_pc
flops_pc = 2.0 * tile_elems * k / grid_k
gm_bytes = k * (2 * math.sqrt(tile_elems)) * dt / grid_k
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=t_reduce,
t_drain=t_reduce, # 归约串行追加
gm_read_bytes=gm_bytes, l2_read_bytes=0.0, dma_cmd_count=0.0,
cube_flops=flops_pc, fixpipe_bytes=fix_bytes,
)