tune_gemm API

Luna-Flow/mare_mark/tune_gemm 是一个完整的调优领域示例:双精度批量矩阵乘法 C=ABC = AB,具有可配置的布局、以缓存块和寄存器块大小为参数的分块标量实现、用于验证的参考实现、确定性的候选网格、工作区统计、可复现的输入矩阵,以及可序列化的调优配置。参见 tune_gemm 设计。

源码:src/tune_gemm/gemm.mbt、src/tune_gemm/versioning.mbt。

import {
  "Luna-Flow/mare_mark/tune_gemm",
}

问题描述

GemmScale

GemmScale 描述一种问题形状。

pub struct GemmScale {
  m : Int
  n : Int
  k : Int
  batch : Int
  reuse_count : Int
  layout_a : Layout
  layout_b : Layout
  layout_c : Layout
}
pub fn GemmScale::new(Int, Int, Int, Int, Int, Layout, Layout, Layout) -> Self

AA 为 m×km \times k,BB 为 k×nk \times n,CC 为 m×nm \times n,batch 个相互独立的乘积依次存放。reuse_count 记录操作数被复用的次数(用于摊销的计时范围);内核不会读取它。

Layout, row_major, column_major

Layout 是一个矩阵的存储顺序。

pub(all) enum Layout {
  RowMajor
  ColumnMajor
}
pub fn row_major() -> Layout
pub fn column_major() -> Layout

r×cr \times c 矩阵第 bb 个批次的元素 (i,j)(i, j) 位于 b rc+i c+jb\,rc + i\,c + j(行主序)或 b rc+j r+ib\,rc + j\,r + i(列主序)。

matrix_length

matrix_length 返回操作数 "a"、"b" 或 "c" 的元素个数。

pub fn matrix_length(GemmScale, String) -> Int

它为 batch * m * k、batch * k * n 或 batch * m * n;对未知操作数或非正的维度或批次数为 0。

GemmProblem

GemmProblem 把形状与其操作数捆绑在一起。

pub struct GemmProblem {
  scale : GemmScale
  a : Array[Double]
  b : Array[Double]
  c : Array[Double]
}
pub fn GemmProblem::new(GemmScale, Array[Double], Array[Double], Array[Double]) -> Self
pub fn GemmProblem::from_inputs(GemmScale, Array[Double], Array[Double]) -> Self

from_inputs 分配一个含 batch * m * n 个元素的全零 c。数组按原样使用;不进行转置。

boundary_shapes

boundary_shapes 返回 m−1m - 1 的形状、形状本身以及 m+1m + 1 的形状。

pub fn boundary_shapes(GemmScale) -> Array[GemmScale]

恰好位于块大小附近的形状会考验内核的尾部处理。对于 m=0m = 0,第一个形状同样是 m=0m = 0。

候选

GemmCandidate

GemmCandidate 是一种内核配置。

pub struct GemmCandidate {
  mc : Int
  nc : Int
  kc : Int
  mr : Int
  nr : Int
  packing : Packing
  loop_order : String
  microkernel_id : String
  allocation_mode : AllocationMode
  timing_scope : TimingScope
}
pub fn GemmCandidate::new(Int, Int, Int, Int, Int, Packing, String, String, AllocationMode, TimingScope) -> Self

mc、nc、kc 是缓存块大小,mr、nr 是寄存器块大小。要使候选有效,loop_order 必须为 "m_n_k",且 microkernel_id 必须是 supported_microkernels() 之一。

Packing, pack_ab

Packing 表示在分块循环之前哪些操作数会被复制成行主序形式。

pub(all) enum Packing {
  None
  A
  B
  AB
}
pub fn pack_ab() -> Packing

AllocationMode, reusable_workspace

AllocationMode 表示缓冲区如何获得。

pub(all) enum AllocationMode {
  FreshPerOperation
  ReuseOutput
  ReusableWorkspace
  PrepackedA
  PrepackedB
  PrepackedAB
}
pub fn reusable_workspace() -> AllocationMode

预打包模式要求相应的打包方式(参见 valid_candidate)。内置内核在每种模式下都分配新的缓冲区;该字段是为测试框架和报告描述实验用的。

TimingScope, compute_only

TimingScope 表示对候选的一次测量包含哪些内容。

pub(all) enum TimingScope {
  ComputeOnly
  EndToEnd
  Amortized(Int)
  SteadyState
}
pub fn compute_only() -> TimingScope

测量时把它映射到夹具的 SetupPolicy:ComputeOnly 不含打包和分配,EndToEnd 包含它们,Amortized(r) 把准备开销分摊到 r 次使用上,SteadyState 测量在预热过的缓冲区上的重复调用。

candidate_id

candidate_id 返回候选的稳定文本 id。

pub fn candidate_id(GemmCandidate) -> String

格式为 mc x nc x kc : mr x nr : microkernel : loop_order : packing : allocation : timing,不含空格。

