use super::*;
impl MetalQwen35State {
#[cfg(test)]
pub(super) fn force_q4_gemm_fallback_for_test(&mut self) -> bool {
self.engine.pipelines.gemm_q4_tiled.take().is_some()
}
#[cfg(test)]
pub(super) fn reset_q4_gemm_fallback_dispatches_for_test() {
Q4_GEMM_FALLBACK_DISPATCHES_FOR_TEST.with(|dispatches| dispatches.set(0));
}
#[cfg(test)]
pub(super) fn q4_gemm_fallback_dispatches_for_test() -> u64 {
Q4_GEMM_FALLBACK_DISPATCHES_FOR_TEST.with(std::cell::Cell::get)
}
pub(super) fn dispatch_matmul_half(
&self,
enc: &ComputeCommandEncoderRef,
a: &Buffer,
b: &Buffer,
c: &Buffer,
m: u32,
n: u32,
k: u32,
) {
let params = GemmParams {
m,
n,
k,
lda: k,
ldb: k,
ldc: n,
};
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_decode);
enc.set_buffer(0, Some(a), 0);
enc.set_buffer(1, Some(b), 0);
enc.set_buffer(2, Some(c), 0);
enc.set_bytes(
3,
std::mem::size_of::<GemmParams>() as u64,
¶ms as *const GemmParams as *const _,
);
enc.dispatch_thread_groups(MTLSize::new(n as u64, 1, 1), MTLSize::new(256, 1, 1));
}
fn dispatch_matmul_q8(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer, qw: &Q4WeightBuf, y: &Buffer, _m: u32, n: u32,
k: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q8);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(&qw.buffer), 0);
enc.set_buffer(2, Some(y), 0);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1),
MTLSize::new(32, 4, 1), );
}
fn dispatch_gemm_q8(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64, qw: &Q4WeightBuf,
y: &Buffer,
y_offset: u64, m: u32,
n: u32,
k: u32,
) {
if m == 0 || n == 0 {
return;
}
assert!(
k > 0 && k.is_multiple_of(32),
"dispatch_gemm_q8 requires K non-zero and divisible by 32, got {k}"
);
if m <= 1 {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q8);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), 0);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1),
MTLSize::new(32, 4, 1),
);
} else if let Some(tiled) = self.engine.pipelines.gemm_q8_tiled.as_ref() {
enc.set_compute_pipeline_state(tiled);
enc.set_buffer(0, Some(&qw.buffer), 0);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(32) as u64, m.div_ceil(64) as u64, 1),
MTLSize::new(32, 4, 1),
);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemm_q8);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), 0);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, m.div_ceil(4) as u64, 1),
MTLSize::new(32, 4, 1),
);
}
}
pub(super) fn dispatch_matmul_q4(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer, qw: &Q4WeightBuf, y: &Buffer, _m: u32,
n: u32,
k: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q4);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(2, Some(y), 0);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1), MTLSize::new(32, 4, 1),
);
}
pub(super) fn dispatch_gemm_q4(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64,
qw: &Q4WeightBuf,
y: &Buffer,
y_offset: u64,
m: u32,
n: u32,
k: u32,
) {
if m == 0 || n == 0 {
return;
}
assert!(
k > 0 && k.is_multiple_of(32),
"dispatch_gemm_q4 requires K to be non-zero and divisible by 32, got {k}"
);
if m == 1 {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q4);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1),
MTLSize::new(32, 4, 1),
);
} else if let Some(tiled) = self.engine.pipelines.gemm_q4_tiled.as_ref() {
enc.set_compute_pipeline_state(tiled);
enc.set_buffer(0, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(32) as u64, m.div_ceil(64) as u64, 1),
MTLSize::new(32, 4, 1),
);
} else {
#[cfg(test)]
Q4_GEMM_FALLBACK_DISPATCHES_FOR_TEST
.with(|dispatches| dispatches.set(dispatches.get() + 1));
enc.set_compute_pipeline_state(&self.engine.pipelines.gemm_q4);
enc.set_buffer(0, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, m.div_ceil(4) as u64, 1),
MTLSize::new(32, 4, 1),
);
}
}
pub(super) fn dispatch_matmul(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
qw: &Q4WeightBuf,
y: &Buffer,
m: u32,
n: u32,
k: u32,
) {
match self.engine.quant_format {
QuantFormat::Q8_0 => self.dispatch_matmul_q8(enc, x, qw, y, m, n, k),
QuantFormat::Q4_0 => self.dispatch_matmul_q4(enc, x, qw, y, m, n, k),
}
}
#[allow(dead_code)] fn dispatch_gemm_q3(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64,
qw: &Q3WeightBuf,
y: &Buffer,
y_offset: u64,
m: u32,
n: u32,
k: u32,
) -> Result<(), String> {
if m == 0 || n == 0 {
return Ok(());
}
assert!(
k > 0 && k.is_multiple_of(32),
"dispatch_gemm_q3 requires K to be non-zero and divisible by 32, got {k}"
);
if m == 1 {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q3);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1),
MTLSize::new(32, 4, 1),
);
} else {
let tiled = self
.engine
.pipelines
.gemm_q3_tiled
.as_ref()
.ok_or_else(|| {
"dispatch_gemm_q3 called with M>1 but gemm_q3_tiled is unavailable \
(device is not Apple7+, or the tiled kernel failed to compile); \
Stage 2 has no naive Q3 GEMM fallback"
.to_string()
})?;
enc.set_compute_pipeline_state(tiled);
enc.set_buffer(0, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(32) as u64, m.div_ceil(64) as u64, 1),
MTLSize::new(32, 4, 1),
);
}
Ok(())
}
pub(super) fn dispatch_gdn_chunked_prefill_layer(
&self,
enc: &ComputeCommandEncoderRef,
linear_idx: usize,
weights: GdnChunkedWeights<'_>,
params: GdnChunkParams,
) {
self.dispatch_gdn_chunk_materialize_c32(enc, linear_idx, weights, params);
self.dispatch_gdn_chunk_conv_buf_update_c32(enc, linear_idx, params);
self.dispatch_gdn_chunk_solve_c32(enc, params);
for chunk in 0..params.num_chunks {
let mut cp = params;
cp.active_chunk = chunk;
self.dispatch_gdn_chunk_residual_output_c32(enc, linear_idx, cp);
self.dispatch_gdn_chunk_state_update_c32(enc, linear_idx, cp);
}
self.dispatch_gdn_chunk_norm_silu_c32(enc, weights.norm_weight, params);
}
fn dispatch_gdn_chunk_materialize_c32(
&self,
enc: &ComputeCommandEncoderRef,
linear_idx: usize,
weights: GdnChunkedWeights<'_>,
params: GdnChunkParams,
) {
let sc = &self.session.activations.gdn_chunk;
let p_bytes = std::mem::size_of::<GdnChunkParams>() as u64;
enc.set_compute_pipeline_state(&self.engine.pipelines.gdn_chunk_materialize_c32);
enc.set_buffer(0, Some(&self.session.gdn_gpu_conv_bufs[linear_idx]), 0);
enc.set_buffer(1, Some(&self.session.activations.gdn_qkv), 0);
enc.set_buffer(2, Some(weights.conv1d_weight), 0);
enc.set_buffer(3, Some(&self.session.activations.hidden), 0);
enc.set_buffer(4, Some(weights.in_proj_b), 0);
enc.set_buffer(5, Some(weights.in_proj_a), 0);
enc.set_buffer(6, Some(weights.a_log), 0);
enc.set_buffer(7, Some(weights.dt_bias), 0);
enc.set_buffer(8, Some(&sc.q), 0);
enc.set_buffer(9, Some(&sc.k), 0);
enc.set_buffer(10, Some(&sc.v), 0);
enc.set_buffer(11, Some(&sc.beta_log_alpha), 0);
enc.set_bytes(12, p_bytes, ¶ms as *const GdnChunkParams as *const _);
enc.dispatch_thread_groups(
MTLSize::new(params.num_chunks as u64, params.num_value_heads as u64, 1),
MTLSize::new(32, 4, 1),
);
}
fn dispatch_gdn_chunk_conv_buf_update_c32(
&self,
enc: &ComputeCommandEncoderRef,
linear_idx: usize,
params: GdnChunkParams,
) {
let p_bytes = std::mem::size_of::<GdnChunkParams>() as u64;
enc.set_compute_pipeline_state(&self.engine.pipelines.gdn_chunk_conv_buf_update_c32);
enc.set_buffer(0, Some(&self.session.gdn_gpu_conv_bufs[linear_idx]), 0);
enc.set_buffer(1, Some(&self.session.activations.gdn_qkv), 0);
enc.set_bytes(2, p_bytes, ¶ms as *const GdnChunkParams as *const _);
enc.dispatch_thread_groups(
MTLSize::new(1, params.num_value_heads as u64, 1),
MTLSize::new(32, 4, 1),
);
}
fn dispatch_gdn_chunk_solve_c32(&self, enc: &ComputeCommandEncoderRef, params: GdnChunkParams) {
let sc = &self.session.activations.gdn_chunk;
let p_bytes = std::mem::size_of::<GdnChunkParams>() as u64;
enc.set_compute_pipeline_state(&self.engine.pipelines.gdn_chunk_solve_c32);
enc.set_buffer(0, Some(&sc.q), 0);
enc.set_buffer(1, Some(&sc.k), 0);
enc.set_buffer(2, Some(&sc.v), 0);
enc.set_buffer(3, Some(&sc.beta_log_alpha), 0);
enc.set_buffer(4, Some(&sc.gamma), 0);
enc.set_buffer(5, Some(&sc.gamma_end), 0);
enc.set_buffer(6, Some(&sc.kkt), 0);
enc.set_buffer(7, Some(&sc.qk_l), 0);
enc.set_buffer(8, Some(&sc.w), 0);
enc.set_buffer(9, Some(&sc.u), 0);
enc.set_buffer(10, Some(&sc.k_right), 0);
enc.set_bytes(11, p_bytes, ¶ms as *const GdnChunkParams as *const _);
enc.dispatch_thread_groups(
MTLSize::new(params.num_chunks as u64, params.num_value_heads as u64, 1),
MTLSize::new(32, 4, 1),
);
}
fn dispatch_gdn_chunk_residual_output_c32(
&self,
enc: &ComputeCommandEncoderRef,
linear_idx: usize,
params: GdnChunkParams,
) {
let sc = &self.session.activations.gdn_chunk;
let vd = params.value_dim as u64;
let num_v_tiles = vd.div_ceil(8);
let p_bytes = std::mem::size_of::<GdnChunkParams>() as u64;
enc.set_compute_pipeline_state(&self.engine.pipelines.gdn_chunk_residual_output_c32);
enc.set_buffer(0, Some(&self.session.gdn_gpu_s_matrices[linear_idx]), 0);
enc.set_buffer(1, Some(&sc.q), 0);
enc.set_buffer(2, Some(&sc.w), 0);
enc.set_buffer(3, Some(&sc.u), 0);
enc.set_buffer(4, Some(&sc.gamma), 0);
enc.set_buffer(5, Some(&sc.qk_l), 0);
enc.set_buffer(6, Some(&sc.r), 0);
enc.set_buffer(7, Some(&sc.raw_out), 0);
enc.set_bytes(8, p_bytes, ¶ms as *const GdnChunkParams as *const _);
enc.dispatch_thread_groups(
MTLSize::new(num_v_tiles, params.num_value_heads as u64, 1),
MTLSize::new(32, 8, 1),
);
}
fn dispatch_gdn_chunk_state_update_c32(
&self,
enc: &ComputeCommandEncoderRef,
linear_idx: usize,
params: GdnChunkParams,
) {
let sc = &self.session.activations.gdn_chunk;
let kd = params.key_dim as u64;
let vd = params.value_dim as u64;
let p_bytes = std::mem::size_of::<GdnChunkParams>() as u64;
enc.set_compute_pipeline_state(&self.engine.pipelines.gdn_chunk_state_update_c32);
enc.set_buffer(0, Some(&self.session.gdn_gpu_s_matrices[linear_idx]), 0);
enc.set_buffer(1, Some(&sc.r), 0);
enc.set_buffer(2, Some(&sc.k_right), 0);
enc.set_buffer(3, Some(&sc.gamma_end), 0);
enc.set_bytes(4, p_bytes, ¶ms as *const GdnChunkParams as *const _);
enc.dispatch_thread_groups(
MTLSize::new(
kd.div_ceil(16),
vd.div_ceil(16),
params.num_value_heads as u64,
),
MTLSize::new(16, 16, 1),
);
}
fn dispatch_gdn_chunk_norm_silu_c32(
&self,
enc: &ComputeCommandEncoderRef,
norm_weight: &Buffer,
params: GdnChunkParams,
) {
let sc = &self.session.activations.gdn_chunk;
let p_bytes = std::mem::size_of::<GdnChunkParams>() as u64;
enc.set_compute_pipeline_state(&self.engine.pipelines.gdn_chunk_norm_silu_c32);
enc.set_buffer(0, Some(&sc.raw_out), 0);
enc.set_buffer(1, Some(&self.session.activations.gdn_z), 0);
enc.set_buffer(2, Some(norm_weight), 0);
enc.set_buffer(3, Some(&self.session.activations.gdn_z), 0);
enc.set_bytes(4, p_bytes, ¶ms as *const GdnChunkParams as *const _);
enc.dispatch_thread_groups(
MTLSize::new(params.n_tokens as u64, params.num_value_heads as u64, 1),
MTLSize::new(128, 1, 1),
);
}
pub(super) fn dispatch_gemm(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64,
qw: &Q4WeightBuf,
y: &Buffer,
y_offset: u64,
m: u32,
n: u32,
k: u32,
) {
match self.engine.quant_format {
QuantFormat::Q8_0 => self.dispatch_gemm_q8(enc, x, x_offset, qw, y, y_offset, m, n, k),
QuantFormat::Q4_0 => self.dispatch_gemm_q4(enc, x, x_offset, qw, y, y_offset, m, n, k),
}
}
pub(super) fn dispatch_rms_norm(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
gamma: &Buffer,
row_len: u32,
num_rows: u32,
eps: f32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.rms_norm);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(gamma), 0);
enc.set_bytes(2, 4, &row_len as *const u32 as *const _);
enc.set_bytes(3, 4, &num_rows as *const u32 as *const _);
enc.set_bytes(4, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(MTLSize::new(num_rows as u64, 1, 1), MTLSize::new(wg, 1, 1));
}
pub(super) fn dispatch_per_head_rms_norm(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
gamma: &Buffer,
num_heads: u32,
head_dim: u32,
eps: f32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.per_head_rms_norm);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(gamma), 0);
enc.set_bytes(2, 4, &num_heads as *const u32 as *const _);
enc.set_bytes(3, 4, &head_dim as *const u32 as *const _);
enc.set_bytes(4, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(MTLSize::new(num_heads as u64, 1, 1), MTLSize::new(wg, 1, 1));
}
pub(super) fn dispatch_partial_rope(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
num_heads: u32,
head_dim: u32,
half_rope_dim: u32,
position: u32,
mrope_override: Option<(&Buffer, &Buffer)>,
) {
let (cos_buf, sin_buf, pos_offset) = match mrope_override {
Some((cos_buf, sin_buf)) => (cos_buf, sin_buf, 0u32),
None => (&self.engine.rope_cos, &self.engine.rope_sin, position),
};
enc.set_compute_pipeline_state(&self.engine.pipelines.partial_rope);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(cos_buf), 0);
enc.set_buffer(2, Some(sin_buf), 0);
enc.set_bytes(3, 4, &num_heads as *const u32 as *const _);
enc.set_bytes(4, 4, &head_dim as *const u32 as *const _);
enc.set_bytes(5, 4, &half_rope_dim as *const u32 as *const _);
enc.set_bytes(6, 4, &pos_offset as *const u32 as *const _);
let total_pairs = num_heads * half_rope_dim;
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(total_pairs as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
#[allow(clippy::too_many_arguments)]
pub(super) fn dispatch_decode_attention(
&self,
enc: &ComputeCommandEncoderRef,
k_cache: &Buffer,
v_cache: &Buffer,
cache_len: u32,
head_dim: u32,
num_q_heads: u32,
num_kv_heads: u32,
q_dim: u32,
kv_dim: u32,
scale: f32,
) {
const PARTITION_TOKENS: u32 = METAL_FLASH_PARTITION_TOKENS as u32;
const DIRECT_THRESHOLD: u32 = 512;
validate_flash_decode_shape(
head_dim as usize,
num_q_heads as usize,
num_kv_heads as usize,
q_dim as usize,
kv_dim as usize,
)
.expect("invalid Metal FlashAttention decode shape");
let kv_f16 = self.use_kv_f16;
let set_common_bufs = |enc: &ComputeCommandEncoderRef, out_or_partials: &Buffer| {
enc.set_buffer(0, Some(&self.session.activations.q_separated), 0);
enc.set_buffer(1, Some(k_cache), 0);
enc.set_buffer(2, Some(v_cache), 0);
enc.set_buffer(3, Some(out_or_partials), 0);
enc.set_bytes(4, 4, &cache_len as *const u32 as *const _);
enc.set_bytes(5, 4, &head_dim as *const u32 as *const _);
enc.set_bytes(6, 4, &num_q_heads as *const u32 as *const _);
enc.set_bytes(7, 4, &num_kv_heads as *const u32 as *const _);
enc.set_bytes(8, 4, &q_dim as *const u32 as *const _);
enc.set_bytes(9, 4, &kv_dim as *const u32 as *const _);
enc.set_bytes(10, 4, &scale as *const f32 as *const _);
};
if cache_len <= DIRECT_THRESHOLD {
if kv_f16 {
enc.set_compute_pipeline_state(&self.engine.pipelines.decode_attention_f16);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.decode_attention);
}
set_common_bufs(enc, &self.session.activations.attn_out);
enc.dispatch_thread_groups(
MTLSize::new(num_kv_heads as u64, 1, 1),
MTLSize::new(256, 1, 1),
);
if self.path_proof_enabled {
self.path_proof
.decode_attn_direct
.fetch_add(1, Ordering::Relaxed);
}
} else {
let num_partitions = cache_len.div_ceil(PARTITION_TOKENS);
if kv_f16 {
enc.set_compute_pipeline_state(&self.engine.pipelines.decode_attn_partial_f16);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.decode_attn_partial);
}
set_common_bufs(enc, &self.session.activations.attn_partials);
enc.set_bytes(11, 4, &PARTITION_TOKENS as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(num_kv_heads as u64, num_partitions as u64, 1),
MTLSize::new(256, 1, 1),
);
if self.path_proof_enabled {
self.path_proof
.decode_attn_split_partial
.fetch_add(1, Ordering::Relaxed);
}
enc.set_compute_pipeline_state(&self.engine.pipelines.decode_attn_reduce);
enc.set_buffer(0, Some(&self.session.activations.attn_partials), 0);
enc.set_buffer(1, Some(&self.session.activations.attn_out), 0);
enc.set_bytes(2, 4, &num_q_heads as *const u32 as *const _);
enc.set_bytes(3, 4, &num_kv_heads as *const u32 as *const _);
enc.set_bytes(4, 4, &num_partitions as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(num_kv_heads as u64, 1, 1),
MTLSize::new(256, 1, 1),
);
if self.path_proof_enabled {
self.path_proof
.decode_attn_split_reduce
.fetch_add(1, Ordering::Relaxed);
}
}
}
pub(super) fn dispatch_sigmoid_gate(&self, enc: &ComputeCommandEncoderRef, count: u32) {
enc.set_compute_pipeline_state(&self.engine.pipelines.sigmoid_gate);
enc.set_buffer(0, Some(&self.session.activations.attn_out), 0);
enc.set_buffer(1, Some(&self.session.activations.gate_z), 0);
enc.set_bytes(2, 4, &count as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_scatter_q_gate(
&self,
enc: &ComputeCommandEncoderRef,
num_heads: u32,
head_dim: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.scatter_q_gate);
enc.set_buffer(0, Some(&self.session.activations.q), 0);
enc.set_buffer(1, Some(&self.session.activations.q_separated), 0);
enc.set_buffer(2, Some(&self.session.activations.gate_z), 0);
enc.set_bytes(3, 4, &num_heads as *const u32 as *const _);
enc.set_bytes(4, 4, &head_dim as *const u32 as *const _);
let total = num_heads * head_dim;
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(total as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_scatter_q_gate_batch(
&self,
enc: &ComputeCommandEncoderRef,
num_tokens: u32,
num_heads: u32,
head_dim: u32,
) {
let total = num_tokens * num_heads * head_dim;
enc.set_compute_pipeline_state(&self.engine.pipelines.scatter_q_gate_batch);
enc.set_buffer(0, Some(&self.session.activations.q), 0);
enc.set_buffer(1, Some(&self.session.activations.q_separated), 0);
enc.set_buffer(2, Some(&self.session.activations.gate_z), 0);
enc.set_bytes(3, 4, &num_tokens as *const u32 as *const _);
enc.set_bytes(4, 4, &num_heads as *const u32 as *const _);
enc.set_bytes(5, 4, &head_dim as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(total as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_per_head_rms_norm_batch(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
gamma: &Buffer,
num_tokens: u32,
num_heads: u32,
head_dim: u32,
eps: f32,
) {
let total_groups = num_tokens * num_heads;
enc.set_compute_pipeline_state(&self.engine.pipelines.per_head_rms_norm_batch);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(gamma), 0);
enc.set_bytes(2, 4, &num_tokens as *const u32 as *const _);
enc.set_bytes(3, 4, &num_heads as *const u32 as *const _);
enc.set_bytes(4, 4, &head_dim as *const u32 as *const _);
enc.set_bytes(5, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(
MTLSize::new(total_groups as u64, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
#[allow(clippy::too_many_arguments)]
pub(super) fn dispatch_partial_rope_batch(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
num_tokens: u32,
num_heads: u32,
head_dim: u32,
half_rope_dim: u32,
base_pos: u32,
) {
let total_pairs = num_tokens * num_heads * half_rope_dim;
enc.set_compute_pipeline_state(&self.engine.pipelines.partial_rope_batch);
enc.set_buffer(0, Some(x), 0);
enc.set_buffer(1, Some(&self.engine.rope_cos), 0);
enc.set_buffer(2, Some(&self.engine.rope_sin), 0);
enc.set_bytes(3, 4, &num_tokens as *const u32 as *const _);
enc.set_bytes(4, 4, &num_heads as *const u32 as *const _);
enc.set_bytes(5, 4, &head_dim as *const u32 as *const _);
enc.set_bytes(6, 4, &half_rope_dim as *const u32 as *const _);
enc.set_bytes(7, 4, &base_pos as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(total_pairs as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
#[allow(clippy::too_many_arguments)]
pub(super) fn dispatch_copy_kv_cache_batch(
&self,
enc: &ComputeCommandEncoderRef,
k_src: &Buffer,
v_src: &Buffer,
k_cache: &Buffer,
v_cache: &Buffer,
num_tokens: u32,
kv_dim: u32,
base_pos: u32,
) {
let total = num_tokens * kv_dim;
if self.use_kv_f16 {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_kv_cache_batch_f16);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_kv_cache_batch);
}
enc.set_buffer(0, Some(k_src), 0);
enc.set_buffer(1, Some(v_src), 0);
enc.set_buffer(2, Some(k_cache), 0);
enc.set_buffer(3, Some(v_cache), 0);
enc.set_bytes(4, 4, &num_tokens as *const u32 as *const _);
enc.set_bytes(5, 4, &kv_dim as *const u32 as *const _);
enc.set_bytes(6, 4, &base_pos as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(total as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
if self.path_proof_enabled {
self.path_proof
.prefill_kv_batch
.fetch_add(1, Ordering::Relaxed);
}
}
#[allow(clippy::too_many_arguments)]
pub(super) fn dispatch_prefill_attention_batched(
&self,
enc: &ComputeCommandEncoderRef,
k_cache: &Buffer,
v_cache: &Buffer,
base_pos: u32,
num_tokens: u32,
head_dim: u32,
num_q_heads: u32,
num_kv_heads: u32,
q_dim: u32,
kv_dim: u32,
scale: f32,
) -> Result<(), String> {
validate_flash_decode_shape(
head_dim as usize,
num_q_heads as usize,
num_kv_heads as usize,
q_dim as usize,
kv_dim as usize,
)?;
let cache_len_total = base_pos.checked_add(num_tokens).ok_or_else(|| {
"prefill_attention_batched: base_pos + num_tokens overflow".to_string()
})?;
if self.use_kv_f16 {
enc.set_compute_pipeline_state(&self.engine.pipelines.prefill_attention_batched_f16);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.prefill_attention_batched);
}
enc.set_buffer(0, Some(&self.session.activations.q_separated), 0);
enc.set_buffer(1, Some(k_cache), 0);
enc.set_buffer(2, Some(v_cache), 0);
enc.set_buffer(3, Some(&self.session.activations.attn_out), 0);
enc.set_bytes(4, 4, &base_pos as *const u32 as *const _);
enc.set_bytes(5, 4, &num_tokens as *const u32 as *const _);
enc.set_bytes(6, 4, &cache_len_total as *const u32 as *const _);
enc.set_bytes(7, 4, &head_dim as *const u32 as *const _);
enc.set_bytes(8, 4, &num_q_heads as *const u32 as *const _);
enc.set_bytes(9, 4, &num_kv_heads as *const u32 as *const _);
enc.set_bytes(10, 4, &q_dim as *const u32 as *const _);
enc.set_bytes(11, 4, &kv_dim as *const u32 as *const _);
enc.set_bytes(12, 4, &scale as *const f32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(num_kv_heads as u64, num_tokens as u64, 1),
MTLSize::new(256, 1, 1),
);
if self.path_proof_enabled {
self.path_proof
.prefill_attn_batched
.fetch_add(1, Ordering::Relaxed);
}
Ok(())
}
pub(super) fn dispatch_silu_mul(&self, enc: &ComputeCommandEncoderRef, count: u32) {
enc.set_compute_pipeline_state(&self.engine.pipelines.silu_mul);
enc.set_buffer(0, Some(&self.session.activations.gate), 0);
enc.set_buffer(1, Some(&self.session.activations.up), 0);
enc.set_bytes(2, 4, &count as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_silu_mul_fused(
&self,
enc: &ComputeCommandEncoderRef,
data: &Buffer,
count: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.silu_mul_fused);
enc.set_buffer(0, Some(data), 0);
enc.set_bytes(2, 4, &count as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_gemm_at(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64,
qw: &Q4WeightBuf,
qw_extra_offset: u64,
y: &Buffer,
y_offset: u64,
m: u32,
n: u32,
k: u32,
) {
match self.engine.quant_format {
QuantFormat::Q8_0 => self.dispatch_gemm_q8_at(
enc,
x,
x_offset,
qw,
qw_extra_offset,
y,
y_offset,
m,
n,
k,
),
QuantFormat::Q4_0 => self.dispatch_gemm_q4_at(
enc,
x,
x_offset,
qw,
qw_extra_offset,
y,
y_offset,
m,
n,
k,
),
}
}
fn dispatch_gemm_q8_at(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64,
qw: &Q4WeightBuf,
qw_extra_offset: u64,
y: &Buffer,
y_offset: u64,
m: u32,
n: u32,
k: u32,
) {
let wq_offset = qw.payload_offset + qw_extra_offset;
if m == 0 || n == 0 {
return;
}
assert!(
k > 0 && k.is_multiple_of(32),
"dispatch_gemm_q8_at requires K non-zero and divisible by 32, got {k}"
);
if m <= 1 {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q8);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), wq_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1),
MTLSize::new(32, 4, 1),
);
} else if let Some(tiled) = self.engine.pipelines.gemm_q8_tiled.as_ref() {
enc.set_compute_pipeline_state(tiled);
enc.set_buffer(0, Some(&qw.buffer), wq_offset);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(32) as u64, m.div_ceil(64) as u64, 1),
MTLSize::new(32, 4, 1),
);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemm_q8);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), wq_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, m.div_ceil(4) as u64, 1),
MTLSize::new(32, 4, 1),
);
}
}
fn dispatch_gemm_q4_at(
&self,
enc: &ComputeCommandEncoderRef,
x: &Buffer,
x_offset: u64,
qw: &Q4WeightBuf,
qw_extra_offset: u64,
y: &Buffer,
y_offset: u64,
m: u32,
n: u32,
k: u32,
) {
if m == 0 || n == 0 {
return;
}
assert!(
k > 0 && k.is_multiple_of(32),
"dispatch_gemm_q4_at requires K to be non-zero and divisible by 32, got {k}"
);
let wq_offset = qw.payload_offset + qw_extra_offset;
if m == 1 {
enc.set_compute_pipeline_state(&self.engine.pipelines.gemv_q4);
enc.set_buffer(0, Some(x), x_offset);
enc.set_buffer(1, Some(&qw.buffer), wq_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &n as *const u32 as *const _);
enc.set_bytes(4, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, 1, 1),
MTLSize::new(32, 4, 1),
);
} else if let Some(tiled) = self.engine.pipelines.gemm_q4_tiled.as_ref() {
enc.set_compute_pipeline_state(tiled);
enc.set_buffer(0, Some(&qw.buffer), wq_offset);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(32) as u64, m.div_ceil(64) as u64, 1),
MTLSize::new(32, 4, 1),
);
} else {
#[cfg(test)]
Q4_GEMM_FALLBACK_DISPATCHES_FOR_TEST
.with(|dispatches| dispatches.set(dispatches.get() + 1));
enc.set_compute_pipeline_state(&self.engine.pipelines.gemm_q4);
enc.set_buffer(0, Some(&qw.buffer), wq_offset);
enc.set_buffer(1, Some(x), x_offset);
enc.set_buffer(2, Some(y), y_offset);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(n.div_ceil(2) as u64, m.div_ceil(4) as u64, 1),
MTLSize::new(32, 4, 1),
);
}
}
pub(super) fn dispatch_copy(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
dst: &Buffer,
count: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(dst), 0);
enc.set_bytes(2, 4, &count as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_copy_offset(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
dst: &Buffer,
count: u32,
dst_offset: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_offset);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(dst), 0);
enc.set_bytes(2, 4, &count as *const u32 as *const _);
enc.set_bytes(3, 4, &dst_offset as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_copy_offset_kv(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
dst: &Buffer,
count: u32,
dst_offset: u32,
) {
if self.use_kv_f16 {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_offset_f16);
} else {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_offset);
}
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(dst), 0);
enc.set_bytes(2, 4, &count as *const u32 as *const _);
enc.set_bytes(3, 4, &dst_offset as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
if self.path_proof_enabled {
self.path_proof
.decode_kv_copy
.fetch_add(1, Ordering::Relaxed);
}
}
pub(super) fn dispatch_fused_residual_add_norm(
&self,
enc: &ComputeCommandEncoderRef,
base: &Buffer,
delta: &Buffer,
residual_out: &Buffer,
normed_out: &Buffer,
gamma: &Buffer,
row_len: u32,
eps: f32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.fused_residual_add_norm);
enc.set_buffer(0, Some(base), 0);
enc.set_buffer(1, Some(delta), 0);
enc.set_buffer(2, Some(residual_out), 0);
enc.set_buffer(3, Some(normed_out), 0);
enc.set_buffer(4, Some(gamma), 0);
enc.set_bytes(5, 4, &row_len as *const u32 as *const _);
enc.set_bytes(6, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(MTLSize::new(1, 1, 1), MTLSize::new(wg, 1, 1));
}
pub(super) fn dispatch_copy_and_rms_norm(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
residual_out: &Buffer,
gamma: &Buffer,
row_len: u32,
eps: f32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_and_rms_norm);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(residual_out), 0);
enc.set_buffer(2, Some(gamma), 0);
enc.set_bytes(3, 4, &row_len as *const u32 as *const _);
enc.set_bytes(4, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(MTLSize::new(1, 1, 1), MTLSize::new(wg, 1, 1));
}
pub(super) fn dispatch_copy_and_rms_norm_batch(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
residual_out: &Buffer,
gamma: &Buffer,
row_len: u32,
num_rows: u32,
eps: f32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.copy_and_rms_norm_batch);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(residual_out), 0);
enc.set_buffer(2, Some(gamma), 0);
enc.set_bytes(3, 4, &row_len as *const u32 as *const _);
enc.set_bytes(4, 4, &num_rows as *const u32 as *const _);
enc.set_bytes(5, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(MTLSize::new(num_rows as u64, 1, 1), MTLSize::new(wg, 1, 1));
}
pub(super) fn dispatch_fused_residual_add_norm_batch(
&self,
enc: &ComputeCommandEncoderRef,
base: &Buffer,
delta: &Buffer,
residual_out: &Buffer,
normed_out: &Buffer,
gamma: &Buffer,
row_len: u32,
num_rows: u32,
eps: f32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.fused_residual_add_norm_batch);
enc.set_buffer(0, Some(base), 0);
enc.set_buffer(1, Some(delta), 0);
enc.set_buffer(2, Some(residual_out), 0);
enc.set_buffer(3, Some(normed_out), 0);
enc.set_buffer(4, Some(gamma), 0);
enc.set_bytes(5, 4, &row_len as *const u32 as *const _);
enc.set_bytes(6, 4, &num_rows as *const u32 as *const _);
enc.set_bytes(7, 4, &eps as *const f32 as *const _);
let wg = 256u64;
enc.dispatch_thread_groups(MTLSize::new(num_rows as u64, 1, 1), MTLSize::new(wg, 1, 1));
}
pub(super) fn dispatch_add_and_copy(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
residual: &Buffer,
dst: &Buffer,
count: u32,
) {
enc.set_compute_pipeline_state(&self.engine.pipelines.add_and_copy);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(residual), 0);
enc.set_buffer(2, Some(dst), 0);
enc.set_bytes(3, 4, &count as *const u32 as *const _);
let wg = 256u64;
enc.dispatch_threads(
MTLSize::new(div_ceil(count as u64, wg) * wg, 1, 1),
MTLSize::new(wg, 1, 1),
);
}
pub(super) fn dispatch_lm_head_block_stage1_enc(
&self,
enc: &ComputeCommandEncoderRef,
cfg: &Qwen35Config,
hidden_offset: u64,
local_k: u32,
slot: usize,
) -> u32 {
let vocab_size = cfg.vocab_size as u32;
let row_groups = vocab_size.div_ceil(LM_HEAD_ROWS_PER_TG);
let candidate_capacity = self.session.activations.topk_scratch_a.length() / 8;
let candidates_written = row_groups as u64 * local_k as u64;
assert!(
candidates_written <= candidate_capacity,
"lm_head block-topk Stage-1 would write {candidates_written} candidates \
({row_groups} groups * local_k {local_k}) but topk_scratch_a only holds \
{candidate_capacity} — buffer-sizing invariant violated"
);
match self.engine.quant_format {
QuantFormat::Q8_0 => {
enc.set_compute_pipeline_state(&self.engine.pipelines.lm_head_block_topk_f16[slot]);
enc.set_buffer(0, Some(&self.session.activations.hidden), hidden_offset);
enc.set_buffer(1, Some(&self.engine.embed_tokens), 0);
enc.set_buffer(2, Some(&self.session.activations.topk_scratch_a), 0);
enc.set_bytes(3, 4, &vocab_size as *const u32 as *const _);
}
QuantFormat::Q4_0 => {
let qw = &self.engine.embed_tokens_q8;
enc.set_compute_pipeline_state(&self.engine.pipelines.lm_head_block_topk_q4[slot]);
enc.set_buffer(0, Some(&self.session.activations.hidden), hidden_offset);
enc.set_buffer(1, Some(&qw.buffer), qw.payload_offset);
enc.set_buffer(2, Some(&self.session.activations.topk_scratch_a), 0);
enc.set_bytes(3, 4, &vocab_size as *const u32 as *const _);
}
}
enc.dispatch_thread_groups(
MTLSize::new(row_groups as u64, 1, 1),
MTLSize::new(256, 1, 1),
);
row_groups
}
pub(super) fn dispatch_block_topk_merge_enc(
&self,
enc: &ComputeCommandEncoderRef,
candidate_groups: u32,
local_k: u32,
) -> u8 {
if local_k == 1 {
enc.set_compute_pipeline_state(&self.engine.pipelines.argmax_merge);
enc.set_buffer(0, Some(&self.session.activations.topk_scratch_a), 0);
enc.set_buffer(1, Some(&self.session.activations.topk_scratch_b), 0);
enc.set_bytes(2, 4, &candidate_groups as *const u32 as *const _);
enc.dispatch_thread_groups(MTLSize::new(1, 1, 1), MTLSize::new(1024, 1, 1));
return 1;
}
let mut current_groups = candidate_groups;
let mut which: u8 = 0;
while current_groups > 1 {
let fan_in: u32 = 16u32.min(current_groups);
let out_groups = current_groups.div_ceil(fan_in);
let (in_buf, out_buf) = if which == 0 {
(
&self.session.activations.topk_scratch_a,
&self.session.activations.topk_scratch_b,
)
} else {
(
&self.session.activations.topk_scratch_b,
&self.session.activations.topk_scratch_a,
)
};
enc.set_compute_pipeline_state(&self.engine.pipelines.topk_merge_pass);
enc.set_buffer(0, Some(in_buf), 0);
enc.set_buffer(1, Some(out_buf), 0);
enc.set_bytes(2, 4, ¤t_groups as *const u32 as *const _);
enc.set_bytes(3, 4, &local_k as *const u32 as *const _);
enc.set_bytes(4, 4, &fan_in as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(out_groups as u64, 1, 1),
MTLSize::new(256, 1, 1),
);
current_groups = out_groups;
which = 1 - which;
}
which
}
pub(super) fn dispatch_topk_enc(
&self,
enc: &ComputeCommandEncoderRef,
vocab_size: u32,
k: u32,
) -> u8 {
debug_assert!(
self.session.compact_route != GpuTopkRoute::CpuFallback,
"dispatch_topk_enc called with CpuFallback route"
);
if k == 1 {
let groups = vocab_size.div_ceil(1024);
enc.set_compute_pipeline_state(&self.engine.pipelines.argmax_first);
enc.set_buffer(0, Some(&self.session.activations.logits), 0);
enc.set_buffer(1, Some(&self.session.activations.topk_scratch_a), 0);
enc.set_bytes(2, 4, &vocab_size as *const u32 as *const _);
enc.dispatch_thread_groups(MTLSize::new(groups as u64, 1, 1), MTLSize::new(1024, 1, 1));
enc.set_compute_pipeline_state(&self.engine.pipelines.argmax_merge);
enc.set_buffer(0, Some(&self.session.activations.topk_scratch_a), 0);
enc.set_buffer(1, Some(&self.session.activations.topk_scratch_b), 0);
enc.set_bytes(2, 4, &groups as *const u32 as *const _);
enc.dispatch_thread_groups(MTLSize::new(1, 1, 1), MTLSize::new(1024, 1, 1));
return 1; }
debug_assert_eq!(k, 50, "only HierarchicalK50 route is supported for k>1");
debug_assert_eq!(self.session.compact_route, GpuTopkRoute::HierarchicalK50);
let tile = 1024u32;
let first_pass_groups = vocab_size.div_ceil(tile);
enc.set_compute_pipeline_state(&self.engine.pipelines.topk_select50_first);
enc.set_buffer(0, Some(&self.session.activations.logits), 0);
enc.set_buffer(1, Some(&self.session.activations.topk_scratch_a), 0);
enc.set_bytes(2, 4, &vocab_size as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(first_pass_groups as u64, 1, 1),
MTLSize::new(256, 1, 1),
);
let mut current_groups = first_pass_groups;
let mut which: u8 = 0;
while current_groups > 1 {
let fan_in: u32 = 16u32.min(current_groups);
let out_groups = current_groups.div_ceil(fan_in);
let (in_buf, out_buf) = if which == 0 {
(
&self.session.activations.topk_scratch_a,
&self.session.activations.topk_scratch_b,
)
} else {
(
&self.session.activations.topk_scratch_b,
&self.session.activations.topk_scratch_a,
)
};
enc.set_compute_pipeline_state(&self.engine.pipelines.topk_select50_merge);
enc.set_buffer(0, Some(in_buf), 0);
enc.set_buffer(1, Some(out_buf), 0);
enc.set_bytes(2, 4, ¤t_groups as *const u32 as *const _);
enc.set_bytes(3, 4, &fan_in as *const u32 as *const _);
enc.dispatch_thread_groups(
MTLSize::new(out_groups as u64, 1, 1),
MTLSize::new(256, 1, 1),
);
current_groups = out_groups;
which = 1 - which;
}
which
}
}