Metal (Apple Silicon) backend for OpenAI Triton [1][2]. Write @triton.jit kernels and run them on your Mac's GPU.
@triton.jit → Triton TTIR → TTGIR → MSL → metallib → Apple GPU
The same @triton.jit source runs on NVIDIA. triton-msl is a Triton backend, not a
dialect: only the final stage (→ MSL) is Apple-specific, so the kernel you develop and
correctness-debug on your Mac is the identical code that runs on a CUDA GPU. Verified on a
real NVIDIA A40: two of the three example fp32 kernels (vector-add, ieee matmul) produced
bit-identical outputs across vendors; softmax matched to fp rounding (~1e-9)
(PORTABILITY.md). Develop kernel logic on the laptop you own, rent a
GPU only for the performance pass.
Alpha: actively developed, not yet production-ready.
- 0 failures across the upstream Triton
test_core.pysuite: 5,560 kernels attempted and correct, 3,782 documented skips. Reconciled test by test, almost all are hardware-bounded (no fp8/mxfp matrix unit, no fp64, no 64-bit atomics, no tf32, or over Metal's 1024-thread / 32 KB threadgroup ceilings); each is a loud refusal or a hardware-impossible case, never silent-wrong, and effectively none are unwritten lowerings. Aligned with Triton [2] release3.7.0. Measured byscripts/run_upstream_tests.py(the single source of truth for this count), which runs--device cpu(torch references compute on CPU while the Metal backend compiles and runs the kernels on the GPU, since upstreamtest_coreotherwise assumes CUDA). Re-run it to reproduce; counts in this file andCHANGELOG.mdare regenerated from it, not hand-maintained. - 1,971 passed / 0 failed in the project suite (codegen, GPU correctness,
integration, FlashAttention, MLX backend, fast-matmul / compile_shader
zero-copy,
torch.compile, and training). FlashAttention: causal + non-causal at HEAD_DIM 32 / 64 / 128 (head_dim 128 fp32 + fp16 via the simdgroup-MMA template; see [4] for the algorithm); 15 / 15 MLX backend tests; the project suite grew from 434 → 603 → 716 → 877 → 1,971 since0.1.0-alpha. (A further ~20 C++-MLIR-backend tests skip unless that optional extension is built.) torch.compileroutes through triton-msl on Python 3.10–3.14 (PyTorch Inductor [12]), inference and training (AOTAutograd backward), static anddynamic=True; 32 / 32torch.compilemodel tests (plus 1 inductor-config regression test = 33 total) and the training suite pass.- Triton tutorials 01–03, 05 passing.
- Built against Triton's
TRITON_EXT_ENABLED=1plugin architecture (upstream PR #9783). - Integrity contract: kernels we can lower run correctly; kernels we
cannot are refused (
MetalNonRecoverableError), never silent-wrong. Seedocs/SUPPORTED_OPS.mdfor the supported ops/dtypes matrix + the loud-refusal catalog, anddocs/ARCHITECTURE.md"Lowering paths and the integrity model" for the lowering paths.
See REFERENCES.md for citations and
docs/superpowers/specs/2026-05-30-triton-msl-roadmap.md
for the active pre-1.0 roadmap.
triton-msl is a backend for Triton, so your @triton.jit source is standard Triton and the same code runs on NVIDIA (only the final codegen stage differs). The three example
kernels, unmodified, were run on an Apple M4 Max and a rented NVIDIA A40, each checked
against the same NumPy reference:
| kernel | Mac vs NumPy | NVIDIA vs NumPy | Mac ↔ NVIDIA Δ |
|---|---|---|---|
vector_add |
0 | 0 | 0, bit-identical |
fused_softmax |
7.45e-9 | 5.59e-9 | 7.45e-9 |
matmul (fp32 / ieee) |
4.58e-5 | 4.58e-5 | 0, bit-identical |
matmul (tf32) |
refused (no tf32 on Metal) | 6.07e-2 | n/a |
Correctness/logic is portable: two of three kernels were bit-identical across vendors,
and it crossed Triton 3.0.0 ↔ 3.7.0. Performance is not (block sizes and fast-path
routing are hardware-specific). So the workflow is: develop and debug kernel logic on the
Mac you own, then rent a GPU for minutes for the performance pass. The one numerical caveat
is NVIDIA's tf32 default for tl.dot (pass input_precision="ieee" to match). Full
receipt + a reproduce harness: PORTABILITY.md.
- Apple Silicon Mac (M1 or later)
- macOS 14 (Sonoma) or later
- Xcode Command Line Tools:
xcode-select --install - Python 3.10+
- PyTorch 2.12+ (2.5+ for the zero-copy MPS fast path;
torch.compileis developed + tested against 2.12) and Triton 3.7.0
pip install triton-msl
# Triton is required but installed separately. There is no official macOS wheel,
# so build it from source (the primary, supported path, a one-time ~12 min build):
pip install git+https://github.com/triton-lang/triton.gitIf you're on the exact platform tuple Python 3.14 / macOS 15+ / Apple Silicon (M1–M5), an unofficial prebuilt Triton wheel is attached to the GitHub releases so you can skip the build:
pip install https://github.com/bledden/triton-msl/releases/download/triton-wheel-3.7.0-cp314-macos-arm64/triton-3.7.0+git4da2e268-cp314-cp314-macosx_15_0_arm64.whlimport torch
import triton
import triton.language as tl
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < n
x = tl.load(x_ptr + offs, mask=mask)
y = tl.load(y_ptr + offs, mask=mask)
tl.store(out_ptr + offs, x + y, mask=mask)
n = 1024
x = torch.randn(n, device="cpu")
y = torch.randn(n, device="cpu")
out = torch.empty(n, device="cpu")
add_kernel[(n + 255) // 256,](x, y, out, n, BLOCK=256)
print(f"Max error: {(out - (x + y)).abs().max():.2e}")Run a fuller, runnable version, vector-add, fused-softmax, and a tiled matmul, each
on a non-multiple shape (so every load/store boundary mask is exercised) and verified
against a NumPy reference, with examples/local_triton_dev.py:
python examples/local_triton_dev.pyIt's the same @triton.jit source you'd run on a CUDA GPU: develop and verify locally
on your Mac, then ship the identical kernels to NVIDIA.
import torch
import triton_msl.inductor
triton_msl.inductor.register_metal_triton_backend()
model = torch.nn.Sequential(
torch.nn.Linear(256, 512),
torch.nn.ReLU(),
torch.nn.Linear(512, 256),
)
compiled = torch.compile(model, backend="metal")
x = torch.randn(32, 256)
out = compiled(x)Performance. The compiled kernels dispatch zero-copy through the same compile_shader
fast path as the hand-written kernels. At small model sizes the compiled latency is roughly
on par with eager MPS, so the value here is coverage, not speed: torch.compile graphs
(inference and training) route through triton-msl and stay correct to eager within
floating-point tolerance, validated across 32 models spanning transformer blocks, a small
GPT, a ViT, ResNets, and LSTMs, forward and training backward.
Coverage. Transformer/attention, RNN, CNN (incl. BatchNorm), normalization, and the
reductions (sum, product, mean, max/min incl. NaN-propagating, var/std, argmax/argmin, softmax,
logsumexp, cumsum/cumprod), including small under-filling reductions and a 2-D reduction
fused with a 2-D scan in one kernel (e.g. x.sum(1) + x.cumprod(1)[:,-1]), compile and match
eager. Anything genuinely beyond the hardware (a reduction/scan tile exceeding Metal's 1024
threads/threadgroup) is refused loudly rather than mis-computed (MetalNonRecoverableError,
never silent-wrong); the rest of the graph is unaffected.
import mlx.core as mx
import triton
import triton.language as tl
from triton_msl.mlx import triton_call
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < n
x = tl.load(x_ptr + offs, mask=mask)
y = tl.load(y_ptr + offs, mask=mask)
tl.store(out_ptr + offs, x + y, mask=mask)
n = 1024
x = mx.random.normal((n,))
y = mx.random.normal((n,))
out = mx.zeros((n,))
results = triton_call(add_kernel, x, y, out, n, grid=(4,), BLOCK=256)The same @triton.jit kernel runs zero-copy on torch MPS tensors: the driver
dispatches the emitted Metal through torch.mps.compile_shader, skipping the host
round-trip (~10× faster on memory-bound kernels; on by default, no code change):
x = torch.randn(n, device="mps")
y = torch.randn(n, device="mps")
out = torch.empty(n, device="mps")
add_kernel[(n + 255) // 256,](x, y, out, n, BLOCK=256) # runs on the GPU, no copyfp16/fp32 matmuls (K%8) on MPS tensors take a direct simdgroup-matrix path, dispatched
zero-copy. A deterministic, occupancy-gated tile selector picks the coarsest blocking the
shape's M and N alignment allow: M%32 + N%32 runs the (4,4) tile at ~11–12 TFLOP/s on
M4 Max; M%8 or N%8 (not %32) runs a finer rescue tile (lower, but still far above the
generic path). M%8 ≠ 0 or N%8 ≠ 0 (e.g. an odd vocab/hidden dim) falls to the ~2.4 TFLOP/s
generic path: a partial simdgroup strip can't be masked without writing past the matrix.
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K,
sam, sak, sbk, sbn, scm, scn,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr):
pid_m = tl.program_id(0); pid_n = tl.program_id(1)
offm = pid_m * BM + tl.arange(0, BM)
offn = pid_n * BN + tl.arange(0, BN)
offk = tl.arange(0, BK)
a = a_ptr + (offm[:, None] * sam + offk[None, :] * sak)
b = b_ptr + (offk[:, None] * sbk + offn[None, :] * sbn)
acc = tl.zeros((BM, BN), dtype=tl.float32)
for k in range(0, K, BK):
acc += tl.dot(tl.load(a), tl.load(b))
a += BK * sak; b += BK * sbk
tl.store(c_ptr + (offm[:, None] * scm + offn[None, :] * scn), acc.to(tl.float16))
M = N = K = 2048
A = torch.randn(M, K, device="mps", dtype=torch.float16)
B = torch.randn(K, N, device="mps", dtype=torch.float16)
C = torch.empty(M, N, device="mps", dtype=torch.float16)
matmul_kernel[(M // 64, N // 64)](
A, B, C, M, N, K,
A.stride(0), A.stride(1), B.stride(0), B.stride(1), C.stride(0), C.stride(1),
BM=64, BN=64, BK=32)A kernel that triton-msl cannot lower correctly raises MetalNonRecoverableError
rather than returning garbage. For example, a pid-tiled matmul that bakes its M/N
dims as constexpr (so the true output strides can't be recovered) is refused:
from triton_msl.errors import MetalNonRecoverableError
@triton.jit
def matmul_baked_dims(a_ptr, b_ptr, c_ptr, K,
M: tl.constexpr, N: tl.constexpr,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr):
pid_m = tl.program_id(0); pid_n = tl.program_id(1)
offm = pid_m * BM + tl.arange(0, BM)
offn = pid_n * BN + tl.arange(0, BN)
offk = tl.arange(0, BK)
a = a_ptr + (offm[:, None] * K + offk[None, :])
b = b_ptr + (offk[:, None] * BN + offn[None, :])
acc = tl.zeros((BM, BN), dtype=tl.float32)
for _k in range(0, K, BK):
acc += tl.dot(tl.load(a), tl.load(b)); a += BK; b += BK * BN
tl.store(c_ptr + (offm[:, None] * BN + offn[None, :]), acc)
try:
matmul_baked_dims[(2, 2)](A, B, C, K, M=64, N=64, BM=32, BN=32, BK=32)
except MetalNonRecoverableError as e:
print("refused (not silent-wrong):", e)See docs/SUPPORTED_OPS.md for the full op/dtype support
matrix and the loud-refusal catalog.
A full FlashAttention v2 forward (causal + non-causal) runs through the standard
@triton.jit path at BLOCK_M = BLOCK_N = 32 for head_dim 32, 64, and 128, see tests/test_flash_attention.py for the kernel and
launch. head_dim 32/64 use the generic lowering; head_dim 128 is routed to an
Apple simdgroup_matrix MMA kernel (fp32 + fp16, causal + non-causal, any N_CTX) that
measures ~5.2× (fp32) / ~6.4× (fp16) faster than the prior scalar path at N=1024
(5.1 / 6.3 TFLOP/s), rising to ~6.1× / ~7.9× at N=2048 (6.8 / 8.8 TFLOP/s), ~45–55% of the in-repo matmul-template peak. This is NOT
competitive with Apple metal-flash-attention or MLX in absolute terms; it is the
practical ceiling for a triton-jit-routed MSL kernel at these tile sizes. Out-of-range
configs are refused loudly (MetalNonRecoverableError, never silent-wrong): head_dim
128, block tiles ≠ 32, bf16 inputs, non-contiguous innermost stride, and any FA-shaped kernel whose strides/scale can't be resolved unambiguously. Larger blocks and head_dim > 128 are on the roadmap.
All default-on; set to 0 to disable (an escape hatch for bisecting a regression):
| Flag | Effect when disabled |
|---|---|
TRITON_MSL_COMPILE_SHADER=0 |
Use the host-copy driver instead of the zero-copy compile_shader dispatch |
TRITON_MSL_FAST_MATMUL=0 |
Use the generic matmul instead of the fast simdgroup-matrix path |
TRITON_MSL_MATMUL_AUTOTUNE=0 |
Pin matmul tile selection to the fixed (4,4) blocking (M%32≠0 / N%32≠0 shapes drop to the generic path instead of the finer M%8/N%8 rescue tiles) |
TRITON_MSL_MEPT=0 |
Disable the multi-element-per-thread register-array model |
TRITON_MSL_LEGACY=1 |
Opt in to the heuristic legacy text parser (off by default, it can be silent-wrong) |
| Category | Operations |
|---|---|
| Elementwise | add, sub, mul, div, exp, log, sqrt, abs, neg, SiLU, GELU, sigmoid, tanh, ReLU, leaky ReLU, clamp, FMA |
| Reductions | sum, max, min, argmax, argmin, xor_sum |
| Dot product | tl.dot with strided matmul template, all epilogues (add, softmax, chain-dot, transpose) |
| Attention | FlashAttention [4] (causal + non-causal) at BLOCK_M=BLOCK_N=32, HEAD_DIM 32 / 64 / 128 via the Python MSL path (head_dim 128 routed to a simdgroup-MMA template, contiguous-stride fast path, scalar fallback otherwise, fp32 + fp16). Out-of-range configs (head_dim > 128; block tiles ≠ 32; bf16) are refused (MetalNonRecoverableError, never silent-wrong); larger blocks/head_dim are on the roadmap. |
| Normalization | Layer norm, RMS norm, batch norm |
| Type casts | FP32, FP16, BF16, INT8, INT16, INT32, bool |
| Control flow | scf.for, scf.if, while loops |
| Atomics | atomic_add, atomic_max, atomic_min, atomic_and, atomic_or, atomic_xor, CAS |
| Tensor ops | cat, join, split, interleave, reshape, permute, transpose, histogram, gather |
| torch.compile | 32 models including MLP, ResBlock, TransformerBlock, SmallGPT, MiniViT, LSTM |
| MLX | Zero-copy dispatch via mx.fast.metal_kernel() |
| Feature | Reason |
|---|---|
| FP64 | Metal has no FP64 support |
| FP8, TF32 | Not available on Apple GPUs |
| Multi-GPU | Apple Silicon is single-GPU |
tl.dot with sizePerThread > 1 |
Requires 2D cooperative execution model (addressed by the register-array spine, WS1) |
Unstructured control flow (cf.cond_br) |
Refused with MetalNonRecoverableError (never silent-wrong); a cf-dialect lowerer is WS2 |
tt.dot_scaled (microscaling matmul) |
No Apple microscaling hardware; refused |
Performance (M4 Max [13])
Measured numbers via the zero-copy compile_shader path (default-on); see
reports/perf_baseline.json. Hardware peak: 546 GB/s memory, 18.4 / 36.9 TFLOP/s
fp32 / fp16.
| Kernel | Size | Throughput | % of peak | vs host-copy path |
|---|---|---|---|---|
| Vector add | 16M | 347 GB/s | 64% | 13× |
| Elementwise | 16M | 315 GB/s | 58% | 13.4× |
| Softmax | 8192×1024 | 232 GB/s | 42% | 17.8× |
| Reduction | 16M | 235 GB/s | 43% | 8.2× |
| Matmul (fp32) | 2048³ | 11.2 TFLOP/s§ | 61% of fp32 peak | ~4× generic |
| Matmul (fp16 in / fp32 out) | 2048³ | 12.3 TFLOP/s | ≈ fp32 rate* | ~4× generic |
| Matmul (fp16 in / fp16 out) | 2048³ | 12.2 TFLOP/s | ≈ fp32 rate* | ~4× generic |
| Matmul (bf16 in / fp32 out)◊ | 2048³ | 12.0 TFLOP/s | ≈ fp32 rate* | ~4.9× generic |
| Matmul (bf16 in / bf16 out)◊ | 2048³ | 11.9 TFLOP/s | ≈ fp32 rate* | ~4.9× generic |
| FlashAttention (fp32, head_dim=128)‡ | Z=1,H=8,N=1024 | 5.1 TFLOP/s | ~28% of fp32 peak | ~5.2× scalar FA |
| FlashAttention (fp16, head_dim=128)‡ | Z=1,H=8,N=1024 | 6.3 TFLOP/s | † | ~6.4× scalar FA |
* fp16 matmul uses fp16 inputs with a float32 accumulator (for precision). The
12.3 figure is the default fp32-output path; the true fp16→fp16 path
(out_dtype=fp16) measures 12.2 TFLOP/s, essentially identical, since both do
the same MACs and the output-cast cost is negligible. Either way it runs at roughly
the fp32 matrix-unit rate: Apple's simdgroup-matrix unit isn't faster for half
accumulation, so the 36.9 TFLOP/s fp16 figure is an unreachable vector-ALU peak. The ~58–64% memory-bound
and ~60% fp32-matmul numbers are near the practical ceilings for these kernel
classes on this hardware (the raw 546 / 18.4 / 36.9 spec peaks are not reachable by
compute), see the Phase-5 readiness audit (docs/audits/).
† FA fp16 accumulates in fp32 (correct); % of the 36.9 fp16 vector-ALU peak is not a useful comparison; the practical reference is the in-repo matmul peak, of which fp32 FA reaches ~46% (5.1 / 11.2) and fp16 ~51% (6.3 / 12.3). The FA numbers are not competitive with Apple metal-flash-attention or MLX in absolute terms.
§ The fp32 matmul figure is the cold-machine peak (~11 TFLOP/s, verified 2026-06-24).
It is thermally sensitive: under sustained GPU load an M4 Max throttles the fp32
path to ~9 TFLOP/s, so reports/perf_baseline.json, re-measured after a long
benchmark session, currently records ~9.2; re-run test_fast_matmul_perf on an idle
machine to see the cold peak. fp16/fp16out throttle less (measured ~12.3/12.2 cold).
The 11.2 figure is the dense row-major (contiguous-innermost) case only.
Non-contiguous / transposed / sliced operands (e.g. x @ w.t(), a column-major
output, or a column slice t[:, :K] whose inner stride ≠ 1 or whose row stride ≠ the
matrix dim) fall off the simdgroup fast path; they are computed correctly by a
fully stride-aware scalar matmul (or refused when un-inferable; never silently
wrong), but that scalar path is ~15–23× slower (below the generic floor).
contiguous() the operand before the kernel to stay on the fast path. (Batched 3-D
matmuls are not yet implemented and refuse loudly.)
◊ bf16 matmul uses Apple's simdgroup_bfloat8x8 matrix unit (float32 accumulate),
verified on M4; on a part without a bfloat matrix unit it falls back to the
(correct) generic float-compute path, never silently wrong.
‡ FA rows are measured on M4 Max (integration microbenchmark of the shipped
make_flash_attention_kernel_simdgroup: warmup + median over 50 iters), not yet a
committed regression-benchmark harness. Throughput scales with sequence length:
~6.8 TFLOP/s (fp32, ~6.1×) / ~8.8 TFLOP/s (fp16, ~7.9×) at N=2048. Correctness
is verified by the 16-case differential gate (tests/test_fa_simdgroup_diff.py: simd
== scalar oracle == torch, all pass). A dedicated FA throughput benchmark is on the
roadmap.
MPS tensors run zero-copy via torch.mps.compile_shader (default-on); the prior
host-round-trip copy bottleneck is gone. CPU tensors and the MLX backend
[7] (mx.fast.metal_kernel) also dispatch zero-copy.
@triton.jit kernel
→ Triton frontend (Python AST → TTIR)
→ Triton optimizer (TTIR → TTGIR)
→ mlir_walker.py: walk TTGIR module → IRGraph
→ generic_lowerer.py: IRGraph → MSL source
→ xcrun metal: MSL → AIR → metallib
→ driver.py: load metallib, dispatch on GPU
See docs/ARCHITECTURE.md for details.
See CONTRIBUTING.md.
If you use triton-msl in research or technical work, see
CITING.md for a suggested BibTeX entry. For citations of
the papers and projects this backend builds on (Triton, FlashAttention,
online softmax, MLX, Asahi/applegpu, the MSL specification, PyTorch
Inductor), see REFERENCES.md.