test "candidate ids" {
  let candidate = @tune_gemm.GemmCandidate::new(
    64, 64, 32, 4, 4, @tune_gemm.pack_ab(), "m_n_k", "scalar",
    @tune_gemm.reusable_workspace(), @tune_gemm.compute_only(),
  )
  inspect(@tune_gemm.candidate_id(candidate), content="64x64x32:4x4:scalar:m_n_k:ab:workspace:compute")
}

supported_microkernels

pub fn supported_microkernels() -> Array[String]

返回 ["scalar", "scalar_f64", "auto"]。三者运行同一个标量循环;这些 id 的存在是为了让测试框架能为内核变体命名。

enumerate_gemm_candidates

enumerate_gemm_candidates 返回默认的候选网格。

pub fn enumerate_gemm_candidates() -> Array[GemmCandidate]

网格为 mc,nc∈{32,64,128}mc, nc \in \{32, 64, 128\}、kc∈{32,64}kc \in \{32, 64\}、mr,nr∈{2,4,8}mr, nr \in \{2, 4, 8\} 以及三个微内核 id,采用 AB 打包、m_n_k 顺序、ReusableWorkspace 和 ComputeOnly:按固定顺序共 3⋅3⋅2⋅3⋅3⋅3=4863 \cdot 3 \cdot 2 \cdot 3 \cdot 3 \cdot 3 = 486 个候选。

workspace_bytes

workspace_bytes 返回候选的打包工作区大小。

pub fn workspace_bytes(GemmCandidate) -> UInt64

它为 8 (mc⋅kc+kc⋅nc+mc⋅nc)8\,(mc \cdot kc + kc \cdot nc + mc \cdot nc) 字节:以双精度数表示的 AA 的一个 mc×kcmc \times kc 块、BB 的一个 kc×nckc \times nc 面板和 CC 的一个 mc×ncmc \times nc 块;若某个块大小不为正,则为 0。

valid_candidate

valid_candidate 依据内核约束和工作区上限检查候选。

pub fn valid_candidate(GemmCandidate, UInt64) -> Bool

需同时满足:块大小为正;mr≤mcmr \le mc 且 nr≤ncnr \le nc;循环顺序为 m_n_k;受支持的微内核 id;Amortized(r) 满足 r>0r > 0;预打包分配只与匹配的打包方式搭配(PrepackedA 需要 A 或 AB,PrepackedB 需要 B 或 AB,PrepackedAB 需要 AB);且 workspace_bytes 不超过上限。

candidate_valid_for_shape

pub fn candidate_valid_for_shape(GemmCandidate, GemmScale, UInt64) -> Bool

valid_candidate,并且 mm、nn、kk、batch 和 reuse_count 均为正。

test "constraints" {
  let candidate = @tune_gemm.GemmCandidate::new(
    128, 128, 64, 8, 8, @tune_gemm.pack_ab(), "m_n_k", "auto",
    @tune_gemm.reusable_workspace(), @tune_gemm.compute_only(),
  )
  inspect(@tune_gemm.workspace_bytes(candidate), content="262144")
  inspect(@tune_gemm.valid_candidate(candidate, 262144UL), content="true")
  inspect(@tune_gemm.valid_candidate(candidate, 262143UL), content="false")
  inspect(@tune_gemm.enumerate_gemm_candidates().length(), content="486")
}

内核

gemm_reference

gemm_reference 用教科书式的三重循环计算 C=ABC = AB。

pub fn gemm_reference(GemmScale, Array[Double], Array[Double]) -> Array[Double]

每个元素从 0.00.0 开始按 kk 递增的顺序累加,并以 layout_c 存储。对非正的维度或长度错误的操作数返回 []。

gemm_blocked

gemm_blocked 用缓存分块和寄存器分块计算 C=ABC = AB。

pub fn gemm_blocked(GemmScale, Array[Double], Array[Double], GemmCandidate) -> Array[Double]

循环先遍历 mc×nc×kcmc \times nc \times kc 块,再遍历 mr×nrmr \times nr 寄存器分片,边缘处为不完整的分片。操作数按 packing 打包。对无效的维度、块大小或操作数长度返回 [];它不检查其他候选约束。

pack_operand

pack_operand 把操作数 "a" 或 "b" 复制为行主序。

pub fn pack_operand(GemmScale, Array[Double], String) -> Array[Double]

结果长度相同;无论源布局如何,每个批次都逐行存储。对未知操作数或无效输入返回 []。

execute_candidate

execute_candidate 在检查候选之后运行它。

pub fn execute_candidate(GemmScale, Array[Double], Array[Double], GemmCandidate, UInt64) -> Array[Double]?

当 candidate_valid_for_shape 失败或某个操作数长度错误时返回 None,否则返回 Some(gemm_blocked(...))。

验证

GemmValidation

pub struct GemmValidation {
  valid : Bool
  mismatches : Int
  max_abs_error : Double
}

validate_gemm

validate_gemm 以绝对容差逐元素比较两个结果数组。

pub fn validate_gemm(Array[Double], Array[Double], Double) -> GemmValidation

