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 は で GFLOP/s に換算します。
- 形状をホールドアウトします。一部の形状で選択し、
boundary_shapesとより大きな形状で確認してください。tune のチュートリアルを参照してください。 - グリッド全体では遅すぎる場合は、
seeded_orderを使って 486 個の候補のランダムな部分集合を計測してください。
よくある落とし穴
- 異なる計時範囲を比較する。 範囲が ID の一部になっているのには理由があります。
- 端数を忘れる。 64 × 64 で正しいカーネルが 63 × 64 では誤っていることがあります。
- 順序を変えるカーネルで等価性を検査する。 の上界を使ってください。
reproducible_matrixの 1 回の呼び出しからバッチ化されたオペランドを作る。 作られるのは 1 つの行列です。バッチごとに連結してください。