How MLX Computation Graph Optimization Works Internally: A Deep Dive into the Compiler Pipeline
MLX executes a five-stage compilation pipeline—tracing, tape construction, simplification, kernel fusion, and caching—to transform lazy computation graphs into optimized device kernels.
MLX (from the ml-explore/mlx repository) is an array framework designed for Apple silicon that employs lazy evaluation to build computation graphs dynamically. When you call a compiled function, MLX doesn't execute operations immediately; instead, it constructs a graph that undergoes aggressive optimization before ever reaching the GPU or CPU. Understanding how MLX computation graph optimization works internally reveals the machinery behind its efficient kernel generation.
Stage 1: Tracing the Function
The process begins when you wrap a function with compile(). MLX creates tracer inputs—placeholder arrays without underlying data—and executes your function on them to record operations without performing actual computation.
In mlx/compile.cpp (lines 3998–4016), the compile_trace function handles this:
// compile_trace – creates tracer inputs, runs the user function on them,
// and returns (tracer_inputs, outputs, extra)
auto [tracer_inputs, trace_outputs, extra] = compile_trace(fun, inputs, shapeless);
During this phase, every real input is replaced with a tracer via array::set_tracer(true). Because tracers contain no data, primitive operations are merely recorded, creating a "toy graph" that represents the computation structure without memory allocation or device execution.
Stage 2: Building the Tape via Depth-First Search
Once tracing completes, MLX converts the toy graph into a linear sequence of operations called a tape. The compile_dfs function (lines 4244–4292 in mlx/compile.cpp) performs a depth-first walk of the graph:
std::tie(entry.tape, parents_map) =
compile_dfs(entry.inputs, entry.outputs, inputs);
This DFS traversal achieves two critical objectives:
- Tape population: It collects every array that has a primitive into the
tapevector, creating an ordered list of operations. - Parent mapping: It builds a
parents_mapthat records which arrays consume each output, enabling later transformations to rewire connections without mutating the original graph.
The function also copies original inputs and outputs during this walk, ensuring subsequent optimization passes can modify the structure while preserving the initial graph semantics.
Stage 3: Simplifying the Tape
With the tape constructed, MLX runs compile_simplify (lines 6670–6720 in mlx/compile.cpp) across multiple passes to eliminate redundancy:
compile_simplify(entry.tape, parents_map, entry.outputs, /*passes=*/3);
The simplification engine performs three key optimizations:
- Scalar folding: Identifies identical zero-dimensional arrays and merges them using the
mergehelper function (lines 6632–6645), reducing duplicate constant computations. - No-op removal: Eliminates trivial operations like
CopyorStopGradientby rewiring their inputs directly to consumers (lines 6548–6563). - Common subexpression elimination: Merges arrays that compute equivalent values when neither is a graph output, checked via
array_equivalent(lines 6845–6864).
After simplification, the tape contains fewer nodes, reducing memory overhead and compilation time for subsequent stages.
Stage 4: Kernel Fusion
The most aggressive optimization occurs in compile_fuse (lines 8080–8110 in mlx/compile.cpp), which groups compatible primitives into single compiled kernels:
compile_fuse(entry.tape, parents_map, entry.inputs, entry.outputs);
This function iterates backward through the tape to identify fusable sub-graphs based on strict criteria:
- Depth limit: Fusion stops at
max_compile_depth(11 levels) to prevent exponential kernel compilation times. - Array count: Inputs to a fused kernel cannot exceed
max_compile_arrays(24 arrays). - Primitive compatibility: Only operations sharing the same stream and classified as "fusable"—unary, binary, ternary, or broadcast primitives—are considered.
When a valid sub-graph is identified, MLX:
- Collects the sub-graph's inputs.
- Encapsulates the operations into a
Compiledprimitive viastd::make_shared<Compiled>(…). - Removes the original primitives from the tape and updates
parents_mapto route consumers to the new fused node.
The result replaces multiple discrete operations with a single Compiled primitive that launches as one kernel on the device.
Stage 5: Caching and Final Replacement
MLX caches compilation results to avoid redundant optimization work. The CompilerCache structure (lines 3020–3038 in mlx/compile.cpp) stores compiled tapes keyed by input shapes, dtypes, and constant IDs:
// Cache lookup during compile() invocation
if (auto entry = compiler_cache().find(...)) {
// Reuse existing compiled tape
}
Before execution, compile_replace (lines 10990–11020) substitutes the tracer placeholders with real arrays, materializing the optimized computation graph while preserving the high-level API semantics.
Practical Example: Visualizing Graph Optimization
You can observe this pipeline in action using MLX's graph export utilities:
import mlx.core as mx
from mlx import compile
def f(a, b, c, d):
return (a + b) * (c - d)
# Compile triggers the five-stage pipeline
cf = compile(f)
# Create inputs and build lazy graph
a, b, c, d = [mx.array(i) for i in range(4)]
graph = cf([a, b, c, d]) # Still lazy—no device execution yet
# Export to DOT format to visualize the optimized graph
with open("graph.dot", "w") as f:
mx.export_to_dot(f, mx.utils.NodeNamer(), graph)
# Evaluate to run the fused kernel
out = graph[0].eval()
The exported DOT file (implemented in mlx/graph_utils.cpp, lines 105–120) reveals a single Compiled node instead of six separate primitives, confirming that the fusion stage successfully collapsed the addition, subtraction, and multiplication into one kernel launch.
Summary
- Lazy tracing: MLX records operations using tracer arrays in
compile_tracewithout allocating device memory. - Linearization:
compile_dfsflattens the graph into a tape and builds a parent map for dependency tracking. - Simplification: Multiple passes in
compile_simplifyfold scalars, remove no-ops, and deduplicate common subexpressions. - Kernel fusion:
compile_fusegroups compatible primitives into singleCompilednodes subject to depth and array count limits. - Caching: The
CompilerCachestores optimized kernels keyed by signature, eliminating recompilation for identical graph structures.
Frequently Asked Questions
What is the "tape" in MLX compilation?
The tape is a linearized vector of primitive operations extracted from the lazy computation graph during the DFS traversal. It represents a sequential execution order that the simplification and fusion stages transform before final code generation.
How does MLX determine which operations can be fused?
MLX fuses operations that belong to the same stream, are classified as fusable (unary, binary, ternary, or broadcast primitives), and fit within constraints of maximum depth (11) and maximum input arrays (24). These limits prevent kernel compilation overhead from dominating execution time.
Why does MLX use tracing instead of executing operations immediately?
Tracing allows MLX to build a complete computation graph before any device execution occurs. This global view enables aggressive optimizations—such as common subexpression elimination and kernel fusion—that would be impossible with eager execution, where each operation completes before the next begins.
Can I disable specific optimizations in the MLX compiler?
The public API in mlx/compile.cpp does not expose granular flags for individual optimization passes. However, you can control fusion behavior through the max_compile_depth and max_compile_arrays constants, and you can bypass compilation entirely by calling functions without the compile() wrapper to execute operations eagerly.
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 →