How MLX's Compile Function Works and When to Use Compile Mode
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 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 (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.
@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.
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.
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:
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:
@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– Comprehensive test suite covering shapeless compilation, RNG handling, buffer donation, and non-finite constant handling.python/mlx/extension.py– Bridges Python calls to the low-level compiler; defines thecompileentry point and tracing logic.python/mlx/__init__.py– Exposescompile,enable_compile, anddisable_compilein the public API.benchmarks/python/compile_bench.py– Performance benchmarks demonstrating fusion speedups.
Summary
- MLX's
compilefunction 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
inputsandoutputsparameters 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()andmx.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.
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 →