Skip to content

Expose precompiled_metal_kernel to load offline-compiled .metallib (-Ofast) for custom kernels #4541

Description

@dahai80

Motivation

mx.fast.metal_kernel compiles Metal source at runtime via MTLDevice::newLibrary(source, options). The runtime compiler does not expose the -O / -ffast-math optimization flags available to the offline xcrun metal compiler. For memory-latency-bound custom kernels (e.g. INT4 GEMV for LLM decode), this limits how aggressively the user's kernel can be optimized relative to MLX's own mlx.metallib (built offline via mlx_build_metallib).

I found that the C++ layer already has the scaffolding for precompiled kernels but it is not wired through:

  • CustomKernel in mlx/fast_primitives.h already has an is_precompiled_ field (L529)
  • metal_kernel() in mlx/backend/common/metal_kernel.cpp passes is_precompiled=false (L377)
  • CustomKernel::eval_gpu in mlx/backend/metal/custom_kernel.cpp silences the field with (void)is_precompiled_; (L17) and always calls d.get_library(name, compile_options, builder) — the JIT path
  • Device::get_library(name, path) in mlx/backend/metal/device.cpp (L631) already supports loading a .metallib from a path, and load_library (L251) handles colocated/path-based lookup

So the machinery is 90% present — is_precompiled_ appears to be a designed-but-unimplemented feature.

Proposal

Expose mx.fast.precompiled_metal_kernel(name, input_names, output_names, metallib_path, ...) that:

  1. Loads a user-provided .metallib (compiled offline with xcrun -sdk macosx metal -Ofast ...)
  2. Looks up the kernel function by the same custom_kernel_<name>_<types> naming convention the JIT path uses

This mirrors the existing mx.fast.precompiled_cuda_kernel (which already exists for the CUDA backend).

Measured impact (M5 Max, MLX 0.32)

INT4 GEMV kernel (same source as MLX's native qmv_fast_impl algorithm), compiled offline with -Ofast vs runtime JIT vs native mx.quantized_matmul. N-op loop (200 ops, 31 trials, paired A/B, median):

K native JIT -Ofast precompiled -Ofast vs native
20480 1.83ms 1.65ms 1.77ms -3.6%
28672 2.68ms 2.64ms 2.52ms -5.9%
32768 2.44ms 2.38ms 2.35ms -3.9%
40960 5.03ms 4.82ms 4.83ms -4.0%
  • -Ofast precompiled is 7.5% faster than the JIT version at K=20480 (fastMath + offline optimizer)
  • At K≥20480 it surpasses native mx.quantized_matmul by 3-6%, confirmed across two independent runs
  • Parity maintained (max diff 9e-4 in fp16, same as JIT)
  • At K<20480 (8K-16K) the op is launch-overhead-bound and the gap is within noise

The native mlx.metallib is compiled with -fno-fast-math (per cmake/extension.cmake L29). Allowing user kernels to opt into -Ofast (fastMath) gives them an advantage native doesn't have, for ops where the precision tradeoff is acceptable (int4 dequant accumulation is fp32 internally; only the final half cast is affected).

Files to change (rough)

  • mlx/backend/common/metal_kernel.cpp: add precompiled_metal_kernel() factory (mirror metal_kernel but pass is_precompiled=true + path as source_)
  • mlx/backend/metal/custom_kernel.cpp: when is_precompiled_, call d.get_library(lib_name, source_) (path-based) instead of the builder path
  • mlx/fast.h: declare precompiled_metal_kernel()
  • python/src/fast.cpp: bind mx.fast.precompiled_metal_kernel

Happy to open a PR if this approach is acceptable.

🤖 Generated with Claude Code

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions