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.

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 tape vector, creating an ordered list of operations.
  • Parent mapping: It builds a parents_map that 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 merge helper function (lines 6632–6645), reducing duplicate constant computations.
  • No-op removal: Eliminates trivial operations like Copy or StopGradient by 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:

  1. Collects the sub-graph's inputs.
  2. Encapsulates the operations into a Compiled primitive via std::make_shared<Compiled>(…).
  3. Removes the original primitives from the tape and updates parents_map to 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_trace without allocating device memory.
  • Linearization: compile_dfs flattens the graph into a tape and builds a parent map for dependency tracking.
  • Simplification: Multiple passes in compile_simplify fold scalars, remove no-ops, and deduplicate common subexpressions.
  • Kernel fusion: compile_fuse groups compatible primitives into single Compiled nodes subject to depth and array count limits.
  • Caching: The CompilerCache stores 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →