use std::ffi::CStr;
use std::sync::{Mutex, OnceLock};
use crate::engine::TranslationContext;
use crate::mlx::{Array, FastMetalKernel, FastMetalKernelConfig, VectorArray, VectorString};
use crate::sys::mlx;
const KERNEL_NAME: &CStr = c"onnxrt_mlx_qmm_bf16_fp16";
const HEADER: &CStr = cr#"
#include <metal_stdlib>
#include <metal_simdgroup_matrix>
using namespace metal;
"#;
const SOURCE: &CStr = cr#"
constexpr int BM = 64;
constexpr int BK = 32;
constexpr int BN = 32;
constexpr int WM = 2;
constexpr int WN = 2;
constexpr int TM = BM / (WM * 8); // 4
constexpr int TN = BN / (WN * 8); // 2
constexpr int BK_PAD = BK + 8; // 40
constexpr int BN_PAD = BN + 8; // 40
const int M = x_shape[0];
const int K = x_shape[1];
const int N = w_shape[0];
const int nblocks = scales_shape[1];
const uint3 tid = threadgroup_position_in_grid;
const uint simd_gid = simdgroup_index_in_threadgroup;
const uint simd_lid = thread_index_in_simdgroup;
threadgroup half Xs[BM * BK_PAD];
threadgroup half Ws[BK * BN_PAD];
const int y_row0 = int(tid.y) * BM;
const int y_col0 = int(tid.x) * BN;
const uint tidx = simd_gid * 32u + simd_lid;
// X (activation) load: 4 threads cover the BK=32-wide K tile (8 half
// elements/thread, vectorized bf16->half4 casts), 32 rows/pass -> 2 passes
// for BM=64.
const uint xrow_local_base = tidx / 4u;
const uint xk0_local = (tidx % 4u) * 8u;
// W (quantized weight) load: 4 threads cover the BK=32-wide K tile (1
// uint32 word = 8 nibbles/thread), 32 cols -> exactly 1 pass for BN=32
// (matches the group_size=32 alignment already required by `eligible`).
const uint wcol_local = tidx / 4u;
const uint wword_local = tidx % 4u;
const int wcol_global = y_col0 + int(wcol_local);
const bool wcol_ok = wcol_global < N;
const device uint32_t* wrow_ptr = wcol_ok ? (w + (long)wcol_global * K / 8) : w;
const long srow_base = (long)wcol_global * nblocks;
const uint sm_row = simd_gid / uint(WN);
const uint sm_col = simd_gid % uint(WN);
simdgroup_matrix<float,8,8> acc[TM][TN];
for (int i = 0; i < TM; i++) {
for (int j = 0; j < TN; j++) { acc[i][j] = simdgroup_matrix<float,8,8>(0.0f); }
}
for (int k0 = 0; k0 < K; k0 += BK) {
// X load: 2 row-passes of 32 rows each, vectorized bf16->half4 cast (2x
// 4-wide vector loads instead of 8 scalar loads+casts) -- mirrors
// mlx::steel::BlockLoaderCast's cast_vec_width=4 technique.
for (int rp = 0; rp < 2; rp++) {
uint xrow_local = xrow_local_base + uint(rp) * 32u;
int xrow_global = y_row0 + int(xrow_local);
threadgroup half* dst = Xs + xrow_local * uint(BK_PAD) + xk0_local;
if (xrow_global < M) {
const device bfloat4* src4 = (const device bfloat4*)(
x + (long)xrow_global * K + k0 + int(xk0_local));
threadgroup half4* dst4 = (threadgroup half4*)dst;
dst4[0] = static_cast<half4>(src4[0]);
dst4[1] = static_cast<half4>(src4[1]);
} else {
threadgroup half4* dst4 = (threadgroup half4*)dst;
dst4[0] = half4(0.0h);
dst4[1] = half4(0.0h);
}
}
// W load: single pass (BN=32 == group_size), nibble unpack stays scalar
// (bit unpacking does not vectorize). N is only guaranteed a multiple of
// 32 (== BN here), so a column tile is always either fully valid or
// fully out of range -- no partial-column case exists for BN=32.
{
int krow0 = int(wword_local) * 8;
threadgroup half* wdst = Ws + uint(krow0) * uint(BN_PAD) + wcol_local;
if (wcol_ok) {
int group = k0 / 32;
half sc = half(float(scales[srow_base + group]));
half bs = half(float(biases[srow_base + group]));
uint32_t word = wrow_ptr[(k0 / 8) + int(wword_local)];
for (int j = 0; j < 8; j++) {
uint32_t nib = (word >> uint(j * 4)) & 0xFu;
half v = half(float(nib)) * sc + bs;
wdst[j * BN_PAD] = v;
}
} else {
for (int j = 0; j < 8; j++) { wdst[j * BN_PAD] = half(0.0h); }
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (int ks = 0; ks < BK; ks += 8) {
simdgroup_matrix<half,8,8> a[TM];
simdgroup_matrix<half,8,8> b[TN];
for (int i = 0; i < TM; i++) {
simdgroup_load(a[i], Xs, BK_PAD, ulong2(uint(ks), sm_row*32u + uint(i)*8u));
}
for (int j = 0; j < TN; j++) {
simdgroup_load(b[j], Ws, BN_PAD, ulong2(sm_col*16u + uint(j)*8u, uint(ks)));
}
for (int i = 0; i < TM; i++) {
for (int j = 0; j < TN; j++) {
simdgroup_multiply_accumulate(acc[i][j], a[i], b[j], acc[i][j]);
}
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
// Store phase: mirrors mlx::steel::BlockMMA::store_result -- each thread
// writes its own 2 accumulator elements per 8x8 fragment straight to
// device memory via simdgroup_matrix::thread_elements(), using the fixed
// Apple-GPU per-lane (row, col-pair) layout for an 8x8 simdgroup_matrix
// (BaseMMAFrag::get_coord in mlx/backend/metal/kernels/steel/gemm/mma.h).
// No threadgroup staging buffer or extra barrier needed for the epilogue.
{
const int qid = int(simd_lid) / 4;
const int frag_row = (qid & 4) + ((int(simd_lid) / 2) % 4);
const int frag_col0 = (qid & 2) * 2 + (int(simd_lid) % 2) * 2;
for (int i = 0; i < TM; i++) {
int row = y_row0 + int(sm_row) * 32 + i * 8 + frag_row;
if (row >= M) continue;
for (int j = 0; j < TN; j++) {
int col0 = y_col0 + int(sm_col) * 16 + j * 8 + frag_col0;
device bfloat16_t* ydst = y + (long)row * N + col0;
thread float2& elems = reinterpret_cast<thread float2&>(acc[i][j].thread_elements());
if (col0 < N) { ydst[0] = bfloat16_t(elems[0]); }
if (col0 + 1 < N) { ydst[1] = bfloat16_t(elems[1]); }
}
}
}
"#;
const TILE_N: i32 = 32;
const TILE_M: i32 = 64;
const THREADGROUP_X: i32 = 128;
fn env_enabled() -> bool {
if std::env::var_os("MLX_BF16_QMM_FP16").is_some_and(|v| v != "0" && !v.is_empty()) {
return false;
}
std::env::var_os("ONNXRUNTIME_EP_MLX_BF16_QMM_FP16")
.map(|v| v != "0" && !v.is_empty())
.unwrap_or(true)
}
fn kernel_singleton() -> Option<&'static FastMetalKernel> {
static KERNEL: OnceLock<Option<FastMetalKernel>> = OnceLock::new();
KERNEL
.get_or_init(|| {
let mut input_names = VectorString::new();
input_names.append(c"w");
input_names.append(c"scales");
input_names.append(c"biases");
input_names.append(c"x");
let mut output_names = VectorString::new();
output_names.append(c"y");
Some(FastMetalKernel::new(
KERNEL_NAME,
&input_names,
&output_names,
SOURCE,
HEADER,
true,
false,
))
})
.as_ref()
}
fn kernel_apply_lock() -> &'static Mutex<()> {
static LOCK: Mutex<()> = Mutex::new(());
&LOCK
}
#[allow(clippy::too_many_arguments)]
pub fn eligible(
ctx: &TranslationContext,
out_ndim: usize,
m: i32,
k: i32,
big_n: i32,
block: i64,
bits: i64,
x_dtype: mlx::mlx_dtype,
scales_dtype: mlx::mlx_dtype,
biases_dtype: mlx::mlx_dtype,
w_dtype: mlx::mlx_dtype,
) -> bool {
if !env_enabled() {
return false;
}
if !ctx.shape_keyed_compile() {
return false;
}
block == 32
&& bits == 4
&& out_ndim == 2
&& m > 1
&& k % TILE_N == 0
&& big_n % TILE_N == 0
&& x_dtype == mlx::mlx_dtype__MLX_BFLOAT16
&& scales_dtype == mlx::mlx_dtype__MLX_BFLOAT16
&& biases_dtype == mlx::mlx_dtype__MLX_BFLOAT16
&& w_dtype == mlx::mlx_dtype__MLX_UINT32
}
pub fn try_apply(
ctx: &mut TranslationContext,
x: mlx::mlx_array,
w: mlx::mlx_array,
scales: mlx::mlx_array,
biases: mlx::mlx_array,
m: i32,
big_n: i32,
) -> Option<mlx::mlx_array> {
let kernel = kernel_singleton()?;
let mut inputs = VectorArray::new();
inputs.append(w);
inputs.append(scales);
inputs.append(biases);
inputs.append(x);
let mut config = FastMetalKernelConfig::new();
config
.add_output_arg(&[m, big_n], mlx::mlx_dtype__MLX_BFLOAT16)
.ok()?;
let grid_x = THREADGROUP_X * (big_n / TILE_N);
let grid_y = (m + TILE_M - 1) / TILE_M;
config.set_grid(grid_x, grid_y, 1).ok()?;
config.set_thread_group(THREADGROUP_X, 1, 1).ok()?;
let _apply_guard = kernel_apply_lock().lock().ok()?;
let outs = kernel.apply(&inputs, &config, ctx.stream()).ok()?;
if outs.size() != 1 {
return None;
}
let y: Array = outs.get(0);
Some(ctx.keep(y))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::mlx::Stream;
struct Rng(u64);
impl Rng {
fn new(seed: u64) -> Self {
Rng(seed | 1)
}
fn next_u32(&mut self) -> u32 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
(x >> 32) as u32
}
fn next_f32(&mut self) -> f32 {
(self.next_u32() as f32) / (u32::MAX as f32)
}
}
fn bf16_bits(v: f32) -> u16 {
half::bf16::from_f32(v).to_bits()
}
struct Case {
w_words: Vec<u32>, scales: Vec<u16>, biases: Vec<u16>, x: Vec<u16>, y_ref: Vec<f32>, m: i32,
k: i32,
big_n: i32,
nblocks: i32,
}
fn make_case(seed: u64, m: i32, k: i32, big_n: i32) -> Case {
let group = 32i32;
let nblocks = k / group;
let mut rng = Rng::new(seed);
let mut q = vec![0u8; (big_n as usize) * (k as usize)];
for v in q.iter_mut() {
*v = (rng.next_u32() % 16) as u8;
}
let words_per_row = (k / 8) as usize;
let mut w_words = vec![0u32; (big_n as usize) * words_per_row];
for row in 0..big_n as usize {
for word_idx in 0..words_per_row {
let mut word = 0u32;
for j in 0..8usize {
let kk = word_idx * 8 + j;
let nib = q[row * (k as usize) + kk] as u32;
word |= nib << (j * 4);
}
w_words[row * words_per_row + word_idx] = word;
}
}
let mut scales_f = vec![0f32; (big_n as usize) * (nblocks as usize)];
let mut zp_f = vec![0f32; (big_n as usize) * (nblocks as usize)];
for i in 0..scales_f.len() {
scales_f[i] = rng.next_f32() * 0.05 + 0.001;
zp_f[i] = (rng.next_u32() % 16) as f32;
}
let scales: Vec<u16> = scales_f.iter().map(|&v| bf16_bits(v)).collect();
let biases: Vec<u16> = scales_f
.iter()
.zip(zp_f.iter())
.map(|(&s, &zp)| bf16_bits(-zp * s))
.collect();
let mut x_f = vec![0f32; (m as usize) * (k as usize)];
for v in x_f.iter_mut() {
*v = (rng.next_f32() - 0.5) * 1.0;
}
let x: Vec<u16> = x_f.iter().map(|&v| bf16_bits(v)).collect();
let x_bf: Vec<f32> = x
.iter()
.map(|&b| half::bf16::from_bits(b).to_f32())
.collect();
let scales_bf: Vec<f32> = scales
.iter()
.map(|&b| half::bf16::from_bits(b).to_f32())
.collect();
let biases_bf: Vec<f32> = biases
.iter()
.map(|&b| half::bf16::from_bits(b).to_f32())
.collect();
let mut y_ref = vec![0f32; (m as usize) * (big_n as usize)];
for mi in 0..m as usize {
for n in 0..big_n as usize {
let mut acc = 0f32;
for kk in 0..k as usize {
let blk = kk / group as usize;
let nib = q[n * (k as usize) + kk] as f32;
let dq = nib * scales_bf[n * nblocks as usize + blk]
+ biases_bf[n * nblocks as usize + blk];
acc += x_bf[mi * (k as usize) + kk] * dq;
}
y_ref[mi * (big_n as usize) + n] = acc;
}
}
Case {
w_words,
scales,
biases,
x,
y_ref,
m,
k,
big_n,
nblocks,
}
}
fn max_err(got: &[f32], want: &[f32]) -> (f32, f32) {
let mut max_abs = 0f32;
let mut max_ref = 0f32;
for (&g, &w) in got.iter().zip(want.iter()) {
max_abs = max_abs.max((g - w).abs());
max_ref = max_ref.max(w.abs());
}
(max_abs, max_abs / (max_ref + 1e-6))
}
fn read_bf16_output(arr: &Array, count: usize) -> Vec<f32> {
arr.eval();
let ptr = arr.data_bytes() as *const u16;
(0..count)
.map(|i| half::bf16::from_bits(unsafe { *ptr.add(i) }).to_f32())
.collect()
}
fn run_case(case: &Case) {
let _stream = Stream::new_gpu();
let stream_raw = _stream.as_raw();
let w = Array::from_data(
case.w_words.as_ptr() as *const std::os::raw::c_void,
&[case.big_n, case.k / 8],
mlx::mlx_dtype__MLX_UINT32,
);
let scales = Array::from_data(
case.scales.as_ptr() as *const std::os::raw::c_void,
&[case.big_n, case.nblocks],
mlx::mlx_dtype__MLX_BFLOAT16,
);
let biases = Array::from_data(
case.biases.as_ptr() as *const std::os::raw::c_void,
&[case.big_n, case.nblocks],
mlx::mlx_dtype__MLX_BFLOAT16,
);
let x = Array::from_data(
case.x.as_ptr() as *const std::os::raw::c_void,
&[case.m, case.k],
mlx::mlx_dtype__MLX_BFLOAT16,
);
let kernel = kernel_singleton().expect("kernel object should always construct");
let mut inputs = VectorArray::new();
inputs.append(w.as_raw());
inputs.append(scales.as_raw());
inputs.append(biases.as_raw());
inputs.append(x.as_raw());
let mut config = FastMetalKernelConfig::new();
config
.add_output_arg(&[case.m, case.big_n], mlx::mlx_dtype__MLX_BFLOAT16)
.expect("add_output_arg");
let grid_x = THREADGROUP_X * (case.big_n / TILE_N);
let grid_y = (case.m + TILE_M - 1) / TILE_M;
config.set_grid(grid_x, grid_y, 1).expect("set_grid");
config
.set_thread_group(THREADGROUP_X, 1, 1)
.expect("set_thread_group");
config.set_verbose(false).expect("set_verbose");
let outs = kernel
.apply(&inputs, &config, stream_raw)
.expect("fast kernel apply should succeed");
assert_eq!(outs.size(), 1);
let y_fast = outs.get(0);
let fast = read_bf16_output(&y_fast, (case.m as usize) * (case.big_n as usize));
let gs = mlx::mlx_optional_int_ {
value: 32,
has_value: true,
};
let bb = mlx::mlx_optional_int_ {
value: 4,
has_value: true,
};
let mode = c"affine".as_ptr();
let mut res = unsafe { mlx::mlx_array_new() };
let rc = unsafe {
mlx::mlx_quantized_matmul(
&mut res,
x.as_raw(),
w.as_raw(),
scales.as_raw(),
biases.as_raw(),
true,
gs,
bb,
mode,
stream_raw,
)
};
assert_eq!(rc, 0, "mlx_quantized_matmul failed");
let y_ref_arr = Array::from_raw(res);
let ref_out = read_bf16_output(&y_ref_arr, (case.m as usize) * (case.big_n as usize));
let (abs_fast_ref, rel_fast_ref) = max_err(&fast, &ref_out);
let (abs_ref_host, _rel_ref_host) = max_err(&ref_out, &case.y_ref);
let (abs_fast_host, _rel_fast_host) = max_err(&fast, &case.y_ref);
assert!(
rel_fast_ref < 0.05,
"fast kernel vs mlx_quantized_matmul diverges: abs={abs_fast_ref} rel={rel_fast_ref} \
(mlx_quantized_matmul vs host fp32 abs={abs_ref_host}, fast vs host fp32 abs={abs_fast_host})"
);
}
#[test]
fn matches_quantized_matmul_exact_tile() {
run_case(&make_case(1, 32, 32, 32));
}
#[test]
fn matches_quantized_matmul_partial_m_tile() {
run_case(&make_case(2, 4, 32, 32));
run_case(&make_case(3, 7, 64, 64));
}
#[test]
fn matches_quantized_matmul_exact_bm_tile() {
run_case(&make_case(6, 64, 32, 32));
}
#[test]
fn matches_quantized_matmul_larger_shapes() {
run_case(&make_case(4, 64, 128, 256));
run_case(&make_case(5, 100, 256, 512));
}
#[test]
#[ignore = "manual perf benchmark, not a correctness check"]
fn bench_qmm_fast_kernel_vs_reference() {
let stream = Stream::new_gpu();
let stream_raw = stream.as_raw();
let iters = 50usize;
let warmup = 5usize;
for &(m, k, big_n) in &[
(512i32, 6656i32, 19968i32), (128, 4096, 4096),
(256, 4096, 11008),
(512, 4096, 4096),
] {
let case = make_case(42, m, k, big_n);
let w = Array::from_data(
case.w_words.as_ptr() as *const std::os::raw::c_void,
&[case.big_n, case.k / 8],
mlx::mlx_dtype__MLX_UINT32,
);
let scales = Array::from_data(
case.scales.as_ptr() as *const std::os::raw::c_void,
&[case.big_n, case.nblocks],
mlx::mlx_dtype__MLX_BFLOAT16,
);
let biases = Array::from_data(
case.biases.as_ptr() as *const std::os::raw::c_void,
&[case.big_n, case.nblocks],
mlx::mlx_dtype__MLX_BFLOAT16,
);
let x = Array::from_data(
case.x.as_ptr() as *const std::os::raw::c_void,
&[case.m, case.k],
mlx::mlx_dtype__MLX_BFLOAT16,
);
let kernel = kernel_singleton().expect("kernel object should always construct");
let grid_x = THREADGROUP_X * (case.big_n / TILE_N);
let grid_y = (case.m + TILE_M - 1) / TILE_M;
let run_fast = || -> Array {
let mut inputs = VectorArray::new();
inputs.append(w.as_raw());
inputs.append(scales.as_raw());
inputs.append(biases.as_raw());
inputs.append(x.as_raw());
let mut config = FastMetalKernelConfig::new();
config
.add_output_arg(&[case.m, case.big_n], mlx::mlx_dtype__MLX_BFLOAT16)
.unwrap();
config.set_grid(grid_x, grid_y, 1).unwrap();
config.set_thread_group(THREADGROUP_X, 1, 1).unwrap();
let outs = kernel.apply(&inputs, &config, stream_raw).unwrap();
outs.get(0)
};
let run_ref = || -> Array {
let gs = mlx::mlx_optional_int_ {
value: 32,
has_value: true,
};
let bb = mlx::mlx_optional_int_ {
value: 4,
has_value: true,
};
let mode = c"affine".as_ptr();
let mut res = unsafe { mlx::mlx_array_new() };
let rc = unsafe {
mlx::mlx_quantized_matmul(
&mut res,
x.as_raw(),
w.as_raw(),
scales.as_raw(),
biases.as_raw(),
true,
gs,
bb,
mode,
stream_raw,
)
};
assert_eq!(rc, 0);
Array::from_raw(res)
};
for _ in 0..warmup {
let y = run_fast();
unsafe { mlx::mlx_array_eval(y.as_raw()) };
drop(y);
}
let t0 = std::time::Instant::now();
for _ in 0..iters {
let y = run_fast();
unsafe { mlx::mlx_array_eval(y.as_raw()) };
drop(y);
}
let fast_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
for _ in 0..warmup {
let y = run_ref();
unsafe { mlx::mlx_array_eval(y.as_raw()) };
drop(y);
}
let t0 = std::time::Instant::now();
for _ in 0..iters {
let y = run_ref();
unsafe { mlx::mlx_array_eval(y.as_raw()) };
drop(y);
}
let ref_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
println!(
"M={m} K={k} N={big_n}: fast_kernel={fast_ms:.4}ms/iter \
mlx_quantized_matmul={ref_ms:.4}ms/iter speedup={:.2}x",
ref_ms / fast_ms
);
}
}
#[test]
fn eligibility_gate_requires_shape_keyed_compile() {
unsafe { std::env::remove_var("ONNXRUNTIME_EP_MLX_BF16_QMM_FP16") };
let mut plan = crate::engine::Plan::new(Vec::new());
let ctx = crate::engine::TranslationContext::new(
&mut plan,
std::ptr::null(),
std::ptr::null_mut(),
Stream::new_gpu().as_raw(),
);
assert!(!ctx.shape_keyed_compile());
assert!(!eligible(
&ctx,
2,
32,
32,
32,
32,
4,
mlx::mlx_dtype__MLX_BFLOAT16,
mlx::mlx_dtype__MLX_BFLOAT16,
mlx::mlx_dtype__MLX_BFLOAT16,
mlx::mlx_dtype__MLX_UINT32,
));
}
}