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 成立的地方复用调优得到的选择。

更进一步

  • 用 2mnk/(t⋅103)2mnk/(t \cdot 10^3) 把每次操作的 µs 换算为 GFLOP/s。
  • 留出部分形状:在一部分形状上选择,然后在 boundary_shapes 和更大的形状上确认;参见 tune 教程。
  • 当完整网格太慢时,用 seeded_order 测量 486 个候选中的一个随机子集。

常见陷阱

  • 比较不同的计时范围。 计时范围成为 id 的一部分是有原因的。
  • 忘记尾部。 对 64 × 64 正确的内核,对 63 × 64 可能是错误的。
  • 对重排序内核使用相等性检查。 请使用 2γkk2\gamma_k k 界。
  • 用一次 reproducible_matrix 调用生成批量操作数。 它只生成一个矩阵;请按批次拼接。

后续步骤