# How Automatic Differentiation Works in MLX's Transforms: A Deep Dive into the Autograd Engine

> Explore how MLX's autograd engine powers automatic differentiation in transforms for machine learning. Understand dynamic computation graphs and gradient computation.

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

---

**MLX implements automatic differentiation through a hybrid autograd engine that records primitive array operations in a dynamic computation graph and executes reverse-mode (backpropagation) or forward-mode (JVP) passes via the `vjp`, `jvp`, and `grad` transforms defined in [`mlx/transforms.h`](https://github.com/ml-explore/mlx/blob/main/mlx/transforms.h).**

The MLX framework (ml-explore/mlx) provides a NumPy-like API with composable function transformations that enable automatic differentiation without manual gradient computation. At the core of this system lies a graph-based autograd engine that tracks every operation on `mlx::core::array` objects, allowing transforms to compute gradients through registered backward and forward functors.

## The Autograd Engine: Recording and Executing Gradients

Every primitive operation in MLX—such as `add`, `matmul`, `conv`, or `exp`—registers its metadata in a global computation graph. This graph stores **Node IDs** (result arrays), **Parents** (input arrays), and **Backward functors** (callables that propagate gradients). When transforms like `grad` or `vjp` are invoked, they trigger a two-phase execution process.

### Building the Computation Graph

During the forward pass, the engine evaluates the supplied function on *primal* inputs, constructing a graph of dependent operations. Each primitive operation in [`mlx/core/array.h`](https://github.com/ml-explore/mlx/blob/main/mlx/core/array.h) registers a backward lambda with the autograd system. When `async_eval` or `eval` (see **[transforms.h:12‑24]** ) is called, the graph is flushed: forward values materialize, and backward functors become available for subsequent gradient computation.

### Reverse-Mode Differentiation (VJP)

**Reverse-mode automatic differentiation** (also known as backpropagation) computes gradients by walking the graph from outputs to inputs. When you call `mlx::core::vjp` (see **[transforms.h:33‑41]** ), the engine:

1. Executes the forward pass to build the graph.
2. Accepts a cotangent (seed), usually an array of ones.
3. Traverses the graph in reverse, invoking each operation’s backward functor to accumulate cotangents for the parents.

This process yields the vector-Jacobian product, commonly used for computing gradients of scalar loss functions.

### Forward-Mode Differentiation (JVP)

**Forward-mode automatic differentiation** propagates tangents forward through the graph. The `mlx::core::jvp` routine (see **[transforms.h:53‑62]** ) evaluates the function while simultaneously applying each operation’s forward-mode rule to produce Jacobian-vector products. This approach is efficient when the number of inputs exceeds the number of outputs, or when computing directional derivatives.

## The Transform API: High-Level Interfaces to Autograd

MLX exposes its autograd capabilities through a unified transform API. Each transform is a thin wrapper around the core engine:

| Transform | Core Routine | Return Value | Typical Use |
|-----------|--------------|--------------|-------------|
| `vjp` | `mlx::core::vjp` **[transforms.h:33‑41]** | `(primal_outputs, cotangents)` | Compute vector-Jacobian products for custom gradients. |
| `jvp` | `mlx::core::jvp` **[transforms.h:53‑62]** | `(primal_outputs, tangents)` | Compute Jacobian-vector products (forward-mode). |
| `grad` | `mlx::core::grad` **[transforms.h:28‑48]** | `gradient` | Returns the gradient by calling `value_and_grad` and extracting the second element. |
| `value_and_grad` | `mlx::core::value_and_grad` **[transforms.h:75‑92]** | `(value, gradient)` | Returns both the function value and its gradient in one call. |
| `vmap` | `mlx::core::vmap` **[transforms.h:59‑88]** | Vectorized function | Automatically batches a unary or binary function over specified axes. |
| `custom_function` | `mlx::core::custom_function` **[transforms.h:90‑109]** | Function with custom AD rules | Allows overriding default `vjp`, `jvp`, and `vmap` behavior. |
| `checkpoint` | `mlx::core::checkpoint` **[transforms.h:124‑130]** | Memory-optimized function | Discards intermediate results during forward pass and recomputes them during backward pass to save memory. |

## Execution Flow: From Function Call to Gradient

Understanding the end-to-end flow clarifies how MLX moves from Python API calls to computed gradients.

1. **Graph Construction**: When you invoke a transformed function, MLX evaluates the forward computation, recording each primitive operation in the graph managed by [`mlx/core/autograd.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/core/autograd.cpp).

2. **Backward Traversal**: For reverse-mode transforms, the engine supplies a cotangent seed and walks the graph in reverse topological order, executing backward functors stored during the forward pass.

3. **Result Materialization**: Gradients accumulate as `mlx::core::array` objects, which are then returned to the Python layer.

### Practical Example

```python
import mlx.core as mx

# Define a scalar function

def f(x):
    return mx.sin(x) * mx.exp(x)

# Create a gradient transform

grad_f = mx.grad(f)  # Internally calls value_and_grad -> vjp

x = mx.array(0.5)
g = grad_f(x)        # Returns array([-0.352])

```

Under the hood, `mx.grad(f)` builds a wrapper that invokes `mx.value_and_grad`, which in turn calls `mx.vjp`. The forward evaluation records `sin`, `exp`, and multiplication operations, while the backward pass applies the chain rule through the stored backward functors to compute `∂f/∂x`.

## Advanced Features: Custom Rules and Memory Optimization

Beyond standard automatic differentiation, MLX provides mechanisms for customizing gradient behavior and optimizing memory usage in deep networks.

### Defining Custom Functions

The `custom_function` transform (see **[transforms.h:90‑109]** ) allows you to implement user-provided `vjp`, `jvp`, and `vmap` implementations. This is essential for integrating external operations or applying numerical stability tricks that the default autograd rules cannot capture.

### Gradient Checkpointing

For memory-intensive models, the `checkpoint` transform (see **[transforms.h:124‑130]** ) modifies the execution strategy: it discards intermediate activations during the forward pass and recomputes them on demand during the backward pass. This trades additional computation for reduced memory footprint, enabling training of larger models on limited hardware.

## Summary

- **MLX automatic differentiation** relies on a dynamic computation graph that records primitive operations in [`mlx/core/array.h`](https://github.com/ml-explore/mlx/blob/main/mlx/core/array.h) and executes them via transforms defined in [`mlx/transforms.h`](https://github.com/ml-explore/mlx/blob/main/mlx/transforms.h).
- **Reverse-mode AD** (`vjp`, `grad`) traverses the graph backward from outputs to inputs, accumulating cotangents to compute gradients efficiently for scalar losses.
- **Forward-mode AD** (`jvp`) propagates tangents forward through the graph, useful for Jacobian-vector products when input dimensionality is high.
- **High-level transforms** like `value_and_grad`, `vmap`, `custom_function`, and `checkpoint` provide composable interfaces for vectorization, custom gradients, and memory management.
- **Source files**: [`mlx/transforms.h`](https://github.com/ml-explore/mlx/blob/main/mlx/transforms.h) declares the API, [`mlx/transforms.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/transforms.cpp) implements the logic, and [`mlx/core/autograd.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/core/autograd.cpp) manages graph execution.

## Frequently Asked Questions

### What is the difference between `grad` and `vjp` in MLX?

`grad` is a convenience wrapper that internally calls `value_and_grad` and returns only the gradient portion, while `vjp` (vector-Jacobian product) returns both the primal outputs and the cotangents, giving you explicit control over the backward seed. According to the source code in [`mlx/transforms.h`](https://github.com/ml-explore/mlx/blob/main/mlx/transforms.h), `grad` is implemented by extracting the second element from the `value_and_grad` result **[transforms.h:28‑48]**.

### When should I use forward-mode (`jvp`) instead of reverse-mode (`grad`)?

Use `jvp` when you need to compute directional derivatives or when your function has many inputs and few outputs, as forward-mode automatic differentiation scales with the number of outputs rather than inputs. The `mlx::core::jvp` implementation **[transforms.h:53‑62]** propagates tangents forward through the computation graph, making it efficient for calculating Jacobian-vector products in these scenarios.

### How does MLX handle memory during gradient computation?

MLX provides the `checkpoint` transform **[transforms.h:124‑130]** to reduce memory usage by discarding intermediate results during the forward pass and recomputing them during the backward pass. This technique, implemented in the core transforms, allows training deeper networks by trading computation for memory savings.

### Can I define custom gradient rules for existing operations?

Yes, MLX supports custom automatic differentiation rules through the `custom_function` transform **[transforms.h:90‑109]**. This allows you to supply user-defined `vjp`, `jvp`, and `vmap` implementations, enabling integration of custom primitives or application of numerical stability fixes that override the default autograd behavior recorded in [`mlx/core/array.h`](https://github.com/ml-explore/mlx/blob/main/mlx/core/array.h).