Files
matmul-analysis/BMM/BMM_Theory/bmm_theory/constraints.py

188 lines
9.0 KiB
Python

"""单一约束源: L0C/L0A/L0B/L1/dValue/min_TileSize/核数 的硬件约束, 生成与校验共用.
修复 issue#5 (生成与校验约束不一致) + issue#6 (dValue 口径矛盾) 的共性根因:
此前约束知识分散在三处 (各分支 check_conditions / make_plan 硬编码 / evaluator._check_constraints),
由不同轮次独立写出、互相不一致. 本模块是唯一事实来源 (single source of truth).
dValue 口径裁定 (issue#6 的理论裁定):
dValue 约束的是 GM->L1 的**连续维搬移效率**. 两种情形分别处理:
(a) **整矩阵单次搬移** (IterBatch 形态 a/b, K 整体驻留不切 K): 搬移的是完整 [M,K]/[K,N] 矩阵,
连续维是 M/N (大维度), K 是分段维而非连续维 —— dValue 由 M/N 方向保证, **K 向不适用 128B 下限**;
真正的搬移效率约束是单次搬移量 >= min_TileSize (16KB), 由进入条件 4 保证.
(b) **K 切分搬移** (IterBatch 形态 c/d, ASW, MergeBatch, StreamK): k_l1 成为连续维的一部分,
须满足 k_l1 * dtype >= 128B (dValue 下限), 否则 DMA 突发效率崩塌.
故 dValue 下限只对"以 K 段为连续维"的方案生效, 由 plan 的 k_l1 语义区分.
"""
from __future__ import annotations
from .hardware import NpuSpec, ASCEND950PR
from .models import BmmCase, ImplPlan
# 走 AIV 通路的分支 (无 Cube tile 概念)
AIV_BRANCHES = {"特殊分支"}
# 折叠转 Matmul 的分支 (Cube tile 由 Matmul 体系承接, 此处不校验)
FOLD_BRANCHES = {"转Matmul"}
# StreamK 中间部分和按 4B (L0C dtype) 写出防精度丢失, 是正确行为不算违规
PARTIAL_SUM_4B_BRANCHES = {"StreamK"}
def check_plan_constraints(case: BmmCase, plan: ImplPlan,
spec: NpuSpec = ASCEND950PR) -> list:
"""统一约束校验, 返回违规列表 (空 = 可行).
生成 (recommend 自检) 与评估 (evaluate) 共用本函数, 保证口径一致.
"""
s = spec
v = []
# 占位/无效方案: used_core_num<1 说明没有真实方案 (如 router._no_plan 的占位),
# recommend 侧 advice 已标注"[无方案]", evaluate 侧必须判不可行 (issue#18),
# 不得当作可行方案给正常时延.
if plan.used_core_num < 1:
v.append("used_core_num=0: 占位/未生成方案, 不可评估")
return v
# AIV 通路: 只校验 AIV 核数
if plan.branch in AIV_BRANCHES:
if plan.used_core_num > s.aiv_num:
v.append(f"used_core_num={plan.used_core_num} 超 AIV 核数 {s.aiv_num}")
return v
# 转Matmul: 折叠后由 Matmul 体系承接, 不校验 Cube tile
if plan.branch in FOLD_BRANCHES:
return v
# --- 核数 ---
if plan.used_core_num > s.aic_num:
v.append(f"used_core_num={plan.used_core_num} 超 AIC 核数 {s.aic_num}")
# --- L0C --- (ASW 降核每核单份, 其他分支双缓冲)
l0c_factor = 1 if plan.branch == "ASW_Basic_降核" else 2
l0c_need = plan.base_m * plan.base_n * 4 * l0c_factor
if l0c_need > s.l0c_bytes:
v.append(f"L0C tile 超容量: BaseM*BaseN*4B*{l0c_factor}={l0c_need}B > {s.l0c_bytes}B")
# --- L0A/L0B --- (双缓冲两份)
if plan.base_m > 0 and plan.base_k > 0:
if plan.base_m * plan.base_k * case.dtype_in_bytes * 2 > s.l0a_bytes:
v.append(f"L0A tile 超容量: BaseM*BaseK*dt*2={plan.base_m*plan.base_k*case.dtype_in_bytes*2}B > {s.l0a_bytes}B")
if plan.base_n * plan.base_k * case.dtype_in_bytes * 2 > s.l0b_bytes:
v.append(f"L0B tile 超容量: BaseN*BaseK*dt*2={plan.base_n*plan.base_k*case.dtype_in_bytes*2}B > {s.l0b_bytes}B")
# --- L1 --- (MergeBatch/IterBatch 的 k_l1 语义为 K 向粒度)
if plan.branch in ("MergeBatch", "IterBatch"):
l1_need = _l1_need_bytes(plan, case, s)
# 预算按形态与 b_core 分档 (与生成侧 l1_form 的 budget 一致):
# 形态 c 且 b_core>=2 时半预算 (另一半预取下一 batch 驻留侧); b_core=1 时全量.
if plan.l1_form.startswith("c_") and plan.b_core >= 2:
l1_cap = s.l1_bytes / 2
else:
l1_cap = s.l1_bytes
if l1_need > l1_cap:
v.append(f"L1 tile 超容量: {l1_need/1024:.0f}KB > {l1_cap/1024:.0f}KB")
# --- dValue --- (issue#6 口径裁定: 只对"以 K 段为连续维"的方案生效)
if _k_segment_is_contiguous(plan, case) and plan.k_l1 > 0:
if plan.branch == "IterBatch" and plan.l1_form.startswith(("c_", "d_")):
# 转置感知判据与生成守卫/条件 4 同源 (issue#19):
# dv_a = M*dt (A 转置) 或 k_l1*dt; dv_b = N*dt (B 不转置) 或 k_l1*dt;
# 两侧连续维 dValue 均低于下限才算违规.
from .models import dvalue_contig_dims
dv_a, dv_b = dvalue_contig_dims(case, plan.k_l1)
if dv_a < s.dvalue_min and dv_b < s.dvalue_min:
v.append(f"dValueA={dv_a:.0f}B 与 dValueB={dv_b:.0f}B 均 < 下限 "
f"{s.dvalue_min}B, 搬移连续维效率崩塌")
else:
dv = plan.k_l1 * case.dtype_in_bytes
if dv < s.dvalue_min:
v.append(f"dValue={dv}B < 下限 {s.dvalue_min}B, K 段连续维搬移效率崩塌")
# --- 写出 dtype --- (StreamK 部分和 4B 是正确行为)
if plan.branch not in PARTIAL_SUM_4B_BRANCHES:
if plan.out_dtype_bytes != case.dtype_out_bytes:
v.append(f"方案写出 dtype ({plan.out_dtype_bytes}B) 与 case C 矩阵 dtype ({case.dtype_c}) 不一致")
return v
def _l1_need_bytes(plan: ImplPlan, case: BmmCase, spec: NpuSpec) -> float:
"""按 IterBatch/MergeBatch 的 L1 驻留形态计算 L1 占用.
- 形态 a/b (整驻留): 整个 (MK+KN)*dt 驻留, b 形态双 batch 乒乓 x2;
- 形态 c (一侧驻留+对侧切K): 驻留侧全量 + 对侧 k_l1 段双缓冲, 占半预算 L1/2;
- 形态 d (两侧切K): 两侧 k_l1 段双缓冲 2*k_l1*(M+N)*dt;
- MergeBatch: 合并后 [b0*M, k_l1] + [k_l1, b0*N] 双缓冲.
"""
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
if plan.branch == "MergeBatch":
return 2 * plan.k_l1 * plan.merge_b0 * (m + n) * dt
# IterBatch
if plan.l1_form.startswith("a_"):
return (m * k + k * n) * dt
if plan.l1_form.startswith("b_"):
return 2 * (m * k + k * n) * dt
if plan.l1_form.startswith("c_"):
# 驻留侧全量 + 对侧 k_l1 双缓冲, 预算为 L1/min(b_core,2)
budget = spec.l1_bytes / min(max(plan.b_core, 1), 2)
resident = min(m * k, k * n) * dt # 驻留较小侧
return resident + 2 * plan.k_l1 * max(m, n) * dt
# d_ 两侧切K
return 2 * plan.k_l1 * (m + n) * dt
def _k_segment_is_contiguous(plan: ImplPlan, case: BmmCase) -> bool:
"""k_l1 是否作为 GM->L1 的连续维 (此时 dValue 下限生效).
口径裁定 (issue#6): dValue 约束的是"连续维的搬移效率". 当 K 整体驻留不切
(k_l1 >= K, 即 K 不分段、整矩阵单次搬入) 时, 连续维是 M/N 而非 K, dValue 由
M/N 方向保证, K 向的 128B 下限不适用 —— 这类方案豁免 K 向 dValue 检查.
仅当 k_l1 < K (K 被切分成段、K 段成为搬移连续维) 时, 下限才生效.
"""
# 整 K 驻留 (k_l1 >= K): K 非连续维, 豁免
if plan.k_l1 >= case.k > 0:
return False
# 以下分支 k_l1 < K 时 K 段是连续维, dValue 下限生效
if plan.branch == "IterBatch":
return plan.l1_form.startswith(("c_", "d_"))
if plan.branch == "MergeBatch":
return True
return plan.branch in ("ASW_Basic", "ASW_Basic_降核", "StreamK")
# ---------------------------------------------------------------------------
# 生成侧辅助: 供各分支 make_plan 调用, 保证生成的 tile 不越界 (与校验同源)
# ---------------------------------------------------------------------------
def clamp_base_k(base_m: int, base_n: int, dtype_bytes: int, k: int,
spec: NpuSpec = ASCEND950PR) -> int:
"""由 L0A/L0B 容量反推 base_k (双缓冲), 向下 16 对齐, 不超 K.
供 make_plan 生成 base_k 时调用, 避免硬编码 (issue#5 ASW 降核 base_k=min(K,64) 的缺陷).
"""
from .models import align_down
bk = min(
spec.l0a_bytes // (2 * max(base_m, 1) * dtype_bytes),
spec.l0b_bytes // (2 * max(base_n, 1) * dtype_bytes),
k,
)
return max(align_down(bk, spec.fractal), spec.fractal)
def clamp_base_mn_l0c(base_m: int, base_n: int, double_buffer: bool,
spec: NpuSpec = ASCEND950PR) -> tuple:
"""把 base_m/base_n 收敛到 L0C 容量内 (4B/元素), 返回 (base_m, base_n)."""
from .models import align_down
factor = 2 if double_buffer else 1
cap = spec.l0c_bytes // (4 * factor)
bm, bn = base_m, base_n
while bm * bn > cap and (bm > spec.fractal or bn > spec.fractal):
if bm >= bn and bm > spec.fractal:
bm = max(align_down(bm // 2, spec.fractal), spec.fractal)
elif bn > spec.fractal:
bn = max(align_down(bn // 2, spec.fractal), spec.fractal)
else:
break
return bm, bn