tune_gemm のチュートリアル

このチュートリアルでは、行列積のブロッキングパラメータをエンドツーエンドでチューニングします。再現可能な入力を生成し、制約で候補をフィルタリングし、それぞれを参照に照らして検証し、残ったものを runner で計測し、tune でスコア付けして選択し、結果をシリアライズします。

クイックスタート

moon add Luna-Flow/mare_mark@0.3.0
import {
  "Luna-Flow/mare_mark/tune_gemm",
  "Luna-Flow/mare_mark/model",
  "Luna-Flow/mare_mark/event",
  "Luna-Flow/mare_mark/runner",
  "Luna-Flow/mare_mark/tune",
  "moonbitlang/async",
}
test "multiply and check" {
  let shape = @tune_gemm.GemmScale::new(
    8, 8, 8, 1, 1, @tune_gemm.row_major(), @tune_gemm.row_major(), @tune_gemm.row_major(),
  )
  let a = @tune_gemm.reproducible_matrix(1UL, 8, 8, @tune_gemm.row_major())
  let b = @tune_gemm.reproducible_matrix(2UL, 8, 8, @tune_gemm.row_major())
  let candidate = @tune_gemm.GemmCandidate::new(
    4, 4, 4, 2, 2, @tune_gemm.pack_ab(), "m_n_k", "scalar",
    @tune_gemm.reusable_workspace(), @tune_gemm.compute_only(),
  )
  let c = @tune_gemm.execute_candidate(shape, a, b, candidate, 4096UL).unwrap()
  inspect(c.length(), content="64")
  inspect(c == @tune_gemm.gemm_reference(shape, a, b), content="true")
}

日常的な作業

候補グリッドをフィルタリングする

test "candidates that fit 64 KiB" {
  let all = @tune_gemm.enumerate_gemm_candidates()
  let fitting = all.filter(c => @tune_gemm.valid_candidate(c, 65536UL))
  inspect(all.length(), content="486")
  inspect(fitting.length(), content="189")
}

端数を含め、すべての候補を検証する

test "validate on boundary shapes" {
  let base = @tune_gemm.GemmScale::new(
    16, 12, 9, 1, 1, @tune_gemm.row_major(), @tune_gemm.column_major(), @tune_gemm.row_major(),
  )
  let candidates = @tune_gemm.enumerate_gemm_candidates().filter(c => c.mc == 32 && c.nc == 32 && c.kc == 32)
  let all_correct = @tune_gemm.boundary_shapes(base).all(shape => {
    let a = @tune_gemm.reproducible_matrix(11UL, shape.m, shape.k, shape.layout_a)
    let b = @tune_gemm.reproducible_matrix(12UL, shape.k, shape.n, shape.layout_b)
    candidates.all(c => @tune_gemm.gemm_is_correct(shape, a, b, c, 0.0, 262144UL))
  })
  inspect(candidates.length(), content="27")
  inspect(all_correct, content="true")
}

組み込みのカーネルは参照と同じ順序で総和を取るため、許容誤差 0.0 が正しい値です。順序を変えるカーネルの許容誤差は 設計ページで導いています。

計測と選択

各候補をランナーの実装に包み、ブロックを共有するよう 1 つのケースで計測し、確認的な中央値に基づいて選択します:

