Files
matmul-analysis/BMM/BMM_Theory/bmm_theory/branches/special.py
admin b9e07edc1d Fix #27-#30: dtype感知算力(Cube/AIV速率表) / GM首读下限与整芯片字节列口径+设计文档 / Fixpipe输出落点R4(整case驻留L2否则直写GM) / MergeBatch Cube公式复核注释
- docs/05_L2驻留GM读写与dtype算力口径_设计分析.md: R1-R6公理、S_A/S_B/S_C场景、输出落点R4、dtype速率表(白皮书出处+待标定假设)、逐分支GM/L2归属表 (issue#29/#30 先文档后代码)
- hardware: CUBE_DTYPE_FACTOR(f16/bf16=1, fp8=2x, fp4=4x, fp32=1/2假设) + AIV_DTYPE_FACTOR + q_cube/aiv_elem_rate (issue#28)
- 全分支 t_mmad/t_comp/drain/尾轮主导项/θ_c/R16 语义按输入dtype取算力; 混精度取慢侧; StreamK归约保持fp32(AIV fp32部分和)
- Fixpipe输出落点R4: to_l2 <=> V_in+V_out(+workspace)<=L2; 否则直写GM计入共享总线; ASW场景改S_A整case全驻留(原单batch驻留判定漏计整case输出累积逐出)
- 字节列统一整芯片口径(gm/l2/fix/cube_flops), dma_cmd_count注明单核; GM>=V_in不变量入测试; MergeBatch每步flops=2(b0M)(b0N)K公式注释显式化(#27复核与CSV一致无数值改动)
- tests 39->49 全过; 双压力seed7/6000+seed2024/4000: 0违规/0占位/0NaN/0GM<输入; examples三件套重生成且可复现0diff
2026-09-04 16:26:10 +08:00

110 lines
4.7 KiB
Python

"""特殊分支: K=0 (纯写值) / K=1 (逐元素乘), 走 AIV 向量通路.
理论依据: 《BMM算子优化分析 v0.98》§九 + docs/02_分支理论/04_特殊分支.md.
K 维是 Cube 存在的意义 (累加深度). K=0/1 时 Cube 的 16x16x16 粒度浪费,
走 AIV 向量通路 (GM->UB->算->GM) 优于 Cube 通路. 与切分正交的前置判断.
"""
from __future__ import annotations
from ..hardware import NpuSpec, ASCEND950PR
from ..models import BmmCase, ImplPlan, HardwareTiming
from ..timing import assemble_timing, output_to_l2
from .base import Branch, ConditionCheck
class SpecialBranch(Branch):
name = "特殊分支"
def __init__(self, spec: NpuSpec = ASCEND950PR):
super().__init__(spec)
def check_conditions(self, case: BmmCase) -> list:
c1 = case.k <= 1
checks = [ConditionCheck("1_K<=1 (Cube 无用)", c1, f"K={case.k}")]
if case.k == 1:
# K=1 的 AIV 通路恒可用 (issue#12/#17): B>=2*AIV 开 UB 乒乓; B<128 退化为
# AIV 单缓冲 (无乒乓, 逐 batch 串行搬入), 不再是无方案空洞.
b = case.batch_c
pingpong = b >= 2 * self.spec.aiv_num
mode = "UB乒乓" if pingpong else "AIV单缓冲(逐batch串行, B<2*AIV)"
checks.append(ConditionCheck(
"2_K=1的AIV通路: 恒可用 (B>=128 开UB乒乓, 否则单缓冲)",
True, f"B={b}, 模式={mode}"))
return checks
# ------------------------------------------------------------------
def make_plan(self, case: BmmCase) -> ImplPlan:
s = self.spec
if case.k == 0:
sub = "K=0纯写值"
mode = ""
note = "无任何计算, C=bias 或 0, 纯 AIV 写值; 按行均分到 AIV 核"
else:
sub = "K=1逐元素乘"
pingpong = case.batch_c >= 2 * s.aiv_num
mode = "UB乒乓" if pingpong else "AIV单缓冲"
note = (f"退化为 C=A⊙B 无累加深度, Cube 16x16x16 粒度浪费 15/16; "
f"走 AIV 通路 GM->UB->Mul->GM, {mode} "
f"({'B>=2*AIV 双batch乒乓流水' if pingpong else 'B<2*AIV 逐batch单缓冲串行'})")
return ImplPlan(
case_id=case.case_id, branch=self.name,
used_core_num=s.aiv_num, # 用 AIV 核
split_b=1, m_cnt=1, n_cnt=1, grid_k=1,
core_map="AIV 核间按行均分 (无 Cube tile 概念)",
b_core=0, merge_b0=1,
single_core_m=0, single_core_n=0, single_core_k=case.k,
k_l1=0, b_l1=1,
l1_form="UB驻留(AIV)" if case.k == 0 else "UB驻留(AIV) " + mode,
base_m=0, base_n=0, base_k=0,
l2_policy_in="allocate", l2_policy_out="direct_gm",
swizzle_w=0, workspace_bytes=0,
tail_strategy="不涉及(AIV逐元素)",
fixpipe_unitflag=False,
out_dtype_bytes=case.dtype_out_bytes,
note=f"{sub}: {note}",
)
# ------------------------------------------------------------------
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
"""AIV 通路时延: 瓶颈在搬移 (AIV 算力远剩).
口径: AIV 逐元素通量按输入 dtype (issue#28); 输出落点 R4 (issue#30):
整 case 输入+输出 <= L2 -> 输出写 L2 (异步回写不占算子时延), 否则直写
GM 与读共享总线 (assemble 的 MTE2 链累加, issue#23).
"""
s = self.spec
b = case.batch_c
m, n, k = case.m, case.n, case.k
dt = case.dtype_in_bytes
out_b = case.dtype_out_bytes
out_l2 = output_to_l2(case, s)
w_fix = s.bw_l2 if out_l2 else s.bw_gm
cube_flops = 0.0
t_compute = 0.0
if k == 0:
# 纯写值: 仅写出 (R1 下 GM 读 = 0, 输入本身为空)
t_in = 0.0
t_out = b * m * n * out_b / w_fix
in_bytes = 0.0
else:
# 逐元素乘: 搬入 A+B (GM 每字节一次), 搬出 C
in_bytes = b * (m * k + k * n) * dt
out_bytes = b * m * n * out_b
t_in = in_bytes / s.bw_gm
t_out = out_bytes / w_fix
# AIV 逐元素吞吐按输入 dtype 取通量 (issue#28: bf16/fp16 x2, int8 x4...)
t_compute = b * m * n / s.aiv_elem_rate(case.dtype_in)
cube_flops = float(b * m * n)
return assemble_timing(
t_mte2_gm=t_in, t_mte2_l2=0.0, t_dma_cmd=0.0,
t_mmad=t_compute, t_fixpipe=t_out, t_reduce=0.0, t_drain=0.0,
gm_read_bytes=in_bytes, l2_read_bytes=0.0, dma_cmd_count=0.0,
cube_flops=cube_flops,
fixpipe_bytes=b * m * n * out_b,
fixpipe_to_gm=(not out_l2),
)