diff --git a/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py b/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py index d6cb910..1800c12 100644 --- a/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py +++ b/BMM/BMM_Theory/bmm_theory/branches/iter_batch.py @@ -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]