Mamba, on a laptop GPU.
Mamba's fast kernels are written for NVIDIA GPUs. On a Mac the fallback is a PyTorch scan that writes out the model's entire hidden state at every step of every sequence, and it runs out of memory long before the model gets big. This is the missing piece: the selective scan as one fused Metal kernel, forward and backward, that any PyTorch Mamba can call in place of its own.
Why a Mac could not train one
A Mamba layer runs a recurrence over time: each token updates a small hidden state, and the output reads from it. Written in plain PyTorch, that state is materialised for every token, every channel and every sequence in the batch, and autograd keeps the scan's intermediates alive on top. The official fix is a CUDA kernel that keeps the state on-chip. Apple silicon has no CUDA, so on a Mac the choice was a tiny model or no model.
The state lives in registers
One GPU thread owns one channel of one sequence and carries its whole hidden state in float4 registers from the first token to the last. Nothing per-token is written to memory except a checkpoint every 64 steps. The backward pass walks those chunks in reverse, recomputes each chunk's states from its checkpoint and carries the gradient in registers across the boundaries; the two gradients that sum over channels are reduced in threadgroup memory and added with Metal's float atomics.
Correct, not just consistent
A kernel that agrees with a buggy reference is still wrong, so the gradients are checked twice: against the pure-PyTorch scan at fp32, on shapes chosen to break the easy assumptions — lengths that are not multiples of 64, channel counts that do not fill a threadgroup, a single timestep — and against finite differences, which do not care what either implementation thinks. End to end, a model's loss and every parameter gradient are the same with the kernel as without it.
The bug the port turned up
Pulling the kernel out into its own repository meant re-running a test that had never actually run: it sliced tokens 20 to 26 out of a 20-token batch, so its loop was empty and it passed. Given real tokens it failed, and the cause was in decoding, not the kernel — at batch 1 the state and the readout broadcast into the right shape by luck, and at batch 4 each sequence was reading every other sequence's output. Fixed, and the test now checks the shape as well as the values.