From 16a84ac45c428872d908db162869ff6af1a03f50 Mon Sep 17 00:00:00 2001 From: admin Date: Thu, 3 Sep 2026 11:34:16 +0000 Subject: [PATCH] Add BMM_Theory: bmm_theory/constraints.py (fix review issues #4-#10) --- BMM/BMM_Theory/bmm_theory/constraints.py | 170 +++++++++++++++++++++++ 1 file changed, 170 insertions(+) create mode 100644 BMM/BMM_Theory/bmm_theory/constraints.py diff --git a/BMM/BMM_Theory/bmm_theory/constraints.py b/BMM/BMM_Theory/bmm_theory/constraints.py new file mode 100644 index 0000000..e93ed46 --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/constraints.py @@ -0,0 +1,170 @@ +"""单一约束源: 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 = [] + + # 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: + 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