Update BMM_Theory: bmm_theory/branches/asw_basic.py (fix review issues #4-#10)
This commit is contained in:
@@ -96,9 +96,11 @@ class AswBasicBranch(Branch):
|
|||||||
l2_out = "direct_gm(输出直写GM不占L2)"
|
l2_out = "direct_gm(输出直写GM不占L2)"
|
||||||
r_in = self._l2_group_r_in(case, single_m, single_n)
|
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 = 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
|
n_blk = b * m_cnt * n_cnt
|
||||||
@@ -130,10 +132,16 @@ class AswBasicBranch(Branch):
|
|||||||
"""降核模式: 只用 ceil(P) 核, 每核一个 L0C 满载输出块."""
|
"""降核模式: 只用 ceil(P) 核, 每核一个 L0C 满载输出块."""
|
||||||
s = self.spec
|
s = self.spec
|
||||||
used = max(1, math.ceil(p))
|
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_mn = int(math.sqrt(s.l0c_bytes // 4)) # L0C/4B 单份
|
||||||
single_m = min(case.m, align_down(single_mn, s.fractal))
|
single_m = min(case.m, align_down(single_mn, s.fractal))
|
||||||
single_n = min(case.n, 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(
|
return ImplPlan(
|
||||||
case_id=case.case_id, branch=self.name + "_降核",
|
case_id=case.case_id, branch=self.name + "_降核",
|
||||||
used_core_num=used,
|
used_core_num=used,
|
||||||
@@ -142,10 +150,10 @@ class AswBasicBranch(Branch):
|
|||||||
core_map=f"降核: 只用{used}核, 每核一个L0C满载输出块, 其余核闲置",
|
core_map=f"降核: 只用{used}核, 每核一个L0C满载输出块, 其余核闲置",
|
||||||
b_core=0, merge_b0=1,
|
b_core=0, merge_b0=1,
|
||||||
single_core_m=single_m, single_core_n=single_n, single_core_k=case.k,
|
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="标准核内流水",
|
l1_form="标准核内流水",
|
||||||
base_m=single_m, base_n=single_n,
|
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",
|
l2_policy_in="allocate", l2_policy_out="direct_gm",
|
||||||
swizzle_w=0, workspace_bytes=0,
|
swizzle_w=0, workspace_bytes=0,
|
||||||
tail_strategy="不涉及(每核一块无尾轮)",
|
tail_strategy="不涉及(每核一块无尾轮)",
|
||||||
|
|||||||
Reference in New Issue
Block a user