MLX is Apple's open-source array framework for machine learning, developed by Apple machine learning research and designed around Apple Silicon. It looks familiar on purpose: a NumPy-like core in mlx.core, neural network layers in mlx.nn and optimisers in mlx.optimizers, with Python, C++, C and Swift front ends. Underneath, three design choices make it different from PyTorch or JAX on a Mac: arrays live in unified memory with no device of their own, computation is lazy, and function transformations such as gradients and compilation are composable, as in JAX.

This article is about the programming model and the tooling, not the chip. The hardware side, memory bandwidth and how to estimate tokens per second, is covered in Apple Silicon, in depth. Here you will learn how MLX evaluates work, how to write a correct and fast training step, how its quantisation works, how to run and fine-tune language models with mlx-lm, how to watch and cap memory, how to write a custom Metal kernel, and what goes wrong. MLX moves fast, with releases every few weeks, so check signatures against the documentation for your installed version.

Arrays without devices: the programming model

In PyTorch a tensor belongs to a device, and moving it is a copy. In MLX an array is a buffer in unified memory that the CPU and GPU can both address, so it has no device. Instead each operation runs on a device, chosen by the default or by a stream argument. That lets you run, say, a tokenisation-side reduction on the CPU and a matrix multiply on the GPU over the same arrays without a transfer; MLX tracks dependencies between streams and inserts the necessary waits.

How an MLX program runs: build a lazy graph, transform it, evaluate it on a streamPython / C++ / Swiftmx.array operationsLazy graphnothing computed yetrecordTransformsgrad, vmapmx.compilefuse, cache tracemx.evalor print, itemGPU streamMetal kernelsCPU streamAccelerate, vectorUnified memory: one buffer per arrayboth streams read and write it in place; no host-to-device copiesArrays have no device. Operations do: each op runs on the default device or the stream you pass it.
Operations build a graph; transforms rewrite it; evaluation schedules it on GPU and CPU streams that share one memory pool.
import mlx.core as mx

a = mx.random.normal((4096, 4096))
b = mx.random.normal((4096, 4096))

c = mx.matmul(a, b, stream=mx.gpu)       # GPU
d = mx.sum(a, axis=0, stream=mx.cpu)     # CPU, same buffer, no copy
e = c + d                                # depends on both; MLX orders them
mx.eval(e)

MLX also has a CUDA backend for NVIDIA GPUs, so the same code can run on a server. Treat it as a portability option and benchmark it rather than assuming parity with the Metal backend, which remains the primary target.

Lazy evaluation and where to call mx.eval

Every MLX operation records a node in a graph and returns immediately. Nothing is computed until something needs the values: an explicit mx.eval, printing an array, .item(), converting to NumPy, saving, or using an array in a Python if. Laziness gives MLX freedom to schedule and fuse work, and it means unused results are never computed, for example loading a full checkpoint and only using some weights. It also creates the two most common performance bugs.

  • Evaluating too often. Calling .item() on the loss every step, or branching on an array value inside the loop, forces a synchronisation each time and starves the GPU of batched work.
  • Evaluating too rarely. If you never evaluate inside a loop, the graph grows across iterations until memory runs out or graph building dominates. The documentation's guidance is to evaluate once per outer training iteration; graphs of tens to many thousands of operations per evaluation are fine.

The rule of thumb: one mx.eval per step, covering the loss, the model parameters and the optimiser state, and read scalar values for logging only every N steps.

Composable transforms: grad, value_and_grad, vmap

Gradients in MLX are function transformations, not a tape. mx.grad(f) returns a new function that computes the gradient of f with respect to its first argument; mx.value_and_grad returns both, and mx.vmap vectorises a function over a batch axis. They compose, so per-example gradients are mx.vmap(mx.grad(loss)). For modules, nn.value_and_grad(model, fn) differentiates with respect to the model's trainable parameters, which are returned as a nested dictionary matching the module tree.

import mlx.core as mx

def loss(w, x, y):
    return mx.mean((x @ w - y) ** 2)

grad_fn = mx.grad(loss)                               # d loss / d w
per_example = mx.vmap(mx.grad(loss), in_axes=(None, 0, 0))

w = mx.zeros((8,))
x = mx.random.normal((32, 8)); y = mx.random.normal((32,))
g = grad_fn(w, x, y)                                  # shape (8,)
G = per_example(w, x[:, None, :], y[:, None])         # shape (32, 8)
mx.eval(g, G)

Worked example: a compiled training step

A complete training step combines these pieces. The worked example is a small two-layer classifier, but the structure is the same one mlx-lm uses for language models. The step is wrapped in mx.compile. Compilation traces the function once, fuses element-wise operations into fewer kernels and caches the result. Compiled functions must be pure, so state that changes, the model parameters, the optimiser state and the random state used by dropout, is declared through the inputs and outputs arguments.

