How Automatic Differentiation Works in MLX's Transforms: A Deep Dive into the Autograd Engine
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.
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 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:
- Executes the forward pass to build the graph.
- Accepts a cotangent (seed), usually an array of ones.
- 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.
-
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. -
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.
-
Result Materialization: Gradients accumulate as
mlx::core::arrayobjects, which are then returned to the Python layer.
Practical Example
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.hand executes them via transforms defined inmlx/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, andcheckpointprovide composable interfaces for vectorization, custom gradients, and memory management. - Source files:
mlx/transforms.hdeclares the API,mlx/transforms.cppimplements the logic, andmlx/core/autograd.cppmanages 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, 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.
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 →