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 は 1 つの問題の形状を記述します。

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 は 1 つの行列の格納順序です。

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 は 1 つのカーネル構成です。

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"] を返します。3 つとも同じスカラーループを実行します。これらの 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\} と 3 つのマイクロカーネル 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 は 2 つの結果の配列を要素ごとに絶対許容誤差で比較します。

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]

引数: シード、行数、列数、レイアウト。要素は 64 ビットの線形合同法生成器から引いた 0.0010.001 の倍数で、要求されたレイアウトで格納されるため、同じシードからはどちらのレイアウトでも同じ論理的な行列が得られます。次元が正でない場合は [] を返します。生成するのは 1 つの行列なので、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")、10 進文字列としての seed と max_workspace_bytes(64 ビット値は JSON の数値に正確には収まらないため)、shapes、candidates(それぞれ id 付き)、および 2 つのサンプル数を持ちます。

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 は 9 つのフィールドがすべて等しいときに真です。チューニング済みの設定は、互換性のある環境でのみ再利用できます。