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

> Learn to write custom Metal kernels with MLX fast.h. Explore a JIT compilation interface for direct Python access to Metal compute shaders and streamline your GPU programming.

- Repository: [ml-explore/mlx](https://github.com/ml-explore/mlx)
- Tags: how-to-guide
- Published: 2026-06-18

---

**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`](https://github.com/ml-explore/mlx/blob/main/mlx/fast.h) and generating signatures automatically in [`mlx/backend/common/metal_kernel.cpp`](https://github.com/ml-explore/mlx/blob/main/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`](https://github.com/ml-explore/mlx/blob/main/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)](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.

```cpp
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., `#include`s 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)](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:

```python
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)](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/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.

```python
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`](https://github.com/ml-explore/mlx/blob/main/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`](https://github.com/ml-explore/mlx/blob/main/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/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/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.