use crate::buffer::MlxBuffer;
use crate::device::MlxDevice;
use crate::dtypes::DType;
use crate::encoder::CommandEncoder;
use crate::error::{MlxError, Result};
use crate::kernel_registry::KernelRegistry;
use crate::ops::quantized_matmul_ggml::GgmlType;
#[derive(Debug, Clone, Copy)]
pub struct MulMvExtParams {
pub m: u32,
pub n: u32,
pub k: u32,
pub batch: u32,
pub ggml_type: GgmlType,
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct MulMvExtGpuArgs {
ne00: i32,
ne01: i32,
ne02: i32,
_pad0: u32,
nb00: u64,
nb01: u64,
nb02: u64,
nb03: u64,
ne10: i32,
ne11: i32,
ne12: i32,
_pad1: u32,
nb10: u64,
nb11: u64,
nb12: u64,
nb13: u64,
ne0: i32,
ne1: i32,
r2: i16,
r3: i16,
_pad2: u32,
}
fn pick_nxpsg(k: u32, m: u32) -> i32 {
if k % 256 == 0 && m < 3 {
16
} else if k % 128 == 0 {
8
} else {
4
}
}
fn pick_r1ptg(m: u32) -> Result<i32> {
match m {
2 => Ok(2),
3 | 6 => Ok(3),
4 | 7 | 8 => Ok(4),
5 => Ok(5),
other => Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: unsupported m {} (peer mapping covers 2..=8 only)",
other
))),
}
}
fn kernel_name(ggml_type: GgmlType, r1ptg: i32) -> Result<&'static str> {
Ok(match (ggml_type, r1ptg) {
(GgmlType::Q5_1, 2) => "kernel_mul_mv_ext_q5_1_f32_r1_2",
(GgmlType::Q5_1, 3) => "kernel_mul_mv_ext_q5_1_f32_r1_3",
(GgmlType::Q5_1, 4) => "kernel_mul_mv_ext_q5_1_f32_r1_4",
(GgmlType::Q5_1, 5) => "kernel_mul_mv_ext_q5_1_f32_r1_5",
(GgmlType::IQ4_NL, 2) => "kernel_mul_mv_ext_iq4_nl_f32_r1_2",
(GgmlType::IQ4_NL, 3) => "kernel_mul_mv_ext_iq4_nl_f32_r1_3",
(GgmlType::IQ4_NL, 4) => "kernel_mul_mv_ext_iq4_nl_f32_r1_4",
(GgmlType::IQ4_NL, 5) => "kernel_mul_mv_ext_iq4_nl_f32_r1_5",
(GgmlType::Q4_0, 2) => "kernel_mul_mv_ext_q4_0_f32_r1_2",
(GgmlType::Q4_0, 3) => "kernel_mul_mv_ext_q4_0_f32_r1_3",
(GgmlType::Q4_0, 4) => "kernel_mul_mv_ext_q4_0_f32_r1_4",
(GgmlType::Q4_0, 5) => "kernel_mul_mv_ext_q4_0_f32_r1_5",
(GgmlType::Q5_0, 2) => "kernel_mul_mv_ext_q5_0_f32_r1_2",
(GgmlType::Q5_0, 3) => "kernel_mul_mv_ext_q5_0_f32_r1_3",
(GgmlType::Q5_0, 4) => "kernel_mul_mv_ext_q5_0_f32_r1_4",
(GgmlType::Q5_0, 5) => "kernel_mul_mv_ext_q5_0_f32_r1_5",
(GgmlType::Q8_0, 2) => "kernel_mul_mv_ext_q8_0_f32_r1_2",
(GgmlType::Q8_0, 3) => "kernel_mul_mv_ext_q8_0_f32_r1_3",
(GgmlType::Q8_0, 4) => "kernel_mul_mv_ext_q8_0_f32_r1_4",
(GgmlType::Q8_0, 5) => "kernel_mul_mv_ext_q8_0_f32_r1_5",
(GgmlType::Q4_K, 2) => "kernel_mul_mv_ext_q4_K_f32_r1_2",
(GgmlType::Q4_K, 3) => "kernel_mul_mv_ext_q4_K_f32_r1_3",
(GgmlType::Q4_K, 4) => "kernel_mul_mv_ext_q4_K_f32_r1_4",
(GgmlType::Q4_K, 5) => "kernel_mul_mv_ext_q4_K_f32_r1_5",
(GgmlType::Q5_K, 2) => "kernel_mul_mv_ext_q5_K_f32_r1_2",
(GgmlType::Q5_K, 3) => "kernel_mul_mv_ext_q5_K_f32_r1_3",
(GgmlType::Q5_K, 4) => "kernel_mul_mv_ext_q5_K_f32_r1_4",
(GgmlType::Q5_K, 5) => "kernel_mul_mv_ext_q5_K_f32_r1_5",
(GgmlType::Q6_K, 2) => "kernel_mul_mv_ext_q6_K_f32_r1_2",
(GgmlType::Q6_K, 3) => "kernel_mul_mv_ext_q6_K_f32_r1_3",
(GgmlType::Q6_K, 4) => "kernel_mul_mv_ext_q6_K_f32_r1_4",
(GgmlType::Q6_K, 5) => "kernel_mul_mv_ext_q6_K_f32_r1_5",
(other_type, other_r1) => {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: no kernel for type {:?} × r1ptg {} (supported: Q4_0/Q5_0/Q8_0/Q4_K/Q5_K/Q6_K/Q5_1/IQ4_NL × r1∈{{2,3,4,5}})",
other_type, other_r1
)));
}
})
}
pub fn mul_mv_ext_dispatch(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weight: &MlxBuffer,
input: &MlxBuffer,
output: &MlxBuffer,
params: &MulMvExtParams,
) -> Result<()> {
if weight.dtype() != DType::U8 || input.dtype() != DType::F32 || output.dtype() != DType::F32 {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext requires native U8 GGUF blocks, F32 input, and F32 output; got {:?}/{:?}/{:?}",
weight.dtype(),
input.dtype(),
output.dtype(),
)));
}
for (label, value) in [
("M", params.m),
("N", params.n),
("K", params.k),
("batch", params.batch),
] {
if value > i32::MAX as u32 {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: {label} exceeds the signed Metal ABI"
)));
}
}
if params.batch > i16::MAX as u32 {
return Err(MlxError::InvalidArgument(
"mul_mv_ext: batch exceeds the signed r2 broadcast ABI".into(),
));
}
if params.m == 0 || params.n == 0 || params.k == 0 || params.batch == 0 {
return Err(MlxError::InvalidArgument(
"mul_mv_ext: m, n, k, batch must all be > 0".into(),
));
}
let block_qk = params.ggml_type.block_values();
if params.k % block_qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: k ({}) must be divisible by block QK ({}) for {:?}",
params.k, block_qk, params.ggml_type
)));
}
let r1ptg = pick_r1ptg(params.m)?;
let nxpsg = pick_nxpsg(params.k, params.m);
let nsg: i32 = 2;
let nypsg = 32 / nxpsg;
let r0ptg = nypsg * nsg;
let kname = kernel_name(params.ggml_type, r1ptg)?;
let pipeline = registry
.get_pipeline_with_constants(
kname,
device.metal_device(),
&[],
&[(600, nsg), (601, nxpsg)],
)?
.clone();
let block_bytes_per_row = (params.k as usize / block_qk as usize)
.checked_mul(params.ggml_type.block_bytes() as usize)
.ok_or_else(|| MlxError::InvalidArgument("mul_mv_ext weight row size overflow".into()))?;
let weight_required = (params.n as usize)
.checked_mul(block_bytes_per_row)
.ok_or_else(|| MlxError::InvalidArgument("mul_mv_ext weight size overflow".into()))?;
if weight.data_byte_len() < weight_required {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: weight buffer too small: {} < {} bytes",
weight.data_byte_len(),
weight_required
)));
}
let input_required = (params.batch as usize)
.checked_mul(params.m as usize)
.and_then(|elements| elements.checked_mul(params.k as usize))
.and_then(|elements| elements.checked_mul(DType::F32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("mul_mv_ext input size overflow".into()))?;
if input.data_byte_len() < input_required {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: input buffer too small: {} < {} bytes",
input.data_byte_len(),
input_required
)));
}
let output_required = (params.batch as usize)
.checked_mul(params.m as usize)
.and_then(|elements| elements.checked_mul(params.n as usize))
.and_then(|elements| elements.checked_mul(DType::F32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("mul_mv_ext output size overflow".into()))?;
if output.data_byte_len() < output_required {
return Err(MlxError::InvalidArgument(format!(
"mul_mv_ext: output buffer too small: {} < {} bytes",
output.data_byte_len(),
output_required
)));
}
let nb00 = params.ggml_type.block_bytes() as u64;
let nb01 = block_bytes_per_row as u64;
let nb02 = nb01 * params.n as u64;
let nb10: u64 = 4;
let nb11 = (params.k as u64) * 4;
let nb12 = nb11 * params.m as u64;
let args = MulMvExtGpuArgs {
ne00: params.k as i32,
ne01: params.n as i32,
ne02: 1,
_pad0: 0,
nb00,
nb01,
nb02,
nb03: nb02, ne10: params.k as i32,
ne11: params.m as i32,
ne12: params.batch as i32,
_pad1: 0,
nb10,
nb11,
nb12,
nb13: nb12, ne0: params.n as i32,
ne1: params.m as i32,
r2: params.batch as i16,
r3: 1,
_pad2: 0,
};
use crate::encoder::{as_bytes, KernelArg};
let args_bytes = as_bytes(&args);
let r0_groups = u64::from(params.n).div_ceil(r0ptg as u64);
let r1_groups = u64::from(params.m).div_ceil(r1ptg as u64);
encoder.encode_threadgroups_with_args(
&pipeline,
&[
(0, KernelArg::Bytes(args_bytes)),
(1, KernelArg::Buffer(weight)),
(2, KernelArg::Buffer(input)),
(3, KernelArg::Buffer(output)),
],
crate::MTLSize::new(r0_groups, r1_groups, params.batch as u64),
crate::MTLSize::new(32, nsg as u64, 1),
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pick_nxpsg_matches_peer_logic() {
assert_eq!(pick_nxpsg(128, 1), 8);
assert_eq!(pick_nxpsg(256, 2), 16);
assert_eq!(pick_nxpsg(256, 3), 8);
assert_eq!(pick_nxpsg(64, 4), 4);
assert_eq!(pick_nxpsg(2816, 2), 16);
assert_eq!(pick_nxpsg(2816, 3), 8);
assert_eq!(pick_nxpsg(512, 4), 8);
}
#[test]
fn pick_r1ptg_matches_peer_switch() {
assert_eq!(pick_r1ptg(2).unwrap(), 2);
assert_eq!(pick_r1ptg(3).unwrap(), 3);
assert_eq!(pick_r1ptg(4).unwrap(), 4);
assert_eq!(pick_r1ptg(5).unwrap(), 5);
assert_eq!(pick_r1ptg(6).unwrap(), 3);
assert_eq!(pick_r1ptg(7).unwrap(), 4);
assert_eq!(pick_r1ptg(8).unwrap(), 4);
assert!(pick_r1ptg(1).is_err());
assert!(pick_r1ptg(9).is_err());
}
#[test]
fn kernel_name_covers_all_phase1_combinations() {
for r1 in 2..=5 {
assert!(kernel_name(GgmlType::Q5_1, r1).is_ok());
assert!(kernel_name(GgmlType::IQ4_NL, r1).is_ok());
}
}
#[test]
fn kernel_name_covers_all_extended_combinations() {
for r1 in 2..=5 {
for ggml_type in [
GgmlType::Q4_0,
GgmlType::Q5_0,
GgmlType::Q8_0,
GgmlType::Q4_K,
GgmlType::Q5_K,
GgmlType::Q6_K,
] {
assert!(
kernel_name(ggml_type, r1).is_ok(),
"{ggml_type:?} r1={r1} must have a kernel"
);
}
}
}
#[test]
fn kernel_name_rejects_unsupported_combinations() {
assert!(
kernel_name(GgmlType::Q5_1, 1).is_err(),
"r1=1 not supported by any phase"
);
assert!(
kernel_name(GgmlType::Q5_1, 6).is_err(),
"r1=6 not supported by any phase"
);
assert!(
kernel_name(GgmlType::Q4_0, 0).is_err(),
"r1=0 not supported"
);
assert!(
kernel_name(GgmlType::Q4_0, -1).is_err(),
"r1=-1 not supported"
);
}
}