metalworkingGitHub
/mlx/how-an-op-becomes-a-kernel

How an Op Becomes a Kernel

Between mx.matmul(a, b) and a simdgroup_matrix instruction sits one C++ hop: the primitive's eval_gpu, which picks a kernel variant, picks tile sizes, binds function constants, and encodes the dispatch.

CUDA equivalent: the dispatcher layer of PyTorch → cuBLAS heuristics, except open, small, and readable. Knowing where these decisions live is what separates "my model is slow" from "my shapes fall off the fast path," and the war stories' question 2 (unlock an existing fast path) is answered by reading exactly these files.

The chain for a matmul, concretely: the lazy graph hands Matmul to the schedulerMatmul::eval_gpu in mlx/backend/metal/matmul.cpp normalizes strides/transposes, then picks the steel template parameters from device class and problem shape:

#define GEMM_TPARAM_MACRO(devc)                                           \
  if (devc == 'g' || devc == 'p') { /* Small device */                    \
    ...
  } else if (devc == 'd') { /* Large device */                            \
    if ((size_t)batch_size_out * M * N >= 1ul << 20) { /* large matmul */ \
      if (out.dtype() != float32) { /* half and bfloat */                 \
        if (2 * std::max(M, N) > K) { /* Reasonable K */                  \
          bm = 64; bn = 64; bk = 16; wm = 1; wn = 2;                      \
        } else if (!transpose_a && transpose_b) { /* nt with large k */   \
          bm = 64; bn = 32; bk = 32; wm = 2; wn = 2;                      \

mlx/backend/metal/matmul.cpp:89-124, abridged and reformatted

Read what that macro is: a hand-tuned lookup from (chip class, size, dtype, transpose pattern) to the tile shape. cuBLAS's secret heuristics, in greppable form. The kernel name is then assembled as a string (steel_gemm_fused_...bm64_bn64...), the pipeline is fetched or JIT-built, alignment function constants are bound, and the dispatch is encoded.

mx.matmul(a, b) lazy graph → Matmul::eval_gpu weights quantized? → QMV / QMM family instead matrix × vector shape? → gemv kernels, not GEMM very skinny / huge K? → split-K and non-steel paths GEMM_TPARAM_MACRO chip class × size × dtype × transpose bm bn bk · wm wn e.g. 64·64·16 · 1·2 "steel_gemm_fused_...bm64_bn64..." kernel name assembled as a string pipeline cache hit → reuse · miss → JIT build with function constants encode dispatch into the shared command buffer falling off the fast path is silent: MTPLX's 2.24× came from noticing M=3-6 fell between the gemv tuning and the GEMM tiles sdpa has its own gate like this (head-dim lists, sequence thresholds) cuBLAS's secret heuristics, in greppable form: mlx/backend/metal/matmul.cpp

The route a matmul takes. The amber branches are where shapes silently leave the steel fast path; the green spine is the path this glossary's case studies read.

Branches before steel, because falling into the wrong one is a classic silent slowdown: matrix-vector shapes route to gemv kernels rather than GEMM; very skinny/small cases have split-K and non-steel paths; quantized weights go to an entirely different kernel family whose fast path is shape-sensitive in its own ways; and mx.fast.scaled_dot_product_attention has its own dispatch gate deciding fused-vs-fallback. The MTPLX war story, 2.24× from ~10 lines, came from noticing that multi-token decode (M=3-6) fell between the matrix-vector path and the tile sizes this macro assumes.