Update BMM_Theory: bmm_theory/branches/asw_basic.py (fix review issues #4-#10)

This commit is contained in:
2026-09-03 11:34:23 +00:00
parent 933b428c06
commit e6aa792cc2

View File

@@ -96,9 +96,11 @@ class AswBasicBranch(Branch):
l2_out = "direct_gm(输出直写GM不占L2)"
r_in = self._l2_group_r_in(case, single_m, single_n)
# Step 6: k_l1 (GM->L1 K 向粒度, 须 >= 256B/dt)
# Step 6: k_l1 (GM->L1 K 向粒度, 须 >= dValue 下限; K 小于下限时整 K 一次搬入不切)
dv_min_elems = max(s.dvalue_hw_min // dt, 1)
k_l1 = min(k, s.dvalue_recommend // dt)
k_l1 = max(k_l1, s.dvalue_hw_min // dt)
if k_l1 < dv_min_elems:
k_l1 = k # K 本身小于 dValue 下限: 不切 K, 整段搬入 (K 非连续维, dValue 由 M/N 保证)
# ---- 尾轮决策 (必要组成环节) ----
n_blk = b * m_cnt * n_cnt
@@ -130,10 +132,16 @@ class AswBasicBranch(Branch):
"""降核模式: 只用 ceil(P) 核, 每核一个 L0C 满载输出块."""
s = self.spec
used = max(1, math.ceil(p))
# SingleCoreM/N 在 L0C 容量内取最大
dt = case.dtype_in_bytes
# SingleCoreM/N 在 L0C 容量内取最大, 收敛到不越界 (issue#5: 降核漏容量反推)
from ..constraints import clamp_base_mn_l0c, clamp_base_k
single_mn = int(math.sqrt(s.l0c_bytes // 4)) # L0C/4B 单份
single_m = min(case.m, align_down(single_mn, s.fractal))
single_n = min(case.n, align_down(single_mn, s.fractal))
# 降核每核单份 L0C (double_buffer=False)
single_m, single_n = clamp_base_mn_l0c(single_m, single_n, False, s)
# base_k 由 L0A/L0B 容量按 dtype 反推 (issue#5: 原硬编码 min(K,64) 在 fp32 下溢出)
base_k = clamp_base_k(single_m, single_n, dt, case.k, s)
return ImplPlan(
case_id=case.case_id, branch=self.name + "_降核",
used_core_num=used,
@@ -142,10 +150,10 @@ class AswBasicBranch(Branch):
core_map=f"降核: 只用{used}核, 每核一个L0C满载输出块, 其余核闲置",
b_core=0, merge_b0=1,
single_core_m=single_m, single_core_n=single_n, single_core_k=case.k,
k_l1=min(case.k, s.dvalue_recommend // case.dtype_in_bytes), b_l1=1,
k_l1=min(case.k, s.dvalue_recommend // dt), b_l1=1,
l1_form="标准核内流水",
base_m=single_m, base_n=single_n,
base_k=min(case.k, 64),
base_k=base_k,
l2_policy_in="allocate", l2_policy_out="direct_gm",
swizzle_w=0, workspace_bytes=0,
tail_strategy="不涉及(每核一块无尾轮)",