async test "tune two candidates" {
  let shape = @tune_gemm.GemmScale::new(
    24, 24, 24, 1, 1, @tune_gemm.row_major(), @tune_gemm.row_major(), @tune_gemm.row_major(),
  )
  let a = @tune_gemm.reproducible_matrix(1UL, 24, 24, @tune_gemm.row_major())
  let b = @tune_gemm.reproducible_matrix(2UL, 24, 24, @tune_gemm.row_major())
  let candidates = [
    @tune_gemm.GemmCandidate::new(8, 8, 8, 2, 2, @tune_gemm.pack_ab(), "m_n_k", "scalar", @tune_gemm.reusable_workspace(), @tune_gemm.compute_only()),
    @tune_gemm.GemmCandidate::new(32, 32, 32, 4, 4, @tune_gemm.pack_ab(), "m_n_k", "scalar", @tune_gemm.reusable_workspace(), @tune_gemm.compute_only()),
  ]
  let implementations = candidates.map(candidate => {
    @runner.Implementation::stateless(@tune_gemm.candidate_id(candidate), "1", (problem : @tune_gemm.GemmProblem) => {
      @model.OperationResult::completed(@tune_gemm.gemm_blocked(problem.scale, problem.a, problem.b, candidate), ())
    })
  })
  let plan = @runner.single_step("gemm-24", [shape])
    .with_immutable_input(context => @tune_gemm.GemmProblem::from_inputs(context.dataset_key.scale, a, b), _ => "gemm-24")
    .compare(implementations)
    .against_equal(problem => @tune_gemm.gemm_reference(problem.scale, problem.a, problem.b), (expected, actual) => {
      @tune_gemm.validate_gemm(expected, actual, 0.0).valid
    })
    .compile()
    .unwrap()
  let memory = @event.InMemorySink::new()
  let environment = @model.EnvironmentSnapshot::new(
    @model.SemanticEnvironment::new(@model.ExecutionTarget::Native, "moonc", "", "f64"),
    @model.PerformanceEnvironment::new("native", "cpu", "default", 1, "monotonic"),
    @model.ProvenanceEnvironment::new("os", "host", "now", "HEAD", "gemm-tuning"),
  )
  let summary = @runner.run(
    plan,
    @runner.RunContext::new(environment, memory.as_sink(), 3UL, @runner.ProtocolPreset::QuickCheck.validated()),
  )
  inspect(summary.passed_count, content="2")
  let scores = candidates.map(candidate => {
    let id = @tune_gemm.candidate_id(candidate)
    let samples = memory.observations
      .filter(o => o.implementation_id == id && o.valid && o.phase is Confirmatory)
      .map(o => o.raw_elapsed_us)
    @tune.score_samples(id, samples, @tune_gemm.workspace_bytes(candidate).to_double())
  })
  let best = @tune.select_best(scores, 2.0, true)
  inspect(best is Some(_), content="true")
}

どの候補が勝つかはマシンに依存しますが、手順は依存しません。

設定を記録する

test "serialize the tuning input" {
  let shapes = [
    @tune_gemm.GemmScale::new(256, 256, 256, 1, 1, @tune_gemm.row_major(), @tune_gemm.row_major(), @tune_gemm.row_major()),
  ]
  let config = @tune_gemm.GemmTuningConfig::new(
    2026UL, shapes, @tune_gemm.enumerate_gemm_candidates(), 262144UL, 3, 20,
  )
  let json = @tune_gemm.config_json(config)
  inspect(json.contains("\"id\":\"32x32x32:2x2:scalar:m_n_k:ab:workspace:compute\""), content="true")
}

JSON、計測の JSONL、GemmEnvironment を一緒に保存してください。チューニングした選択は、environment_compatible が成り立つ場合にのみ再利用してください。

さらに先へ

  • 1 操作あたりの µs は 2mnk/(t⋅103)2mnk/(t \cdot 10^3) で GFLOP/s に換算します。
  • 形状をホールドアウトします。一部の形状で選択し、boundary_shapes とより大きな形状で確認してください。tune のチュートリアルを参照してください。
  • グリッド全体では遅すぎる場合は、seeded_order を使って 486 個の候補のランダムな部分集合を計測してください。

よくある落とし穴

  • 異なる計時範囲を比較する。 範囲が ID の一部になっているのには理由があります。
  • 端数を忘れる。 64 × 64 で正しいカーネルが 63 × 64 では誤っていることがあります。
  • 順序を変えるカーネルで等価性を検査する。 2γkk2\gamma_k k の上界を使ってください。
  • reproducible_matrix の 1 回の呼び出しからバッチ化されたオペランドを作る。 作られるのは 1 つの行列です。バッチごとに連結してください。

次のステップ