Update BMM_Theory: bmm_theory/branches/iter_batch.py (fix review issues #4-#10)
This commit is contained in:
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user