不匹配包括:超出较短长度的每个元素、含有非有限值的每一对,以及 ∣e−a∣\lvert e - a\rvert 超过容差的每一对。max_abs_error 是最大的有限差值。NaN、无穷大或负的容差会得到 valid = false。

gemm_is_correct

gemm_is_correct 执行候选,并用 gemm_reference 验证它。

pub fn gemm_is_correct(GemmScale, Array[Double], Array[Double], GemmCandidate, Double, UInt64) -> Bool

参数:形状、AA、BB、候选、容差、工作区上限。

test "blocked equals reference" {
  let shape = @tune_gemm.GemmScale::new(
    5, 7, 3, 1, 1, @tune_gemm.row_major(), @tune_gemm.column_major(), @tune_gemm.row_major(),
  )
  let a = @tune_gemm.reproducible_matrix(1UL, 5, 3, @tune_gemm.row_major())
  let b = @tune_gemm.reproducible_matrix(2UL, 3, 7, @tune_gemm.column_major())
  let candidate = @tune_gemm.GemmCandidate::new(
    4, 4, 2, 2, 2, @tune_gemm.pack_ab(), "m_n_k", "scalar",
    @tune_gemm.reusable_workspace(), @tune_gemm.compute_only(),
  )
  inspect(@tune_gemm.gemm_is_correct(shape, a, b, candidate, 0.0, 4096UL), content="true")
  let check = @tune_gemm.validate_gemm([1.0, 2.0], [1.0, 2.5], 0.1)
  inspect(check.mismatches, content="1")
  inspect(check.max_abs_error, content="0.5")
}

输入与配置

reproducible_matrix

reproducible_matrix 生成元素位于 [−1,1][-1, 1] 的确定性矩阵。

pub fn reproducible_matrix(UInt64, Int, Int, Layout) -> Array[Double]

参数:种子、行数、列数、布局。元素是 0.0010.001 的倍数,由 64 位线性同余生成器产生,并以请求的布局存储,因此相同的种子在两种布局下给出相同的逻辑矩阵。对非正的维度返回 []。它只生成一个矩阵;当 batch > 1 时请拼接多个矩阵。

test "the same logical matrix in both layouts" {
  let rows = @tune_gemm.reproducible_matrix(5UL, 2, 3, @tune_gemm.row_major())
  let cols = @tune_gemm.reproducible_matrix(5UL, 2, 3, @tune_gemm.column_major())
  inspect(rows[1 * 3 + 2] == cols[2 * 2 + 1], content="true")
  inspect(rows.all(x => x >= -1.0 && x <= 1.0), content="true")
}

GemmTuningConfig

GemmTuningConfig 是一次调优运行完整且可序列化的输入。

pub struct GemmTuningConfig {
  schema_version : TuningSchemaVersion
  seed : UInt64
  shapes : Array[GemmScale]
  candidates : Array[GemmCandidate]
  max_workspace_bytes : UInt64
  exploration_samples : Int
  confirmation_samples : Int
}
pub fn GemmTuningConfig::new(UInt64, Array[GemmScale], Array[GemmCandidate], UInt64, Int, Int) -> Self

new 把 schema_version 设为 TuningSchemaVersion::V1。

config_json

config_json 序列化调优配置。

pub fn config_json(GemmTuningConfig) -> String

该 JSON 包含 schema_version("mmkts_1")、以十进制字符串表示的 seed 和 max_workspace_bytes(64 位值无法精确放入 JSON 数值)、shapes、candidates(每项带有其 id)以及两个样本数。

test "configuration JSON" {
  let shape = @tune_gemm.GemmScale::new(
    64, 64, 64, 1, 1, @tune_gemm.row_major(), @tune_gemm.row_major(), @tune_gemm.row_major(),
  )
  let config = @tune_gemm.GemmTuningConfig::new(42UL, [shape], [], 262144UL, 3, 10)
  let json = @tune_gemm.config_json(config)
  inspect(json.contains("\"schema_version\":\"mmkts_1\""), content="true")
  inspect(json.contains("\"seed\":\"42\""), content="true")
}

TuningSchemaVersion

pub(all) enum TuningSchemaVersion {
  V1
}
pub fn TuningSchemaVersion::identifier(Self) -> String
pub fn TuningSchemaVersion::implementation(Self) -> String
pub fn TuningSchemaVersion::lifecycle(Self) -> @model.VersionLifecycle
pub fn TuningSchemaVersion::version(Self) -> Int

identifier() 为 "mmkts_1";V1 为 Supported。

GemmEnvironment, environment_compatible

GemmEnvironment 是用于调优结果的扁平环境记录。

pub struct GemmEnvironment {
  target : String
  toolchain : String
  compiler_flags : String
  dtype_abi : String
  runtime : String
  cpu : String
  gc : String
  concurrency : Int
  clock : String
}
pub fn GemmEnvironment::new(String, String, String, String, String, String, String, Int, String) -> Self
pub fn environment_compatible(GemmEnvironment, GemmEnvironment) -> Bool

当九个字段全部相等时,environment_compatible 为真。调优得到的配置只能在兼容的环境中复用。