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 对内置内核是正确的,因为它与参考实现以相同的顺序求和;设计页面推导了重排序内核的容差。
测量与选择
把每个候选包装成一个运行器实现,在同一个用例中测量它们使其共享区组,并基于验证性中位数进行选择:
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 成立的地方复用调优得到的选择。
更进一步
- 用 把每次操作的 µs 换算为 GFLOP/s。
- 留出部分形状:在一部分形状上选择,然后在
boundary_shapes和更大的形状上确认;参见 tune 教程。 - 当完整网格太慢时,用
seeded_order测量 486 个候选中的一个随机子集。
常见陷阱
- 比较不同的计时范围。 计时范围成为 id 的一部分是有原因的。
- 忘记尾部。 对 64 × 64 正确的内核,对 63 × 64 可能是错误的。
- 对重排序内核使用相等性检查。 请使用 界。
- 用一次
reproducible_matrix调用生成批量操作数。 它只生成一个矩阵;请按批次拼接。