- 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
110 lines
4.7 KiB
Python
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),
|
|
)
|