Add BMM_Theory: bmm_theory/branches/iter_batch.py

This commit is contained in:
2026-09-03 08:09:31 +00:00
parent 87e945c5a8
commit d93ae35729

View File

@@ -0,0 +1,203 @@
"""IterBatch 分支: 核间切 B, 核内逐 batch 标准 Matmul 分块计算.
理论依据:
- 《BMM算子优化分析 v0.98》§六 (进入条件四形态 a/b/c/d + 实现方案)
- 《MergeBatch_vs_IterBatch分析 v1.1》§二/§四 (执行模型 + 分界)
核心思想: "切B"最朴素形态, 无算力浪费、无跨 batch 依赖.
进入的核心要求是单 batch 计算核内零重复读 (由 L1 驻留四形态刻画);
batch 间流水靠 L1 双 buffer / 驻留侧预取 (c 形态半预算) 掩盖.
"""
from __future__ import annotations
from ..hardware import NpuSpec, ASCEND950PR
from ..models import BmmCase, ImplPlan, HardwareTiming, align_down
from ..timing import assemble_timing
from .base import Branch, BranchResult, ConditionCheck
class IterBatchBranch(Branch):
name = "IterBatch"
def __init__(self, spec: NpuSpec = ASCEND950PR):
super().__init__(spec)
# ------------------------------------------------------------------
# L1 驻留形态判定 (v0.98 §六 条件 3, 四选一)
# ------------------------------------------------------------------
def l1_form(self, case: BmmCase) -> tuple:
"""返回 (形态 'a'/'b'/'c'/'d'/None, k_l1, 说明)."""
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
b_core = -(-case.batch_c // s.aic_num) # ceil
single = (m * k + k * n) * dt
# a) 单 batch 全驻留
if b_core == 1 and single <= s.l1_bytes:
return "a", k, f"单batch全驻留: (MK+KN)*dtype={single/1024:.0f}KB <= L1"
# b) 双 batch 乒乓
if b_core > 1 and 2 * single <= s.l1_bytes:
return "b", k, f"双batch乒乓: 2*(MK+KN)*dtype={2*single/1024:.0f}KB <= L1"
# c) 一侧驻留 + 对侧切 K, 预算按 b_core 分档
l1_budget = s.l1_bytes / min(b_core, 2)
for resident, side in ((m * k * dt, "A"), (k * n * dt, "B")):
other = n * dt if side == "A" else m * dt
if resident <= l1_budget:
k_l1 = min(int((l1_budget - resident) / other / 2), k)
k_l1 = align_down(max(k_l1, s.fractal), s.fractal)
if k_l1 >= s.fractal and k_l1 * dt >= s.dvalue_min:
return "c", k_l1, (
f"一侧驻留({side})+对侧切K: {side}驻留{resident/1024:.0f}KB, "
f"预算L1/{min(b_core,2)}, k_L1={k_l1}")
# b_core>=2 时另一半 L1 预取下一 batch 驻留侧, 边界无气泡
# d) 两侧都切 K (兜底)
k_l1 = align_down(int(s.l1_bytes / (2 * (m + n) * dt)), s.fractal)
if k_l1 >= s.fractal and k_l1 * dt >= s.dvalue_min:
return "d", k_l1, f"两侧都切K: k_L1={k_l1}, K段成对流水, batch边界天然无缝"
return None, 0, "L1 四形态均不满足 (M/N 相对 L1 过大)"
# ------------------------------------------------------------------
# 进入条件 (v0.98 §六, 四条同时满足)
# ------------------------------------------------------------------
def check_conditions(self, case: BmmCase) -> list:
s = self.spec
b = case.batch_c
b_core = -(-b // s.aic_num)
checks = []
c1 = (case.batch_a == case.batch_b) and (b_core >= 1)
checks.append(ConditionCheck(
"1_batch关系与每核份额: BatchA==BatchB 且 b_core>=1",
c1, f"batchA={case.batch_a}, batchB={case.batch_b}, b_core={b_core}"))
# 条件 2: 负载均衡 (整除 或 尾波活跃核数 >= minCoreNum)
rem = b % s.aic_num
c2 = (rem == 0) or (rem >= s.min_core_num)
checks.append(ConditionCheck(
"2_负载均衡: B mod C == 0 或 >= minCoreNum",
c2, f"B mod C={rem}, minCoreNum={s.min_core_num}"))
form, k_l1, form_desc = self.l1_form(case)
checks.append(ConditionCheck(
"3_L1驻留形态(四选一, 核心要求: 单batch核内零重复读)",
form is not None, form_desc))
# 条件 4: 搬移效率下限 (c/d 切分后)
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
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"))
else:
checks.append(ConditionCheck(
"4_搬移效率(a/b形态不切分, 恒满足)", True, ""))
return checks
# ------------------------------------------------------------------
# 实现方案 (v0.98 §六 实现方案 a/b/c/d)
# ------------------------------------------------------------------
def make_plan(self, case: BmmCase) -> ImplPlan:
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
b_core = -(-case.batch_c // s.aic_num)
form, k_l1, form_desc = self.l1_form(case)
form = form or "d"
# L0 级 tile (v0.98 §六 (a) 伪代码):
# L0C 放得下完整输出 -> BaseM=M, BaseN=N, BaseK 由 L0A/L0B 决定 (双缓冲, 故除 2);
# 放不下 -> 按较小维切, 再定 BaseK.
# 注意 L0A/L0B 双缓冲两份, 且 base_k 不超 k_l1 (L1 tile 的 K 粒度) 与 K 本身.
if s.l0c_bytes >= m * n * 4 * 2:
base_m, base_n = m, n
else:
if m < n:
base_m = align_down(max(m, s.fractal), s.fractal)
base_n = align_down(int(s.l0c_bytes / (2 * 4 * base_m)), s.fractal)
else:
base_n = align_down(max(n, s.fractal), s.fractal)
base_m = align_down(int(s.l0c_bytes / (2 * 4 * base_n)), s.fractal)
base_m = max(base_m, s.fractal)
base_n = max(base_n, s.fractal)
base_k = align_down(int(min(
s.l0a_bytes / (2 * base_m * dt),
s.l0b_bytes / (2 * base_n * dt),
)), s.fractal)
base_k = max(min(base_k, k_l1 if k_l1 > 0 else k, k), s.fractal)
form_name = {"a": "a_单batch全驻留", "b": "b_双batch乒乓",
"c": "c_一侧驻留+对侧切K", "d": "d_两侧都切K"}[form]
return ImplPlan(
case_id=case.case_id, branch=self.name,
used_core_num=s.aic_num,
split_b=s.aic_num, m_cnt=1, n_cnt=1, grid_k=1,
core_map="切B轮转分配(核间零重复读零依赖)",
b_core=b_core, merge_b0=1,
single_core_m=m, single_core_n=n, single_core_k=k,
k_l1=k_l1, b_l1=2 if form == "b" else 1, l1_form=form_name,
base_m=base_m, base_n=base_n, base_k=base_k,
l2_policy_in="allocate(GM->L1随路驻留L2)",
l2_policy_out="direct_gm(输出仅写一次,直写GM不占L2)",
swizzle_w=0, workspace_bytes=0,
tail_strategy="不涉及(核内不切M/N)",
fixpipe_unitflag=True,
out_dtype_bytes=case.dtype_out_bytes,
note=form_desc,
)
# ------------------------------------------------------------------
# 时延评估 (v1.1 §4 端到端模型, IterBatch 侧)
# ------------------------------------------------------------------
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
s = self.spec
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
out_b = case.dtype_out_bytes
b_core, k_l1 = plan.b_core, plan.k_l1
k_truncated = k_l1 >= k
n_k = 1 if k_truncated else -(-k // k_l1)
t_load = min(k_l1, k) * (m + n) * dt / s.bw_pc
t_comp_chunk = 2.0 * m * n * min(k_l1, k) / s.q16
t_write = m * n * out_b / s.bw_pc
# 搬移: 每 batch n_K 次 GM->L1, 每次含 T_cmd
dma_cmds = b_core * n_k
t_mte2_data = dma_cmds * t_load
t_dma_cmd = dma_cmds * s.t_cmd
t_mte2 = t_mte2_data + t_dma_cmd
# Cube: 无冗余
flops_pc = b_core * 2.0 * m * n * k
t_mmad = flops_pc / s.q16
# Fixpipe: b_core 个 batch 输出, 按 C dtype
fix_bytes_pc = b_core * m * n * out_b
t_fix = fix_bytes_pc / s.bw_pc
# drain: 末 batch 排空 = T_comp + T_write
t_drain = t_comp_chunk + t_write
gm_bytes_pc = b_core * (m * k + k * n) * dt
return assemble_timing(
t_mte2_gm=t_mte2_data, t_mte2_l2=0.0, t_dma_cmd=t_dma_cmd,
t_mmad=t_mmad, t_fixpipe=t_fix, t_reduce=0.0, t_drain=t_drain,
gm_read_bytes=gm_bytes_pc, l2_read_bytes=0.0,
dma_cmd_count=dma_cmds, cube_flops=flops_pc,
fixpipe_bytes=fix_bytes_pc,
)