use crate::buffer::MlxBuffer;
use crate::device::MlxDevice;
use crate::dtypes::DType;
use crate::encoder::{as_bytes, CapturedOpKind, CommandEncoder, DispatchRecord, KernelArg};
use crate::env_flags::{cached_env_default_true, cached_env_eq_one};
use crate::ggml_capability::{
ggml_expert_bytes, plan_expert_auto_route, ExpertAutoPlan, GgmlRoutingPolicy,
GgmlTensorMmPreference,
};
use crate::ggml_routing_policy::ggml_routing_policy_from_environment;
use std::sync::atomic::AtomicI8;
static CACHED_Q6K_ID_MV_NR2: AtomicI8 = AtomicI8::new(-1);
static CACHED_Q8_0_ID_MV_NR2: AtomicI8 = AtomicI8::new(-1);
use crate::error::{MlxError, Result};
use crate::kernel_registry::KernelRegistry;
use crate::ops::dense_mm_capability::is_unavailable_tensor_header;
use crate::ops::quantized_matmul_ggml::GgmlType;
fn checked_byte_extent(label: &str, factors: &[usize]) -> Result<usize> {
factors.iter().try_fold(1usize, |total, factor| {
total.checked_mul(*factor).ok_or_else(|| {
MlxError::InvalidArgument(format!("quantized_matmul_id_ggml: {label} size overflow"))
})
})
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct GgmlMatvecIdGpuParams {
ne00: i64, ne01: i64, ne02: i64, ne10: i64, ne12: i64, ne0: i64, ne1: i64, r2: u32, r3: u32, top_k: u32, n_tokens: u32, expert_stride: i64, }
#[derive(Debug, Clone, Copy)]
pub struct GgmlQuantizedMatmulIdParams {
pub n_tokens: u32,
pub top_k: u32,
pub n: u32,
pub k: u32,
pub n_experts: u32,
pub expert_stride: u64,
pub ggml_type: GgmlType,
}
impl GgmlType {
fn id_kernel_name(self) -> &'static str {
match self {
GgmlType::Q4_0 => "kernel_mul_mv_id_q4_0_f32",
GgmlType::Q8_0 => "kernel_mul_mv_id_q8_0_f32",
GgmlType::Q2_K => "kernel_mul_mv_id_q2_K_f32",
GgmlType::Q3_K => "kernel_mul_mv_id_q3_K_f32",
GgmlType::Q4_K => "kernel_mul_mv_id_q4_K_f32",
GgmlType::Q5_K => "kernel_mul_mv_id_q5_K_f32",
GgmlType::Q6_K => "kernel_mul_mv_id_q6_K_f32",
GgmlType::Q5_1 => "kernel_mul_mv_id_q5_1_f32",
GgmlType::IQ4_NL => "kernel_mul_mv_id_iq4_nl_f32",
GgmlType::F32 | GgmlType::F16 | GgmlType::I16 | GgmlType::I32 => "unsupported",
GgmlType::IQ4_XS => "kernel_mul_mv_id_iq4_xs_f32",
}
}
fn id_mm_kernel_name(self) -> &'static str {
match self {
GgmlType::Q4_0 => "kernel_mul_mm_id_q4_0_f32",
GgmlType::Q8_0 => "kernel_mul_mm_id_q8_0_f32",
GgmlType::Q2_K => "kernel_mul_mm_id_q2_K_f32",
GgmlType::Q3_K => "kernel_mul_mm_id_q3_K_f32",
GgmlType::Q5_K => "kernel_mul_mm_id_q5_K_f32",
GgmlType::Q6_K => "kernel_mul_mm_id_q6_K_f32",
GgmlType::Q4_K => "kernel_mul_mm_id_q4_K_f32",
GgmlType::Q5_1 => "kernel_mul_mm_id_q5_1_f32",
GgmlType::IQ4_NL => "kernel_mul_mm_id_iq4_nl_f32",
GgmlType::F32 | GgmlType::F16 | GgmlType::I16 | GgmlType::I32 => "unsupported",
GgmlType::IQ4_XS => "kernel_mul_mm_id_iq4_xs_f32",
}
}
fn id_mm_tensor_kernel_name(self) -> &'static str {
match self {
GgmlType::Q4_0 => "kernel_mul_mm_id_q4_0_tensor_f32",
GgmlType::Q8_0 => "kernel_mul_mm_id_q8_0_tensor_f32",
GgmlType::Q2_K => "kernel_mul_mm_id_q2_K_tensor_f32",
GgmlType::Q3_K => "kernel_mul_mm_id_q3_K_tensor_f32",
GgmlType::Q5_K => "kernel_mul_mm_id_q5_K_tensor_f32",
GgmlType::Q6_K => "kernel_mul_mm_id_q6_K_tensor_f32",
GgmlType::Q4_K => "kernel_mul_mm_id_q4_K_tensor_f32",
GgmlType::Q5_1 => "kernel_mul_mm_id_q5_1_tensor_f32",
GgmlType::IQ4_NL => "kernel_mul_mm_id_iq4_nl_tensor_f32",
GgmlType::F32 | GgmlType::F16 | GgmlType::I16 | GgmlType::I32 => "unsupported",
GgmlType::IQ4_XS => "kernel_mul_mm_id_iq4_xs_tensor_f32",
}
}
}
#[inline]
fn has_mm_id_map0(top_k: u32) -> bool {
matches!(top_k, 1 | 6 | 8)
}
fn required_expert_weight_bytes(
ggml_type: GgmlType,
n_experts: u32,
n: u32,
k: u32,
expert_stride: u64,
) -> Result<usize> {
if expert_stride > i64::MAX as u64 {
return Err(MlxError::InvalidArgument(
"expert stride exceeds the signed Metal kernel ABI".into(),
));
}
let bytes = ggml_expert_bytes(ggml_type, n_experts, n, k, expert_stride)?;
usize::try_from(bytes)
.map_err(|_| MlxError::InvalidArgument("expert GGUF bytes exceed usize".into()))
}
fn probe_tensor_mm_id(registry: &mut KernelRegistry, device: &MlxDevice) -> Result<bool> {
let probe = registry.probe_optional_pipeline(
"kernel_mul_mm_id_q4_0_tensor_f32",
device.metal_device(),
device.registry_id(),
is_unavailable_tensor_header,
)?;
if probe.newly_probed && std::env::var("MLX_LOG_TENSOR_PROBE").is_ok() {
eprintln!(
"[mlx-native] tensor_mm_id probe: {}",
if probe.available {
"OK (using tensor variant for MoE)"
} else {
"FAILED (falling back to simdgroup MMA)"
}
);
}
Ok(probe.available)
}
pub(crate) fn expert_routing_policy_from_environment() -> GgmlRoutingPolicy {
GgmlRoutingPolicy {
expert_mm_threshold: mm_id_routing_threshold(),
expert_q6k_mv_nr2: cached_env_default_true(&CACHED_Q6K_ID_MV_NR2, "HF2Q_Q6K_ID_MV_NR2"),
expert_q8_0_mv_nr2: cached_env_eq_one(&CACHED_Q8_0_ID_MV_NR2, "HF2Q_Q8_0_ID_MV_NR2"),
expert_tensor_mm: if std::env::var("HF2Q_DISABLE_TENSOR_MM_ID").is_ok() {
GgmlTensorMmPreference::ForceSimd
} else {
GgmlTensorMmPreference::AutoProbe
},
..GgmlRoutingPolicy::default()
}
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
) -> Result<()> {
let routing = ggml_routing_policy_from_environment();
quantized_matmul_id_ggml_with_policy(
encoder, registry, device, input, weight, ids, output, params, &routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_with_policy(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
quantized_matmul_id_ggml_impl(
encoder, registry, device, input, weight, ids, output, params, false, routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_mv(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
) -> Result<()> {
let routing = ggml_routing_policy_from_environment();
quantized_matmul_id_ggml_mv_with_policy(
encoder, registry, device, input, weight, ids, output, params, &routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_mv_with_policy(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
quantized_matmul_id_ggml_impl(
encoder, registry, device, input, weight, ids, output, params, true, routing,
)
}
#[allow(clippy::too_many_arguments)]
fn quantized_matmul_id_ggml_impl(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
force_mv: bool,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
let qk = params.ggml_type.block_values();
if params.n_tokens == 0 || params.k == 0 || params.n == 0 {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml: n_tokens, K, and N must all be > 0".into(),
));
}
if params.top_k == 0 || params.top_k > params.n_experts {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml: top_k must be in 1..=n_experts".into(),
));
}
if params.n_experts == 0 {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml: n_experts must be > 0".into(),
));
}
if params.k % qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml: K ({}) must be divisible by block QK ({})",
params.k, qk
)));
}
let expected_input_bytes = checked_byte_extent(
"input",
&[
params.n_tokens as usize,
params.k as usize,
DType::F32.size_of(),
],
)?;
if input.data_byte_len() < expected_input_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml: input buffer too small: expected {} bytes for [{} x {}] f32, got {}",
expected_input_bytes, params.n_tokens, params.k, input.data_byte_len()
)));
}
let total_weight_bytes = required_expert_weight_bytes(
params.ggml_type,
params.n_experts,
params.n,
params.k,
params.expert_stride,
)?;
if weight.data_byte_len() < total_weight_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml: weight buffer too small: expected {} bytes for {} experts, got {}",
total_weight_bytes, params.n_experts, weight.data_byte_len()
)));
}
let total_rows = (params.n_tokens as usize)
.checked_mul(params.top_k as usize)
.ok_or_else(|| MlxError::InvalidArgument("expert row count overflow".into()))?;
let expected_ids_bytes = checked_byte_extent("ids", &[total_rows, DType::U32.size_of()])?;
if ids.data_byte_len() < expected_ids_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml: ids buffer too small: expected {} bytes for [{} * {}] u32, got {}",
expected_ids_bytes, params.n_tokens, params.top_k, ids.data_byte_len()
)));
}
let expected_output_bytes = checked_byte_extent(
"output",
&[total_rows, params.n as usize, DType::F32.size_of()],
)?;
if output.data_byte_len() < expected_output_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml: output buffer too small: expected {} bytes for [{} x {}] f32, got {}",
expected_output_bytes, total_rows, params.n, output.data_byte_len()
)));
}
if plan_expert_auto_route(params.n_tokens, params.top_k, params.k, force_mv, routing)
== ExpertAutoPlan::Mm
{
if std::env::var("HF2Q_LOG_MM_ID_ROUTE").is_ok() {
eprintln!(
"[mlx-native adr-022 AC-4] dispatch_id_mm engaged: type={:?} \
n_tokens={} top_k={} k={} n={} n_experts={}",
params.ggml_type,
params.n_tokens,
params.top_k,
params.k,
params.n,
params.n_experts,
);
}
return dispatch_id_mm(
encoder, registry, device, input, weight, ids, output, params, routing,
);
}
dispatch_id_mv(
encoder, registry, device, input, weight, ids, output, params, routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_pooled(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
) -> Result<()> {
let routing = ggml_routing_policy_from_environment();
quantized_matmul_id_ggml_pooled_with_policy(
encoder, registry, device, input, weight, ids, output, scratch, params, &routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_pooled_with_policy(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
quantized_matmul_id_ggml_pooled_impl(
encoder,
registry,
device,
input,
weight,
ids,
output,
scratch,
params,
IdMmInputLayout::SharedPerToken,
routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_pooled_pair(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
first_weight: &MlxBuffer,
second_weight: &MlxBuffer,
ids: &MlxBuffer,
first_output: &MlxBuffer,
second_output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
) -> Result<()> {
let routing = ggml_routing_policy_from_environment();
quantized_matmul_id_ggml_pooled_pair_with_policy(
encoder,
registry,
device,
input,
first_weight,
second_weight,
ids,
first_output,
second_output,
scratch,
params,
&routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_pooled_pair_with_policy(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
first_weight: &MlxBuffer,
second_weight: &MlxBuffer,
ids: &MlxBuffer,
first_output: &MlxBuffer,
second_output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
if plan_expert_auto_route(params.n_tokens, params.top_k, params.k, false, routing)
!= ExpertAutoPlan::Mm
{
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair requires the mm_id route: n_tokens={} threshold={} top_k={} k={}",
params.n_tokens,
routing.expert_mm_threshold,
params.top_k,
params.k,
)));
}
if params.n == 0
|| params.n_experts == 0
|| params.top_k == 0
|| params.top_k > params.n_experts
{
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml_pooled_pair requires N > 0 and top_k in 1..=n_experts".into(),
));
}
let qk = params.ggml_type.block_values();
if params.k % qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"K ({}) must be divisible by block QK ({})",
params.k, qk,
)));
}
let total_weight_bytes = required_expert_weight_bytes(
params.ggml_type,
params.n_experts,
params.n,
params.k,
params.expert_stride,
)?;
for (name, weight) in [("first", first_weight), ("second", second_weight)] {
if weight.data_byte_len() < total_weight_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair {name} weight buffer too small: expected {} bytes, got {}",
total_weight_bytes,
weight.data_byte_len(),
)));
}
}
let expected_output_bytes = (params.n_tokens as usize)
.checked_mul(params.top_k as usize)
.and_then(|rows| rows.checked_mul(params.n as usize))
.and_then(|elements| elements.checked_mul(DType::F32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("pair output byte count overflow".into()))?;
for (name, output) in [("first", first_output), ("second", second_output)] {
if output.data_byte_len() < expected_output_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair {name} output buffer too small: expected {} bytes, got {}",
expected_output_bytes,
output.data_byte_len(),
)));
}
}
let expected_input_bytes = (params.n_tokens as usize)
.checked_mul(params.k as usize)
.and_then(|elements| elements.checked_mul(DType::F32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("pair input byte count overflow".into()))?;
if input.data_byte_len() < expected_input_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair input buffer too small: expected {} bytes, got {}",
expected_input_bytes,
input.data_byte_len(),
)));
}
let expected_ids_bytes = (params.n_tokens as usize)
.checked_mul(params.top_k as usize)
.and_then(|elements| elements.checked_mul(DType::U32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("pair ids byte count overflow".into()))?;
if ids.data_byte_len() < expected_ids_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair ids buffer too small: expected {} bytes, got {}",
expected_ids_bytes,
ids.data_byte_len(),
)));
}
scratch.check_capacity(params.n_experts, params.n_tokens)?;
let expected_htpe_bytes = (params.n_experts as usize)
.checked_mul(DType::U32.size_of())
.ok_or_else(|| MlxError::InvalidArgument("pair htpe byte count overflow".into()))?;
let expected_hids_bytes = (params.n_experts as usize)
.checked_mul(params.n_tokens as usize)
.and_then(|elements| elements.checked_mul(DType::U32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("pair hids byte count overflow".into()))?;
for (name, buffer, expected) in [
("htpe", &scratch.htpe, expected_htpe_bytes),
("hids", &scratch.hids, expected_hids_bytes),
] {
if buffer.data_byte_len() < expected {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair {name} scratch buffer too small: expected {} bytes, got {}",
expected,
buffer.data_byte_len(),
)));
}
}
let range = |buffer: &MlxBuffer, extent: usize| {
let start = (buffer.contents_ptr() as usize).saturating_add(buffer.byte_offset() as usize);
(start, start.saturating_add(extent))
};
let overlaps =
|left: (usize, usize), right: (usize, usize)| left.0 < right.1 && right.0 < left.1;
let first_output_range = range(first_output, expected_output_bytes);
let second_output_range = range(second_output, expected_output_bytes);
if overlaps(first_output_range, second_output_range) {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml_pooled_pair output ranges must not overlap".into(),
));
}
let scratch_ranges = [
("htpe", range(&scratch.htpe, expected_htpe_bytes)),
("hids", range(&scratch.hids, expected_hids_bytes)),
];
if overlaps(scratch_ranges[0].1, scratch_ranges[1].1) {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml_pooled_pair scratch ranges must not overlap".into(),
));
}
let immutable_read_ranges = [
("input", range(input, expected_input_bytes)),
("first weight", range(first_weight, total_weight_bytes)),
("second weight", range(second_weight, total_weight_bytes)),
("ids", range(ids, expected_ids_bytes)),
];
for (scratch_name, scratch_range) in scratch_ranges {
if let Some((read_name, _)) = immutable_read_ranges
.iter()
.find(|(_, read_range)| overlaps(scratch_range, *read_range))
{
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair {scratch_name} scratch range must not overlap {read_name}",
)));
}
}
let read_ranges = [
immutable_read_ranges[0],
immutable_read_ranges[1],
immutable_read_ranges[2],
immutable_read_ranges[3],
("htpe scratch", scratch_ranges[0].1),
("hids scratch", scratch_ranges[1].1),
];
for (output_name, output_range) in [
("first", first_output_range),
("second", second_output_range),
] {
if let Some((read_name, _)) = read_ranges
.iter()
.find(|(_, read_range)| overlaps(output_range, *read_range))
{
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_pair {output_name} output range must not overlap {read_name}",
)));
}
}
dispatch_id_mm_pooled_with_layout(
encoder,
registry,
device,
input,
first_weight,
ids,
first_output,
scratch,
params,
IdMmInputLayout::SharedPerToken,
false,
routing,
)?;
dispatch_id_mm_pooled_with_layout(
encoder,
registry,
device,
input,
second_weight,
ids,
second_output,
scratch,
params,
IdMmInputLayout::SharedPerToken,
true,
routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_pooled_slotted(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
) -> Result<()> {
let routing = ggml_routing_policy_from_environment();
quantized_matmul_id_ggml_pooled_slotted_with_policy(
encoder, registry, device, input, weight, ids, output, scratch, params, &routing,
)
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_ggml_pooled_slotted_with_policy(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
quantized_matmul_id_ggml_pooled_impl(
encoder,
registry,
device,
input,
weight,
ids,
output,
scratch,
params,
IdMmInputLayout::Slotted,
routing,
)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum IdMmInputLayout {
SharedPerToken,
Slotted,
}
#[allow(clippy::too_many_arguments)]
fn quantized_matmul_id_ggml_pooled_impl(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
input_layout: IdMmInputLayout,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
let qk = params.ggml_type.block_values();
if params.n_tokens == 0 || params.k == 0 || params.n == 0 {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml_pooled: n_tokens, K, and N must all be > 0".into(),
));
}
if params.top_k == 0 || params.top_k > params.n_experts {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml_pooled: top_k must be in 1..=n_experts".into(),
));
}
if params.n_experts == 0 {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_ggml_pooled: n_experts must be > 0".into(),
));
}
if params.k % qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled: K ({}) must be divisible by block QK ({})",
params.k, qk
)));
}
let input_rows = match input_layout {
IdMmInputLayout::SharedPerToken => params.n_tokens as usize,
IdMmInputLayout::Slotted => (params.n_tokens as usize)
.checked_mul(params.top_k as usize)
.ok_or_else(|| MlxError::InvalidArgument("slotted input row count overflow".into()))?,
};
let expected_input_bytes = input_rows
.checked_mul(params.k as usize)
.and_then(|elements| elements.checked_mul(DType::F32.size_of()))
.ok_or_else(|| MlxError::InvalidArgument("input byte count overflow".into()))?;
if input.data_byte_len() < expected_input_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled: input buffer too small: expected {} bytes for [{} x {}] f32 {:?} input, got {}",
expected_input_bytes, input_rows, params.k, input_layout, input.data_byte_len()
)));
}
let total_weight_bytes = required_expert_weight_bytes(
params.ggml_type,
params.n_experts,
params.n,
params.k,
params.expert_stride,
)?;
if weight.data_byte_len() < total_weight_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled: weight buffer too small: expected {} bytes for {} experts, got {}",
total_weight_bytes, params.n_experts, weight.data_byte_len()
)));
}
let total_rows = (params.n_tokens as usize)
.checked_mul(params.top_k as usize)
.ok_or_else(|| MlxError::InvalidArgument("expert row count overflow".into()))?;
let expected_ids_bytes = checked_byte_extent("ids", &[total_rows, DType::U32.size_of()])?;
if ids.data_byte_len() < expected_ids_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled: ids buffer too small: expected {} bytes for [{} * {}] u32, got {}",
expected_ids_bytes, params.n_tokens, params.top_k, ids.data_byte_len()
)));
}
let expected_output_bytes = checked_byte_extent(
"output",
&[total_rows, params.n as usize, DType::F32.size_of()],
)?;
if output.data_byte_len() < expected_output_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled: output buffer too small: expected {} bytes for [{} x {}] f32, got {}",
expected_output_bytes, total_rows, params.n, output.data_byte_len()
)));
}
if plan_expert_auto_route(params.n_tokens, params.top_k, params.k, false, routing)
== ExpertAutoPlan::Mm
{
if std::env::var("HF2Q_LOG_MM_ID_ROUTE").is_ok() {
eprintln!(
"[mlx-native adr-022 AC-4 pooled] dispatch_id_mm_pooled engaged: \
type={:?} n_tokens={} top_k={} k={} n={} n_experts={}",
params.ggml_type,
params.n_tokens,
params.top_k,
params.k,
params.n,
params.n_experts,
);
}
return dispatch_id_mm_pooled_with_layout(
encoder,
registry,
device,
input,
weight,
ids,
output,
scratch,
params,
input_layout,
false,
routing,
);
}
if input_layout == IdMmInputLayout::Slotted {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml_pooled_slotted requires the mm_id route: n_tokens={} threshold={} top_k={} k={}",
params.n_tokens, routing.expert_mm_threshold, params.top_k, params.k,
)));
}
dispatch_id_mv(
encoder, registry, device, input, weight, ids, output, params, routing,
)
}
pub const MM_ID_ROUTING_THRESHOLD: u32 = 32;
fn mm_id_routing_threshold() -> u32 {
static CACHED: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
*CACHED.get_or_init(|| {
std::env::var("HF2Q_MM_ID_ROUTING_THRESHOLD")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.map(|v| {
if std::env::var("MLX_LOG_TENSOR_PROBE").is_ok() {
eprintln!(
"[mlx-native] mm_id_routing_threshold: OVERRIDE via HF2Q_MM_ID_ROUTING_THRESHOLD={v} (default {})",
MM_ID_ROUTING_THRESHOLD
);
}
v
})
.unwrap_or(MM_ID_ROUTING_THRESHOLD)
})
}
#[allow(clippy::too_many_arguments)]
fn dispatch_id_mv(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
let total_rows = (params.n_tokens as usize) * (params.top_k as usize);
let use_q6k_id_nr2 = matches!(params.ggml_type, GgmlType::Q6_K) && routing.expert_q6k_mv_nr2;
let use_q8_0_id_nr2 = matches!(params.ggml_type, GgmlType::Q8_0) && routing.expert_q8_0_mv_nr2;
let kernel_name = if use_q6k_id_nr2 {
"kernel_mul_mv_id_q6_K_f32_nr2"
} else if use_q8_0_id_nr2 {
"kernel_mul_mv_id_q8_0_f32_nr2"
} else {
params.ggml_type.id_kernel_name()
};
let pipeline = registry.get_pipeline(kernel_name, device.metal_device())?;
let gpu_params = GgmlMatvecIdGpuParams {
ne00: params.k as i64,
ne01: params.n as i64,
ne02: 1,
ne10: params.k as i64,
ne12: 1,
ne0: params.n as i64,
ne1: total_rows as i64,
r2: 1,
r3: 1,
top_k: params.top_k,
n_tokens: params.n_tokens,
expert_stride: params.expert_stride as i64,
};
let (nth0, nth1, align) = match params.ggml_type {
GgmlType::Q4_0
| GgmlType::Q8_0
| GgmlType::Q5_1
| GgmlType::IQ4_NL
| GgmlType::IQ4_XS => (8u64, 8u64, 8usize),
GgmlType::Q2_K => (2u64, 32u64, 8usize),
GgmlType::Q3_K => (2u64, 32u64, 4usize),
GgmlType::Q4_K | GgmlType::Q5_K | GgmlType::Q6_K => (2u64, 32u64, 2usize),
GgmlType::F32
| GgmlType::F16
| GgmlType::I16
| GgmlType::I32 => {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_ggml does not support {:?}",
params.ggml_type
)));
}
};
let align = if use_q6k_id_nr2 { 4usize } else { align };
let (nth0, nth1, align) = if use_q8_0_id_nr2 {
(32u64, 4u64, 2usize)
} else {
(nth0, nth1, align)
};
let n = params.n as usize;
let m = total_rows;
let threadgroups = metal::MTLSize::new(div_ceil(n, align) as u64, m as u64, 1);
let threads_per_tg = metal::MTLSize::new(nth0, nth1, 1);
if use_q8_0_id_nr2 {
let smem_bytes: u64 = 2 * 32 * std::mem::size_of::<f32>() as u64;
encoder.encode_threadgroups_with_args_and_shared(
pipeline,
&[
(0, KernelArg::Buffer(weight)),
(1, KernelArg::Buffer(input)),
(2, KernelArg::Buffer(output)),
(3, KernelArg::Buffer(ids)),
(4, KernelArg::Bytes(as_bytes(&gpu_params))),
],
&[(0, smem_bytes)],
threadgroups,
threads_per_tg,
);
} else {
encoder.encode_threadgroups_with_args(
pipeline,
&[
(0, KernelArg::Buffer(weight)),
(1, KernelArg::Buffer(input)),
(2, KernelArg::Buffer(output)),
(3, KernelArg::Buffer(ids)),
(4, KernelArg::Bytes(as_bytes(&gpu_params))),
],
threadgroups,
threads_per_tg,
);
}
Ok(())
}
pub fn build_q6k_id_nr2_m1_record(
registry: &mut KernelRegistry,
device: &metal::DeviceRef,
n: u32,
k: u32,
top_k: u32,
expert_stride: u64,
) -> Result<Option<DispatchRecord>> {
let routing = ggml_routing_policy_from_environment();
build_q6k_id_nr2_m1_record_with_policy(registry, device, n, k, top_k, expert_stride, &routing)
}
pub fn build_q6k_id_nr2_m1_record_with_policy(
registry: &mut KernelRegistry,
device: &metal::DeviceRef,
n: u32,
k: u32,
top_k: u32,
expert_stride: u64,
routing: &GgmlRoutingPolicy,
) -> Result<Option<DispatchRecord>> {
if expert_stride > i64::MAX as u64 {
return Err(MlxError::InvalidArgument(
"expert stride exceeds the signed Metal kernel ABI".into(),
));
}
if !routing.expert_q6k_mv_nr2 {
return Ok(None);
}
let pipeline = registry
.get_pipeline("kernel_mul_mv_id_q6_K_f32_nr2", device)?
.clone();
let gpu_params = GgmlMatvecIdGpuParams {
ne00: k as i64,
ne01: n as i64,
ne02: 1,
ne10: k as i64,
ne12: 1,
ne0: n as i64,
ne1: top_k as i64,
r2: 1,
r3: 1,
top_k,
n_tokens: 1,
expert_stride: expert_stride as i64,
};
let params_bytes = as_bytes(&gpu_params).to_vec();
const ALIGN: u32 = 4;
let threadgroups =
metal::MTLSize::new(div_ceil(n as usize, ALIGN as usize) as u64, top_k as u64, 1);
let threads_per_tg = metal::MTLSize::new(2, 32, 1);
Ok(Some(DispatchRecord {
pipeline,
threadgroups,
threads_per_tg,
threadgroup_mem: Vec::new(), params_bytes,
params_slot: 4,
buffer_slots: vec![0, 1, 2, 3], op_kind: CapturedOpKind::Other,
kernel_name: "kernel_mul_mv_id_q6_K_f32_nr2".to_string(),
}))
}
pub fn build_q8_0_id_decode_record(
registry: &mut KernelRegistry,
device: &metal::DeviceRef,
n: u32,
k: u32,
real_top_k: u32,
expert_stride: u64,
) -> Result<Option<DispatchRecord>> {
let routing = ggml_routing_policy_from_environment();
build_q8_0_id_decode_record_with_policy(
registry,
device,
n,
k,
real_top_k,
expert_stride,
&routing,
)
}
pub fn build_q8_0_id_decode_record_with_policy(
registry: &mut KernelRegistry,
device: &metal::DeviceRef,
n: u32,
k: u32,
real_top_k: u32,
expert_stride: u64,
routing: &GgmlRoutingPolicy,
) -> Result<Option<DispatchRecord>> {
if expert_stride > i64::MAX as u64 {
return Err(MlxError::InvalidArgument(
"expert stride exceeds the signed Metal kernel ABI".into(),
));
}
if routing.expert_q8_0_mv_nr2 {
return Ok(None);
}
let pipeline = registry
.get_pipeline("kernel_mul_mv_id_q8_0_f32", device)?
.clone();
let gpu_params = GgmlMatvecIdGpuParams {
ne00: k as i64,
ne01: n as i64,
ne02: 1,
ne10: k as i64,
ne12: 1,
ne0: n as i64,
ne1: real_top_k as i64,
r2: 1,
r3: 1,
top_k: 1,
n_tokens: real_top_k,
expert_stride: expert_stride as i64,
};
let params_bytes = as_bytes(&gpu_params).to_vec();
const ALIGN: u32 = 8;
let threadgroups = metal::MTLSize::new(
div_ceil(n as usize, ALIGN as usize) as u64,
real_top_k as u64,
1,
);
let threads_per_tg = metal::MTLSize::new(8, 8, 1);
Ok(Some(DispatchRecord {
pipeline,
threadgroups,
threads_per_tg,
threadgroup_mem: Vec::new(), params_bytes,
params_slot: 4,
buffer_slots: vec![0, 1, 2, 3], op_kind: CapturedOpKind::Other,
kernel_name: "kernel_mul_mv_id_q8_0_f32".to_string(),
}))
}
#[allow(clippy::too_many_arguments)]
pub fn quantized_matmul_id_swiglu_q4_0(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
gate: &MlxBuffer,
up: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
) -> Result<()> {
if params.ggml_type != GgmlType::Q4_0 {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: expected Q4_0, got {:?}",
params.ggml_type
)));
}
let qk = GgmlType::Q4_0.block_values();
if params.n_tokens == 0 || params.k == 0 || params.n == 0 || params.n_experts == 0 {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_swiglu_q4_0: n_tokens, K, N, and n_experts must all be > 0".into(),
));
}
if params.top_k == 0 || params.top_k > params.n_experts {
return Err(MlxError::InvalidArgument(
"quantized_matmul_id_swiglu_q4_0: top_k must be in 1..=n_experts".into(),
));
}
if params.k % qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: K ({}) must be divisible by block QK ({})",
params.k, qk
)));
}
let expected_weight_bytes = required_expert_weight_bytes(
params.ggml_type,
params.n_experts,
params.n,
params.k,
params.expert_stride,
)?;
if weight.data_byte_len() < expected_weight_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: weight buffer too small: expected {} bytes, got {}",
expected_weight_bytes,
weight.data_byte_len(),
)));
}
let total_rows = (params.n_tokens as usize)
.checked_mul(params.top_k as usize)
.ok_or_else(|| MlxError::InvalidArgument("expert row count overflow".into()))?;
let expected_in_bytes = checked_byte_extent(
"swiglu input",
&[total_rows, params.k as usize, DType::F32.size_of()],
)?;
if gate.data_byte_len() < expected_in_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: gate buffer too small: expected {} bytes, got {}",
expected_in_bytes,
gate.data_byte_len()
)));
}
if up.data_byte_len() < expected_in_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: up buffer too small: expected {} bytes, got {}",
expected_in_bytes,
up.data_byte_len()
)));
}
let expected_out_bytes = checked_byte_extent(
"swiglu output",
&[total_rows, params.n as usize, DType::F32.size_of()],
)?;
if output.data_byte_len() < expected_out_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: output buffer too small: expected {} bytes, got {}",
expected_out_bytes,
output.data_byte_len()
)));
}
let expected_ids_bytes =
checked_byte_extent("swiglu ids", &[total_rows, DType::U32.size_of()])?;
if ids.data_byte_len() < expected_ids_bytes {
return Err(MlxError::InvalidArgument(format!(
"quantized_matmul_id_swiglu_q4_0: ids buffer too small: expected {} bytes, got {}",
expected_ids_bytes,
ids.data_byte_len(),
)));
}
let pipeline =
registry.get_pipeline("kernel_mul_mv_id_q4_0_f32_swiglu", device.metal_device())?;
let gpu_params = GgmlMatvecIdGpuParams {
ne00: params.k as i64,
ne01: params.n as i64,
ne02: 1,
ne10: params.k as i64,
ne12: 1,
ne0: params.n as i64,
ne1: total_rows as i64,
r2: 1,
r3: 1,
top_k: params.top_k,
n_tokens: params.n_tokens,
expert_stride: params.expert_stride as i64,
};
let (nth0, nth1, align) = (8u64, 8u64, 8usize);
let n = params.n as usize;
let m = total_rows;
let threadgroups = metal::MTLSize::new(div_ceil(n, align) as u64, m as u64, 1);
let threads_per_tg = metal::MTLSize::new(nth0, nth1, 1);
encoder.encode_threadgroups_with_args(
pipeline,
&[
(0, KernelArg::Buffer(weight)),
(1, KernelArg::Buffer(gate)),
(2, KernelArg::Buffer(up)),
(3, KernelArg::Buffer(output)),
(4, KernelArg::Buffer(ids)),
(5, KernelArg::Bytes(as_bytes(&gpu_params))),
],
threadgroups,
threads_per_tg,
);
Ok(())
}
pub struct IdMmScratch {
pub htpe: MlxBuffer,
pub hids: MlxBuffer,
n_experts_cap: u32,
n_tokens_cap: u32,
}
impl IdMmScratch {
pub fn alloc(device: &MlxDevice, n_experts: u32, max_n_tokens: u32) -> Result<Self> {
let htpe = device.alloc_buffer(
(n_experts as usize) * DType::U32.size_of(),
DType::U32,
vec![n_experts as usize],
)?;
let hids = device.alloc_buffer(
(n_experts as usize) * (max_n_tokens as usize) * DType::U32.size_of(),
DType::U32,
vec![n_experts as usize, max_n_tokens as usize],
)?;
Ok(Self {
htpe,
hids,
n_experts_cap: n_experts,
n_tokens_cap: max_n_tokens,
})
}
fn check_capacity(&self, n_experts: u32, n_tokens: u32) -> Result<()> {
if n_experts > self.n_experts_cap {
return Err(MlxError::InvalidArgument(format!(
"IdMmScratch: n_experts ({}) > cap ({})",
n_experts, self.n_experts_cap,
)));
}
if n_tokens > self.n_tokens_cap {
return Err(MlxError::InvalidArgument(format!(
"IdMmScratch: n_tokens ({}) > cap ({})",
n_tokens, self.n_tokens_cap,
)));
}
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
fn dispatch_id_mm_pooled(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
dispatch_id_mm_pooled_with_layout(
encoder,
registry,
device,
input,
weight,
ids,
output,
scratch,
params,
IdMmInputLayout::SharedPerToken,
false,
routing,
)
}
#[allow(clippy::too_many_arguments)]
fn dispatch_id_mm_pooled_with_layout(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
scratch: &mut IdMmScratch,
params: &GgmlQuantizedMatmulIdParams,
input_layout: IdMmInputLayout,
schedule_prepared: bool,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
scratch.check_capacity(params.n_experts, params.n_tokens)?;
let dispatch = GgmlIdMmDispatchParams {
n_tokens: params.n_tokens,
top_k: params.top_k,
n: params.n,
k: params.k,
n_experts: params.n_experts,
expert_stride: params.expert_stride,
ggml_type: params.ggml_type,
};
dispatch_id_mm_with_layout(
encoder,
registry,
device,
input,
weight,
ids,
&mut scratch.htpe,
&mut scratch.hids,
output,
&dispatch,
input_layout,
schedule_prepared,
routing,
)
}
#[allow(clippy::too_many_arguments)]
fn dispatch_id_mm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlQuantizedMatmulIdParams,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
let mut scratch = IdMmScratch::alloc(device, params.n_experts, params.n_tokens)?;
dispatch_id_mm_pooled(
encoder,
registry,
device,
input,
weight,
ids,
output,
&mut scratch,
params,
routing,
)
}
fn div_ceil(a: usize, b: usize) -> usize {
(a + b - 1) / b
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct GgmlIdMmMap0GpuParams {
ne10: i32, ne11: i32, nb11: u64, nb12: u64, ne21: i32, ne20: i32, nb21: u64, }
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct GgmlIdMmMmGpuParams {
ne00: i32, ne02: i32, nb01: u64, nb02: u64, nb03: u64,
ne11: i32, _pad0: u32,
nb10: u64, nb11: u64, nb12: u64, nb13: u64,
ne20: i32, ne21: i32, ne0: i32, ne1: i32, r2: i16,
r3: i16,
_pad1: u32,
}
#[derive(Debug, Clone, Copy)]
pub struct GgmlIdMmDispatchParams {
pub n_tokens: u32,
pub top_k: u32,
pub n: u32,
pub k: u32,
pub n_experts: u32,
pub expert_stride: u64,
pub ggml_type: GgmlType,
}
impl GgmlIdMmDispatchParams {
pub fn htpe_bytes(&self) -> usize {
(self.n_experts as usize) * DType::U32.size_of()
}
pub fn hids_bytes(&self) -> usize {
(self.n_experts as usize) * (self.n_tokens as usize) * DType::U32.size_of()
}
}
#[doc(hidden)]
#[allow(clippy::too_many_arguments)]
pub fn dispatch_id_mm_for_test(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
htpe: &MlxBuffer,
hids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlIdMmDispatchParams,
) -> Result<()> {
let routing = GgmlRoutingPolicy::default();
dispatch_id_mm_with_layout(
encoder,
registry,
device,
input,
weight,
ids,
htpe,
hids,
output,
params,
IdMmInputLayout::SharedPerToken,
false,
&routing,
)
}
#[doc(hidden)]
#[allow(clippy::too_many_arguments)]
pub fn dispatch_id_mm_with_prepared_schedule_for_test(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
htpe: &MlxBuffer,
hids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlIdMmDispatchParams,
) -> Result<()> {
let routing = GgmlRoutingPolicy::default();
dispatch_id_mm_with_layout(
encoder,
registry,
device,
input,
weight,
ids,
htpe,
hids,
output,
params,
IdMmInputLayout::SharedPerToken,
true,
&routing,
)
}
#[allow(clippy::too_many_arguments)]
fn dispatch_id_mm_with_layout(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
ids: &MlxBuffer,
htpe: &MlxBuffer,
hids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlIdMmDispatchParams,
input_layout: IdMmInputLayout,
schedule_prepared: bool,
routing: &GgmlRoutingPolicy,
) -> Result<()> {
let qk = params.ggml_type.block_values();
match params.ggml_type {
GgmlType::Q4_0
| GgmlType::Q8_0
| GgmlType::Q2_K
| GgmlType::Q3_K
| GgmlType::Q4_K
| GgmlType::Q5_K
| GgmlType::Q6_K
| GgmlType::Q5_1
| GgmlType::IQ4_NL
| GgmlType::IQ4_XS => {}
other => {
return Err(MlxError::InvalidArgument(format!(
"dispatch_id_mm_for_test does not support {:?}",
other
)));
}
}
if params.n_tokens == 0
|| params.k == 0
|| params.n == 0
|| params.top_k == 0
|| params.n_experts == 0
{
return Err(MlxError::InvalidArgument(
"n_tokens, K, N, top_k, n_experts must all be > 0".into(),
));
}
if params.k % qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"K ({}) must be divisible by block QK ({})",
params.k, qk
)));
}
if !has_mm_id_map0(params.top_k) {
return Err(MlxError::InvalidArgument(format!(
"dispatch_id_mm_for_test: top_k {} has no map0 instantiation (need 1, 6, or 8)",
params.top_k
)));
}
let blocks_per_row = params.k / qk;
let block_bytes = params.ggml_type.block_bytes();
let total_weight_bytes = required_expert_weight_bytes(
params.ggml_type,
params.n_experts,
params.n,
params.k,
params.expert_stride,
)?;
if weight.data_byte_len() < total_weight_bytes {
return Err(MlxError::InvalidArgument(format!(
"dispatch_id_mm_for_test: weight buffer too small: expected {} bytes, got {}",
total_weight_bytes,
weight.data_byte_len(),
)));
}
if input.data_byte_len()
< (params.n_tokens as usize) * (params.k as usize) * DType::F32.size_of()
{
return Err(MlxError::InvalidArgument(
"dispatch_id_mm_for_test: input buffer too small".into(),
));
}
let total_rows = (params.n_tokens as usize) * (params.top_k as usize);
if ids.data_byte_len() < total_rows * DType::U32.size_of() {
return Err(MlxError::InvalidArgument(
"dispatch_id_mm_for_test: ids buffer too small".into(),
));
}
if output.data_byte_len() < total_rows * (params.n as usize) * DType::F32.size_of() {
return Err(MlxError::InvalidArgument(
"dispatch_id_mm_for_test: output buffer too small".into(),
));
}
if htpe.data_byte_len() < params.htpe_bytes() {
return Err(MlxError::InvalidArgument(
"dispatch_id_mm_for_test: htpe buffer too small".into(),
));
}
if hids.data_byte_len() < params.hids_bytes() {
return Err(MlxError::InvalidArgument(
"dispatch_id_mm_for_test: hids buffer too small".into(),
));
}
if !schedule_prepared {
let map0_kernel_name = match params.top_k {
1 => "kernel_mul_mm_id_map0_ne20_1",
6 => "kernel_mul_mm_id_map0_ne20_6",
8 => "kernel_mul_mm_id_map0_ne20_8",
other => {
return Err(MlxError::InvalidArgument(format!(
"dispatch_id_mm_for_test: no map0 instantiation for top_k={}",
other
)))
}
};
let map0_pipeline = registry.get_pipeline(map0_kernel_name, device.metal_device())?;
if u64::from(params.n_experts) > map0_pipeline.max_total_threads_per_threadgroup() {
return Err(MlxError::InvalidArgument(format!(
"n_experts ({}) exceeds map0 pipeline threadgroup limit ({})",
params.n_experts,
map0_pipeline.max_total_threads_per_threadgroup(),
)));
}
let map0_params = GgmlIdMmMap0GpuParams {
ne10: params
.n
.try_into()
.map_err(|_| MlxError::InvalidArgument("N out of i32 range".into()))?,
ne11: params.top_k as i32,
nb11: 0,
nb12: 0,
ne21: params.n_tokens as i32,
ne20: params.top_k as i32,
nb21: (params.top_k as u64) * (DType::U32.size_of() as u64),
};
let map0_shmem =
(params.n_experts as u64) * (params.top_k as u64) * std::mem::size_of::<u16>() as u64;
let map0_threadgroups = metal::MTLSize::new(1, 1, 1);
let map0_threads = metal::MTLSize::new(params.n_experts as u64, 1, 1);
if encoder.is_capturing() {
let range = |buffer: &MlxBuffer| {
let start =
(buffer.contents_ptr() as usize).saturating_add(buffer.byte_offset() as usize);
(start, start.saturating_add(buffer.data_byte_len()))
};
encoder.set_pending_buffer_ranges(vec![range(ids)], vec![range(htpe), range(hids)]);
}
encoder.encode_threadgroups_with_args_and_shared(
map0_pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&map0_params))),
(1, KernelArg::Buffer(ids)),
(2, KernelArg::Buffer(htpe)),
(3, KernelArg::Buffer(hids)),
],
&[(0, map0_shmem)],
map0_threadgroups,
map0_threads,
);
encoder.memory_barrier();
}
let use_tensor = routing.expert_tensor_mm == GgmlTensorMmPreference::AutoProbe
&& probe_tensor_mm_id(registry, device)?;
let tensor_name = params.ggml_type.id_mm_tensor_kernel_name();
let mm_kernel_name = if use_tensor && tensor_name != "unsupported" {
tensor_name
} else {
params.ggml_type.id_mm_kernel_name()
};
let mm_pipeline = registry.get_pipeline(mm_kernel_name, device.metal_device())?;
let nb01 = (blocks_per_row as u64) * (block_bytes as u64);
let row_bytes = (params.k as u64) * (DType::F32.size_of() as u64);
let (slot_stride, token_stride) = match input_layout {
IdMmInputLayout::SharedPerToken => (0, row_bytes),
IdMmInputLayout::Slotted => (
row_bytes,
row_bytes
.checked_mul(params.top_k as u64)
.ok_or_else(|| MlxError::InvalidArgument("slotted token stride overflow".into()))?,
),
};
let mm_params = GgmlIdMmMmGpuParams {
ne00: params.k as i32,
ne02: params.n_experts as i32,
nb01,
nb02: params.expert_stride,
nb03: 0,
ne11: params.top_k as i32,
_pad0: 0,
nb10: DType::F32.size_of() as u64,
nb11: slot_stride,
nb12: token_stride,
nb13: 0,
ne20: params.top_k as i32,
ne21: params.n_tokens as i32,
ne0: params.n as i32,
ne1: params.top_k as i32,
r2: 1,
r3: 1,
_pad1: 0,
};
const NR0: u64 = 64;
const NR1: u64 = 32;
const THREADS_PER_TG: u64 = 128;
let mm_threadgroups = metal::MTLSize::new(
(params.n_tokens as u64 + NR1 - 1) / NR1,
(params.n as u64 + NR0 - 1) / NR0,
params.n_experts as u64,
);
let mm_threads = metal::MTLSize::new(THREADS_PER_TG, 1, 1);
const MM_SHMEM_BYTES: u64 = 8192;
if encoder.is_capturing() {
let range = |buffer: &MlxBuffer| {
let start =
(buffer.contents_ptr() as usize).saturating_add(buffer.byte_offset() as usize);
(start, start.saturating_add(buffer.data_byte_len()))
};
encoder.set_pending_buffer_ranges(
vec![range(weight), range(input), range(htpe), range(hids)],
vec![range(output)],
);
}
encoder.encode_threadgroups_with_args_and_shared(
mm_pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&mm_params))),
(1, KernelArg::Buffer(weight)),
(2, KernelArg::Buffer(input)),
(3, KernelArg::Buffer(htpe)),
(4, KernelArg::Buffer(hids)),
(5, KernelArg::Buffer(output)),
],
&[(0, MM_SHMEM_BYTES)],
mm_threadgroups,
mm_threads,
);
Ok(())
}
#[doc(hidden)]
#[allow(clippy::too_many_arguments)]
pub fn dispatch_id_mm_fused_gate_up_silu_for_test(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
gate_w: &MlxBuffer,
up_w: &MlxBuffer,
ids: &MlxBuffer,
htpe: &MlxBuffer,
hids: &MlxBuffer,
output: &MlxBuffer,
params: &GgmlIdMmDispatchParams,
) -> Result<()> {
match params.ggml_type {
GgmlType::Q6_K => {}
other => {
return Err(MlxError::InvalidArgument(format!(
"dispatch_id_mm_fused_gate_up_silu_for_test does not support {:?} (Q6_K only in iter 4)",
other
)));
}
}
if params.n_tokens == 0
|| params.k == 0
|| params.n == 0
|| params.top_k == 0
|| params.n_experts == 0
{
return Err(MlxError::InvalidArgument(
"n_tokens, K, N, top_k, n_experts must all be > 0".into(),
));
}
let qk = params.ggml_type.block_values();
if params.k % qk != 0 {
return Err(MlxError::InvalidArgument(format!(
"K ({}) must be divisible by block QK ({})",
params.k, qk
)));
}
if params.top_k != 1 && params.top_k != 8 {
return Err(MlxError::InvalidArgument(format!(
"fused dispatch: no map0 instantiation for top_k={} (need 1 or 8)",
params.top_k
)));
}
let blocks_per_row = params.k / qk;
let block_bytes = params.ggml_type.block_bytes();
let expected_w = required_expert_weight_bytes(
params.ggml_type,
params.n_experts,
params.n,
params.k,
params.expert_stride,
)?;
if gate_w.data_byte_len() < expected_w {
return Err(MlxError::InvalidArgument(
"fused dispatch: gate_w buffer too small".into(),
));
}
if up_w.data_byte_len() < expected_w {
return Err(MlxError::InvalidArgument(
"fused dispatch: up_w buffer too small".into(),
));
}
if input.data_byte_len()
< (params.n_tokens as usize) * (params.k as usize) * DType::F32.size_of()
{
return Err(MlxError::InvalidArgument(
"fused dispatch: input buffer too small".into(),
));
}
let total_rows = (params.n_tokens as usize) * (params.top_k as usize);
if ids.data_byte_len() < total_rows * DType::U32.size_of() {
return Err(MlxError::InvalidArgument(
"fused dispatch: ids buffer too small".into(),
));
}
if output.data_byte_len() < total_rows * (params.n as usize) * DType::F32.size_of() {
return Err(MlxError::InvalidArgument(
"fused dispatch: output buffer too small".into(),
));
}
if htpe.data_byte_len() < params.htpe_bytes() {
return Err(MlxError::InvalidArgument(
"fused dispatch: htpe buffer too small".into(),
));
}
if hids.data_byte_len() < params.hids_bytes() {
return Err(MlxError::InvalidArgument(
"fused dispatch: hids buffer too small".into(),
));
}
let map0_kernel_name = match params.top_k {
1 => "kernel_mul_mm_id_map0_ne20_1",
8 => "kernel_mul_mm_id_map0_ne20_8",
other => {
return Err(MlxError::InvalidArgument(format!(
"fused dispatch: no map0 instantiation for top_k={}",
other
)));
}
};
let map0_pipeline = registry.get_pipeline(map0_kernel_name, device.metal_device())?;
let map0_params = GgmlIdMmMap0GpuParams {
ne10: params
.n
.try_into()
.map_err(|_| MlxError::InvalidArgument("N out of i32 range".into()))?,
ne11: params.top_k as i32,
nb11: 0,
nb12: 0,
ne21: params.n_tokens as i32,
ne20: params.top_k as i32,
nb21: (params.top_k as u64) * (DType::U32.size_of() as u64),
};
let map0_shmem =
(params.n_experts as u64) * (params.top_k as u64) * std::mem::size_of::<u16>() as u64;
let map0_threadgroups = metal::MTLSize::new(1, 1, 1);
let map0_threads = metal::MTLSize::new(params.n_experts as u64, 1, 1);
encoder.encode_threadgroups_with_args_and_shared(
map0_pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&map0_params))),
(1, KernelArg::Buffer(ids)),
(2, KernelArg::Buffer(htpe)),
(3, KernelArg::Buffer(hids)),
],
&[(0, map0_shmem)],
map0_threadgroups,
map0_threads,
);
encoder.memory_barrier();
let mm_pipeline = registry.get_pipeline(
"kernel_fused_gate_up_silu_mm_id_q6_K_f32",
device.metal_device(),
)?;
let nb01 = (blocks_per_row as u64) * (block_bytes as u64);
let row_bytes = (params.k as u64) * (DType::F32.size_of() as u64);
let mm_params = GgmlIdMmMmGpuParams {
ne00: params.k as i32,
ne02: params.n_experts as i32,
nb01,
nb02: params.expert_stride,
nb03: 0,
ne11: params.top_k as i32,
_pad0: 0,
nb10: DType::F32.size_of() as u64,
nb11: 0,
nb12: row_bytes,
nb13: 0,
ne20: params.top_k as i32,
ne21: params.n_tokens as i32,
ne0: params.n as i32,
ne1: params.top_k as i32,
r2: 1,
r3: 1,
_pad1: 0,
};
const NR0: u64 = 64;
const NR1: u64 = 32;
const THREADS_PER_TG: u64 = 128;
const FUSED_SHMEM_BYTES: u64 = 16384;
let mm_threadgroups = metal::MTLSize::new(
(params.n_tokens as u64 + NR1 - 1) / NR1,
(params.n as u64 + NR0 - 1) / NR0,
params.n_experts as u64,
);
let mm_threads = metal::MTLSize::new(THREADS_PER_TG, 1, 1);
encoder.encode_threadgroups_with_args_and_shared(
mm_pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&mm_params))),
(1, KernelArg::Buffer(gate_w)),
(2, KernelArg::Buffer(up_w)),
(3, KernelArg::Buffer(input)),
(4, KernelArg::Buffer(htpe)),
(5, KernelArg::Buffer(hids)),
(6, KernelArg::Buffer(output)),
],
&[(0, FUSED_SHMEM_BYTES)],
mm_threadgroups,
mm_threads,
);
Ok(())
}
#[cfg(test)]
mod validation_tests {
use super::*;
#[test]
fn padded_expert_stack_requires_last_stride_plus_matrix() {
let matrix = ggml_expert_bytes(GgmlType::Q4_0, 1, 64, 256, 9_216).unwrap();
assert_eq!(matrix, 9_216);
assert_eq!(
required_expert_weight_bytes(GgmlType::Q4_0, 3, 64, 256, 12_288).unwrap(),
33_792
);
}
#[test]
fn expert_stack_rejects_stride_smaller_than_one_matrix() {
let error = required_expert_weight_bytes(GgmlType::Q4_0, 3, 64, 256, 9_215)
.expect_err("undersized stride must fail");
assert!(error.to_string().contains("smaller than one matrix"));
}
}