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
This commit is contained in:
@@ -80,22 +80,14 @@ class AswBasicBranch(Branch):
|
||||
# Step 4: swizzle 窗口 W = max{d | d|C, d <= floor(sqrt(C))}
|
||||
swizzle_w = self._swizzle_w()
|
||||
|
||||
# Step 5: L2 分组判断 (issue#24: 按单 batch 工作集判定, 同一时刻单 batch 激活)
|
||||
a_b = m * k * dt
|
||||
bb_b = k * n * dt
|
||||
out_batch = m * n * case.dtype_out_bytes
|
||||
if a_b + bb_b + out_batch <= s.l2_bytes:
|
||||
l2_scene = "A_单batch全驻留(输入+输出)"
|
||||
l2_out = "resident(输出驻留L2异步回写)"
|
||||
r_in = 1.0
|
||||
elif a_b + bb_b <= s.l2_bytes:
|
||||
l2_scene = "B_单batch输入驻留输出直写GM"
|
||||
l2_out = "direct_gm(输出直写GM不占L2)"
|
||||
r_in = 1.0
|
||||
else:
|
||||
l2_scene = "C_单batch输入超L2分组执行"
|
||||
l2_out = "direct_gm(输出直写GM不占L2)"
|
||||
r_in = self._l2_group_r_gm(case, single_m, single_n)
|
||||
# Step 5: L2 场景判定 (issue#24/#30, 设计文档 docs/05 §4, 芯片 L2 128MB)
|
||||
# S_A 整case全驻留: V_in+V_out <= L2 -> 输出驻留 L2 (R4, GM 写=0);
|
||||
# S_B 整case超L2但单batch输入可驻留: batch 内共享块重复读命中 L2, 输出直写 GM;
|
||||
# S_C 单batch输入超L2: 分组执行 (工作集), 组间共享块落空回 GM (r_gm).
|
||||
scene = self._l2_scene(case, single_m, single_n)
|
||||
l2_out = ("resident(整case全驻留S_A: 输出驻留L2异步回写, GM写=0)"
|
||||
if scene["to_l2"] else
|
||||
"direct_gm(输出直写GM, 输入优先驻留L2)")
|
||||
|
||||
# Step 6: k_l1 (GM->L1 K 向粒度, 须 >= dValue 下限; K 小于下限时整 K 一次搬入不切)
|
||||
dv_min_elems = max(s.dvalue_hw_min // dt, 1)
|
||||
@@ -125,7 +117,9 @@ class AswBasicBranch(Branch):
|
||||
tail_block_cnt=tail["r"], tail_wave_num=tail["n_wave"],
|
||||
fixpipe_unitflag=True,
|
||||
out_dtype_bytes=case.dtype_out_bytes,
|
||||
note=f"L2场景{l2_scene}, GM首读倍率={r_in:.2f}; 尾轮: {tail['reason']}",
|
||||
note=(f"L2场景: {scene['label']} (V_in={case.input_bytes/1048576:.1f}MB, "
|
||||
f"V_out={case.output_bytes/1048576:.1f}MB, L2={s.l2_bytes/1048576:.0f}MB), "
|
||||
f"GM首读倍率={scene['r_gm']:.2f}; 尾轮: {tail['reason']}"),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -143,6 +137,7 @@ class AswBasicBranch(Branch):
|
||||
single_m, single_n = clamp_base_mn_l0c(single_m, single_n, False, s)
|
||||
# base_k 由 L0A/L0B 容量按 dtype 反推 (issue#5: 原硬编码 min(K,64) 在 fp32 下溢出)
|
||||
base_k = clamp_base_k(single_m, single_n, dt, case.k, s)
|
||||
out_l2 = case.input_bytes + case.output_bytes <= s.l2_bytes # R4 (issue#30)
|
||||
return ImplPlan(
|
||||
case_id=case.case_id, branch=self.name + "_降核",
|
||||
used_core_num=used,
|
||||
@@ -155,7 +150,9 @@ class AswBasicBranch(Branch):
|
||||
l1_form="标准核内流水",
|
||||
base_m=single_m, base_n=single_n,
|
||||
base_k=base_k,
|
||||
l2_policy_in="allocate", l2_policy_out="direct_gm",
|
||||
l2_policy_in="allocate",
|
||||
l2_policy_out=("resident(整case全驻留S_A)"
|
||||
if out_l2 else "direct_gm(输出直写GM)"),
|
||||
swizzle_w=0, workspace_bytes=0,
|
||||
tail_strategy="不涉及(每核一块无尾轮)",
|
||||
fixpipe_unitflag=True, out_dtype_bytes=case.dtype_out_bytes,
|
||||
@@ -190,6 +187,31 @@ class AswBasicBranch(Branch):
|
||||
w = d
|
||||
return w
|
||||
|
||||
def _l2_scene(self, case: BmmCase, sm: int, sn: int) -> dict:
|
||||
"""L2 场景判定 (make_plan 与 evaluate 共用, 设计文档 docs/05 §4.1).
|
||||
|
||||
判定顺序 S_A -> S_B -> S_C (L2 为整芯片 128MB):
|
||||
S_A 整case全驻留: V_in + V_out <= L2 -> 输出驻留 L2 (R4);
|
||||
S_B 单batch输入驻留: a_b + bb_b <= L2 -> 输入共享块重复读命中 L2;
|
||||
S_C 单batch输入超L2: 分组执行, 组间落空回 GM (r_gm).
|
||||
"""
|
||||
s = self.spec
|
||||
m, n, k = case.m, case.n, case.k
|
||||
dt = case.dtype_in_bytes
|
||||
a_b = m * k * dt
|
||||
bb_b = k * n * dt
|
||||
if case.input_bytes + case.output_bytes <= s.l2_bytes:
|
||||
return {"code": "S_A",
|
||||
"label": "A_整case全驻留(输入+输出<=L2)",
|
||||
"to_l2": True, "r_gm": 1.0}
|
||||
if a_b + bb_b <= s.l2_bytes:
|
||||
return {"code": "S_B",
|
||||
"label": "B_单batch输入驻留(整case超L2, 输出直写GM)",
|
||||
"to_l2": False, "r_gm": 1.0}
|
||||
return {"code": "S_C",
|
||||
"label": "C_单batch输入超L2分组执行(组间落空回GM)",
|
||||
"to_l2": False, "r_gm": self._l2_group_r_gm(case, sm, sn)}
|
||||
|
||||
def _l2_group_r_gm(self, case, sm, sn) -> float:
|
||||
"""场景 C (单 batch 输入仍超 L2): 分组执行的 GM 重复读倍率.
|
||||
|
||||
@@ -227,7 +249,8 @@ class AswBasicBranch(Branch):
|
||||
|
||||
# 主导项判定
|
||||
bw_eff = s.bw_l2_pc # L2 命中
|
||||
t_mmad = 2 * sm * sn * case.k / s.q16
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b) # issue#28
|
||||
t_mmad = 2 * sm * sn * case.k / qc
|
||||
t_mte2 = case.k * (sm + sn) * dt / bw_eff
|
||||
t_fix = sm * sn * out_b / s.bw_pc
|
||||
t_block = max(t_mmad, t_mte2, t_fix)
|
||||
@@ -268,15 +291,17 @@ class AswBasicBranch(Branch):
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
|
||||
"""MTE2 两段块级模型 (issue#24, 用户澄清):
|
||||
"""MTE2 两段块级模型 (issue#24 用户口径 + #28/#29/#30):
|
||||
|
||||
- 首读走 GM (按 GM 带宽, 不叠加 L2); 共享块 (A 行块被 n_cnt 个 tile 读、
|
||||
B 列块被 m_cnt 个 tile 读) 驻留 L2 后其余 (n_cnt-1)/(m_cnt-1) 次读走
|
||||
L2 读口 (5.2TB/s 独享) —— 场景 A/B (单 batch 工作集可驻留);
|
||||
- 单 batch 输入超 L2 -> 场景 C 分组执行, GM 重复按组间落空计 (r_gm), 窗口内
|
||||
复用未另计 (保守);
|
||||
- 输出: 场景 A 驻留 L2 (5.2 写口) / 场景 B,C 直写 GM —— GM 读写共享总线
|
||||
累加由 assemble 的 MTE2 链处理 (issue#23).
|
||||
L2 读口 (5.2TB/s 独享) —— 场景 S_A/S_B (单 batch 工作集可驻留);
|
||||
- 单 batch 输入超 L2 -> 场景 S_C 分组执行, GM 重复按组间落空计 (r_gm),
|
||||
窗口内复用未另计 (保守);
|
||||
- 输出落点 R4 (issue#30): 仅 S_A (整 case 输入+输出 <= L2) 驻留 L2
|
||||
(5.2 写口, GM 写 = 0); S_B/S_C 直写 GM —— GM 读写共享总线累加由
|
||||
assemble 的 MTE2 链处理 (issue#23);
|
||||
- 字节列整芯片口径 (issue#29); Cube 算力按输入 dtype (issue#28).
|
||||
"""
|
||||
s = self.spec
|
||||
b = case.batch_c
|
||||
@@ -284,37 +309,33 @@ class AswBasicBranch(Branch):
|
||||
dt = case.dtype_in_bytes
|
||||
out_b = case.dtype_out_bytes
|
||||
used = max(plan.used_core_num, 1)
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
|
||||
flops = 2.0 * b * m * n * k
|
||||
t_mmad = flops / (used * s.q16)
|
||||
t_mmad = flops / (used * qc)
|
||||
|
||||
# ---- 字节量 (每 batch 口径) ----
|
||||
a_b = m * k * dt # 每 batch A 字节
|
||||
bb_b = k * n * dt # 每 batch B 字节
|
||||
# ---- 字节量 (整芯片口径) ----
|
||||
a_b = m * k * dt # 单 batch A 字节
|
||||
bb_b = k * n * dt # 单 batch B 字节
|
||||
out_all = b * m * n * out_b # 输出总字节
|
||||
m_cnt = max(plan.m_cnt, 1)
|
||||
n_cnt = max(plan.n_cnt, 1)
|
||||
|
||||
# ---- 场景判定: 按单 batch 工作集 (同一时刻单 batch 激活) ----
|
||||
if a_b + bb_b + m * n * out_b <= s.l2_bytes:
|
||||
scene, to_l2, r_gm = "A_单batch全驻留(输入+输出)", True, 1.0
|
||||
elif a_b + bb_b <= s.l2_bytes:
|
||||
scene, to_l2, r_gm = "B_单batch输入驻留,输出直写GM", False, 1.0
|
||||
else:
|
||||
scene, to_l2, r_gm = "C_单batch输入超L2分组执行", False, \
|
||||
self._l2_group_r_gm(case, plan.single_core_m, plan.single_core_n)
|
||||
# ---- L2 场景 (S_A/S_B/S_C, 与 make_plan 同源) ----
|
||||
scene = self._l2_scene(case, plan.single_core_m, plan.single_core_n)
|
||||
to_l2, r_gm = scene["to_l2"], scene["r_gm"]
|
||||
|
||||
# ---- MTE2: GM 段 (首读, 每字节一次) + L2 段 (共享块重复读) ----
|
||||
gm_read = r_gm * b * (a_b + bb_b)
|
||||
if scene.startswith(("A_", "B_")):
|
||||
if scene["code"] in ("S_A", "S_B"):
|
||||
# 驻留命中: A 行块被 n_cnt 个 tile 复用 -> 其余 (n_cnt-1) 次走 L2 读口
|
||||
l2_read = b * ((n_cnt - 1) * a_b + (m_cnt - 1) * bb_b)
|
||||
else:
|
||||
l2_read = 0.0 # 场景 C: 组间复用落空(已计 GM), 窗口内复用保守不计
|
||||
l2_read = 0.0 # 场景 S_C: 组间复用落空(已计 GM), 窗口内复用保守不计
|
||||
t_gm = gm_read / (used * s.bw_pc)
|
||||
t_l2 = l2_read / (used * s.bw_l2_pc)
|
||||
|
||||
# ---- Fixpipe ----
|
||||
# ---- Fixpipe (R4) ----
|
||||
t_fix = out_all / (used * (s.bw_l2_pc if to_l2 else s.bw_pc))
|
||||
|
||||
# drain: 尾轮暴露 (方案 B 已均匀重切, drain 小; A1b 尾轮凑满, drain 小; A0 尾轮 r 核空转)
|
||||
|
||||
@@ -13,7 +13,7 @@ from __future__ import annotations
|
||||
|
||||
from ..hardware import NpuSpec, ASCEND950PR
|
||||
from ..models import BmmCase, ImplPlan, HardwareTiming, align_down
|
||||
from ..timing import assemble_timing
|
||||
from ..timing import assemble_timing, output_to_l2
|
||||
from .base import Branch, BranchResult, ConditionCheck
|
||||
|
||||
|
||||
@@ -171,6 +171,12 @@ class IterBatchBranch(Branch):
|
||||
form_name = {"a": "a_单batch全驻留", "b": "b_双batch乒乓",
|
||||
"c": "c_一侧驻留+对侧切K", "d": "d_两侧都切K"}[form]
|
||||
|
||||
# 输出落点 (issue#30, R4): 整 case 输入+输出 <= L2 才驻留 L2
|
||||
out_l2 = output_to_l2(case, s)
|
||||
l2_out = ("resident(整case输入+输出<=L2: 输出驻留L2异步回写, GM写=0)"
|
||||
if out_l2 else
|
||||
"direct_gm(整case超L2: 输入优先驻留L2, 输出直写GM不占L2)")
|
||||
|
||||
return ImplPlan(
|
||||
case_id=case.case_id, branch=self.name,
|
||||
used_core_num=s.aic_num,
|
||||
@@ -181,54 +187,63 @@ class IterBatchBranch(Branch):
|
||||
k_l1=k_l1, b_l1=2 if form == "b" else 1, l1_form=form_name,
|
||||
base_m=base_m, base_n=base_n, base_k=base_k,
|
||||
l2_policy_in="allocate(GM->L1随路驻留L2)",
|
||||
l2_policy_out="direct_gm(输出仅写一次,直写GM不占L2)",
|
||||
l2_policy_out=l2_out,
|
||||
swizzle_w=0, workspace_bytes=0,
|
||||
tail_strategy="不涉及(核内不切M/N)",
|
||||
fixpipe_unitflag=True,
|
||||
out_dtype_bytes=case.dtype_out_bytes,
|
||||
note=form_desc,
|
||||
note=form_desc + f"; 输出落点: {'L2驻留' if out_l2 else '直写GM'}",
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 时延评估 (v1.1 §4 端到端模型, IterBatch 侧)
|
||||
# 口径: dtype 感知算力 (issue#28); 字节列整芯片 (issue#29);
|
||||
# 输出落点 R4 (issue#30).
|
||||
# ------------------------------------------------------------------
|
||||
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
|
||||
s = self.spec
|
||||
m, n, k = case.m, case.n, case.k
|
||||
dt = case.dtype_in_bytes
|
||||
out_b = case.dtype_out_bytes
|
||||
b = case.batch_c
|
||||
b_core, k_l1 = plan.b_core, plan.k_l1
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
out_l2 = output_to_l2(case, s)
|
||||
w_fix = s.bw_l2_pc if out_l2 else s.bw_pc
|
||||
|
||||
k_truncated = k_l1 >= k
|
||||
n_k = 1 if k_truncated else -(-k // k_l1)
|
||||
|
||||
t_load = min(k_l1, k) * (m + n) * dt / s.bw_pc
|
||||
t_comp_chunk = 2.0 * m * n * min(k_l1, k) / s.q16
|
||||
t_write = m * n * out_b / s.bw_pc
|
||||
t_comp_chunk = 2.0 * m * n * min(k_l1, k) / qc
|
||||
t_write = m * n * out_b / w_fix
|
||||
|
||||
# 搬移: 每 batch n_K 次 GM->L1, 每次含 T_cmd
|
||||
# 搬移: 每 batch n_K 次 GM->L1 (各 (batch,K段) 数据互不重叠, GM 每字节一次),
|
||||
# 每次含 T_cmd; dma_cmds 为单核命令数 (各核并行, issue#29 口径)
|
||||
dma_cmds = b_core * n_k
|
||||
t_mte2_data = dma_cmds * t_load
|
||||
t_dma_cmd = dma_cmds * s.t_cmd
|
||||
t_mte2 = t_mte2_data + t_dma_cmd
|
||||
|
||||
# Cube: 无冗余
|
||||
# Cube: 无冗余; 每核 flops = b_core*2MNK (整芯片列 = b*2MNK)
|
||||
flops_pc = b_core * 2.0 * m * n * k
|
||||
t_mmad = flops_pc / s.q16
|
||||
flops_chip = b * 2.0 * m * n * k
|
||||
t_mmad = flops_pc / qc
|
||||
|
||||
# Fixpipe: b_core 个 batch 输出, 按 C dtype
|
||||
# Fixpipe (R4): b_core 个 batch 输出, 按 C dtype
|
||||
fix_bytes_pc = b_core * m * n * out_b
|
||||
t_fix = fix_bytes_pc / s.bw_pc
|
||||
fix_bytes_chip = b * m * n * out_b
|
||||
t_fix = fix_bytes_pc / w_fix
|
||||
|
||||
# drain: 末 batch 排空 = T_comp + T_write
|
||||
t_drain = t_comp_chunk + t_write
|
||||
|
||||
gm_bytes_pc = b_core * (m * k + k * n) * dt
|
||||
|
||||
return assemble_timing(
|
||||
t_mte2_gm=t_mte2_data, t_mte2_l2=0.0, t_dma_cmd=t_dma_cmd,
|
||||
t_mmad=t_mmad, t_fixpipe=t_fix, t_reduce=0.0, t_drain=t_drain,
|
||||
gm_read_bytes=gm_bytes_pc, l2_read_bytes=0.0,
|
||||
dma_cmd_count=dma_cmds, cube_flops=flops_pc,
|
||||
fixpipe_bytes=fix_bytes_pc,
|
||||
gm_read_bytes=b * n_k * min(k_l1, k) * (m + n) * dt,
|
||||
l2_read_bytes=0.0,
|
||||
dma_cmd_count=dma_cmds, cube_flops=flops_chip,
|
||||
fixpipe_bytes=fix_bytes_chip,
|
||||
fixpipe_to_gm=(not out_l2),
|
||||
)
|
||||
|
||||
@@ -18,7 +18,7 @@ import math
|
||||
|
||||
from ..hardware import NpuSpec, ASCEND950PR
|
||||
from ..models import BmmCase, ImplPlan, HardwareTiming, align_down
|
||||
from ..timing import MoveInPlan, assemble_timing
|
||||
from ..timing import assemble_timing, output_to_l2
|
||||
from .base import Branch, BranchResult, ConditionCheck
|
||||
|
||||
MIN_B0 = 2 # b0: 合并搬移有收益的最小合并数
|
||||
@@ -142,9 +142,17 @@ class MergeBatchBranch(Branch):
|
||||
))
|
||||
b_l1 = max(b_l1, b0)
|
||||
|
||||
# Step 5: 输出落点 (issue#30): 整 case 输入+输出可驻留 L2 才写 L2 (R4)
|
||||
out_l2 = output_to_l2(case, s)
|
||||
l2_out = ("resident(整case输入+输出<=L2: 输出驻留L2异步回写, GM写=0)"
|
||||
if out_l2 else
|
||||
"direct_gm(整case超L2: 输入优先驻留L2, 输出直写GM不占L2)")
|
||||
|
||||
note = (f"b0={b0} (L0C上限{b0_l0c:.1f}/算存比上限{b0_ai:.1f}/b_core={b_core}); "
|
||||
f"{'K截断' if k_truncated else 'L1绑定'}; "
|
||||
f"合并后单次DMA搬入 A'[{b0*m},{k_l1}]+B'[{k_l1},{b0*n}]")
|
||||
f"合并后单次DMA搬入 A'[{b0*m},{k_l1}]+B'[{k_l1},{b0*n}]; "
|
||||
f"输出落点: {'L2驻留' if out_l2 else '直写GM'} (整case V_in+V_out"
|
||||
f"={case.input_bytes/1048576:.1f}MB vs L2={s.l2_bytes/1048576:.0f}MB)")
|
||||
|
||||
return ImplPlan(
|
||||
case_id=case.case_id, branch=self.name,
|
||||
@@ -157,7 +165,7 @@ class MergeBatchBranch(Branch):
|
||||
# L0C 双缓冲约束已由进入条件 2 保证 (2*(b0*M)*(b0*N)*4B <= L0C)
|
||||
base_m=b0 * m, base_n=b0 * n, base_k=max(min(k_l0, k_l1, k), s.fractal),
|
||||
l2_policy_in="allocate(GM->L1随路驻留L2)",
|
||||
l2_policy_out="direct_gm(输出仅写一次,直写GM不占L2)" if not case.out_nd else "direct_gm",
|
||||
l2_policy_out=l2_out,
|
||||
swizzle_w=0, workspace_bytes=0,
|
||||
tail_strategy="不涉及(核内不切M/N)",
|
||||
fixpipe_unitflag=True,
|
||||
@@ -166,55 +174,71 @@ class MergeBatchBranch(Branch):
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 时延评估 (v1.1 §4 端到端模型)
|
||||
# 时延评估 (v1.1 §4 端到端模型; issue#27/#28/#29/#30 口径)
|
||||
# - Cube flops (issue#27 复核): 每合并步计算全网格 (b0*M)x(b0*N)xK
|
||||
# 含交叉项冗余, flops_step = 2*(b0M)*(b0N)*K, 每核步数 = b_core/b0
|
||||
# -> 每核 flops = b_core*b0*2MNK (与 CSV 及文档公式一致, 无修改);
|
||||
# - 算力按输入 dtype (issue#28): q_cube(dtype_a, dtype_b);
|
||||
# - 字节列 = 整芯片口径 (issue#29); GM 首读 = V_in 一次 (合并按组搬入
|
||||
# 各 (batch,K段) 数据互不重叠, 无 L2 重复读);
|
||||
# - 输出落点按整 case 驻留判定 (issue#30, R4).
|
||||
# ------------------------------------------------------------------
|
||||
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
|
||||
s = self.spec
|
||||
m, n, k = case.m, case.n, case.k
|
||||
dt = case.dtype_in_bytes
|
||||
out_b = case.dtype_out_bytes
|
||||
b = case.batch_c
|
||||
b_core, b0, k_l1 = plan.b_core, plan.merge_b0, plan.k_l1
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
out_l2 = output_to_l2(case, s, 0.0)
|
||||
w_fix = s.bw_l2_pc if out_l2 else s.bw_pc
|
||||
|
||||
k_truncated = k_l1 >= k
|
||||
# 每 K 分块搬移/计算时延 (未合并基准, v1.1 §4.1 符号)
|
||||
t_load = k_l1 * (m + n) * dt / s.bw_pc
|
||||
t_comp_chunk = 2.0 * m * n * k_l1 / s.q16
|
||||
t_write = m * n * out_b / s.bw_pc # 单 batch 输出写回 (单核带宽份额)
|
||||
t_comp_chunk = 2.0 * m * n * k_l1 / qc
|
||||
t_write = m * n * out_b / w_fix # 单 batch 输出写回 (R4 落点带宽)
|
||||
|
||||
if k_truncated:
|
||||
# K 截断: n_K=1, 合并后单次搬移量 b0 倍
|
||||
# K 截断: n_K=1, 每核 b_core/b0 次合并搬入, 每次搬 b0 个 batch 全 K
|
||||
n_move = b_core / b0
|
||||
t_mte2_data = n_move * b0 * t_load
|
||||
dma_cmds = n_move
|
||||
gm_chip = b * (m + n) * k * dt # = V_in, GM 每字节一次
|
||||
else:
|
||||
# L1 绑定: k_L1^m = k_L1/b0, n_K^m = b0*n_K, 搬移次数与 IterBatch 相同
|
||||
# (每 (合并组, K段) 数据互不重叠, GM 仍每字节一次, 末段含 padding 上取)
|
||||
n_k = -(-k // k_l1)
|
||||
n_move = b_core * n_k
|
||||
t_mte2_data = n_move * t_load
|
||||
dma_cmds = n_move
|
||||
gm_chip = b * (m + n) * n_k * k_l1 * dt # >= V_in (padding 上取)
|
||||
|
||||
t_dma_cmd = dma_cmds * s.t_cmd
|
||||
t_mte2 = t_mte2_data + t_dma_cmd
|
||||
|
||||
# Cube: 合并计算含冗余 (b0^2 输出, 有效 b0) -> 计算量 = b_core*b0*2MNK
|
||||
# Cube (每核口径): 合并计算含冗余 (b0^2 输出, 有效 b0), 每核 b_core/b0 步,
|
||||
# 每步 flops = 2*(b0*M)*(b0*N)*K -> 每核 = b_core*b0*2MNK (issue#27 复核一致)
|
||||
flops_pc = b_core * b0 * 2.0 * m * n * k
|
||||
t_mmad = flops_pc / s.q16
|
||||
flops_chip = b * b0 * 2.0 * m * n * k # 整芯片口径列
|
||||
t_mmad = flops_pc / qc
|
||||
|
||||
# Fixpipe: 只写对角块, 写出量 = b_core*MN*outB (C 矩阵 dtype, 随路转换)
|
||||
# Fixpipe (R4): 只写对角块, 写出量 = b_core*MN*outB (C 矩阵 dtype, 随路转换)
|
||||
fix_bytes_pc = b_core * m * n * out_b
|
||||
t_fix = fix_bytes_pc / s.bw_pc
|
||||
fix_bytes_chip = b * m * n * out_b
|
||||
t_fix = fix_bytes_pc / w_fix
|
||||
|
||||
# drain: 末合并 batch 排空 = b0*(T_comp + T_write)
|
||||
t_drain = b0 * (t_comp_chunk + t_write)
|
||||
|
||||
gm_bytes_pc = b_core * (m * k + k * n) * dt
|
||||
|
||||
return assemble_timing(
|
||||
t_mte2_gm=t_mte2_data, t_mte2_l2=0.0, t_dma_cmd=t_dma_cmd,
|
||||
t_mmad=t_mmad, t_fixpipe=t_fix, t_reduce=0.0, t_drain=t_drain,
|
||||
gm_read_bytes=gm_bytes_pc, l2_read_bytes=0.0,
|
||||
dma_cmd_count=dma_cmds, cube_flops=flops_pc,
|
||||
fixpipe_bytes=fix_bytes_pc,
|
||||
gm_read_bytes=gm_chip, l2_read_bytes=0.0,
|
||||
dma_cmd_count=dma_cmds, cube_flops=flops_chip,
|
||||
fixpipe_bytes=fix_bytes_chip,
|
||||
fixpipe_to_gm=(not out_l2),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -243,7 +267,9 @@ class MergeBatchBranch(Branch):
|
||||
|
||||
plan = self.make_plan(case)
|
||||
b0 = plan.merge_b0
|
||||
t_comp = 2.0 * m * n * min(k_l1_iter, k) / s.q16
|
||||
# T_comp 按输入 dtype 算力 (issue#28); T_write 保持 v1.1 直写 GM 语义
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
t_comp = 2.0 * m * n * min(k_l1_iter, k) / qc
|
||||
t_write = m * n * out_b / s.bw_pc
|
||||
drain_pen = (b0 - 1) * (t_comp + t_write)
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ from __future__ import annotations
|
||||
|
||||
from ..hardware import NpuSpec, ASCEND950PR
|
||||
from ..models import BmmCase, ImplPlan, HardwareTiming
|
||||
from ..timing import assemble_timing
|
||||
from ..timing import assemble_timing, output_to_l2
|
||||
from .base import Branch, ConditionCheck
|
||||
|
||||
|
||||
@@ -68,32 +68,42 @@ class SpecialBranch(Branch):
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
|
||||
"""AIV 通路时延: 瓶颈在搬移 (AIV 算力远剩)."""
|
||||
"""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:
|
||||
# 纯写值: 仅写出
|
||||
t_in, t_compute, t_out = 0.0, 0.0, b * m * n * out_b / s.bw_gm
|
||||
in_bytes, cube_flops = 0.0, 0.0
|
||||
# 纯写值: 仅写出 (R1 下 GM 读 = 0, 输入本身为空)
|
||||
t_in = 0.0
|
||||
t_out = b * m * n * out_b / w_fix
|
||||
in_bytes = 0.0
|
||||
else:
|
||||
# 逐元素乘: 搬入 A+B, 搬出 C, AIV 算力远剩
|
||||
# 逐元素乘: 搬入 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 / s.bw_gm
|
||||
# AIV 求积吞吐 (近似按 Q_AIV)
|
||||
t_compute = b * m * n / s.q_aiv
|
||||
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)
|
||||
|
||||
t_total = max(t_in, t_out, t_compute)
|
||||
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),
|
||||
)
|
||||
|
||||
@@ -12,7 +12,7 @@ import math
|
||||
|
||||
from ..hardware import NpuSpec, ASCEND950PR
|
||||
from ..models import BmmCase, ImplPlan, HardwareTiming, ceil_div
|
||||
from ..timing import assemble_timing, eval_streamk_reduce
|
||||
from ..timing import assemble_timing, eval_streamk_reduce, output_to_l2
|
||||
from .base import Branch, ConditionCheck
|
||||
|
||||
|
||||
@@ -32,10 +32,14 @@ class StreamKBranch(Branch):
|
||||
p = self._p_value(case)
|
||||
return int(self.spec.aic_num // max(1, math.ceil(p)))
|
||||
|
||||
def _theta_c(self) -> float:
|
||||
"""归约代价系数 theta_c = Q16/2 * (8B/W_L2 + 1/Q_AIV) ≈ 12."""
|
||||
def _theta_c(self, case: BmmCase) -> float:
|
||||
"""归约代价系数 theta_c = Q_cube(dtype)/2 * (8B/W_L2 + 1/Q_AIV_fp32).
|
||||
|
||||
Q_AIV 用 fp32 档 (部分和为 fp32, issue#28); Cube 侧按输入 dtype.
|
||||
"""
|
||||
s = self.spec
|
||||
return s.q16 / 2 * (8 / s.bw_l2 + 1 / s.q_aiv)
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
return qc / 2 * (8 / s.bw_l2 + 1 / s.q_aiv)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
def check_conditions(self, case: BmmCase) -> list:
|
||||
@@ -61,7 +65,7 @@ class StreamKBranch(Branch):
|
||||
|
||||
# 条件 3: K > grid_K^2/(grid_K-1) * theta_c (归约代价可接受)
|
||||
if grid_k >= 2:
|
||||
theta_c = self._theta_c()
|
||||
theta_c = self._theta_c(case)
|
||||
k_thresh = grid_k * grid_k / (grid_k - 1) * theta_c
|
||||
c3 = case.k > k_thresh
|
||||
checks.append(ConditionCheck(
|
||||
@@ -139,7 +143,9 @@ class StreamKBranch(Branch):
|
||||
(工作集超 L2, 保守同 ASW 场景 C 公式);
|
||||
- MMAD: 芯片总量 2*b*m*n*k / (C核) —— 全核忙稳态;
|
||||
- 归约 (部分和 4B 写 L2/AIV 读回求和/最终写回): 每 tile-组一次、组间串行
|
||||
追加为 drain: t_drain = waves * eval_streamk_reduce(tile, grid_k, out).
|
||||
追加为 drain: t_drain = waves * eval_streamk_reduce(tile, grid_k, out);
|
||||
最终输出写回落点按 R4 (issue#30): 整 case 输入+输出+workspace <= L2
|
||||
时写 L2, 否则直写 GM (drain 内按 GM 带宽).
|
||||
"""
|
||||
s = self.spec
|
||||
b = case.batch_c
|
||||
@@ -152,6 +158,8 @@ class StreamKBranch(Branch):
|
||||
tile_m = max(plan.single_core_m, 1)
|
||||
tile_n = max(plan.single_core_n, 1)
|
||||
tile = tile_m * tile_n # 每核实际 tile 元素数
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
out_l2 = output_to_l2(case, s, float(plan.workspace_bytes))
|
||||
|
||||
# ---- MTE2: GM 段 + L2 段 (同 ASW 块级规则) ----
|
||||
a_b = m * k * dt
|
||||
@@ -166,13 +174,14 @@ class StreamKBranch(Branch):
|
||||
t_l2 = l2_read / s.bw_l2
|
||||
t_mte2 = t_gm + t_l2
|
||||
|
||||
# ---- MMAD: 芯片总量 ----
|
||||
# ---- MMAD: 芯片总量, dtype 感知算力 (issue#28) ----
|
||||
flops = 2.0 * b * m * n * k
|
||||
t_mmad = flops / (s.aic_num * s.q16)
|
||||
t_mmad = flops / (s.aic_num * qc)
|
||||
|
||||
# ---- 归约: 每 tile-组一次, 组间分波串行追加 ----
|
||||
waves = ceil_div(b * m_cnt * n_cnt * grid_k, s.aic_num)
|
||||
t_reduce_group = eval_streamk_reduce(tile, grid_k, out_b, s)
|
||||
t_reduce_group = eval_streamk_reduce(tile, grid_k, out_b, s,
|
||||
out_to_gm=not out_l2)
|
||||
t_reduce = waves * t_reduce_group
|
||||
fix_bytes = 0.0
|
||||
t_fix = 0.0
|
||||
|
||||
@@ -12,7 +12,7 @@ from __future__ import annotations
|
||||
|
||||
from ..hardware import NpuSpec, ASCEND950PR
|
||||
from ..models import BmmCase, ImplPlan, HardwareTiming
|
||||
from ..timing import assemble_timing
|
||||
from ..timing import assemble_timing, output_to_l2
|
||||
from .base import Branch, ConditionCheck
|
||||
|
||||
|
||||
@@ -71,7 +71,12 @@ class ToMatmulBranch(Branch):
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
def evaluate(self, case: BmmCase, plan: ImplPlan) -> HardwareTiming:
|
||||
"""折叠后按 Matmul 粗估 (详细切分待 MM 理论体系打通后接入)."""
|
||||
"""折叠后按 Matmul 粗估 (详细切分待 MM 理论体系打通后接入).
|
||||
|
||||
口径: Cube 算力按输入 dtype (issue#28); GM 首读 = V_in 一次 (R1);
|
||||
输出落点 R4 (issue#30): 整 case 输入+输出 <= L2 时写 L2 写口, 否则
|
||||
直写 GM 与读共享总线 (assemble 累加).
|
||||
"""
|
||||
s = self.spec
|
||||
if case.batch_b == 1:
|
||||
fold_m, fold_n, kk = case.batch_a * case.m, case.n, case.k
|
||||
@@ -79,13 +84,15 @@ class ToMatmulBranch(Branch):
|
||||
fold_m, fold_n, kk = case.m, case.batch_b * case.n, case.k
|
||||
dt = case.dtype_in_bytes
|
||||
out_b = case.dtype_out_bytes
|
||||
qc = s.q_cube(case.dtype_a, case.dtype_b)
|
||||
out_l2 = output_to_l2(case, s)
|
||||
|
||||
flops = 2.0 * fold_m * fold_n * kk
|
||||
t_mmad = flops / (s.aic_num * s.q16)
|
||||
t_mmad = flops / (s.aic_num * qc)
|
||||
in_bytes = (fold_m * kk + kk * fold_n) * dt
|
||||
out_bytes = fold_m * fold_n * out_b
|
||||
t_mte2 = in_bytes / s.bw_gm
|
||||
t_fix = out_bytes / s.bw_gm
|
||||
t_fix = out_bytes / (s.bw_l2 if out_l2 else s.bw_gm)
|
||||
|
||||
return assemble_timing(
|
||||
t_mte2_gm=t_mte2, t_mte2_l2=0.0, t_dma_cmd=0.0,
|
||||
@@ -93,4 +100,5 @@ class ToMatmulBranch(Branch):
|
||||
t_drain=0.0,
|
||||
gm_read_bytes=in_bytes, l2_read_bytes=0.0, dma_cmd_count=0.0,
|
||||
cube_flops=flops, fixpipe_bytes=out_bytes,
|
||||
fixpipe_to_gm=(not out_l2),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user