How MLX Primitives Form the Backend of All Operations

MLX primitives form the backend of all operations by serving as low-level computational units that encapsulate device-specific execution, gradient propagation, and batching logic through a unified C++ interface.

The MLX machine learning framework from the ml-explore/mlx repository implements a primitive-based computational graph where every tensor operation—from element-wise addition to complex convolutions—decomposes into these fundamental building blocks. Understanding how MLX primitives form the backend of all operations reveals the architectural decisions that enable seamless CPU/GPU portability, JAX-style automatic differentiation, and efficient vectorization without kernel rewrites.

The Primitive Interface

At the core of MLX's architecture lies the abstract base class mlx::core::Primitive, defined in [mlx/primitives.h](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.h#L48-L66). This interface mandates that every concrete primitive must implement three fundamental capabilities:

  1. Device execution via eval_cpu and eval_gpu methods
  2. Gradient propagation via jvp (forward-mode) and vjp (reverse-mode) methods
  3. Batching via the vmap method

The base class declaration establishes this contract:

class Primitive {
public:
    explicit Primitive(Stream stream) : stream_(stream) {}
    virtual void eval_cpu(const std::vector<array>& inputs,
                          std::vector<array>& outputs) = 0;
    virtual void eval_gpu(const std::vector<array>& inputs,
                          std::vector<array>& outputs) = 0;
    virtual std::vector<array> jvp(const std::vector<array>& primals,
                                   const std::vector<array>& tangents,
                                   const std::vector<int>& argnums);
    virtual std::vector<array> vjp(const std::vector<array>& primals,
                                   const std::vector<array>& cotangents,
                                   const std::vector<int>& argnums,
                                   const std::vector<array>& outputs);
    virtual std::pair<std::vector<array>, std::vector<int>> vmap(
            const std::vector<array>& inputs,
            const std::vector<int>& axes);
    virtual const char* name() const = 0;
    …
};

Each primitive stores a Stream object at construction time, binding the operation to a specific device (CPU or GPU) and determining which evaluation path executes.

Concrete Implementation: The Add Primitive

Concrete primitives inherit from Primitive (or intermediate classes like UnaryPrimitive) and implement the virtual interface. The Add primitive, defined in [mlx/primitives.h](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.h#L76-L88), demonstrates this pattern:

class Add : public UnaryPrimitive {
public:
    explicit Add(Stream stream) : UnaryPrimitive(stream) {}
    void eval_cpu(const std::vector<array>& inputs, array& out) override;
    void eval_gpu(const std::vector<array>& inputs, array& out) override;
    DEFINE_VMAP()
    DEFINE_GRADS()
    DEFINE_NAME(Add)
    DEFINE_DEFAULT_IS_EQUIVALENT()
    DEFINE_INPUT_OUTPUT_SHAPE()
};

The eval_cpu implementation forwards to the element-wise add kernel in mlx/ops.cpp, while eval_gpu dispatches to the Metal GPU implementation. For automatic differentiation, Add::jvp adds the tangents of the arguments, and Add::vjp propagates the cotangent unchanged (or splits it when both arguments require gradients). The vmap implementation uses vmap_binary_op to align inputs on a common batch axis before invoking the standard kernel.

How High-Level Operations Use Primitives

All high-level MLX operations—such as mx.add, mx.matmul, or mx.conv—are thin wrappers that instantiate the appropriate primitive and invoke its evaluation methods. When you call mx.add(x, y) in Python, the binding layer creates an Add primitive, passes the input arrays, and executes add_op.eval_* on the active device.

This design provides a uniform entry point for:

  • Device selection – The primitive's stored Stream determines execution target at construction
  • Autodiff – The same primitive object supplies jvp/vjp, enabling composable differentiation
  • Batching – vmap implementations reshape or transpose inputs so kernels apply over new batch dimensions without rewrites

When the computational graph compiles via mlx.compile, the compiler walks the tree of primitive objects, collects operations, and generates fused kernels respecting the selected device and data layout.

Code Examples

Direct Primitive Usage in C++

You can instantiate and invoke primitives directly in C++:

#include "mlx/array.h"
#include "mlx/primitives.h"

int main() {
    mlx::core::Stream s;                     // default device (CPU)
    auto a = mlx::core::array({1, 2, 3}, s);
    auto b = mlx::core::array({4, 5, 6}, s);

    mlx::core::Add add_op(s);                // instantiate the Add primitive
    std::vector<mlx::core::array> inputs = {a, b};
    std::vector<mlx::core::array> outputs(1);   // one output for binary ops
    add_op.eval_cpu(inputs, outputs);        // compute a + b on CPU

    // `outputs[0]` now holds {5, 7, 9}
}

This example demonstrates explicit primitive construction, device stream handling, and direct invocation of eval_cpu.

Python High-Level API

The Python frontend provides ergonomic wrappers that hide primitive instantiation:

import mlx.core as mx

x = mx.array([1, 2, 3])
y = mx.array([4, 5, 6])

z = mx.add(x, y)          # high‑level call

print(z)                  # → [5 7 9]

mx.add internally creates an Add primitive and executes eval_cpu or eval_gpu based on the current default device.

Automatic Differentiation Through Primitives

Primitives enable reverse-mode automatic differentiation via their vjp implementations:

import mlx.core as mx

def f(a, b):
    return mx.add(a, b) * mx.square(a)

a = mx.array([1., 2.])
b = mx.array([3., 4.])
y = f(a, b)

grad_a = mx.grad(f, argnums=0)(a, b)   # uses Add.jvp / Add.vjp, Square.jvp, etc.

print(grad_a)                         # → derivative of f w.r.t. a

During mx.grad, each primitive in the graph contributes its vector-Jacobian product implementation, guaranteeing correct derivative propagation through the computational tree.

Key Source Files

Summary

  • MLX primitives are the foundational backend objects that implement every tensor operation through a unified C++ interface.
  • Each primitive implements device execution (eval_cpu/eval_gpu), gradient rules (jvp/vjp), and vectorization (vmap).
  • The Primitive base class in mlx/primitives.h enforces this contract, ensuring consistent behavior across CPU and GPU.
  • High-level APIs like mx.add are thin wrappers that instantiate primitives and delegate to their evaluation methods.
  • Automatic differentiation works by traversing the primitive graph and invoking each node's vjp method during reverse-mode passes.

Frequently Asked Questions

What is the difference between MLX primitives and high-level operations?

High-level operations like mx.add or mx.matmul are user-facing functions that handle array validation, broadcasting, and device placement before instantiating the corresponding primitive. The primitive is the actual C++ object that holds the implementation logic for execution, gradients, and batching, while the high-level operation provides the convenient API wrapper.

How do primitives handle device selection between CPU and GPU?

Primitives store a Stream object at construction time, which encapsulates the target device. When evaluation triggers, the primitive calls eval_cpu for CPU streams or eval_gpu for GPU streams, ensuring the correct backend kernel executes without runtime device checks in the hot path.

How does automatic differentiation work with primitives?

During forward passes, MLX builds a computational graph of primitive objects. When computing gradients via mx.grad, the system traverses this graph in reverse, calling each primitive's vjp (vector-Jacobian product) method to propagate cotangents backward. Similarly, forward-mode differentiation uses jvp to push tangents forward through the graph.

Can users define custom primitives in MLX?

Currently, custom primitives require implementing the Primitive interface in C++ and exposing it through the Python bindings. Users must provide implementations for eval_cpu, eval_gpu, and the differentiation methods (jvp, vjp, vmap) to ensure the custom operation integrates with MLX's compilation and autodiff systems.

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 →