用户澄清: MergeBatch vs IterBatch 的本质区别不只是 DMA 命令数 —— 合并 b0 个 batch 的左/右矩阵一起搬入 L1, 使单块 tile = nValue*dValue*dt 放大 b0 倍 (堆叠 方向视转置: A ND 非转置沿 M(nValue), B ND 非转置沿 N(dValue)), 搬移效率更高, 即便 T_cmd=0 也有效益。 - models.move_eff: 单命令搬移效率 eff = min(1, tile/min_TileSize) (16KB 饱和, 与进入条件4效率下限语义同源); gm_move_time 按 A/B 两侧字节加权 t = (V_A/eff_A + V_B/eff_B)/BW_gm; 只影响时间列, GM 字节量仍 = V_in - IterBatch: l1_form 补驻留侧返回; move_tiles 分侧口径 (a/b 双侧整K, c 驻留侧 整K+对侧k_l1, d 双侧k_l1), evaluate 接入效率加权 - MergeBatch: 合并 tile 放大 b0 倍接入效率加权; beats_iterbatch 净收益 = 命令节省(cmds差×T_cmd) + 效率节省(t_data差) − drain惩罚, K截断且效率打平且 T_cmd>0 时严格退化为 v1.1 §4.5 闭式; 退役 T_cmd<=0 策略特判 - router: 退役 "T_cmd<=0 策略优先 MergeBatch" 覆盖, 时延模型统一终审 - hardware: t_cmd_ns 50 -> 0 (未标定按 0; 合并收益不再依赖 T_cmd 估计值) - 作用域: 仅切B 两分支接入 (逐命令 tile 小、效率差显著); ASW/StreamK 单命令 tile 通常已饱和, 极端小 tile 走 issue#34 效率降级标注通道 - 用户 case 家族 B=128,M=1~16,N=128,K=512: m=1~8 -> MergeBatch (效率节省 ~0.61us > drain), m=16 -> IterBatch (iter A tile 恰达 16KB 饱和, 效率打平, drain 决定); 分界与时延全家族一致 - demo: merge_demo_k_trunc 形状 (2048,32,32,256)->(2048,16,64,128) (原形状 两侧 tile 均已 16KB 饱和, t_cmd=0 下无收益转 IterBatch; 新形状 iter A tile 4KB eff=0.25 vs 合并 16KB eff=1.0, 保持 MergeBatch 胜出演示且仍 K截断) - 测试: 74/74 (新增 TestIssue36 5 例: 效率曲线/字节不变/效率差胜出/家族; TestArbitration/TestZeroCmdHandling 按 t_cmd=0+效率语义重写; TestIssue35 家族期望更新) - 文档: 01_MergeBatch §4/§5 效率模型+泛化净收益; 02_IterBatch 口径注; 00_总纲胜出条件; 01_软件架构 T_cmd 标定说明; 05 时间列效率口径注; README 要点 - 验证: examples 重生成可复现 0 diff; 压力 10000 例 0 崩溃/0 NaN/0 违规/ 0 GM<V_in, 七分支覆盖 (MergeBatch 386 例)
363 lines
14 KiB
Python
363 lines
14 KiB
Python
"""核心数据模型: Case 输入 / 实现方案结构体 / 评估结果.
|
||
|
||
约定:
|
||
- 所有"时延"内部统一用秒 (float), 输出 csv 时转 us;
|
||
- 所有"字节数"用 int (Byte);
|
||
- dtype 一律归一化为小写字符串, 如 "bf16"/"fp16"/"fp8"/"fp32".
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import math
|
||
from dataclasses import dataclass, field, asdict
|
||
from typing import Optional
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# dtype 工具
|
||
# ---------------------------------------------------------------------------
|
||
|
||
DTYPE_BYTES = {
|
||
"fp32": 4,
|
||
"f32": 4,
|
||
"tf32": 4,
|
||
"fp16": 2,
|
||
"f16": 2,
|
||
"bf16": 2,
|
||
"fp8": 1,
|
||
"fp8_e4m3": 1,
|
||
"fp8_e5m2": 1,
|
||
"int8": 1,
|
||
"fp4": 0.5, # 半字节 (对齐 bmmv3, 2026-09-03)
|
||
"fp4_e2m1": 0.5,
|
||
}
|
||
|
||
# Cube 累加器 (L0C) 中元素字节数: 16bit 输入 -> fp32 累加; fp8/fp4 输入 -> fp32 累加
|
||
L0C_DTYPE_BYTES = 4
|
||
|
||
|
||
def dtype_bytes(dtype: str) -> float:
|
||
key = dtype.strip().lower()
|
||
if key not in DTYPE_BYTES:
|
||
raise ValueError(f"不支持的 dtype: {dtype!r}, 支持 {sorted(DTYPE_BYTES)}")
|
||
return DTYPE_BYTES[key]
|
||
|
||
|
||
def ceil_div(a: int, b: int) -> int:
|
||
return -(-a // b)
|
||
|
||
|
||
def align_up(x: int, align: int) -> int:
|
||
return ceil_div(x, align) * align
|
||
|
||
|
||
def align_down(x: int, align: int) -> int:
|
||
return (x // align) * align
|
||
|
||
|
||
def dvalue_contig_dims(case: "BmmCase", k_l1: float) -> tuple:
|
||
"""A/B 两侧 GM->L1 搬移的连续维 dValue (Byte) —— 依转置定连续维 (ND 排布).
|
||
|
||
A 不转置存储 [M,K]: 行内 K 连续, 切 K 后每行连续段 k_l1 -> k_l1*dt;
|
||
A 转置存储 [K,M]: M 连续 -> m*dt;
|
||
B 不转置存储 [K,N]: N 连续 -> n*dt;
|
||
B 转置存储 [N,K]: 行内 K 连续, 切 K 后连续段 k_l1 -> k_l1*dt.
|
||
|
||
供三处共用 (issue#19): l1_form 生成守卫 / IterBatch 条件 4 / constraints 校验,
|
||
保证"生成说可行、条件说可行、校验说可行"口径一致.
|
||
"""
|
||
dt = case.dtype_in_bytes
|
||
dv_a = case.m * dt if case.trans_a else k_l1 * dt
|
||
dv_b = k_l1 * dt if case.trans_b else case.n * dt
|
||
return dv_a, dv_b
|
||
|
||
|
||
def move_eff(tile_bytes: float, min_tile_size: int) -> float:
|
||
"""GM->L1 单命令搬移效率 (issue#36, 用户澄清口径).
|
||
|
||
单命令搬移效率由单块 tile = nValue × dValue × input_dtype 决定: tile 越大,
|
||
有效带宽越高; 达到 min_TileSize (16KB, 即进入条件 4 的效率下限常数) 饱和,
|
||
之下线性退化 —— 与"max(MK,KN)·dt >= min_TileSize 才保证搬移效率"的语义同源
|
||
(该条件只要求一侧达标, 另一侧小 tile 的低效由此函数量化)。
|
||
|
||
MergeBatch 合并 b0 个 batch 后单块 tile 放大 b0 倍 (A 侧: b0·M·k_l1·dt,
|
||
B 侧: b0·N·k_l1·dt; 堆叠方向视转置/排布而定 —— A ND 非转置时沿 M(nValue)
|
||
堆叠, B ND 非转置时沿 N(dValue) 堆叠, 乘积口径不变), 因此即便 T_cmd=0
|
||
合并仍有搬移效率收益。
|
||
"""
|
||
if tile_bytes <= 0:
|
||
return 1.0
|
||
return min(1.0, tile_bytes / min_tile_size)
|
||
|
||
|
||
def gm_move_time(v_a: float, v_b: float, tile_a: float, tile_b: float,
|
||
spec) -> float:
|
||
"""效率加权的 GM->L1 数据时延 (秒, 整芯片口径, issue#36).
|
||
|
||
t = (V_A/eff_A + V_B/eff_B) / BW_gm; 只影响**时间**列, GM 字节量不变
|
||
(仍按每输入字节恰读一次 = V_in 计, issue#31)。
|
||
"""
|
||
ea = move_eff(tile_a, spec.min_tile_size)
|
||
eb = move_eff(tile_b, spec.min_tile_size)
|
||
return (v_a / ea + v_b / eb) / spec.bw_gm
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Case 输入
|
||
# ---------------------------------------------------------------------------
|
||
|
||
@dataclass
|
||
class BmmCase:
|
||
"""一个 batch_mat_mul_v3 case 的输入描述.
|
||
|
||
对应算子接口:
|
||
A: [batch_a, M, K] (可带转置: trans_a=True 表示 [batch_a, K, M])
|
||
B: [batch_b, K, N] (可带转置: trans_b=True 表示 [batch_b, N, K])
|
||
bias: [B, 1, N] 可选
|
||
C: [batch_c, M, N], batch_c = broadcast(batch_a, batch_b)
|
||
"""
|
||
|
||
case_id: str = ""
|
||
batch_a: int = 1
|
||
batch_b: int = 1
|
||
m: int = 1
|
||
n: int = 1
|
||
k: int = 1
|
||
dtype_a: str = "bf16"
|
||
dtype_b: str = "bf16"
|
||
dtype_c: str = "bf16" # 输出 C 的 dtype (fp16/fp8 时 fixpipe 随路转换)
|
||
trans_a: bool = False
|
||
trans_b: bool = False
|
||
has_bias: bool = False
|
||
out_nd: bool = True # 输出是否 ND 格式 (StreamK 要求 ND)
|
||
deterministic_level: int = 0 # 确定性等级, >=2 禁用 StreamK
|
||
|
||
# ---- 派生属性 ----
|
||
def __post_init__(self):
|
||
"""输入合法性校验 (issue#15): 非法维度/负值静默产出伪方案, 必须明确报错."""
|
||
bad = []
|
||
for nm, v, lo, ok0 in (("batch_a", self.batch_a, 1, False),
|
||
("batch_b", self.batch_b, 1, False),
|
||
("m", self.m, 1, False),
|
||
("n", self.n, 1, False),
|
||
("k", self.k, 0, True)):
|
||
if not isinstance(v, int):
|
||
bad.append(f"{nm}={v!r} 非整数")
|
||
elif v < lo or (v == 0 and not ok0):
|
||
bad.append(f"{nm}={v} 非法 (需 >= {lo})")
|
||
if bad:
|
||
raise ValueError("case 维度非法: " + "; ".join(bad) +
|
||
" (m/n/batch 必须为正, k 可为 0)")
|
||
for nm, dt in (("dtype_a", self.dtype_a), ("dtype_b", self.dtype_b),
|
||
("dtype_c", self.dtype_c)):
|
||
if str(dt).strip().lower() not in DTYPE_BYTES:
|
||
raise ValueError(f"不支持的 dtype: {dt!r}, 支持 {sorted(DTYPE_BYTES)}")
|
||
|
||
@property
|
||
def batch_c(self) -> int:
|
||
return max(self.batch_a, self.batch_b)
|
||
|
||
@property
|
||
def dtype_in_bytes(self) -> float:
|
||
# A/B 输入元素字节数 (要求 A/B 同 dtype, 不一致时取较大者并在校验中报 warning)
|
||
# 返回 float 以兼容 fp4 (0.5B)
|
||
return max(dtype_bytes(self.dtype_a), dtype_bytes(self.dtype_b))
|
||
|
||
@property
|
||
def dtype_in(self) -> str:
|
||
"""输入侧"较慢"dtype (元素字节较大者) — Cube/AIV 速率取慢侧 (issue#28)."""
|
||
return (self.dtype_a if dtype_bytes(self.dtype_a) >= dtype_bytes(self.dtype_b)
|
||
else self.dtype_b)
|
||
|
||
@property
|
||
def dtype_out_bytes(self) -> float:
|
||
return dtype_bytes(self.dtype_c)
|
||
|
||
@property
|
||
def flops(self) -> float:
|
||
"""总计算量 (乘加各计一次)."""
|
||
return 2.0 * self.batch_c * self.m * self.n * self.k
|
||
|
||
@property
|
||
def input_bytes(self) -> float:
|
||
"""输入总数据量 (按广播前实际存储计)."""
|
||
dt = self.dtype_in_bytes
|
||
return self.batch_a * self.m * self.k * dt + self.batch_b * self.k * self.n * dt
|
||
|
||
@property
|
||
def output_bytes(self) -> float:
|
||
return self.batch_c * self.m * self.n * self.dtype_out_bytes
|
||
|
||
@property
|
||
def ai_full(self) -> float:
|
||
"""case 固有全量算存比 AI_full = 2MNK / (MK + KN + MN) (单 batch)."""
|
||
m, n, k = self.m, self.n, self.k
|
||
denom = m * k + k * n + m * n
|
||
return (2.0 * m * n * k / denom) if denom > 0 else 0.0
|
||
|
||
def to_row(self) -> dict:
|
||
d = asdict(self)
|
||
return d
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 实现方案结构体 (标准结构体定义)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
@dataclass
|
||
class ImplPlan:
|
||
"""BMM 实现方案 (理论分析输出 / 用户评估输入 共用的标准结构体).
|
||
|
||
字段分四组:
|
||
1) 分支与核间切分; 2) 核内 tiling; 3) 存储/Cache 策略; 4) 尾轮与流水策略.
|
||
|
||
说明: 本结构体对齐 batch_mat_mul_v3 tiling 的概念层级
|
||
(核间 grid 划分 -> singleCoreM/N/K -> L1 tile -> L0 tile),
|
||
字段名采用理论文档符号, 便于与文档公式直接对照.
|
||
"""
|
||
|
||
# --- 0) 基本信息 ---
|
||
case_id: str = ""
|
||
branch: str = "" # 转Matmul / 特殊分支 / MergeBatch / IterBatch /
|
||
# StreamK / ASW_Basic / ASW_Basic_降核
|
||
npu: str = "Ascend950PR"
|
||
op: str = "batch_mat_mul_v3"
|
||
|
||
# --- 1) 核间切分 (grid 级) ---
|
||
used_core_num: int = 0 # 实际使用 AIC 核数 (降核时 < C)
|
||
split_b: int = 1 # 核间 B 维切分数
|
||
m_cnt: int = 1 # 核间 M 维切分数
|
||
n_cnt: int = 1 # 核间 N 维切分数
|
||
grid_k: int = 1 # 核间 K 维切分数 (StreamK > 1)
|
||
core_map: str = "" # 核间分配策略, 如 "B->M->N线性映射+ASW滑窗蛇形(W=4)"
|
||
|
||
# --- 2) 核内 tiling ---
|
||
b_core: int = 0 # 每核 batch 数 (切 B 分支)
|
||
merge_b0: int = 1 # MergeBatch 合并数 (IterBatch/其他 = 1)
|
||
single_core_m: int = 0 # 每核输出 tile M
|
||
single_core_n: int = 0 # 每核输出 tile N
|
||
single_core_k: int = 0 # 每核 K 段长度 (核间切 K 时 < K)
|
||
k_l1: int = 0 # GM->L1 的 K 向粒度
|
||
b_l1: int = 1 # L1 内驻留 batch 数 (MergeBatch/IterBatch)
|
||
l1_form: str = "" # IterBatch L1 形态: a/b/c/d
|
||
base_m: int = 0 # L0 级 tile
|
||
base_n: int = 0
|
||
base_k: int = 0
|
||
|
||
# --- 3) 存储/Cache 策略 ---
|
||
l2_policy_in: str = "" # 输入 L2 策略: allocate(随路驻留) / non_allocate
|
||
l2_policy_out: str = "" # 输出: resident(驻留L2异步回写) / direct_gm(直写GM)
|
||
swizzle_w: int = 0 # ASW 滑窗宽度 (0 = 不用)
|
||
workspace_bytes: int = 0 # StreamK 中间结果 workspace (驻留 L2)
|
||
|
||
# --- 4) 尾轮与流水策略 ---
|
||
# 尾轮切分成员, 参考源码 MatMulV3TailInfo{mCnt,nCnt,kCnt,mTailMain,nTailMain}
|
||
# (mat_mul_v3/op_host/op_tiling/arch35/matmul_v3_common_advanced.h:120)
|
||
tail_strategy: str = "" # 尾轮策略: A0 / A1a / A1b / 方案B (仅切 M/N 类分支)
|
||
tail_m_cnt: int = 1 # 尾轮 M 向切分数 (tailInfo.mCnt)
|
||
tail_n_cnt: int = 1 # 尾轮 N 向切分数 (tailInfo.nCnt)
|
||
tail_k_cnt: int = 1 # 尾轮 K 向切分数 (tailInfo.kCnt, StreamK 时 = grid_K)
|
||
tail_m_main: int = 0 # 尾轮 M 向主体块数 (tailInfo.mTailMain)
|
||
tail_n_main: int = 0 # 尾轮 N 向主体块数 (tailInfo.nTailMain)
|
||
tail_block_cnt: int = 0 # 尾轮块数 r = N_blk mod C (0 = 无尾轮)
|
||
tail_wave_num: int = 0 # 总轮次 n_wave = ceil(N_blk / C)
|
||
fixpipe_unitflag: bool = True # fixpipe 开 unitflag 随路搬出
|
||
out_dtype_bytes: int = 2 # fixpipe 写出元素字节数 (C 矩阵 dtype)
|
||
|
||
# --- 备注 ---
|
||
note: str = ""
|
||
|
||
def to_row(self) -> dict:
|
||
return asdict(self)
|
||
|
||
@staticmethod
|
||
def csv_fields() -> list:
|
||
return list(ImplPlan.__dataclass_fields__.keys())
|
||
|
||
@staticmethod
|
||
def from_row(row: dict) -> "ImplPlan":
|
||
"""从 csv 行 (字符串字典) 恢复 ImplPlan."""
|
||
kw = {}
|
||
for name, f in ImplPlan.__dataclass_fields__.items():
|
||
if name not in row or row[name] is None or str(row[name]).strip() == "":
|
||
continue
|
||
v = row[name]
|
||
if f.type == "int":
|
||
kw[name] = int(float(v))
|
||
elif f.type == "bool":
|
||
kw[name] = str(v).strip().lower() in ("1", "true", "yes", "y")
|
||
else:
|
||
kw[name] = v
|
||
return ImplPlan(**kw)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 评估结果
|
||
# ---------------------------------------------------------------------------
|
||
|
||
@dataclass
|
||
class HardwareTiming:
|
||
"""各硬件流水级时延 (秒) 与数据量明细.
|
||
|
||
数据量列口径 (issue#29, 设计文档 docs/05 §4.4):
|
||
gm_read_bytes / l2_read_bytes / fixpipe_bytes / cube_flops = **整芯片**总量
|
||
(跨分支可比, GM 首读下限断言: gm_read_bytes >= case.input_bytes);
|
||
dma_cmd_count = **单核** GM->L1 DMA 命令数 (各核引擎并行, 墙钟 T_cmd =
|
||
单核命令数 x T_cmd), 与字节列口径不同属.
|
||
"""
|
||
|
||
# 搬入 (MTE2)
|
||
gm_read_bytes: float = 0.0 # GM->L1 直读数据量 (整芯片, 首读/落空重读)
|
||
l2_read_bytes: float = 0.0 # L2->L1 数据量 (整芯片, 驻留 L2 后重复读命中部分)
|
||
t_mte2_gm: float = 0.0 # GM->L1 时延 (按 GM 带宽, 不累加 L2->L1)
|
||
t_mte2_l2: float = 0.0 # L2->L1 时延 (按 L2 带宽)
|
||
t_mte2: float = 0.0 # 搬入合计 = t_mte2_gm + t_mte2_l2 (两者发生在不同数据上)
|
||
dma_cmd_count: float = 0.0 # 单核 GM->L1 DMA 命令次数 (T_cmd 分析用)
|
||
t_dma_cmd: float = 0.0 # DMA 命令固定开销合计 (墙钟)
|
||
|
||
# 计算 (Cube MMAD)
|
||
cube_flops: float = 0.0 # Cube 实际计算量 (整芯片, MergeBatch 含冗余)
|
||
t_mmad: float = 0.0
|
||
|
||
# 搬出 (Fixpipe)
|
||
fixpipe_bytes: float = 0.0 # 写出数据量 (整芯片, 按 C 矩阵 dtype / StreamK 临时矩阵按 4B)
|
||
t_fixpipe: float = 0.0
|
||
|
||
# 归约 (StreamK 专用)
|
||
t_reduce: float = 0.0
|
||
|
||
# 汇总
|
||
t_steady: float = 0.0 # 稳态流水时延 = max(各级)
|
||
t_drain: float = 0.0 # 流水排空暴露
|
||
t_total: float = 0.0 # 端到端时延
|
||
bottleneck: str = "" # 瓶颈级: MTE2_GM / MTE2_L2 / MMAD / FIXPIPE / REDUCE
|
||
|
||
def to_row(self, prefix: str = "") -> dict:
|
||
return {prefix + k: v for k, v in asdict(self).items()}
|
||
|
||
|
||
@dataclass
|
||
class EvalResult:
|
||
"""单个 case 的完整评估输出 (csv 一行的内容)."""
|
||
|
||
case: BmmCase = field(default_factory=BmmCase)
|
||
plan: ImplPlan = field(default_factory=ImplPlan)
|
||
timing: HardwareTiming = field(default_factory=HardwareTiming)
|
||
feasible: bool = True # 方案是否满足硬件约束
|
||
violations: str = "" # 违反的约束列表 (";" 分隔)
|
||
bound_type: str = "" # 计算Bound / 访存Bound / 写出Bound / 归约Bound
|
||
advice: str = "" # 瓶颈分析与优化建议
|
||
|
||
def to_row(self) -> dict:
|
||
row = {}
|
||
row.update(self.case.to_row())
|
||
row.update({"plan_" + k: v for k, v in self.plan.to_row().items()})
|
||
row.update(self.timing.to_row())
|
||
row.update({
|
||
"feasible": self.feasible,
|
||
"violations": self.violations,
|
||
"bound_type": self.bound_type,
|
||
"advice": self.advice,
|
||
})
|
||
return row
|