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)]
pub struct TqParams {
pub n_tokens: u32,
pub n_heads: u32,
pub head_dim: u32,
pub src_stride: u32,
pub dst_pos: u32,
pub max_seq_len: u32,
pub sign_off: u32,
pub q_cap: u32,
pub c0: f32,
pub c1: f32,
pub c2: f32,
pub c3: f32,
pub b0: f32,
pub b1: f32,
pub b2: f32,
pub _pad: u32,
}
const _: () = assert!(size_of::<TqParams>() == 64);
impl MetalParams for TqParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct TqAttnParams {
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub max_seq: u32,
pub start_pos: u32,
pub scale: f32,
pub q_cap: u32,
pub out_stride: u32,
pub qjl_scale: f32,
pub sign_off: u32,
pub c0: f32,
pub c1: f32,
pub c2: f32,
pub c3: f32,
pub q_base: u32,
pub cache_cap: u32,
}
const _: () = assert!(size_of::<TqAttnParams>() == 64);
impl MetalParams for TqAttnParams {}
#[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 {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ArgmaxParams {
pub n: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<ArgmaxParams>() == 8);
impl MetalParams for ArgmaxParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct NormParams {
pub n: u32,
pub eps_bits: u32,
pub _pad0: u32,
pub _pad1: u32,
}
const _: () = assert!(size_of::<NormParams>() == 16);
impl MetalParams for NormParams {}
impl NormParams {
pub fn words(n: u32, eps_bits: u32) -> [u32; 4] {
[n, eps_bits, 0, 0]
}
}
const _: () = assert!(size_of::<[u32; 4]>() == size_of::<NormParams>());
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Conv1dParams {
pub hs: u32,
pub kernel_size: u32,
pub d_conv: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<Conv1dParams>() == 16);
impl MetalParams for Conv1dParams {}
impl Conv1dParams {
pub fn words(hs: u32, kernel_size: u32, d_conv: u32) -> [u32; 4] {
[hs, kernel_size, d_conv, 0]
}
}
const _: () = assert!(size_of::<[u32; 4]>() == size_of::<Conv1dParams>());
impl ArgmaxParams {
pub fn words(n: u32) -> [u32; 2] {
[n, 0]
}
}
const _: () = assert!(size_of::<[u32; 2]>() == size_of::<ArgmaxParams>());
impl ElementwiseParams {
pub fn new(n: u32) -> Self {
Self { n, _pad: 0 }
}
pub fn words(n: u32) -> [u32; 2] {
[n, 0]
}
}
const _: () = assert!(size_of::<[u32; 2]>() == size_of::<ElementwiseParams>());
#[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 {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Conv2dDirectParams {
pub in_ch: u32,
pub out_ch: u32,
pub h_in: u32,
pub w_in: u32,
pub kh: u32,
pub kw: u32,
pub stride_h: u32,
pub stride_w: u32,
pub pad_h: u32,
pub pad_w: u32,
pub h_out: u32,
pub w_out: u32,
pub groups: u32,
pub _pad0: u32,
pub _pad1: u32,
pub _pad2: u32,
}
const _: () = assert!(size_of::<Conv2dDirectParams>() == 64);
impl MetalParams for Conv2dDirectParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TransposeBlockedParams {
pub a: u32,
pub b: u32,
pub k: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<TransposeBlockedParams>() == 16);
impl MetalParams for TransposeBlockedParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Batch2dParams {
pub outer: u32,
pub inner: u32,
pub _pad0: u32,
pub _pad1: u32,
}
const _: () = assert!(size_of::<Batch2dParams>() == 16);
impl MetalParams for Batch2dParams {}
impl Batch2dParams {
pub fn new(outer: u32, inner: u32) -> Self {
Self {
outer,
inner,
_pad0: 0,
_pad1: 0,
}
}
}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct AudioXlAttnParams {
pub tokens: u32,
pub n_head: u32,
pub head_dim: u32,
pub scale_bits: u32,
}
const _: () = assert!(size_of::<AudioXlAttnParams>() == 16);
impl MetalParams for AudioXlAttnParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct StftFrameParams {
pub n_frames: u32,
pub n_fft: u32,
pub hop: u32,
pub center_pad: u32,
pub n_samples: u32,
pub preemph_bits: u32,
pub _pad0: u32,
pub _pad1: u32,
}
const _: () = assert!(size_of::<StftFrameParams>() == 32);
impl MetalParams for StftFrameParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct PowerSpecParams {
pub n_frames: u32,
pub n_fft: u32,
pub n_bins: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<PowerSpecParams>() == 16);
impl MetalParams for PowerSpecParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MelProjectParams {
pub n_mel: u32,
pub n_frames: u32,
pub n_bins: u32,
pub eps_bits: u32,
}
const _: () = assert!(size_of::<MelProjectParams>() == 16);
impl MetalParams for MelProjectParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MelNormParams {
pub n_mel: u32,
pub n_frames: u32,
pub effective_n_len: u32,
pub eps_bits: u32,
}
const _: () = assert!(size_of::<MelNormParams>() == 16);
impl MetalParams for MelNormParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MoeRouteParams {
pub n_expert: u32,
pub n_used: u32,
pub n_tokens: u32,
pub _pad: u32,
}
const _: () = assert!(size_of::<MoeRouteParams>() == 16);
impl MetalParams for MoeRouteParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MoeGemvParams {
pub m: u32,
pub k: u32,
pub n_used: u32,
pub n_entries: u32,
pub expert_stride: u32,
pub x_by_entry: u32,
pub _pad0: u32,
pub _pad1: u32,
}
const _: () = assert!(size_of::<MoeGemvParams>() == 32);
impl MetalParams for MoeGemvParams {}
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MoeCombineParams {
pub hidden: u32,
pub n_used: u32,
pub n_tokens: u32,
pub accumulate: u32,
}
const _: () = assert!(size_of::<MoeCombineParams>() == 16);
impl MetalParams for MoeCombineParams {}