# How MLX Primitives Form the Backend of All Operations

> Discover how MLX primitives power all MLX operations. Learn about their role as low-level computational units handling execution, gradients, and batching via a unified C++ interface.

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

---

**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)](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:

```cpp
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)](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.h#L76-L88), demonstrates this pattern:

```cpp
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`](https://github.com/ml-explore/mlx/blob/main/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++:

```cpp
#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:

```python
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:

```python
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

- **[`mlx/primitives.h`](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.h)** – Declares the `Primitive` base class and all concrete primitives (e.g., `Add`, `Matmul`, `Convolution`). Source: [[`mlx/primitives.h`](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.h)](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.h)

- **[`mlx/primitives.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.cpp)** – Implements primitive methods (`eval_*`, `jvp`, `vjp`, `vmap`) and helper utilities. Source: [[`mlx/primitives.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.cpp)](https://github.com/ml-explore/mlx/blob/main/mlx/primitives.cpp)

- **[`mlx/ops.h`](https://github.com/ml-explore/mlx/blob/main/mlx/ops.h) / [`mlx/ops.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/ops.cpp)** – High-level operation wrappers that instantiate primitives and expose the user-friendly API. Source: [[`mlx/ops.h`](https://github.com/ml-explore/mlx/blob/main/mlx/ops.h)](https://github.com/ml-explore/mlx/blob/main/mlx/ops.h)

- **`mlx/backend/*/`** – Device-specific kernels (`cpu`, `gpu`, `metal`) that primitives call within `eval_cpu` and `eval_gpu`. Example: [[`mlx/backend/cpu/unary_ops.h`](https://github.com/ml-explore/mlx/blob/main/mlx/backend/cpu/unary_ops.h)](https://github.com/ml-explore/mlx/blob/main/mlx/backend/cpu/unary_ops.h)

## 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`](https://github.com/ml-explore/mlx/blob/main/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.