use metal::{Buffer, ComputeCommandEncoderRef};
pub trait MetalParams: Sized {
fn set(&self, enc: &ComputeCommandEncoderRef, slot: u64) {
enc.set_bytes(
slot,
std::mem::size_of_val(self) as u64,
self as *const Self as *const _,
);
}
}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct QkNormRopeParams {
pub pos: u32,
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub eps_bits: u32,
pub freq_base_bits: u32,
pub rope_type: u32,
pub has_freq_factors: u32,
pub has_qk_norm: u32,
}
const _: () = assert!(size_of::<QkNormRopeParams>() == 36);
impl QkNormRopeParams {
pub fn bind(&self, enc: &ComputeCommandEncoderRef, freq_factors: &Buffer) {
self.set(enc, 4);
enc.set_buffer(5, Some(freq_factors), 0);
}
}
impl MetalParams for QkNormRopeParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct QkNormRopeBatchParams {
pub start_pos: u32,
pub n_tokens: u32,
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub eps_bits: u32,
pub freq_base_bits: u32,
pub rope_type: u32,
pub q_stride: u32,
pub k_stride: u32,
pub has_freq_factors: u32,
pub has_qk_norm: u32,
}
const _: () = assert!(size_of::<QkNormRopeBatchParams>() == 48);
impl QkNormRopeBatchParams {
pub fn bind(&self, enc: &ComputeCommandEncoderRef, freq_factors: &Buffer) {
self.set(enc, 4);
enc.set_buffer(5, Some(freq_factors), 0);
}
}
impl MetalParams for QkNormRopeBatchParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct KvShiftKParams {
pub n_keep: u32,
pub shift: u32,
pub new_seq_len: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub freq_base_bits: u32,
pub delta_pos: i32,
pub rope_type: u32,
pub has_freq_factors: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<KvShiftKParams>() == 40);
impl KvShiftKParams {
pub fn bind(&self, enc: &ComputeCommandEncoderRef, freq_factors: &Buffer) {
self.set(enc, 2);
enc.set_buffer(3, Some(freq_factors), 0);
}
}
impl MetalParams for KvShiftKParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RopeParams {
pub pos: u32,
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub freq_base_bits: u32,
}
const _: () = assert!(size_of::<RopeParams>() == 20);
impl MetalParams for RopeParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct GemmF32Params {
pub m: u32,
pub n: u32,
pub k: u32,
}
const _: () = assert!(size_of::<GemmF32Params>() == 12);
impl MetalParams for GemmF32Params {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct QuantGemmParams {
pub m: u32,
pub k: u32,
pub n: u32,
pub x_stride: u32,
pub y_stride: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<QuantGemmParams>() == 24);
impl MetalParams for QuantGemmParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct GemvBatchParams {
pub m: u32,
pub k: u32,
pub n: u32,
pub x_stride: u32,
pub y_stride: u32,
pub accum: u32,
}
const _: () = assert!(size_of::<GemvBatchParams>() == 24);
impl MetalParams for GemvBatchParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct GemvQkvParams {
pub m_q: u32,
pub m_kv: u32,
pub k: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<GemvQkvParams>() == 16);
impl MetalParams for GemvQkvParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct GemvRmsParams {
pub m: u32,
pub k: u32,
pub eps_bits: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<GemvRmsParams>() == 16);
impl MetalParams for GemvRmsParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct GemvSplitKParams {
pub m: u32,
pub k: u32,
pub n_splits: u32,
}
const _: () = assert!(size_of::<GemvSplitKParams>() == 12);
impl MetalParams for GemvSplitKParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct FlashAttnParams {
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub kv_dim: u32,
pub seq_len: u32,
pub scale_bits: u32,
pub _pad0: u32,
pub _pad1: u32,
}
const _: () = assert!(size_of::<FlashAttnParams>() == 32);
impl MetalParams for FlashAttnParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct SplitAttnParams {
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub kv_dim: u32,
pub seq_len: u32,
pub scale_bits: u32,
pub n_splits: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<SplitAttnParams>() == 32);
impl MetalParams for SplitAttnParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct PrefillAttnParams {
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub kv_dim: u32,
pub start_pos: u32,
pub n_queries: u32,
pub scale_bits: u32,
pub q_stride: u32,
pub out_stride: u32,
}
const _: () = assert!(size_of::<PrefillAttnParams>() == 36);
impl MetalParams for PrefillAttnParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ElementwiseParams {
pub n: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<ElementwiseParams>() == 8);
impl MetalParams for ElementwiseParams {}
impl ElementwiseParams {
pub fn new(n: u32) -> Self {
Self { n, _pad: 0 }
}
}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ScaleParams {
pub n: u32,
pub scale_bits: u32,
}
const _: () = assert!(size_of::<ScaleParams>() == 8);
impl MetalParams for ScaleParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct BiasAddParams {
pub total: u32,
pub dim: u32,
}
const _: () = assert!(size_of::<BiasAddParams>() == 8);
impl MetalParams for BiasAddParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RmsNormBatchParams {
pub n: u32,
pub eps_bits: u32,
pub src_stride: u32,
pub dst_stride: u32,
pub res_scale_bits: u32,
}
const _: () = assert!(size_of::<RmsNormBatchParams>() == 20);
impl MetalParams for RmsNormBatchParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Conv1dBatchParams {
pub hidden_size: u32,
pub kernel_size: u32,
pub d_conv: u32,
pub n_tokens: u32,
pub proj_stride: u32,
pub out_stride: u32,
}
const _: () = assert!(size_of::<Conv1dBatchParams>() == 24);
impl MetalParams for Conv1dBatchParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct KvCopyParams {
pub n_elements: u32,
pub src_offset_elements: u32,
pub dst_offset_elements: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<KvCopyParams>() == 16);
impl MetalParams for KvCopyParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct VitLinearParams {
pub m: u32,
pub k: u32,
pub n: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<VitLinearParams>() == 16);
impl MetalParams for VitLinearParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct VitAttnParams {
pub tokens: u32,
pub n_head: u32,
pub head_dim: u32,
pub scale_bits: u32,
}
const _: () = assert!(size_of::<VitAttnParams>() == 16);
impl MetalParams for VitAttnParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LayerNormBatchParams {
pub n: u32,
pub eps_bits: u32,
pub src_stride: u32,
pub dst_stride: u32,
}
const _: () = assert!(size_of::<LayerNormBatchParams>() == 16);
impl MetalParams for LayerNormBatchParams {}