# How MLX's Compile Function Works and When to Use Compile Mode

> Discover how MLX's compile function optimizes Python code into efficient kernels. Learn when to use CompileMode for performance gains and memory savings in your ML projects.

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

---

**MLX's `compile` function transforms pure-Python operations into fused, statically-typed kernels that eliminate Python overhead and intermediate memory allocations, with `shapeless` mode enabling single-kernel reuse across varying input shapes.**

MLX's just-in-time (JIT) compilation system, available in the `ml-explore/mlx` repository, provides a `compile` function that converts regular Python functions into optimized compute kernels. This compilation process traces the function's primitive operations, fuses them into unified kernels, and caches the results for repeated execution. Understanding when to use standard compilation versus shapeless mode—or when to disable compilation entirely—is essential for maximizing performance in machine learning workflows.

## How MLX's Compile Function Works

The compilation pipeline in [`python/mlx/extension.py`](https://github.com/ml-explore/mlx/blob/main/python/mlx/extension.py) follows a six-stage process that transforms Python code into platform-specific kernels for Metal, CUDA, or CPU backends.

### Function Registration and Tracing

When you call `mx.compile(f)` or apply the `@mx.compile` decorator, MLX registers the function and returns a new callable that intercepts subsequent invocations. According to the test suite in [`python/tests/test_compile.py`](https://github.com/ml-explore/mlx/blob/main/python/tests/test_compile.py) (lines 21-23), this registration creates a compiled wrapper that forwards calls to the JIT engine. On the first invocation, MLX traces the function by recording every primitive operation—such as `mx.add` or `mx.exp`—to build a static data flow graph.

### Graph Lowering and Kernel Fusion

The traced graph undergoes lowering to the backend target, where MLX fuses compatible primitives into a single kernel. This fusion eliminates intermediate array allocations and reduces kernel launch overhead. The `test_enable_disable` test case (lines 38-50) verifies that compiled execution produces fewer primitives than uncompiled execution (`n_compiled < n_uncompiled`), confirming the fusion optimization. The compiler also handles edge cases like non-finite literals (`nan`, `inf`) by special-casing them to avoid invalid identifiers in generated source code, as verified in `test_compile_nonfinite_constants`.

### Caching and Execution

Compiled kernels are cached based on shape-type signatures, unless shapeless mode is enabled. When a compiled function is invoked, the kernel executes directly on the device with minimal Python overhead. The `test_compile_donates_input_buffer` test (lines 14-17) demonstrates that outputs can reuse input buffers to minimize memory copies, returning MLX `array` objects without unnecessary data duplication.

## When to Use Compile Mode

Compile mode encompasses both the standard JIT compilation via `mx.compile` and configuration options including the `shapeless` parameter and global enable/disable flags. Choose your configuration based on workload characteristics.

### Standard Compilation for Repeated Execution

Use the default `mx.compile` decorator for functions called repeatedly in training loops or inference pipelines. The one-time tracing cost amortizes over many executions, while kernel fusion provides significant speedups for computational graphs with dozens or hundreds of primitives.

### Shapeless Mode for Dynamic Shapes

Enable `shapeless=True` when your function must accept varying input dimensions without triggering recompilation. This mode skips shape-based caching, allowing a single kernel to handle arbitrary compatible shapes. Avoid this when the function body contains shape-dependent control flow that would produce different graphs.

```python
@mx.compile(shapeless=True)
def scale(x, factor=2):
    return x * factor

# Works for any shape with same rank

a = mx.ones((2, 3))
b = mx.arange(6).reshape(2, 3)
print(scale(a))            # → all 1’s

print(scale(b, factor=5))  # → b*5

```

### Capturing Mutable State

When functions capture mutable Python objects—such as dictionaries or random number generator states—pass them explicitly via the `inputs` and `outputs` parameters. For reproducible randomness in compiled code, capture `mx.random.state` as both input and output, as demonstrated in `test_compile_rng`. Changing captured state without declaring it in `inputs`/`outputs` forces unnecessary recompilation.

```python
state = {"bias": mx.array(1.0)}

@mx.compile(inputs=state, outputs=state)
def bias_add(x):
    return x + state["bias"]

print(bias_add(mx.array(2.0)))   # → 3.0

state["bias"] = mx.array(3.0)    # triggers recompilation on next call

print(bias_add(mx.array(2.0)))   # → 5.0

```

### Disabling Compilation for Debugging

Use `mx.disable_compile()` to run the uncompiled Python implementation for profiling or debugging, then restore compilation with `mx.enable_compile()`. This comparison helps isolate whether performance bottlenecks stem from kernel fusion or Python overhead.

```python
mx.disable_compile()

# Run uncompiled version for profiling

mx.enable_compile()

# Run compiled version for comparison

```

### Avoiding Compilation Pitfalls

Do not compile functions containing dynamic control flow dependent on input values (e.g., `if` statements on scalar arrays), as the tracer only supports static control flow. Such branches execute in Python before tracing begins, potentially causing unexpected behavior.

## Practical Code Examples

The following patterns demonstrate typical usage for the `mlx` compile function:

**Basic compilation with decorator syntax:**

```python
import mlx.core as mx

@mx.compile
def add(x, y):
    return x + y

x = mx.array([1., 2.])
y = mx.array([3., 4.])
print(add(x, y))  # → [4., 6.]

```

**Using compiled functions with transformations:**

```python
@mx.compile
def unary(x):
    return -mx.exp(x)

# Vectorized mapping

batched = mx.vmap(unary)(mx.arange(4.0))

# Gradient computation

grad_fn = mx.grad(unary)
print(grad_fn(mx.array(1.0)))

```

## Key Implementation Files

Understanding the compile system requires familiarity with these source locations:

- **[`python/tests/test_compile.py`](https://github.com/ml-explore/mlx/blob/main/python/tests/test_compile.py)** – Comprehensive test suite covering shapeless compilation, RNG handling, buffer donation, and non-finite constant handling.
- **[`python/mlx/extension.py`](https://github.com/ml-explore/mlx/blob/main/python/mlx/extension.py)** – Bridges Python calls to the low-level compiler; defines the `compile` entry point and tracing logic.
- **[`python/mlx/__init__.py`](https://github.com/ml-explore/mlx/blob/main/python/mlx/__init__.py)** – Exposes `compile`, `enable_compile`, and `disable_compile` in the public API.
- **[`benchmarks/python/compile_bench.py`](https://github.com/ml-explore/mlx/blob/main/benchmarks/python/compile_bench.py)** – Performance benchmarks demonstrating fusion speedups.

## Summary

- **MLX's `compile` function** traces Python operations, fuses them into single kernels, and caches the results for repeated execution with minimal overhead.
- **Shapeless mode** (`shapeless=True`) enables a single compiled kernel to handle varying input shapes, ideal for generic utilities but unsuitable for shape-dependent logic.
- **Mutable state** must be explicitly declared via `inputs` and `outputs` parameters to avoid unnecessary recompilation when values change.
- **Dynamic control flow** dependent on tensor values cannot be compiled; only static data flow graphs are supported.
- **Debugging** is facilitated by `mx.disable_compile()` and `mx.enable_compile()` to compare compiled versus uncompiled performance.

## Frequently Asked Questions

### What is the difference between `mx.compile` and `mx.compile(shapeless=True)`?

Standard `mx.compile` caches kernels based on specific input shapes and types, triggering recompilation when dimensions change. **Shapeless mode** generates a single kernel that handles any compatible shape, eliminating cache misses for varying dimensions but requiring the function logic to be shape-independent.

### Can I use random number generation inside a compiled function?

Yes, but you must pass `mx.random.state` as both an input and output to maintain reproducible streams within the compiled kernel. According to `test_compile_rng` in the test suite, this ensures the RNG state updates correctly without forcing recompilation.

### Why does my compiled function recompile on every call?

Recompilation occurs when captured Python state changes between invocations without being declared in the `inputs`/`outputs` parameters, or when input shapes change (in standard mode). Ensure mutable objects are explicitly captured and consider `shapeless=True` if dimensions vary.

### How do I verify that compilation is actually improving performance?

Use `mx.disable_compile()` to run the uncompiled version and compare primitive counts. The `test_enable_disable` test confirms that compiled graphs contain fewer primitives (`n_compiled < n_uncompiled`), indicating successful fusion and reduced kernel launch overhead.