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