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

This commit is contained in:
2026-09-03 11:34:25 +00:00
parent e6aa792cc2
commit c0b358fa94

View File

@@ -116,26 +116,39 @@ class IterBatchBranch(Branch):
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:
# L0 级 tile (BaseM x BaseN): 核内 L1->L0 的切分, 与核间切分无关.
# 原则 (v0.98 §六 + ASW Step0): 把 L0C 双缓冲用满 (32768 元素), 长宽比跟随 M/N.
# M*N 大时切多个 L0 tile 流水计算, L1->L0 由 MTE1 搬运与 Cube 掩盖, 无额外搬移代价.
# 收敛到同时满足 L0C 与 L0A/L0B (base_k>=16 前提) 容量, 用同源辅助 (issue#5).
from ..constraints import clamp_base_mn_l0c, clamp_base_k
l0c_cap = s.l0c_bytes // (4 * 2) # 32768 元素 (双缓冲)
# 初始: 尽量跟随 M/N 但不超 L0C 容量, 方形附近取
if m * n <= l0c_cap:
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)
# 长宽比跟随 M/N, 面积 = l0c_cap
ratio = m / max(n, 1)
base_n = align_down(max(int((l0c_cap / max(ratio, 1e-9)) ** 0.5), s.fractal), s.fractal)
base_m = align_down(max(int(l0c_cap / base_n), s.fractal), s.fractal)
base_m = min(base_m, align_down(max(m, s.fractal), s.fractal))
base_n = min(base_n, align_down(max(n, s.fractal), s.fractal))
base_m, base_n = clamp_base_mn_l0c(base_m, base_n, True, s)
# 迭代收敛: base_k>=16 且 L0A/L0B 不溢出 (fp32 大 M/N 时继续缩)
for _ in range(32):
bk = clamp_base_k(base_m, base_n, dt, k, s)
if base_m * bk * dt * 2 <= s.l0a_bytes and base_n * bk * dt * 2 <= s.l0b_bytes:
base_k = bk
break
if base_m >= base_n and base_m > s.fractal:
base_m = max(align_down(base_m // 2, s.fractal), s.fractal)
elif base_n > s.fractal:
base_n = max(align_down(base_n // 2, s.fractal), 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)
base_k = bk
break
else:
base_k = clamp_base_k(base_m, base_n, dt, k, s)
base_k = max(min(base_k, k_l1 if k_l1 > 0 else k), s.fractal)
form_name = {"a": "a_单batch全驻留", "b": "b_双batch乒乓",
"c": "c_一侧驻留+对侧切K", "d": "d_两侧都切K"}[form]