crabml 0.1.0

crabml core package
use bytemuck;

#[derive(Debug, Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
#[repr(C, align(16))]
pub struct RmsNormMeta {
    pub m: u32,
    pub n: u32,
    pub eps: f32,
    pub _padding: f32,
}

// (M, N) x (N, K) = (M, K), now we only support K = 1
#[derive(Debug, Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
#[repr(C, align(16))]
pub struct MatmulMeta {
    pub m: u32,
    pub k: u32,
    pub n: u32,
    pub _padding: u32,
}

// (M, N, K) x (N, K) = (M, N)
#[derive(Debug, Copy, Clone, bytemuck::Pod, bytemuck::Zeroable, Default)]
#[repr(C, align(16))]
pub struct BatchMatmulMeta {
    pub m: u32,
    pub n: u32,
    pub k: u32,
    pub _padding_0: u32,
    pub strides_0: [u32; 3],
    pub _padding_1: u32,
}

#[derive(Debug, Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
#[repr(C, align(16))]
pub struct RopeMeta {
    pub m: u32,
    pub n: u32,
    pub pos: u32,
    pub n_heads: u32,
    pub rope_dims: u32,
    pub _padding: [u32; 7],
}