from functools import partial
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim

class MLP(nn.Module):
    def __init__(self, d_in, d_hidden, n_classes):
        super().__init__()
        self.l1 = nn.Linear(d_in, d_hidden)
        self.drop = nn.Dropout(0.1)
        self.l2 = nn.Linear(d_hidden, n_classes)
    def __call__(self, x):
        return self.l2(self.drop(nn.gelu(self.l1(x))))

model = MLP(784, 512, 10)
optimizer = optim.AdamW(learning_rate=3e-4, weight_decay=0.01)

def loss_fn(model, x, y):
    return nn.losses.cross_entropy(model(x), y, reduction="mean")

state = [model.state, optimizer.state, mx.random.state]

@partial(mx.compile, inputs=state, outputs=state)
def step(x, y):
    loss, grads = nn.value_and_grad(model, loss_fn)(model, x, y)
    optimizer.update(model, grads)
    return loss

for it, (x, y) in enumerate(batches()):         # your data iterator yielding mx.arrays
    loss = step(x, y)
    mx.eval(state)                              # one evaluation per step
    if it % 100 == 0:
        print(it, loss.item())                  # sync only when logging

Compilation re-traces when input shapes, dtypes or the number of arguments change. Variable-length batches therefore cause repeated recompiles; pad to a few bucket lengths, or pass shapeless=True when the function has no shape-dependent logic such as a hard-coded reshape. When something looks wrong, mx.disable_compile() turns compilation off globally so you can print intermediate arrays.

Quantisation and what it costs in memory

MLX's default quantisation is affine and group-wise. Weights are split into groups of group_size consecutive values along the last axis. For each group MLX stores a scale s = (max - min) / (2^bits - 1) and a bias equal to the group minimum, and each weight becomes round((w - min) / s). mx.quantize returns the packed weights, scales and biases, and nn.quantize(model, ...) swaps eligible Linear and Embedding layers for quantised versions whose matrix multiply dequantises on the fly. Recent releases also offer mxfp4, mxfp8 and nvfp4 modes. Defaults for group size have differed between releases and modes, so always pass group_size and bits explicitly.

import mlx.core as mx
import mlx.nn as nn

w = mx.random.normal((4096, 4096)).astype(mx.float16)
wq, scales, biases = mx.quantize(w, group_size=64, bits=4)
w_hat = mx.dequantize(wq, scales, biases, group_size=64, bits=4)
print(mx.abs(w - w_hat).max())                  # worst-case reconstruction error

nn.quantize(model, group_size=64, bits=4)       # in place, on a loaded model

Worked sizing: with 4-bit weights, a group size of 64 and a 16-bit scale and bias per group, each group costs 64 times 4 bits plus 32 bits, which is 4.5 bits per weight. A 7-billion-parameter model is therefore about 3.9 GB of weights instead of 14 GB in float16, before the KV cache and activations. Smaller groups improve accuracy and raise the overhead; group size 32 costs 5 bits per weight. For the wider landscape of formats compare GGUF and llama.cpp, which make similar trade-offs with different block layouts.

Running and fine-tuning LLMs with mlx-lm

Most people meet MLX through mlx-lm, the companion package for language models. It loads Hugging Face models, converts and quantises them, generates text, serves an HTTP endpoint and fine-tunes with LoRA, DoRA or full fine-tuning.

pip install mlx-lm

# Convert and quantise a Hugging Face model to MLX format
mlx_lm.convert --model mistralai/Mistral-7B-Instruct-v0.3 -q

# Generate, or chat with preserved context
mlx_lm.generate --model mlx-community/Mistral-7B-Instruct-v0.3-4bit --prompt "Explain KV caching"
mlx_lm.chat --model mlx-community/Mistral-7B-Instruct-v0.3-4bit

# LoRA fine-tune: --data points to a folder with train.jsonl and valid.jsonl
mlx_lm.lora --model mlx-community/Mistral-7B-Instruct-v0.3-4bit --train \
    --data ./data --iters 600 --batch-size 1 --num-layers 8 \
    --mask-prompt --grad-checkpoint --adapter-path ./adapters

# Merge the adapter into the base weights for serving
mlx_lm.fuse --model mlx-community/Mistral-7B-Instruct-v0.3-4bit --adapter-path ./adapters

Training on a quantised base is the QLoRA pattern: the frozen base stays 4-bit and only the small adapter matrices are trained in higher precision. The memory levers, in order of effect, are batch size, the number of layers that receive adapters, sequence length and gradient checkpointing, which trades recomputation for activation memory. Each line in the JSONL files must be one complete example in the chat, completions, text or tools format. From Python, from mlx_lm import load, generate gives the same functionality, and stream_generate yields tokens as they are produced.

Watching and capping memory

Because the GPU shares system RAM, running out of memory does not raise a clean device error first; the system starts compressing and swapping and everything slows down. MLX gives you tools to see and cap its usage. In current releases they live directly under mlx.core; older code calls them through mx.metal.

import mlx.core as mx

mx.set_memory_limit(48 * 1024**3)      # cap MLX allocations
mx.set_cache_limit(4 * 1024**3)        # cap freed buffers kept for reuse
mx.reset_peak_memory()

run_one_step()
print("active", mx.get_active_memory() / 1e9, "GB")
print("peak  ", mx.get_peak_memory() / 1e9, "GB")
print("cache ", mx.get_cache_memory() / 1e9, "GB")
mx.clear_cache()

The allocator keeps freed buffers in a cache so the next step can reuse them without asking the operating system. That is why memory appears not to drop after a large step; the cache limit bounds it. mx.set_wired_limit controls how much memory MLX asks the system to keep resident, which helps large models avoid being paged out between tokens. Measure peak memory for one step at your real batch and sequence length before starting a long run.

Custom Metal kernels

When an operation is missing or too slow, mx.fast.metal_kernel lets you write the body of a Metal kernel as a string and call it like any other MLX function. MLX generates the function signature from the input and output names, supplies template parameters, and launches it over the grid you specify. The mx.fast namespace also contains tuned built-ins worth checking first, such as RMS norm, layer norm, rotary embeddings and scaled dot-product attention.

import mlx.core as mx

source = '''
    uint i = thread_position_in_grid.x;
    T v = inp[i];
    out[i] = v / (T(1) + metal::exp(-v));      // SiLU
'''
silu_kernel = mx.fast.metal_kernel(
    name="silu_v1", input_names=["inp"], output_names=["out"], source=source)

def silu(a):
    return silu_kernel(
        inputs=[a], template=[("T", a.dtype)],
        grid=(a.size, 1, 1), threadgroup=(256, 1, 1),
        output_shapes=[a.shape], output_dtypes=[a.dtype])[0]

Give each distinct source a distinct name. A reported issue showed that two kernels with the same name and different sources could run stale code within one evaluation batch, which is the kind of bug that produces plausible but wrong numbers. Always compare a custom kernel against a reference implementation in a test.

Distributed MLX

MLX can spread work over several Macs. mx.distributed.init() returns a group with a rank and size; collectives such as mx.distributed.all_sum become no-ops when there is one process, so the same script runs alone or in a group. For data-parallel training, nn.average_gradients batches many small gradient reductions into fewer large ones. The mlx.launch helper starts processes locally with -n or across machines with --hosts or --hostfile, and the documented backends are MPI, ring, JACCL and NCCL. Interconnect bandwidth between Macs is far below what a GPU server's fabric offers, so distribution is most useful for fitting a model that does not fit one machine, not for speeding up a model that does.

Failure modes

FailureSymptomFix
.item() or array branching every stepGPU utilisation low, step time highEvaluate once per step, log every N steps
No eval inside a loopMemory climbs until swapAdd mx.eval(state) each iteration
Variable shapes under compileSteps slow, recompilesBucket lengths or shapeless=True
Mutable state not declared to compileParameters do not updatePass state through inputs and outputs
Model larger than free RAMSystem swaps, tokens per second collapsesSmaller quantisation, set memory limits, close apps
Same kernel name, new sourceStale resultsVersion kernel names; test against reference

Trade-offs against other Mac stacks

Choose MLX when you work on a Mac and want both inference and training in one framework that uses unified memory natively, or when you want JAX-style transforms with a NumPy feel. Choose PyTorch with the MPS backend when you need the PyTorch ecosystem and portability to CUDA servers matters more than peak Mac performance. Choose llama.cpp or a wrapper such as Ollama for packaged inference across many platforms. Choose Core ML when shipping a model inside an app, especially if you want the Neural Engine, which MLX does not program directly.

What to do next

  1. Install mlx and mlx-lm, run mlx_lm.generate on a 4-bit model, and note tokens per second and peak memory.
  2. Port one small training loop: nn.value_and_grad, one mx.eval(state) per step, then wrap the step in mx.compile with declared state and compare step time.
  3. Quantise a model with explicit group_size and bits, and measure quality on your own prompts at 4 and 8 bits.
  4. Run a LoRA fine-tune on a few hundred examples with --mask-prompt, evaluate on a held-out file, then fuse and serve it.
  5. Set memory and cache limits for long runs and log peak memory each epoch.
  6. Before writing a custom kernel, check mx.fast for a built-in, and test any kernel you write against a reference.
Key takeaway: MLX arrays live in unified memory and operations choose the device, so CPU and GPU share data without copies. Computation is lazy: evaluate once per training step, compile the step with its state declared, and avoid reading scalars every iteration. Quantise with explicit group size and bits, fine-tune with mlx-lm LoRA on a quantised base, cap and measure memory, and test any custom Metal kernel against a reference.