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– Whentrue, MLX automatically rearranges inputs to be row-contiguous before kernel launch.atomic_outputs– Whentrue, output buffers are declared asdevice atomic<T>instead ofdevice 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:
- Validates input and output counts (lines 66-78).
- Resolves a Metal GPU stream (lines 21-34).
- Generates a unique kernel name combining the logical name, input dtypes, and template arguments (lines 89-118).
- Builds the full Metal source by emitting a function signature (
write_signature) that incorporates buffer bindings and constant-array handling (lines 51-74). - Optionally prints the generated source when
verbose=True(lines 42-47). - Creates a
CustomKernelobject 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:
metal_kernelinstantiation – Creates aCustomKernelFunctionthat captures your source code and parameter names.- Kernel invocation – The backend generates a complete Metal function by combining the signature (buffer declarations) with your
sourcebody. - Template substitution –
("T", mx.float32)causes the generated Metal code to usefloatwhereverTappears. - 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 frommlx/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 inmlx/backend/common/metal_kernel.cpp. - Type specialization – Use the
templateparameter to bind generic types likeTto 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
gridin total threads andthreadgroupin 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →