QR, symmetric eigendecomposition and SVD for batches of matrices on Apple GPUs, on PyTorch tensors.
import torch
import metal_linalg_torch as mlt
a = torch.randn(1000, 64, 32, device="mps") # a batch of 1000 matrices; "cpu" works too
Q, R = mlt.qr(a) # Q (1000, 64, 32), R (1000, 32, 32)
U, S, Vh = mlt.svd(a) # thin: U (1000, 64, 32), S (1000, 32), Vh (1000, 32, 32)
L, V = mlt.eigh(a.mT @ a) # L ascending (1000, 32), V (1000, 32, 32)
The functions take the arguments of their torch.linalg namesakes and return
the same result types (Q, R = ..., result.eigenvalues, …), on the
device of their input. Each call is routed to the fastest Metal kernel for
its shape and batch, or to LAPACK on the CPU, by a policy measured on the
Mac it runs on:
mlt.device_name() # 'Apple M5 Pro'
mlt.eigh_policy_source() # 'tuned:Apple M5 Pro'
mlt.eigh_backend(32, 4096) # 'ql': 4096 matrices of 32x32 go to the GPU
mlt.svd_backend(4096, 4096) # 'bidiag': GPU bidiagonalization, LAPACK's solver
mlt.svdvals_backend(4096, 4096) # 'band': the two-stage reduction, values alone
mlt.set_eigh_policy(gpu_min_batch=1) # override the measured policy
The kernels, the routing and the measurements behind them are described in the main README.
pip install metal-linalg-torch
On an Apple Silicon Mac with macOS 14 or later, Python 3.10 or later and
PyTorch 2.4 or later. This is the PyTorch package; for
MLX arrays install metal-linalg
instead. The two are independent: this one does not install or load MLX.
One wheel serves every PyTorch and every Python: the package calls the
library’s C API (on plain float buffers) and is not compiled against
either, so upgrading torch never breaks it.
To build from source instead (needs Xcode’s command line tools):
pip install ./python-torch # in a clone
pip install "git+https://github.com/c0rmac/metal-linalg.git#subdirectory=python-torch"
| function | like | returns |
|---|---|---|
qr(A, mode="reduced") |
torch.linalg.qr |
(Q, R), thin; mode="r" gives an empty Q |
eigh(A, UPLO="L") |
torch.linalg.eigh |
(eigenvalues, eigenvectors), ascending |
eigvalsh(A, UPLO="L") |
torch.linalg.eigvalsh |
eigenvalues, ascending; less work than eigh |
svd(A, full_matrices=False) |
torch.linalg.svd |
(U, S, Vh), thin, S descending |
svdvals(A) |
torch.linalg.svdvals |
singular values, descending; about half the work of svd |
A is [..., M, N] with any number of batch dimensions, on "cpu" or
"mps". What differs from torch.linalg:
TypeError.svd defaults to full_matrices=False (torch’s
default is True), and full_matrices=True or qr(mode="complete") on a
non-square matrix raises NotImplementedError.Underneath, the functions are custom operators,
torch.ops.metal_linalg.{qr, eigh, eigvalsh, svd, svdvals}, with fake
implementations and the backward formulas torch.linalg uses:
A = torch.randn(64, 32, 32, device="mps", requires_grad=True)
S = mlt.svdvals(A)
S.sum().backward() # A.grad: the gradient of the nuclear norm
f = torch.compile(lambda x: mlt.eigh(x @ x.mT).eigenvalues.sum())
f(torch.randn(8, 16, 16))
As in torch, the gradients of eigh and svd are defined only for distinct
eigenvalues or singular values (they divide by their differences), and are
not unique for a loss that depends on the signs of the vectors. Second
derivatives are not supported. eigvalsh and svdvals compute the vectors
as well when their input requires grad, since the gradient needs them.
Against torch.linalg on an M5 Pro with PyTorch 2.13 (conda-forge’s, its CPU
LAPACK from Accelerate) and metal-linalg 2.16 (best of five,
benchmarks/benchmark_torch.py;
the same tensors on MPS for torch’s MPS path and for this package, which uses
them in place, see MPS tensors):
| torch, CPU | torch, MPS | metal-linalg-torch | |
|---|---|---|---|
| QR, 1024 × 128×128 | 202 ms | 32 ms | 4.4 ms |
| SVD, 256 × 128×64 | 69 ms | 71 ms | 5.1 ms |
| SVD, 4096 × 32×32 | 223 ms | 232 ms | 6.3 ms |
| eigh, 4096 × 16×16 | 33 ms | 35 ms | 2.0 ms |
| eigh, one 2048×2048 | 263 ms | 272 ms | 56 ms |
| SVD, one 4096×4096 | 3.60 s | 3.64 s | 339 ms |
| eigvalsh, one 4096×4096 | 1.78 s | 1.80 s | 125 ms |
| svdvals, one 4096×4096 | 1.95 s | 1.99 s | 193 ms |
It is ahead on every row. Of these calls PyTorch 2.13 runs only QR on the GPU
for MPS tensors, and this is 7x faster there (the QR batch runs on the
library’s Householder kernels); its SVD takes as long on MPS
as on the CPU, and eigh, eigvalsh and svdvals have no MPS kernels and go
through its CPU fallback. Against those, 13-35x for the other batches of
small matrices (the SVD of 256 matrices of 128×64 is a QR and then
golub_kahan on the GPU; the two batches of 4096 run on the GPU and the CPU
at once), 4.7x for eigh of one 2048×2048
and 10x for the SVD of one 4096×4096 with its vectors, and 10-14x for its
eigenvalues or singular values alone (the last three by a two-stage
reduction). Which
backend a shape gets on your Mac: mlt.svd_backend(m, n, batch) and its
siblings.
An MPS tensor is used in place. On Apple Silicon PyTorch keeps MPS tensors in
Metal buffers in shared storage, which the CPU can address too: the library
reads its input there and writes its results into new MPS tensors, with no
copy to the CPU and back. Its kernels run on a Metal command queue of its
own, so a call first waits for the work PyTorch has queued on the GPU
(torch.mps.synchronize()), which may still be writing the input, and
returns when its results are written. mlt.mps_in_place() says whether this
applies; where it does not (a torch that keeps MPS tensors in private
storage), or with METAL_LINALG_TORCH_MPS_COPY=1, an MPS tensor is copied to
the CPU and the results back. Since 2.16 the tensors’ own Metal buffers are
also handed to the library for the call (metal_linalg_know_buffer), so
that its GPU kernels use them rather than wrapping the memory in new
buffers, whose pages the GPU maps on first use.
Up to 2.13 every MPS call made those copies. They cost most where the
decomposition costs least: 9-12x for a handful of matrices, 1.2-1.8x for
large batches, a few percent for one large matrix (M5 Pro, PyTorch 2.13, best
of five with the two alternating, benchmarks/benchmark_torch.py --mps-ab):
| copied (2.13) | in place | |
|---|---|---|
| eigh, 16 × 16×16 | 1.2 ms | 0.10 ms |
| QR, 16 × 48×16 | 0.69 ms | 0.08 ms |
| SVD, 1024 × 32×32 | 4.4 ms | 3.0 ms |
| QR, 16384 × 64×32 | 25 ms | 14 ms |
| eigh, 16384 × 32×32 | 22 ms | 17 ms |
| eigh, one 2048×2048 | 93 ms | 90 ms |
| svdvals, one 4096×4096 | 231 ms | 227 ms |
A CPU tensor is used in place too, and its results stay on the CPU, though the work still runs on the GPU where that is faster.
Calls are thread-safe; they are serialised, and release the GIL.
Each Mac’s routing comes from measurements of that Mac. On one that has
none (or older ones), the import raises a CalibrationWarning once per
decomposition, and mlt.calibration_status() says where each stands. The
library works either way, with settings estimated from a measured Mac and
published benchmarks, which lean toward the CPU; measuring the Mac takes one
command and improves it for everyone with that Mac: see
which Macs are measured
and how to contribute.
METAL_LINALG_NO_CALIBRATION_NOTICE=1 silences the warning.