How to Write Custom Metal Kernels Using MLX's fast.h Module

MLX provides a just-in-time (JIT) compilation interface for custom GPU kernels via mx.fast.metal_kernel, which exposes Metal compute shaders directly to Python by wrapping the C++ API declared in mlx/fast.h and generating signatures automatically in mlx/backend/common/metal_kernel.cpp.

MLX (Machine Learning for Apple Silicon) enables developers to write custom GPU kernels in Apple Metal without leaving Python. The fast.h module provides the low-level bridge between MLX's array infrastructure and the Metal Performance Shaders runtime. This guide explains how to use the metal_kernel API to JIT-compile custom Metal code and execute it on GPU tensors.

C++ API Foundation in mlx/fast.h

The entry point for custom Metal kernels is declared in [mlx/fast.h](https://github.com/ml-explore/mlx/blob/main/mlx/fast.h#L71-L78). The metal_kernel function returns a CustomKernelFunction object that captures your kernel configuration for later execution.

MLX_API CustomKernelFunction metal_kernel(
    const std::string& name,
    const std::vector<std::string>& input_names,
    const std::vector<std::string>& output_names,
    const std::string& source,
    const std::string& header = "",
    bool ensure_row_contiguous = true,
    bool atomic_outputs = false);

Parameter breakdown:

  • name – Logical identifier for the kernel (used in caching and debugging).
  • input_names / output_names – Parameter names that appear in the generated Metal function signature.
  • source – The body of the Metal function (the code that appears after the generated signature).
  • header – Optional helper code (e.g., #includes or utility functions) prepended to the generated source.
  • ensure_row_contiguous – When true, MLX automatically rearranges inputs to be row-contiguous before kernel launch.
  • atomic_outputs – When true, output buffers are declared as device atomic<T> instead of device T*, enabling atomic operations for custom reductions.

Python Wrapper: mx.fast.metal_kernel

The C++ API is exposed to Python in [python/src/fast.cpp](https://github.com/ml-explore/mlx/blob/main/python/src/fast.cpp#L99-L112). This wrapper creates a callable Python object that forwards arguments to the underlying CustomKernelFunction.

When invoked, the callable accepts the following arguments:

outputs = kernel(
    inputs=[a, b],                    # List of mlx arrays

    template=[("T", mx.float32)],     # Template arguments for Metal types

    grid=(a.size, 1, 1),              # Total threads in each dimension

    threadgroup=(256, 1, 1),          # Threads per threadgroup

    output_shapes=[a.shape],          # Shape of each output

    output_dtypes=[a.dtype],          # Dtype of each output

    verbose=False,                    # Print generated Metal source

    stream=None,                      # Optional GPU stream

)

The template parameter binds compile-time constants (such as int, bool, or mlx.Dtype) to template parameters in your Metal source. For example, ("T", mx.float32) causes the backend to substitute T with float in the generated Metal code.

Backend Implementation and Code Generation

The actual compilation and launch logic resides in [mlx/backend/common/metal_kernel.cpp](https://github.com/ml-explore/mlx/blob/main/mlx/backend/common/metal_kernel.cpp). When you call the kernel object, MLX performs the following steps:

  1. Validates input and output counts (lines 66-78).
  2. Resolves a Metal GPU stream (lines 21-34).
  3. Generates a unique kernel name combining the logical name, input dtypes, and template arguments (lines 89-118).
  4. Builds the full Metal source by emitting a function signature (write_signature) that incorporates buffer bindings and constant-array handling (lines 51-74).
  5. Optionally prints the generated source when verbose=True (lines 42-47).
  6. Creates a CustomKernel object that compiles and launches the Metal kernel (lines 49-64).

The generated signature automatically handles buffer bindings for inputs and outputs, making the source parameter contain only the function body logic.

Key Concepts for Custom Metal Kernels

Template Arguments and Type Specialization

Template arguments allow you to write generic Metal code that works across multiple dtypes. In [metal_kernel.cpp](https://github.com/ml-explore/mlx/blob/main/mlx/backend/common/metal_kernel.cpp#L78-L98), the write_template function converts Python tuples like ("T", mx.float32) into Metal template syntax kernel_name<float>.

Grid vs. Threadgroup Configuration

The grid parameter specifies the total number of threads to dispatch (e.g., (a.size, 1, 1) for element-wise operations), while threadgroup specifies how threads are grouped into blocks (e.g., (256, 1, 1)). This matches Metal's dispatchThreads API and determines GPU occupancy.

Row Contiguity Guarantees

Setting ensure_row_contiguous=True (the default) forces MLX to transpose or copy non-contiguous arrays before kernel execution. This ensures that your Metal code can use simple linear indexing (thread_position_in_grid.x) without handling strided memory layouts.

Atomic Outputs for Reductions

Enable atomic_outputs=True when your kernel performs scatter operations or reductions. This changes the output buffer declaration from device T* to device atomic<T>, allowing use of Metal's atomic_fetch_add and similar functions.

Complete Working Example: Element-wise Exponential

This example defines a custom kernel that computes the exponential of each element using Metal's standard library functions.

import mlx as mx

# Define the Metal kernel body. The generated signature exposes

# an input buffer named 'inp' and an output buffer named 'out'.

source = """
    uint elem = thread_position_in_grid.x;  // each thread processes one element
    T value = inp[elem];                    // T is a template parameter (dtype)
    out[elem] = metal::exp(value);          // Metal's exponential function
"""

# Create the JIT-compiled kernel

kernel = mx.fast.metal_kernel(
    name="exp_elemwise",
    input_names=["inp"],
    output_names=["out"],
    source=source,
    header="",                # no extra headers needed

    ensure_row_contiguous=True,
    atomic_outputs=False,
)

# Allocate input array

a = mx.random.normal(shape=(4, 16)).astype(mx.float32)

# Launch the kernel

outputs = kernel(
    inputs=[a],
    template=[("T", mx.float32)],  # bind template parameter T to float

    grid=(a.size, 1, 1),           # one thread per element

    threadgroup=(256, 1, 1),       # 256 threads per threadgroup

    output_shapes=[a.shape],
    output_dtypes=[a.dtype],
    verbose=False,
)

# Retrieve result

b = outputs[0]

# Verify correctness

assert mx.allclose(b, mx.exp(a)).item()

How the execution works:

  1. metal_kernel instantiation – Creates a CustomKernelFunction that captures your source code and parameter names.
  2. Kernel invocation – The backend generates a complete Metal function by combining the signature (buffer declarations) with your source body.
  3. Template substitution – ("T", mx.float32) causes the generated Metal code to use float wherever T appears.
  4. Launch configuration – The grid of 64 threads (4×16) is divided into threadgroups of 256, executing in parallel on the GPU.

Summary

  • Entry point – Use mx.fast.metal_kernel (exposed from mlx/fast.h) to create JIT-compiled Metal kernels from Python.
  • Source composition – Provide only the function body in source; MLX automatically generates buffer bindings and function signatures in mlx/backend/common/metal_kernel.cpp.
  • Type specialization – Use the template parameter to bind generic types like T to specific Metal types (float, half, etc.).
  • Memory layout – Enable ensure_row_contiguous (default) for simple linear indexing, or disable it and handle strides manually for zero-copy operations.
  • Launch geometry – Specify grid in total threads and threadgroup in threads per group to control GPU parallelism.

Frequently Asked Questions

How does MLX generate the Metal function signature automatically?

MLX constructs the signature in [metal_kernel.cpp](https://github.com/ml-explore/mlx/blob/main/mlx/backend/common/metal_kernel.cpp#L51-L74) by analyzing the input_names and output_names you provide. It emits device const T* for inputs and device T* (or device atomic<T>) for outputs, then appends your source as the function body. This means you write only the algorithmic logic, not the boilerplate buffer declarations.

What is the difference between grid and threadgroup in MLX Metal kernels?

The grid parameter defines the total number of threads to dispatch across the entire computation (e.g., one thread per array element), while threadgroup defines how many threads execute together in a single block (e.g., 256 threads sharing a local cache). Metal uses this to schedule work across GPU cores. In the MLX source, these map directly to the dispatchThreads API in the Metal Performance Shaders framework.

When should I use atomic_outputs versus standard output buffers?

Enable atomic_outputs=True when your kernel performs scatter operations or custom reductions where multiple threads write to the same memory location. This declares outputs as device atomic<T> instead of device T*, allowing you to use Metal atomic functions like atomic_fetch_add_explicit. For standard element-wise operations where each thread writes to a unique location, leave atomic_outputs=False (the default) for better performance.

Can I use template arguments other than dtypes in custom Metal kernels?

Yes. The template parameter accepts tuples of (name, value) where value can be an int, bool, or mlx.Dtype. In [metal_kernel.cpp](https://github.com/ml-explore/mlx/blob/main/mlx/backend/common/metal_kernel.cpp#L78-L98), MLX's write_template function emits these as template arguments in the generated Metal source (e.g., template <typename T, int N>). This allows you to compile specialized versions of kernels with unrolled loops or fixed-size arrays optimized for specific sizes.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →