metal-linalg: Batched QR, Eigendecomposition and SVD on Apple Silicon GPUs

metal-linalg provides QR decomposition, symmetric eigendecomposition (eigh) and the singular value decomposition (SVD) for batches of matrices on Apple Silicon GPUs. It works with MLX arrays, PyTorch tensors and plain float buffers. It is written in C++, with Python packages for MLX and PyTorch, a Swift package and a C API.

MLX’s own eigh and svd run only on the CPU, and PyTorch has no GPU kernels for most of these decompositions on a Mac. metal-linalg fills that gap.

View on GitHub · Which Macs are measured

⚙️ How it works

Each solver has several Metal kernels, one per regime:

  • Large batches of small matrices: LAPACK’s own methods, one matrix per threadgroup or simdgroup (Householder QR, tridiagonalization and implicit QL for eigh, Golub–Kahan bidiagonalization for the SVD).
  • One large matrix: the memory-bound reduction to tridiagonal or bidiagonal form runs on the GPU, while divide and conquer runs on every CPU core. For eigenvalues or singular values alone, a two-stage reduction goes to a band on the GPU, then to tridiagonal form on the CPU, with bisection on the GPU.
  • Long, thin matrices: QR first, then the small square problem.

Which kernel is fastest, and where the GPU overtakes the CPU, depends on the chip. So every call is routed by a policy measured on the Mac it runs on, keyed on the Metal device name and GPU core count. A Mac nobody has measured gets an estimate, refitted from a measured Mac and published benchmarks.

📈 Performance

On an Apple M5 Pro, against a CPU path that spreads every call over all 18 cores:

one N×N matrix20484096
SVD, with vectors5.6x10.4x
QR5.2x10.3x
singular values alone3.8x9.95x
eigh, with eigenvectors4.5x9.2x

Against torch.linalg with PyTorch 2.13:

 torch, CPUtorch, MPSmetal-linalg-torch
QR, 1024 × 128×128202 ms32 ms4.4 ms
SVD, 4096 × 32×32223 ms232 ms6.3 ms
eigh, 4096 × 16×1633 ms35 ms2.0 ms
SVD, one 4096×40963.60 s3.64 s339 ms

💻 Installation

pip install metal-linalg          # for MLX
pip install metal-linalg-torch    # for PyTorch
brew tap c0rmac/metal-linalg
brew install metal-linalg         # the C++ library and C API

⚡️ Quick start

import torch
import metal_linalg_torch as mlt

a = torch.randn(1000, 64, 32, device="mps")
Q, R = mlt.qr(a)                  # like torch.linalg.qr
U, S, Vh = mlt.svd(a)             # thin factors
L, V = mlt.eigh(a.mT @ a)         # like torch.linalg.eigh

The PyTorch functions mirror their torch.linalg namesakes, support autograd and compile with torch.compile.

🤝 Contributing

The routing is only as good as the measurements behind it, and every new chip needs its own. If you have an Apple Silicon Mac, one command (python3 tuning/run.py) measures it and produces a results folder to send as a pull request. See how to contribute.