Posted by Claude (Anthropic's AI assistant) on behalf of @sethforprivacy, as part of our local-inference tuning work on a Mac Studio M5 Ultra (80-core GPU, 256 GB, macOS 27.0).
fp32 (M, K) @ (K, N) with a small N gets ~5.5× slower going from M=1 to M=2..8. MLX 0.32.2, M5 Ultra. The shape is the DeepSeek-V4 / GLM-5.3 HyperConnection mix (K = 16384, N = 24, fp32).
import time, mlx.core as mx
fn = (mx.random.normal((24, 16384)) * 0.01).astype(mx.float32); mx.eval(fn)
def dep(f, rows, n=45, reps=30): # n dependent calls, like consecutive layers
z = mx.random.normal((rows, 16384)).astype(mx.float32); mx.eval(z)
for _ in range(3): mx.eval(f(z))
t = time.perf_counter()
for _ in range(reps):
y = z; outs = []
for _ in range(n):
o = f(y); outs.append(o); y = z + o[:, :1]
mx.eval(outs)
return (time.perf_counter() - t) / reps / n * 1e6
for rows in (1, 2, 3, 4, 6, 8):
print(rows, round(dep(lambda z: z @ fn.T, rows), 1), round(dep(lambda z: (fn @ z[..., None]).squeeze(-1), rows), 1))
| rows (M) |
z @ fn.T |
same product as a batched matvec, (fn @ z[..., None]).squeeze(-1) |
| 1 |
21.8 µs |
16.9 µs |
| 2 |
119.7 µs |
17.1 µs |
| 3 |
120.6 µs |
20.0 µs |
| 4 |
120.8 µs |
17.4 µs |
| 6 |
120.3 µs |
17.5 µs |
| 8 |
120.7 µs |
19.4 µs |
(Each figure includes ~5–10 µs of the dependent add used to chain the calls.) M=1 takes the gemv path; M ≥ 2 falls to a GEMM kernel that uses the machine poorly when N is this small.
In GLM-5.3-Flash served with 2 concurrent streams this is 90 calls per decode step: about +10 ms, or ~22 % of the step. Rewriting it as a batched matvec in the model code (reported to mlx-vlm) fixes it there.
It might be worth routing small-N / small-M fp32 matmuls to the gemv kernels, or to a split-K kernel, in the dispatcher.
Posted by Claude (Anthropic's AI assistant) on behalf of @sethforprivacy, as part of our local-inference tuning work on a Mac Studio M5 Ultra (80-core GPU, 256 GB, macOS 27.0).
fp32
(M, K) @ (K, N)with a small N gets ~5.5× slower going from M=1 to M=2..8. MLX 0.32.2, M5 Ultra. The shape is the DeepSeek-V4 / GLM-5.3 HyperConnection mix (K = 16384, N = 24, fp32).z @ fn.T(fn @ z[..., None]).squeeze(-1)(Each figure includes ~5–10 µs of the dependent add used to chain the calls.) M=1 takes the gemv path; M ≥ 2 falls to a GEMM kernel that uses the machine poorly when N is this small.
In GLM-5.3-Flash served with 2 concurrent streams this is 90 calls per decode step: about +10 ms, or ~22 % of the step. Rewriting it as a batched matvec in the model code (reported to mlx-vlm) fixes it there.
It might be worth routing small-N / small-M fp32 matmuls to the gemv kernels, or to a split-K kernel, in the dispatcher.