use crate::{
Q8Activations, Q8KActivations, Q4_0_BLOCK_BYTES, Q4_0_BLOCK_ELEMS, Q4_K_BLOCK_BYTES,
Q4_K_BLOCK_ELEMS, Q5_K_BLOCK_BYTES, Q5_K_BLOCK_ELEMS, Q6_K_BLOCK_BYTES, Q6_K_BLOCK_ELEMS,
Q8_0_BLOCK_BYTES, Q8_0_BLOCK_ELEMS,
};
use half::f16;
pub const Q4_KX8_BLOCK_BYTES: usize = 1152;
pub const Q4_KX8_NROWS: usize = 8;
const KMASK1: u32 = 0x3f3f_3f3f;
const KMASK2: u32 = 0x0f0f_0f0f;
const KMASK3: u32 = 0x0303_0303;
#[inline]
pub fn q4_kx8_interleave() -> usize {
if cfg!(target_arch = "x86_64") {
return 8;
}
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return 8;
}
}
4
}
#[inline]
fn f16_from_bytes(b: &[u8]) -> f32 {
f16::from_le_bytes([b[0], b[1]]).to_f32()
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum AccelX4 {
NeonI8mm,
Portable,
}
impl AccelX4 {
#[inline]
pub fn detect() -> Self {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return AccelX4::NeonI8mm;
}
}
AccelX4::Portable
}
}
pub fn make_block_q4_kx8(
rows: [&[u8]; Q4_KX8_NROWS],
interleave: usize,
) -> [u8; Q4_KX8_BLOCK_BYTES] {
debug_assert!(interleave == 4 || interleave == 8);
for r in &rows {
debug_assert_eq!(r.len(), Q4_K_BLOCK_BYTES);
}
let mut out = [0u8; Q4_KX8_BLOCK_BYTES];
for (i, row) in rows.iter().enumerate() {
out[i * 2] = row[0];
out[i * 2 + 1] = row[1];
out[16 + i * 2] = row[2];
out[16 + i * 2 + 1] = row[3];
}
let end = (Q4_K_BLOCK_ELEMS * 4) / interleave; let qs_out = &mut out[128..];
for i in 0..end {
let src_id = i % Q4_KX8_NROWS;
let src_offset = (i / Q4_KX8_NROWS) * interleave;
let dst_offset = i * interleave;
let src_qs = &rows[src_id][16..144];
qs_out[dst_offset..dst_offset + interleave]
.copy_from_slice(&src_qs[src_offset..src_offset + interleave]);
}
let mut s = [0u8; 8];
let mut m = [0u8; 8];
let scales_out = &mut out[32..128];
for i in 0..4 {
for j in 0..8 {
let sc = &rows[j][4..16];
s[j] = sc[i] & 63;
m[j] = sc[i + 4] & 63;
}
let base = i * 12;
scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
}
for i in 0..4 {
for j in 0..8 {
let sc = &rows[j][4..16];
s[j] = ((sc[i] & 192) >> 2) | (sc[i + 8] & 15);
m[j] = ((sc[i + 4] & 192) >> 2) | ((sc[i + 8] & 240) >> 4);
}
let base = 48 + i * 12;
scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
}
out
}
pub fn pack_q4_k_matrix_x8(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
assert!(cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
let n_blocks = cols / Q4_K_BLOCK_ELEMS;
let row_bytes = n_blocks * Q4_K_BLOCK_BYTES;
assert_eq!(data.len(), rows * row_bytes);
let n_groups = rows / Q4_KX8_NROWS;
let mut out = Vec::with_capacity(n_groups * n_blocks * Q4_KX8_BLOCK_BYTES);
for g in 0..n_groups {
for b in 0..n_blocks {
let mut row_refs: [&[u8]; Q4_KX8_NROWS] = [&[]; Q4_KX8_NROWS];
for (r, slot) in row_refs.iter_mut().enumerate() {
let base = (g * Q4_KX8_NROWS + r) * row_bytes + b * Q4_K_BLOCK_BYTES;
*slot = &data[base..base + Q4_K_BLOCK_BYTES];
}
out.extend_from_slice(&make_block_q4_kx8(row_refs, interleave));
}
}
out
}
#[inline]
fn decode_scales_mins(scales12: &[u8], scales_out: &mut [u8; 8], mins_out: &mut [u8; 8]) {
debug_assert!(scales12.len() >= 12);
let mut utmp = [0u32; 4];
utmp[0] = u32::from_le_bytes(scales12[0..4].try_into().unwrap());
utmp[1] = u32::from_le_bytes(scales12[4..8].try_into().unwrap());
utmp[2] = u32::from_le_bytes(scales12[8..12].try_into().unwrap());
utmp[3] = ((utmp[2] >> 4) & KMASK2) | (((utmp[1] >> 6) & KMASK3) << 4);
let uaux_0 = utmp[1] & KMASK1;
utmp[1] = (utmp[2] & KMASK2) | (((utmp[0] >> 6) & KMASK3) << 4);
utmp[2] = uaux_0;
utmp[0] &= KMASK1;
let bytes = unsafe { std::slice::from_raw_parts(utmp.as_ptr() as *const u8, 16) };
scales_out.copy_from_slice(&bytes[0..8]);
mins_out.copy_from_slice(&bytes[8..16]);
}
fn gemv_q4_kx8_q8_k_scalar_4(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let blocklen = 4;
let ncols_interleaved = Q4_KX8_NROWS;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
debug_assert_eq!(packed.len(), n_row_groups * nb * Q4_KX8_BLOCK_BYTES);
for x in 0..n_row_groups {
let mut sumf = [0f32; 8];
let mut sum_minf = [0f32; 8];
let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q4_KX8_BLOCK_BYTES..][..Q4_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qs = &blk[128..];
let da = act.d[l];
let q8 = &act.q[l * Q4_K_BLOCK_ELEMS..(l + 1) * Q4_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen); for k in 0..n_k {
let sb_pair = k / 8;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
for j in 0..ncols_interleaved {
let mut sumi = 0i32;
for i in 0..blocklen {
let qbyte = qs[k * ncols_interleaved * blocklen + j * blocklen + i];
let v0 = (qbyte & 0x0F) as i32;
let v1 = (qbyte >> 4) as i32;
let a0 = q8[(k / 8) * 64 + (k % 8) * blocklen + i] as i32;
let a1 = q8[(k / 8) * 64 + (k % 8) * blocklen + i + 32] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for sb in 0..8 {
let mins = &all_mins[sb];
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols_interleaved {
sum_minf[j] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
let base = x * ncols_interleaved;
for j in 0..ncols_interleaved {
out[base + j] = sumf[j] - sum_minf[j];
}
}
}
fn gemv_q4_kx8_q8_k_scalar_8(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols_interleaved = Q4_KX8_NROWS;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
for x in 0..n_row_groups {
let mut sumf = [0f32; 8];
let mut sum_minf = [0f32; 8];
let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q4_KX8_BLOCK_BYTES..][..Q4_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qs = &blk[128..];
let da = act.d[l];
let q8 = &act.q[l * Q4_K_BLOCK_ELEMS..(l + 1) * Q4_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen); for k in 0..n_k {
let sb_pair = k / 4;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
for j in 0..ncols_interleaved {
let mut sumi = 0i32;
for i in 0..blocklen {
let qbyte = qs[k * ncols_interleaved * blocklen + j * blocklen + i];
let v0 = (qbyte & 0x0F) as i32;
let v1 = (qbyte >> 4) as i32;
let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for sb in 0..8 {
let mins = &all_mins[sb];
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols_interleaved {
sum_minf[j] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
let base = x * ncols_interleaved;
for j in 0..ncols_interleaved {
out[base + j] = sumf[j] - sum_minf[j];
}
}
}
pub fn gemv_q4_kx8_q8_k(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
interleave: usize,
out: &mut [f32],
) {
assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
assert_eq!(out.len(), n_row_groups * Q4_KX8_NROWS);
match interleave {
4 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q4_kx8_q8_k_neon_sdot(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q4_kx8_q8_k_scalar_4(packed, act, n_cols, n_row_groups, out);
}
8 => {
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
unsafe {
avx2::gemv_q4_kx8_q8_k_avx2(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q4_kx8_q8_k_neon_8x8(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q4_kx8_q8_k_scalar_8(packed, act, n_cols, n_row_groups, out);
}
_ => panic!("q4_kx8 interleave must be 4 or 8, got {interleave}"),
}
}
#[inline]
pub fn gemv_q4_kx8_group(
packed: &[u8],
group: usize,
act: &Q8KActivations,
n_cols: usize,
interleave: usize,
out8: &mut [f32],
) {
debug_assert_eq!(out8.len(), Q4_KX8_NROWS);
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let off = group * nb * Q4_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q4_KX8_BLOCK_BYTES];
gemv_q4_kx8_q8_k(slice, act, n_cols, 1, interleave, out8);
}
pub const Q8_0X4_BLOCK_BYTES: usize = 136;
pub const Q8_0X4_NROWS: usize = 4;
pub const Q8_0X4_INTERLEAVE: usize = 4;
#[inline]
pub fn q8_0x4_interleave() -> usize {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return 8;
}
}
Q8_0X4_INTERLEAVE
}
pub fn make_block_q8_0x4(
rows: [&[u8]; Q8_0X4_NROWS],
interleave: usize,
) -> [u8; Q8_0X4_BLOCK_BYTES] {
debug_assert!(interleave == 4 || interleave == 8);
for r in &rows {
debug_assert_eq!(r.len(), Q8_0_BLOCK_BYTES);
}
let mut out = [0u8; Q8_0X4_BLOCK_BYTES];
for (i, row) in rows.iter().enumerate() {
out[i * 2] = row[0];
out[i * 2 + 1] = row[1];
}
let end = (Q8_0_BLOCK_ELEMS * Q8_0X4_NROWS) / interleave;
let qs_out = &mut out[8..];
for i in 0..end {
let src_id = i % Q8_0X4_NROWS;
let src_offset = (i / Q8_0X4_NROWS) * interleave;
let dst_offset = i * interleave;
let src_qs = &rows[src_id][2..34];
qs_out[dst_offset..dst_offset + interleave]
.copy_from_slice(&src_qs[src_offset..src_offset + interleave]);
}
out
}
pub fn pack_q8_0_matrix_x4(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
assert!(cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
let n_blocks = cols / Q8_0_BLOCK_ELEMS;
let row_bytes = n_blocks * Q8_0_BLOCK_BYTES;
assert_eq!(data.len(), rows * row_bytes);
let n_groups = rows / Q8_0X4_NROWS;
let mut out = Vec::with_capacity(n_groups * n_blocks * Q8_0X4_BLOCK_BYTES);
for g in 0..n_groups {
for b in 0..n_blocks {
let mut row_refs: [&[u8]; Q8_0X4_NROWS] = [&[]; Q8_0X4_NROWS];
for (r, slot) in row_refs.iter_mut().enumerate() {
let base = (g * Q8_0X4_NROWS + r) * row_bytes + b * Q8_0_BLOCK_BYTES;
*slot = &data[base..base + Q8_0_BLOCK_BYTES];
}
out.extend_from_slice(&make_block_q8_0x4(row_refs, interleave));
}
}
out
}
fn gemv_q8_0x4_q8_0_scalar(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
blocklen: usize,
out: &mut [f32],
) {
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let ncols = Q8_0X4_NROWS;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols);
debug_assert_eq!(packed.len(), n_row_groups * nb * Q8_0X4_BLOCK_BYTES);
for x in 0..n_row_groups {
let mut sumf = [0f32; 4];
let group_off = x * nb * Q8_0X4_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q8_0X4_BLOCK_BYTES..][..Q8_0X4_BLOCK_BYTES];
let qs = &blk[8..];
let da = act.d[l];
let q8 = &act.q[l * Q8_0_BLOCK_ELEMS..(l + 1) * Q8_0_BLOCK_ELEMS];
for k in 0..(Q8_0_BLOCK_ELEMS / blocklen) {
for j in 0..ncols {
let mut sumi = 0i32;
for i in 0..blocklen {
let v0 = qs[k * ncols * blocklen + j * blocklen + i] as i8 as i32;
sumi += v0 * q8[k * blocklen + i] as i32;
}
sumf[j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
}
}
}
let base = x * ncols;
out[base..base + ncols].copy_from_slice(&sumf);
}
}
pub fn gemv_q8_0x4_q8_0(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
interleave: usize,
out: &mut [f32],
) {
assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
assert_eq!(out.len(), n_row_groups * Q8_0X4_NROWS);
match interleave {
4 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q8_0x4_q8_0_neon_sdot(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q8_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 4, out);
}
8 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q8_0x4_q8_0_neon_4x8(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q8_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 8, out);
}
_ => panic!("q8_0x4 interleave must be 4 or 8, got {interleave}"),
}
}
pub const Q8_0X4_GEMM_NC: usize = 8;
pub fn gemm_q8_0x4_group(
packed: &[u8],
group: usize,
acts: &[Q8Activations],
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
assert_eq!(out.len(), Q8_0X4_NROWS * acts.len());
assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
if acts.is_empty() {
return;
}
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let off = group * nb * Q8_0X4_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q8_0X4_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
{
if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
for (t, chunk) in acts.chunks(Q8K_ACTS_X4_NC).enumerate() {
let tile = prepare_q8_acts_x4(chunk, n_cols);
let mut tmp = [0f32; Q8_0X4_NROWS * Q8K_ACTS_X4_NC];
let n = chunk.len();
unsafe {
neon::gemm_q8_0x4_q8_0_neon_i8mm(
slice,
&tile,
n_cols,
&mut tmp[..Q8_0X4_NROWS * n],
);
}
for r in 0..Q8_0X4_NROWS {
for j in 0..n {
out[r * acts.len() + t * Q8K_ACTS_X4_NC + j] = tmp[r * n + j];
}
}
}
return;
}
if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemm_q8_0x4_q8_0_neon_sdot(slice, acts, n_cols, out);
}
return;
}
}
let mut tmp = [0f32; Q8_0X4_NROWS];
for (j, act) in acts.iter().enumerate() {
gemv_q8_0x4_q8_0(slice, act, n_cols, 1, interleave, &mut tmp);
for (r, v) in tmp.iter().enumerate() {
out[r * acts.len() + j] = *v;
}
}
}
pub struct Q8ActsX4 {
pub na: usize,
pub n_blocks: usize,
pub qs: Vec<i8>,
pub d: Vec<f32>,
}
pub fn prepare_q8_acts_x4(acts: &[Q8Activations], n_cols: usize) -> Q8ActsX4 {
assert!(acts.len() <= Q8K_ACTS_X4_NC);
assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
let na = acts.len();
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let mut qs = vec![0i8; nb * Q8_0_BLOCK_ELEMS * 4];
let mut d = vec![0f32; nb * 4];
for (a, act) in acts.iter().enumerate() {
debug_assert_eq!(act.d.len(), nb);
for b in 0..nb {
let src = &act.q[b * Q8_0_BLOCK_ELEMS..(b + 1) * Q8_0_BLOCK_ELEMS];
let dst = &mut qs[b * Q8_0_BLOCK_ELEMS * 4..(b + 1) * Q8_0_BLOCK_ELEMS * 4];
for (c, run) in src.as_chunks::<8>().0.iter().enumerate() {
dst[c * 32 + a * 8..c * 32 + a * 8 + 8].copy_from_slice(run);
}
d[b * 4 + a] = act.d[b];
}
}
Q8ActsX4 {
na,
n_blocks: nb,
qs,
d,
}
}
#[inline]
pub fn q8_0x4_gemm_uses_acts_x4(interleave: usize) -> bool {
#[cfg(target_arch = "aarch64")]
{
interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
}
#[cfg(not(target_arch = "aarch64"))]
{
let _ = interleave;
false
}
}
pub fn gemm_q8_0x4_group_x4(
packed: &[u8],
group: usize,
tile: &Q8ActsX4,
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
gemm_q8_0x4_group_x4_on(
packed,
group,
tile,
n_cols,
interleave,
AccelX4::detect(),
out,
);
}
#[inline]
pub fn gemm_q8_0x4_group_x4_on(
packed: &[u8],
group: usize,
tile: &Q8ActsX4,
n_cols: usize,
interleave: usize,
accel: AccelX4,
out: &mut [f32],
) {
assert_eq!(
interleave, 8,
"the x4 GEMM only exists for interleave-8 packing"
);
assert_eq!(out.len(), Q8_0X4_NROWS * tile.na);
assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
debug_assert_eq!(tile.n_blocks, n_cols / Q8_0_BLOCK_ELEMS);
if tile.na == 0 {
return;
}
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let off = group * nb * Q8_0X4_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q8_0X4_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
if accel == AccelX4::NeonI8mm {
unsafe {
neon::gemm_q8_0x4_q8_0_neon_i8mm(slice, tile, n_cols, out);
}
return;
}
let _ = accel;
gemm_q8_0x4_acts_x4_scalar_8(slice, tile, n_cols, out);
}
fn gemm_q8_0x4_acts_x4_scalar_8(packed: &[u8], tile: &Q8ActsX4, n_cols: usize, out: &mut [f32]) {
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let blocklen = 8;
let ncols = Q8_0X4_NROWS;
let na = tile.na;
let mut sumf = [[0f32; Q8_0X4_NROWS]; Q8K_ACTS_X4_NC];
for l in 0..nb {
let blk = &packed[l * Q8_0X4_BLOCK_BYTES..][..Q8_0X4_BLOCK_BYTES];
let qs = &blk[8..];
let q8 = &tile.qs[l * Q8_0_BLOCK_ELEMS * 4..][..Q8_0_BLOCK_ELEMS * 4];
for a in 0..na {
let da = tile.d[l * 4 + a];
for k in 0..(Q8_0_BLOCK_ELEMS / blocklen) {
for j in 0..ncols {
let mut sumi = 0i32;
for i in 0..blocklen {
let v0 = qs[k * ncols * blocklen + j * blocklen + i] as i8 as i32;
let e = k * blocklen + i;
sumi += v0 * q8[(e / 8) * 32 + a * 8 + (e % 8)] as i32;
}
sumf[a][j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
}
}
}
}
for j in 0..ncols {
for (a, row) in sumf.iter().take(na).enumerate() {
out[j * na + a] = row[j];
}
}
}
pub const Q4_KX8_GEMM_NC: usize = 4;
pub fn gemm_q4_kx8_group(
packed: &[u8],
group: usize,
acts: &[Q8KActivations],
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
assert_eq!(out.len(), Q4_KX8_NROWS * acts.len());
assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
if acts.is_empty() {
return;
}
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let off = group * nb * Q4_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q4_KX8_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
if acts.len() <= Q4_KX8_GEMM_NC {
if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
let tile = prepare_q8_k_acts_x4(acts, n_cols);
unsafe {
neon::gemm_q4_kx8_q8_k_neon_i8mm(slice, &tile, n_cols, out);
}
return;
}
if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemm_q4_kx8_q8_k_neon_sdot(slice, acts, n_cols, out);
}
return;
}
}
let mut tmp = [0f32; Q4_KX8_NROWS];
for (j, act) in acts.iter().enumerate() {
gemv_q4_kx8_q8_k(slice, act, n_cols, 1, interleave, &mut tmp);
for (r, v) in tmp.iter().enumerate() {
out[r * acts.len() + j] = *v;
}
}
}
pub const Q8K_ACTS_X4_NC: usize = 4;
pub struct Q8KActsX4 {
pub na: usize,
pub n_blocks: usize,
pub qs: Vec<i8>,
pub bsums: Vec<i16>,
pub d: Vec<f32>,
}
pub fn prepare_q8_k_acts_x4(acts: &[Q8KActivations], n_cols: usize) -> Q8KActsX4 {
assert!(acts.len() <= Q8K_ACTS_X4_NC);
assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
let na = acts.len();
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let mut qs = vec![0i8; nb * Q4_K_BLOCK_ELEMS * 4];
let mut bsums = vec![0i16; nb * 4 * 8];
let mut d = vec![0f32; nb * 4];
for (a, act) in acts.iter().enumerate() {
debug_assert_eq!(act.n_blocks(), nb);
for b in 0..nb {
let src = &act.q[b * Q4_K_BLOCK_ELEMS..(b + 1) * Q4_K_BLOCK_ELEMS];
let dst = &mut qs[b * Q4_K_BLOCK_ELEMS * 4..(b + 1) * Q4_K_BLOCK_ELEMS * 4];
for (c, run) in src.as_chunks::<8>().0.iter().enumerate() {
dst[c * 32 + a * 8..c * 32 + a * 8 + 8].copy_from_slice(run);
}
let src_bs = &act.bsums[b * 16..(b + 1) * 16];
let dst_bs = &mut bsums[(b * 4 + a) * 8..(b * 4 + a) * 8 + 8];
for (slot, pair) in dst_bs.iter_mut().zip(src_bs.as_chunks::<2>().0) {
*slot = pair[0] + pair[1];
}
d[b * 4 + a] = act.d[b];
}
}
Q8KActsX4 {
na,
n_blocks: nb,
qs,
bsums,
d,
}
}
#[inline]
pub fn q4_kx8_gemm_uses_acts_x4(interleave: usize) -> bool {
#[cfg(target_arch = "aarch64")]
{
interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
}
#[cfg(not(target_arch = "aarch64"))]
{
let _ = interleave;
false
}
}
pub fn gemm_q4_kx8_group_x4(
packed: &[u8],
group: usize,
tile: &Q8KActsX4,
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
gemm_q4_kx8_group_x4_on(
packed,
group,
tile,
n_cols,
interleave,
AccelX4::detect(),
out,
);
}
#[inline]
pub fn gemm_q4_kx8_group_x4_on(
packed: &[u8],
group: usize,
tile: &Q8KActsX4,
n_cols: usize,
interleave: usize,
accel: AccelX4,
out: &mut [f32],
) {
assert_eq!(
interleave, 8,
"the x4 GEMM only exists for interleave-8 packing"
);
assert_eq!(out.len(), Q4_KX8_NROWS * tile.na);
assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
debug_assert_eq!(tile.n_blocks, n_cols / Q4_K_BLOCK_ELEMS);
if tile.na == 0 {
return;
}
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let off = group * nb * Q4_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q4_KX8_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
if accel == AccelX4::NeonI8mm {
unsafe {
neon::gemm_q4_kx8_q8_k_neon_i8mm(slice, tile, n_cols, out);
}
return;
}
let _ = accel;
gemm_q4_kx8_acts_x4_scalar_8(slice, tile, n_cols, out);
}
fn gemm_q4_kx8_acts_x4_scalar_8(packed: &[u8], tile: &Q8KActsX4, n_cols: usize, out: &mut [f32]) {
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols_interleaved = Q4_KX8_NROWS;
let na = tile.na;
let mut sumf = [[0f32; Q4_KX8_NROWS]; Q4_KX8_GEMM_NC];
let mut sum_minf = [[0f32; Q4_KX8_NROWS]; Q4_KX8_GEMM_NC];
for l in 0..nb {
let blk = &packed[l * Q4_KX8_BLOCK_BYTES..][..Q4_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qs = &blk[128..];
let q8 = &tile.qs[l * Q4_K_BLOCK_ELEMS * 4..][..Q4_K_BLOCK_ELEMS * 4];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
for a in 0..na {
let da = tile.d[l * 4 + a];
let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen); for k in 0..n_k {
let sb_pair = k / 4;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
for j in 0..ncols_interleaved {
let mut sumi = 0i32;
for i in 0..blocklen {
let qbyte = qs[k * ncols_interleaved * blocklen + j * blocklen + i];
let v0 = (qbyte & 0x0F) as i32;
let v1 = (qbyte >> 4) as i32;
let e0 = (k >> 2) * 64 + (k % 4) * blocklen + i;
let e1 = e0 + 32;
let a0 = q8[(e0 / 8) * 32 + a * 8 + (e0 % 8)] as i32;
let a1 = q8[(e1 / 8) * 32 + a * 8 + (e1 % 8)] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
sumf[a][j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for (sb, mins) in all_mins.iter().enumerate() {
let bsum = tile.bsums[(l * 4 + a) * 8 + sb] as i32;
for j in 0..ncols_interleaved {
sum_minf[a][j] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
}
for j in 0..ncols_interleaved {
for (a, row) in sumf.iter().take(na).enumerate() {
out[j * na + a] = row[j] - sum_minf[a][j];
}
}
}
pub const Q5_KX8_BLOCK_BYTES: usize = 1408;
pub const Q5_KX8_NROWS: usize = 8;
#[inline]
pub fn q5_kx8_interleave() -> usize {
if cfg!(target_arch = "x86_64") {
return 8;
}
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return 8;
}
}
4
}
pub fn make_block_q5_kx8(
rows: [&[u8]; Q5_KX8_NROWS],
interleave: usize,
) -> [u8; Q5_KX8_BLOCK_BYTES] {
debug_assert!(interleave == 4 || interleave == 8);
for r in &rows {
debug_assert_eq!(r.len(), Q5_K_BLOCK_BYTES);
}
let mut out = [0u8; Q5_KX8_BLOCK_BYTES];
for (i, row) in rows.iter().enumerate() {
out[i * 2] = row[0];
out[i * 2 + 1] = row[1];
out[16 + i * 2] = row[2];
out[16 + i * 2 + 1] = row[3];
}
let end = (Q5_K_BLOCK_ELEMS * 4) / interleave;
let qs_out = &mut out[384..];
for i in 0..end {
let src_id = i % Q5_KX8_NROWS;
let src_offset = (i / Q5_KX8_NROWS) * interleave;
let dst_offset = i * interleave;
let src_qs = &rows[src_id][48..176];
qs_out[dst_offset..dst_offset + interleave]
.copy_from_slice(&src_qs[src_offset..src_offset + interleave]);
}
let qh_end = end / 4;
let qh_out = &mut out[128..384];
for i in 0..qh_end {
let src_id = i % Q5_KX8_NROWS;
let src_offset = (i / Q5_KX8_NROWS) * interleave;
let dst_offset = i * interleave;
let src_qh = &rows[src_id][16..48];
qh_out[dst_offset..dst_offset + interleave]
.copy_from_slice(&src_qh[src_offset..src_offset + interleave]);
}
let mut s = [0u8; 8];
let mut m = [0u8; 8];
let scales_out = &mut out[32..128];
for i in 0..4 {
for j in 0..8 {
let sc = &rows[j][4..16];
s[j] = sc[i] & 63;
m[j] = sc[i + 4] & 63;
}
let base = i * 12;
scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
}
for i in 0..4 {
for j in 0..8 {
let sc = &rows[j][4..16];
s[j] = ((sc[i] & 192) >> 2) | (sc[i + 8] & 15);
m[j] = ((sc[i + 4] & 192) >> 2) | ((sc[i + 8] & 240) >> 4);
}
let base = 48 + i * 12;
scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
}
out
}
pub fn pack_q5_k_matrix_x8(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
assert!(cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
let n_blocks = cols / Q5_K_BLOCK_ELEMS;
let row_bytes = n_blocks * Q5_K_BLOCK_BYTES;
assert_eq!(data.len(), rows * row_bytes);
let n_groups = rows / Q5_KX8_NROWS;
let mut out = Vec::with_capacity(n_groups * n_blocks * Q5_KX8_BLOCK_BYTES);
for g in 0..n_groups {
for b in 0..n_blocks {
let mut row_refs: [&[u8]; Q5_KX8_NROWS] = [&[]; Q5_KX8_NROWS];
for (r, slot) in row_refs.iter_mut().enumerate() {
let base = (g * Q5_KX8_NROWS + r) * row_bytes + b * Q5_K_BLOCK_BYTES;
*slot = &data[base..base + Q5_K_BLOCK_BYTES];
}
out.extend_from_slice(&make_block_q5_kx8(row_refs, interleave));
}
}
out
}
fn gemv_q5_kx8_q8_k_scalar_4(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let blocklen = 4;
let ncols_interleaved = Q5_KX8_NROWS;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
debug_assert_eq!(packed.len(), n_row_groups * nb * Q5_KX8_BLOCK_BYTES);
for x in 0..n_row_groups {
let mut sumf = [0f32; 8];
let mut sum_minf = [0f32; 8];
let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qh = &blk[128..384];
let qs = &blk[384..];
let da = act.d[l];
let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen); for k in 0..n_k {
let sb_pair = k / 8;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
let qh_shift = sb_pair * 2;
for j in 0..ncols_interleaved {
let mut sumi = 0i32;
for i in 0..blocklen {
let b_qs_offset = k * ncols_interleaved * blocklen + j * blocklen + i;
let qh_idx = (k * blocklen + i) % 32;
let qh_chunk = qh_idx / blocklen;
let qh_pos = qh_idx % blocklen;
let b_qh_offset =
qh_chunk * (blocklen * ncols_interleaved) + j * blocklen + qh_pos;
let qh_val = qh[b_qh_offset];
let h0 = (qh_val >> qh_shift) & 1;
let h1 = (qh_val >> (qh_shift + 1)) & 1;
let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
let a0 = q8[(k / 8) * 64 + (k % 8) * blocklen + i] as i32;
let a1 = q8[(k / 8) * 64 + (k % 8) * blocklen + i + 32] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for sb in 0..8 {
let mins = &all_mins[sb];
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols_interleaved {
sum_minf[j] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
let base = x * ncols_interleaved;
for j in 0..ncols_interleaved {
out[base + j] = sumf[j] - sum_minf[j];
}
}
}
fn gemv_q5_kx8_q8_k_scalar_8(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols_interleaved = Q5_KX8_NROWS;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
for x in 0..n_row_groups {
let mut sumf = [0f32; 8];
let mut sum_minf = [0f32; 8];
let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qh = &blk[128..384];
let qs = &blk[384..];
let da = act.d[l];
let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen); for k in 0..n_k {
let sb_pair = k / 4;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
let qh_shift = sb_pair * 2;
for j in 0..ncols_interleaved {
let mut sumi = 0i32;
for i in 0..blocklen {
let b_qs_offset = k * ncols_interleaved * blocklen + j * blocklen + i;
let qh_idx = (k * blocklen + i) % 32;
let qh_chunk = qh_idx / blocklen;
let qh_pos = qh_idx % blocklen;
let b_qh_offset =
qh_chunk * (blocklen * ncols_interleaved) + j * blocklen + qh_pos;
let qh_val = qh[b_qh_offset];
let h0 = (qh_val >> qh_shift) & 1;
let h1 = (qh_val >> (qh_shift + 1)) & 1;
let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for sb in 0..8 {
let mins = &all_mins[sb];
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols_interleaved {
sum_minf[j] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
let base = x * ncols_interleaved;
for j in 0..ncols_interleaved {
out[base + j] = sumf[j] - sum_minf[j];
}
}
}
pub fn gemv_q5_kx8_q8_k(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
interleave: usize,
out: &mut [f32],
) {
assert!(n_cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
assert_eq!(out.len(), n_row_groups * Q5_KX8_NROWS);
match interleave {
4 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q5_kx8_q8_k_neon_sdot(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q5_kx8_q8_k_scalar_4(packed, act, n_cols, n_row_groups, out);
}
8 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q5_kx8_q8_k_neon_8x8(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q5_kx8_q8_k_scalar_8(packed, act, n_cols, n_row_groups, out);
}
_ => panic!("q5_kx8 interleave must be 4 or 8, got {interleave}"),
}
}
#[inline]
pub fn gemv_q5_kx8_group(
packed: &[u8],
group: usize,
act: &Q8KActivations,
n_cols: usize,
interleave: usize,
out8: &mut [f32],
) {
debug_assert_eq!(out8.len(), Q5_KX8_NROWS);
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let off = group * nb * Q5_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q5_KX8_BLOCK_BYTES];
gemv_q5_kx8_q8_k(slice, act, n_cols, 1, interleave, out8);
}
pub const Q5_KX8_GEMM_NC: usize = 4;
pub fn gemm_q5_kx8_group(
packed: &[u8],
group: usize,
acts: &[Q8KActivations],
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
assert_eq!(out.len(), Q5_KX8_NROWS * acts.len());
assert!(n_cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
if acts.is_empty() {
return;
}
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let off = group * nb * Q5_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q5_KX8_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
{
if acts.len() <= Q5_KX8_GEMM_NC {
if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
let tile = prepare_q8_k_acts_x4(acts, n_cols);
unsafe {
neon::gemm_q5_kx8_q8_k_neon_i8mm(slice, &tile, n_cols, out);
}
return;
}
if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemm_q5_kx8_q8_k_neon_sdot(slice, acts, n_cols, out);
}
return;
}
}
}
match interleave {
4 => gemm_q5_kx8_q8_k_scalar_4(slice, acts, n_cols, out),
8 => gemm_q5_kx8_q8_k_scalar_8(slice, acts, n_cols, out),
_ => panic!("q5_kx8 interleave must be 4 or 8, got {interleave}"),
}
}
#[inline]
pub fn q5_kx8_gemm_uses_acts_x4(interleave: usize) -> bool {
#[cfg(target_arch = "aarch64")]
{
interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
}
#[cfg(not(target_arch = "aarch64"))]
{
let _ = interleave;
false
}
}
pub fn gemm_q5_kx8_group_x4(
packed: &[u8],
group: usize,
tile: &Q8KActsX4,
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
gemm_q5_kx8_group_x4_on(
packed,
group,
tile,
n_cols,
interleave,
AccelX4::detect(),
out,
);
}
#[inline]
pub fn gemm_q5_kx8_group_x4_on(
packed: &[u8],
group: usize,
tile: &Q8KActsX4,
n_cols: usize,
interleave: usize,
accel: AccelX4,
out: &mut [f32],
) {
assert_eq!(
interleave, 8,
"the x4 GEMM only exists for interleave-8 packing"
);
assert_eq!(out.len(), Q5_KX8_NROWS * tile.na);
assert!(n_cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
debug_assert_eq!(tile.n_blocks, n_cols / Q5_K_BLOCK_ELEMS);
if tile.na == 0 {
return;
}
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let off = group * nb * Q5_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q5_KX8_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
if accel == AccelX4::NeonI8mm {
unsafe {
neon::gemm_q5_kx8_q8_k_neon_i8mm(slice, tile, n_cols, out);
}
return;
}
let _ = accel;
gemm_q5_kx8_acts_x4_scalar_8(slice, tile, n_cols, out);
}
fn gemm_q5_kx8_acts_x4_scalar_8(packed: &[u8], tile: &Q8KActsX4, n_cols: usize, out: &mut [f32]) {
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols_interleaved = Q5_KX8_NROWS;
let na = tile.na;
let mut sumf = [[0f32; Q5_KX8_NROWS]; Q5_KX8_GEMM_NC];
let mut sum_minf = [[0f32; Q5_KX8_NROWS]; Q5_KX8_GEMM_NC];
for l in 0..nb {
let blk = &packed[l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qh = &blk[128..384];
let qs = &blk[384..];
let q8 = &tile.qs[l * Q5_K_BLOCK_ELEMS * 4..][..Q5_K_BLOCK_ELEMS * 4];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
for a in 0..na {
let da = tile.d[l * 4 + a];
let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen); for k in 0..n_k {
let sb_pair = k / 4;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
let qh_shift = sb_pair * 2;
for j in 0..ncols_interleaved {
let mut sumi = 0i32;
for i in 0..blocklen {
let b_qs_offset = k * ncols_interleaved * blocklen + j * blocklen + i;
let qh_idx = (k * blocklen + i) % 32;
let qh_chunk = qh_idx / blocklen;
let qh_pos = qh_idx % blocklen;
let b_qh_offset =
qh_chunk * (blocklen * ncols_interleaved) + j * blocklen + qh_pos;
let qh_val = qh[b_qh_offset];
let h0 = (qh_val >> qh_shift) & 1;
let h1 = (qh_val >> (qh_shift + 1)) & 1;
let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
let e0 = (k >> 2) * 64 + (k % 4) * blocklen + i;
let e1 = e0 + 32;
let a0 = q8[(e0 / 8) * 32 + a * 8 + (e0 % 8)] as i32;
let a1 = q8[(e1 / 8) * 32 + a * 8 + (e1 % 8)] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
sumf[a][j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for (sb, mins) in all_mins.iter().enumerate() {
let bsum = tile.bsums[(l * 4 + a) * 8 + sb] as i32;
for j in 0..ncols_interleaved {
sum_minf[a][j] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
}
for j in 0..ncols_interleaved {
for (a, row) in sumf.iter().take(na).enumerate() {
out[j * na + a] = row[j] - sum_minf[a][j];
}
}
}
fn gemm_q5_kx8_q8_k_scalar_4(
packed: &[u8],
acts: &[Q8KActivations],
n_cols: usize,
out: &mut [f32],
) {
let na = acts.len();
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let blocklen = 4;
let ncols = Q5_KX8_NROWS;
out.fill(0.0);
let mut sum_minf = vec![0f32; ncols * na];
for l in 0..nb {
let blk = &packed[l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qh = &blk[128..384];
let qs = &blk[384..];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen);
for (a, act) in acts.iter().enumerate() {
let da = act.d[l];
let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
for k in 0..n_k {
let sb_pair = k / 8;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
let qh_shift = sb_pair * 2;
for j in 0..ncols {
let mut sumi = 0i32;
for i in 0..blocklen {
let b_qs_offset = k * ncols * blocklen + j * blocklen + i;
let qh_idx = (k * blocklen + i) % 32;
let qh_chunk = qh_idx / blocklen;
let qh_pos = qh_idx % blocklen;
let b_qh_offset = qh_chunk * (blocklen * ncols) + j * blocklen + qh_pos;
let qh_val = qh[b_qh_offset];
let h0 = (qh_val >> qh_shift) & 1;
let h1 = (qh_val >> (qh_shift + 1)) & 1;
let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
let a0 = q8[(k / 8) * 64 + (k % 8) * blocklen + i] as i32;
let a1 = q8[(k / 8) * 64 + (k % 8) * blocklen + i + 32] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
out[j * na + a] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for sb in 0..8 {
let mins = &all_mins[sb];
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols {
sum_minf[j * na + a] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
}
for i in 0..ncols * na {
out[i] -= sum_minf[i];
}
}
fn gemm_q5_kx8_q8_k_scalar_8(
packed: &[u8],
acts: &[Q8KActivations],
n_cols: usize,
out: &mut [f32],
) {
let na = acts.len();
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols = Q5_KX8_NROWS;
out.fill(0.0);
let mut sum_minf = vec![0f32; ncols * na];
for l in 0..nb {
let blk = &packed[l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let dmin = &blk[16..32];
let scales = &blk[32..128];
let qh = &blk[128..384];
let qs = &blk[384..];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen);
for (a, act) in acts.iter().enumerate() {
let da = act.d[l];
let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
for k in 0..n_k {
let sb_pair = k / 4;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
let qh_shift = sb_pair * 2;
for j in 0..ncols {
let mut sumi = 0i32;
for i in 0..blocklen {
let b_qs_offset = k * ncols * blocklen + j * blocklen + i;
let qh_idx = (k * blocklen + i) % 32;
let qh_chunk = qh_idx / blocklen;
let qh_pos = qh_idx % blocklen;
let b_qh_offset = qh_chunk * (blocklen * ncols) + j * blocklen + qh_pos;
let qh_val = qh[b_qh_offset];
let h0 = (qh_val >> qh_shift) & 1;
let h1 = (qh_val >> (qh_shift + 1)) & 1;
let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
out[j * na + a] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
}
}
for sb in 0..8 {
let mins = &all_mins[sb];
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols {
sum_minf[j * na + a] +=
mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
}
}
}
}
for i in 0..ncols * na {
out[i] -= sum_minf[i];
}
}
pub const Q6_KX8_BLOCK_BYTES: usize = 1680;
pub const Q6_KX8_NROWS: usize = 8;
#[inline]
pub fn q6_kx8_interleave() -> usize {
if cfg!(target_arch = "x86_64") {
return 8;
}
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return 8;
}
}
4
}
pub fn make_block_q6_kx8(
rows: [&[u8]; Q6_KX8_NROWS],
interleave: usize,
) -> [u8; Q6_KX8_BLOCK_BYTES] {
debug_assert!(interleave == 4 || interleave == 8);
for r in &rows {
debug_assert_eq!(r.len(), Q6_K_BLOCK_BYTES);
}
let mut out = [0u8; Q6_KX8_BLOCK_BYTES];
for (i, row) in rows.iter().enumerate() {
out[i * 2] = row[208];
out[i * 2 + 1] = row[209];
}
let end_ls = (Q6_K_BLOCK_ELEMS * 4) / interleave;
let ql_out = &mut out[144..1168];
for i in 0..end_ls {
let src_id = i % Q6_KX8_NROWS;
let src_offset = (i / Q6_KX8_NROWS) * interleave;
let dst_offset = i * interleave;
let src_ql = &rows[src_id][0..128];
ql_out[dst_offset..dst_offset + interleave]
.copy_from_slice(&src_ql[src_offset..src_offset + interleave]);
}
let end_hs = end_ls / 2;
let qh_out = &mut out[1168..];
for i in 0..end_hs {
let src_id = i % Q6_KX8_NROWS;
let src_offset = (i / Q6_KX8_NROWS) * interleave;
let dst_offset = i * interleave;
let src_qh = &rows[src_id][128..192];
qh_out[dst_offset..dst_offset + interleave]
.copy_from_slice(&src_qh[src_offset..src_offset + interleave]);
}
let n_scales = Q6_K_BLOCK_ELEMS / 16;
let scales_out = &mut out[16..144];
for i in 0..Q6_KX8_NROWS {
let src_sc = &rows[i][192..208];
for j in 0..n_scales {
scales_out[j * Q6_KX8_NROWS + i] = src_sc[j];
}
}
out
}
pub fn pack_q6_k_matrix_x8(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
assert!(cols.is_multiple_of(Q6_K_BLOCK_ELEMS));
let row_bytes = (cols / Q6_K_BLOCK_ELEMS) * Q6_K_BLOCK_BYTES;
assert_eq!(data.len(), rows * row_bytes);
let n_blocks = cols / Q6_K_BLOCK_ELEMS;
let n_groups = rows / Q6_KX8_NROWS;
let mut out = Vec::with_capacity(n_groups * n_blocks * Q6_KX8_BLOCK_BYTES);
for g in 0..n_groups {
for b in 0..n_blocks {
let mut row_refs: [&[u8]; Q6_KX8_NROWS] = [&[]; Q6_KX8_NROWS];
for (r, slot) in row_refs.iter_mut().enumerate() {
let base = (g * Q6_KX8_NROWS + r) * row_bytes + b * Q6_K_BLOCK_BYTES;
*slot = &data[base..base + Q6_K_BLOCK_BYTES];
}
out.extend_from_slice(&make_block_q6_kx8(row_refs, interleave));
}
}
out
}
fn gemv_q6_kx8_q8_k_scalar(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
blocklen: usize,
out: &mut [f32],
) {
let nb = n_cols / Q6_K_BLOCK_ELEMS;
let ncols = Q6_KX8_NROWS;
let blocks_per_half = 64 / blocklen;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols);
for x in 0..n_row_groups {
let mut sumf = [0f32; 8];
let group_off = x * nb * Q6_KX8_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q6_KX8_BLOCK_BYTES..][..Q6_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let scales = &blk[16..144];
let ql = &blk[144..1168];
let qh = &blk[1168..];
let da = act.d[l];
let q8 = &act.q[l * Q6_K_BLOCK_ELEMS..(l + 1) * Q6_K_BLOCK_ELEMS];
for k in 0..(Q6_K_BLOCK_ELEMS / (2 * blocklen)) {
let base_l = (k / blocks_per_half) * 128 + (k % blocks_per_half) * blocklen;
let base_h = base_l + 64;
let scale_idx_l = base_l / 16;
let scale_idx_h = base_h / 16;
let qh_shift_l = ((base_l % 128) / 32) * 2;
let qh_shift_h = ((base_h % 128) / 32) * 2;
let qh_half_l = (base_l / 128) * 32;
let qh_half_h = (base_h / 128) * 32;
for j in 0..ncols {
let scale_l = scales[scale_idx_l * ncols + j] as i8 as i32;
let scale_h = scales[scale_idx_h * ncols + j] as i8 as i32;
let mut sumi_l = 0i32;
let mut sumi_h = 0i32;
for i in 0..blocklen {
let ql_pos = k * ncols * blocklen + j * blocklen + i;
let l_4 = (ql[ql_pos] & 0x0F) as i32;
let hi_4 = ((ql[ql_pos] >> 4) & 0x0F) as i32;
let qh_idx_l = qh_half_l + ((base_l + i) % 32);
let qh_chunk_l = qh_idx_l / blocklen;
let qh_pos_l = qh_idx_l % blocklen;
let qh_offset_l = qh_chunk_l * (blocklen * ncols) + j * blocklen + qh_pos_l;
let hi_2_l = ((qh[qh_offset_l] >> qh_shift_l) & 0x3) as i32;
let qh_idx_h = qh_half_h + ((base_h + i) % 32);
let qh_chunk_h = qh_idx_h / blocklen;
let qh_pos_h = qh_idx_h % blocklen;
let qh_offset_h = qh_chunk_h * (blocklen * ncols) + j * blocklen + qh_pos_h;
let hi_2_h = ((qh[qh_offset_h] >> qh_shift_h) & 0x3) as i32;
let q_l = ((hi_2_l << 4) | l_4) - 32;
let q_h = ((hi_2_h << 4) | hi_4) - 32;
sumi_l += q_l * (q8[base_l + i] as i32);
sumi_h += q_h * (q8[base_h + i] as i32);
}
sumf[j] += (sumi_l * scale_l + sumi_h * scale_h) as f32
* f16_from_bytes(&d[j * 2..])
* da;
}
}
}
let base = x * ncols;
out[base..base + ncols].copy_from_slice(&sumf);
}
}
pub fn gemv_q6_kx8_q8_k(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
interleave: usize,
out: &mut [f32],
) {
assert_eq!(out.len(), n_row_groups * Q6_KX8_NROWS);
match interleave {
4 => gemv_q6_kx8_q8_k_scalar(packed, act, n_cols, n_row_groups, 4, out),
8 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q6_kx8_q8_k_neon_8x8(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q6_kx8_q8_k_scalar(packed, act, n_cols, n_row_groups, 8, out)
}
_ => panic!("q6_kx8 interleave must be 4 or 8, got {interleave}"),
}
}
pub fn gemv_q6_kx8_group(
packed: &[u8],
group: usize,
act: &Q8KActivations,
n_cols: usize,
interleave: usize,
out8: &mut [f32],
) {
debug_assert_eq!(out8.len(), Q6_KX8_NROWS);
let nb = n_cols / Q6_K_BLOCK_ELEMS;
let off = group * nb * Q6_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q6_KX8_BLOCK_BYTES];
gemv_q6_kx8_q8_k(slice, act, n_cols, 1, interleave, out8);
}
pub const Q6_KX8_GEMM_NC: usize = 8;
pub fn gemm_q6_kx8_group(
packed: &[u8],
group: usize,
acts: &[Q8KActivations],
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
assert_eq!(out.len(), Q6_KX8_NROWS * acts.len());
assert!(n_cols.is_multiple_of(Q6_K_BLOCK_ELEMS));
if acts.is_empty() {
return;
}
let nb = n_cols / Q6_K_BLOCK_ELEMS;
let off = group * nb * Q6_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q6_KX8_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
{
if interleave == 8
&& acts.len() <= Q8K_ACTS_X4_NC
&& std::arch::is_aarch64_feature_detected!("i8mm")
{
let tile = prepare_q8_k_acts_x4(acts, n_cols);
unsafe {
neon::gemm_q6_kx8_q8_k_neon_i8mm(slice, &tile, n_cols, out);
}
return;
}
}
let blocklen = interleave;
assert!(blocklen == 4 || blocklen == 8);
let na = acts.len();
let ncols = Q6_KX8_NROWS;
let blocks_per_half = 64 / blocklen;
out.fill(0.0);
for l in 0..nb {
let blk = &slice[l * Q6_KX8_BLOCK_BYTES..][..Q6_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let scales = &blk[16..144];
let ql = &blk[144..1168];
let qh = &blk[1168..];
let mut d_f = [0f32; 8];
for j in 0..8 {
d_f[j] = f16_from_bytes(&d[j * 2..]);
}
for k in 0..(Q6_K_BLOCK_ELEMS / (2 * blocklen)) {
let base_l = (k / blocks_per_half) * 128 + (k % blocks_per_half) * blocklen;
let base_h = base_l + 64;
let scale_idx_l = base_l / 16;
let scale_idx_h = base_h / 16;
let qh_shift_l = ((base_l % 128) / 32) * 2;
let qh_shift_h = ((base_h % 128) / 32) * 2;
let qh_half_l = (base_l / 128) * 32;
let qh_half_h = (base_h / 128) * 32;
for j in 0..ncols {
let scale_l = scales[scale_idx_l * ncols + j] as i8 as i32;
let scale_h = scales[scale_idx_h * ncols + j] as i8 as i32;
let mut q_l = [0i32; 8];
let mut q_h = [0i32; 8];
for i in 0..blocklen {
let ql_pos = k * ncols * blocklen + j * blocklen + i;
let l_4 = (ql[ql_pos] & 0x0F) as i32;
let hi_4 = ((ql[ql_pos] >> 4) & 0x0F) as i32;
let qh_idx_l = qh_half_l + ((base_l + i) % 32);
let qh_chunk_l = qh_idx_l / blocklen;
let qh_pos_l = qh_idx_l % blocklen;
let qh_offset_l = qh_chunk_l * (blocklen * ncols) + j * blocklen + qh_pos_l;
let hi_2_l = ((qh[qh_offset_l] >> qh_shift_l) & 0x3) as i32;
let qh_idx_h = qh_half_h + ((base_h + i) % 32);
let qh_chunk_h = qh_idx_h / blocklen;
let qh_pos_h = qh_idx_h % blocklen;
let qh_offset_h = qh_chunk_h * (blocklen * ncols) + j * blocklen + qh_pos_h;
let hi_2_h = ((qh[qh_offset_h] >> qh_shift_h) & 0x3) as i32;
q_l[i] = ((hi_2_l << 4) | l_4) - 32;
q_h[i] = ((hi_2_h << 4) | hi_4) - 32;
}
for (a, act) in acts.iter().enumerate() {
let da = act.d[l];
let q8 = &act.q[l * Q6_K_BLOCK_ELEMS..(l + 1) * Q6_K_BLOCK_ELEMS];
let mut sumi_l = 0i32;
let mut sumi_h = 0i32;
for i in 0..blocklen {
sumi_l += q_l[i] * (q8[base_l + i] as i32);
sumi_h += q_h[i] * (q8[base_h + i] as i32);
}
out[j * na + a] += (sumi_l * scale_l + sumi_h * scale_h) as f32 * d_f[j] * da;
}
}
}
}
}
#[inline]
pub fn q6_kx8_gemm_uses_acts_x4(interleave: usize) -> bool {
#[cfg(target_arch = "aarch64")]
{
interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
}
#[cfg(not(target_arch = "aarch64"))]
{
let _ = interleave;
false
}
}
pub fn gemm_q6_kx8_group_x4(
packed: &[u8],
group: usize,
tile: &Q8KActsX4,
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
gemm_q6_kx8_group_x4_on(
packed,
group,
tile,
n_cols,
interleave,
AccelX4::detect(),
out,
);
}
#[inline]
pub fn gemm_q6_kx8_group_x4_on(
packed: &[u8],
group: usize,
tile: &Q8KActsX4,
n_cols: usize,
interleave: usize,
accel: AccelX4,
out: &mut [f32],
) {
assert_eq!(
interleave, 8,
"the x4 GEMM only exists for interleave-8 packing"
);
assert_eq!(out.len(), Q6_KX8_NROWS * tile.na);
assert!(n_cols.is_multiple_of(Q6_K_BLOCK_ELEMS));
debug_assert_eq!(tile.n_blocks, n_cols / Q6_K_BLOCK_ELEMS);
if tile.na == 0 {
return;
}
let nb = n_cols / Q6_K_BLOCK_ELEMS;
let off = group * nb * Q6_KX8_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q6_KX8_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
if accel == AccelX4::NeonI8mm {
unsafe {
neon::gemm_q6_kx8_q8_k_neon_i8mm(slice, tile, n_cols, out);
}
return;
}
let _ = accel;
gemm_q6_kx8_acts_x4_scalar_8(slice, tile, n_cols, out);
}
fn gemm_q6_kx8_acts_x4_scalar_8(packed: &[u8], tile: &Q8KActsX4, n_cols: usize, out: &mut [f32]) {
let nb = n_cols / Q6_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols = Q6_KX8_NROWS;
let blocks_per_half = 64 / blocklen;
let na = tile.na;
let mut sumf = [[0f32; Q6_KX8_NROWS]; Q8K_ACTS_X4_NC];
for l in 0..nb {
let blk = &packed[l * Q6_KX8_BLOCK_BYTES..][..Q6_KX8_BLOCK_BYTES];
let d = &blk[0..16];
let scales = &blk[16..144];
let ql = &blk[144..1168];
let qh = &blk[1168..];
let q8 = &tile.qs[l * Q6_K_BLOCK_ELEMS * 4..][..Q6_K_BLOCK_ELEMS * 4];
for a in 0..na {
let da = tile.d[l * 4 + a];
for k in 0..(Q6_K_BLOCK_ELEMS / (2 * blocklen)) {
let base_l = (k / blocks_per_half) * 128 + (k % blocks_per_half) * blocklen;
let base_h = base_l + 64;
let scale_idx_l = base_l / 16;
let scale_idx_h = base_h / 16;
let qh_shift_l = ((base_l % 128) / 32) * 2;
let qh_shift_h = ((base_h % 128) / 32) * 2;
let qh_half_l = (base_l / 128) * 32;
let qh_half_h = (base_h / 128) * 32;
for j in 0..ncols {
let scale_l = scales[scale_idx_l * ncols + j] as i8 as i32;
let scale_h = scales[scale_idx_h * ncols + j] as i8 as i32;
let mut sumi_l = 0i32;
let mut sumi_h = 0i32;
for i in 0..blocklen {
let ql_pos = k * ncols * blocklen + j * blocklen + i;
let l_4 = (ql[ql_pos] & 0x0F) as i32;
let hi_4 = ((ql[ql_pos] >> 4) & 0x0F) as i32;
let qh_idx_l = qh_half_l + ((base_l + i) % 32);
let qh_chunk_l = qh_idx_l / blocklen;
let qh_pos_l = qh_idx_l % blocklen;
let qh_offset_l = qh_chunk_l * (blocklen * ncols) + j * blocklen + qh_pos_l;
let hi_2_l = ((qh[qh_offset_l] >> qh_shift_l) & 0x3) as i32;
let qh_idx_h = qh_half_h + ((base_h + i) % 32);
let qh_chunk_h = qh_idx_h / blocklen;
let qh_pos_h = qh_idx_h % blocklen;
let qh_offset_h = qh_chunk_h * (blocklen * ncols) + j * blocklen + qh_pos_h;
let hi_2_h = ((qh[qh_offset_h] >> qh_shift_h) & 0x3) as i32;
let q_l = ((hi_2_l << 4) | l_4) - 32;
let q_h = ((hi_2_h << 4) | hi_4) - 32;
let e_l = base_l + i;
let e_h = base_h + i;
sumi_l += q_l * (q8[(e_l / 8) * 32 + a * 8 + (e_l % 8)] as i32);
sumi_h += q_h * (q8[(e_h / 8) * 32 + a * 8 + (e_h % 8)] as i32);
}
sumf[a][j] += (sumi_l * scale_l + sumi_h * scale_h) as f32
* f16_from_bytes(&d[j * 2..])
* da;
}
}
}
}
for j in 0..ncols {
for (a, row) in sumf.iter().take(na).enumerate() {
out[j * na + a] = row[j];
}
}
}
pub const Q4_0X4_BLOCK_BYTES: usize = 72;
pub const Q4_0X4_NROWS: usize = 4;
pub const Q4_0X4_INTERLEAVE: usize = 4;
#[inline]
pub fn q4_0x4_interleave() -> usize {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return 8;
}
}
Q4_0X4_INTERLEAVE
}
const Q4_0X4_XOR_MASK_U32: u32 = 0x8888_8888;
const Q4_0X4_XOR_MASK_U64: u64 = 0x8888_8888_8888_8888;
pub fn make_block_q4_0x4(
rows: [&[u8]; Q4_0X4_NROWS],
interleave: usize,
) -> [u8; Q4_0X4_BLOCK_BYTES] {
debug_assert!(interleave == 4 || interleave == 8);
for r in &rows {
debug_assert_eq!(r.len(), Q4_0_BLOCK_BYTES);
}
let mut out = [0u8; Q4_0X4_BLOCK_BYTES];
for (i, row) in rows.iter().enumerate() {
out[i * 2] = row[0];
out[i * 2 + 1] = row[1];
}
let end = (Q4_0_BLOCK_ELEMS * 2) / interleave;
let qs_out = &mut out[8..];
for i in 0..end {
let src_id = i % Q4_0X4_NROWS;
let src_offset = (i / Q4_0X4_NROWS) * interleave;
let dst_offset = i * interleave;
let src_qs = &rows[src_id][2..18];
if interleave == 4 {
let mut elems = u32::from_le_bytes(
src_qs[src_offset..src_offset + 4]
.try_into()
.expect("4-byte interleave chunk"),
);
elems ^= Q4_0X4_XOR_MASK_U32;
qs_out[dst_offset..dst_offset + 4].copy_from_slice(&elems.to_le_bytes());
} else {
let mut elems = u64::from_le_bytes(
src_qs[src_offset..src_offset + 8]
.try_into()
.expect("8-byte interleave chunk"),
);
elems ^= Q4_0X4_XOR_MASK_U64;
qs_out[dst_offset..dst_offset + 8].copy_from_slice(&elems.to_le_bytes());
}
}
out
}
pub fn pack_q4_0_matrix_x4(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
assert!(cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
let n_blocks = cols / Q4_0_BLOCK_ELEMS;
let row_bytes = n_blocks * Q4_0_BLOCK_BYTES;
assert_eq!(data.len(), rows * row_bytes);
let n_groups = rows / Q4_0X4_NROWS;
let mut out = Vec::with_capacity(n_groups * n_blocks * Q4_0X4_BLOCK_BYTES);
for g in 0..n_groups {
for b in 0..n_blocks {
let mut row_refs: [&[u8]; Q4_0X4_NROWS] = [&[]; Q4_0X4_NROWS];
for (r, slot) in row_refs.iter_mut().enumerate() {
let base = (g * Q4_0X4_NROWS + r) * row_bytes + b * Q4_0_BLOCK_BYTES;
*slot = &data[base..base + Q4_0_BLOCK_BYTES];
}
out.extend_from_slice(&make_block_q4_0x4(row_refs, interleave));
}
}
out
}
#[inline]
fn q4_0x4_nibble_dot(byte: u8, q8_lo: i32, q8_hi: i32) -> i32 {
let v0 = ((byte << 4) as i8) as i32;
let v1 = ((byte & 0xF0) as i8) as i32;
((v0 * q8_lo) + (v1 * q8_hi)) >> 4
}
fn gemv_q4_0x4_q8_0_scalar(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
blocklen: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let ncols = Q4_0X4_NROWS;
debug_assert_eq!(act.n_blocks(), nb);
debug_assert_eq!(out.len(), n_row_groups * ncols);
debug_assert_eq!(packed.len(), n_row_groups * nb * Q4_0X4_BLOCK_BYTES);
for x in 0..n_row_groups {
let mut sumf = [0f32; 4];
let group_off = x * nb * Q4_0X4_BLOCK_BYTES;
for l in 0..nb {
let blk = &packed[group_off + l * Q4_0X4_BLOCK_BYTES..][..Q4_0X4_BLOCK_BYTES];
let da = act.d[l];
let q8 = &act.q[l * Q4_0_BLOCK_ELEMS..(l + 1) * Q4_0_BLOCK_ELEMS];
for k in 0..(Q4_0_BLOCK_ELEMS / (2 * blocklen)) {
for j in 0..ncols {
let mut sumi = 0i32;
for i in 0..blocklen {
let byte = blk[8 + k * ncols * blocklen + j * blocklen + i];
sumi += q4_0x4_nibble_dot(
byte,
q8[k * blocklen + i] as i32,
q8[k * blocklen + i + Q4_0_BLOCK_ELEMS / 2] as i32,
);
}
sumf[j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
}
}
}
let base = x * ncols;
out[base..base + ncols].copy_from_slice(&sumf);
}
}
pub fn gemv_q4_0x4_q8_0(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
interleave: usize,
out: &mut [f32],
) {
assert!(n_cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
assert_eq!(out.len(), n_row_groups * Q4_0X4_NROWS);
match interleave {
4 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q4_0x4_q8_0_neon_sdot(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q4_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 4, out);
}
8 => {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemv_q4_0x4_q8_0_neon_4x8(packed, act, n_cols, n_row_groups, out);
}
return;
}
}
gemv_q4_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 8, out);
}
_ => panic!("q4_0x4 interleave must be 4 or 8, got {interleave}"),
}
}
pub const Q4_0X4_GEMM_NC: usize = 4;
pub fn gemm_q4_0x4_group(
packed: &[u8],
group: usize,
acts: &[Q8Activations],
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
assert_eq!(out.len(), Q4_0X4_NROWS * acts.len());
assert!(n_cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
if acts.is_empty() {
return;
}
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let off = group * nb * Q4_0X4_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q4_0X4_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
{
if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
for (t, chunk) in acts.chunks(Q8K_ACTS_X4_NC).enumerate() {
let tile = prepare_q8_acts_x4(chunk, n_cols);
let mut tmp = [0f32; Q4_0X4_NROWS * Q8K_ACTS_X4_NC];
let n = chunk.len();
unsafe {
neon::gemm_q4_0x4_q8_0_neon_i8mm(
slice,
&tile,
n_cols,
&mut tmp[..Q4_0X4_NROWS * n],
);
}
for r in 0..Q4_0X4_NROWS {
for j in 0..n {
out[r * acts.len() + t * Q8K_ACTS_X4_NC + j] = tmp[r * n + j];
}
}
}
return;
}
if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
unsafe {
neon::gemm_q4_0x4_q8_0_neon_sdot(slice, acts, n_cols, out);
}
return;
}
}
let mut tmp = [0f32; Q4_0X4_NROWS];
for (j, act) in acts.iter().enumerate() {
gemv_q4_0x4_q8_0(slice, act, n_cols, 1, interleave, &mut tmp);
for (r, v) in tmp.iter().enumerate() {
out[r * acts.len() + j] = *v;
}
}
}
#[inline]
pub fn q4_0x4_gemm_uses_acts_x4(interleave: usize) -> bool {
#[cfg(target_arch = "aarch64")]
{
interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
}
#[cfg(not(target_arch = "aarch64"))]
{
let _ = interleave;
false
}
}
pub fn gemm_q4_0x4_group_x4(
packed: &[u8],
group: usize,
tile: &Q8ActsX4,
n_cols: usize,
interleave: usize,
out: &mut [f32],
) {
gemm_q4_0x4_group_x4_on(
packed,
group,
tile,
n_cols,
interleave,
AccelX4::detect(),
out,
);
}
#[inline]
pub fn gemm_q4_0x4_group_x4_on(
packed: &[u8],
group: usize,
tile: &Q8ActsX4,
n_cols: usize,
interleave: usize,
accel: AccelX4,
out: &mut [f32],
) {
assert_eq!(
interleave, 8,
"the x4 GEMM only exists for interleave-8 packing"
);
assert_eq!(out.len(), Q4_0X4_NROWS * tile.na);
assert!(n_cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
debug_assert_eq!(tile.n_blocks, n_cols / Q4_0_BLOCK_ELEMS);
if tile.na == 0 {
return;
}
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let off = group * nb * Q4_0X4_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q4_0X4_BLOCK_BYTES];
#[cfg(target_arch = "aarch64")]
if accel == AccelX4::NeonI8mm {
unsafe {
neon::gemm_q4_0x4_q8_0_neon_i8mm(slice, tile, n_cols, out);
}
return;
}
let _ = accel;
gemm_q4_0x4_acts_x4_scalar_8(slice, tile, n_cols, out);
}
fn gemm_q4_0x4_acts_x4_scalar_8(packed: &[u8], tile: &Q8ActsX4, n_cols: usize, out: &mut [f32]) {
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let blocklen = 8;
let ncols = Q4_0X4_NROWS;
let na = tile.na;
let mut sumf = [[0f32; Q4_0X4_NROWS]; Q8K_ACTS_X4_NC];
for l in 0..nb {
let blk = &packed[l * Q4_0X4_BLOCK_BYTES..][..Q4_0X4_BLOCK_BYTES];
let q8 = &tile.qs[l * Q4_0_BLOCK_ELEMS * 4..][..Q4_0_BLOCK_ELEMS * 4];
for a in 0..na {
let da = tile.d[l * 4 + a];
for k in 0..(Q4_0_BLOCK_ELEMS / (2 * blocklen)) {
for j in 0..ncols {
let mut sumi = 0i32;
for i in 0..blocklen {
let byte = blk[8 + k * ncols * blocklen + j * blocklen + i];
let e0 = k * blocklen + i;
let e1 = e0 + Q4_0_BLOCK_ELEMS / 2;
sumi += q4_0x4_nibble_dot(
byte,
q8[(e0 / 8) * 32 + a * 8 + (e0 % 8)] as i32,
q8[(e1 / 8) * 32 + a * 8 + (e1 % 8)] as i32,
);
}
sumf[a][j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
}
}
}
}
for j in 0..ncols {
for (a, row) in sumf.iter().take(na).enumerate() {
out[j * na + a] = row[j];
}
}
}
#[inline]
pub fn gemv_q4_0x4_group(
packed: &[u8],
group: usize,
act: &Q8Activations,
n_cols: usize,
interleave: usize,
out4: &mut [f32],
) {
debug_assert_eq!(out4.len(), Q4_0X4_NROWS);
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let off = group * nb * Q4_0X4_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q4_0X4_BLOCK_BYTES];
gemv_q4_0x4_q8_0(slice, act, n_cols, 1, interleave, out4);
}
#[inline]
pub fn gemv_q8_0x4_group(
packed: &[u8],
group: usize,
act: &Q8Activations,
n_cols: usize,
interleave: usize,
out4: &mut [f32],
) {
debug_assert_eq!(out4.len(), Q8_0X4_NROWS);
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let off = group * nb * Q8_0X4_BLOCK_BYTES;
let slice = &packed[off..off + nb * Q8_0X4_BLOCK_BYTES];
gemv_q8_0x4_q8_0(slice, act, n_cols, 1, interleave, out4);
}
#[cfg(target_arch = "aarch64")]
mod neon {
use super::*;
use std::arch::aarch64::*;
#[target_feature(enable = "neon,i8mm")]
unsafe fn vmmla_s32(mut acc: int32x4_t, a: int8x16_t, b: int8x16_t) -> int32x4_t {
std::arch::asm!(
"smmla {acc:v}.4s, {a:v}.16b, {b:v}.16b",
acc = inout(vreg) acc,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
);
acc
}
#[target_feature(enable = "neon,dotprod")]
unsafe fn sdot_lane(mut acc: int32x4_t, a: int8x16_t, b: int8x16_t, lane: u32) -> int32x4_t {
match lane {
0 => std::arch::asm!(
"sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[0]",
acc = inout(vreg) acc,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
),
1 => std::arch::asm!(
"sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[1]",
acc = inout(vreg) acc,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
),
2 => std::arch::asm!(
"sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[2]",
acc = inout(vreg) acc,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
),
3 => std::arch::asm!(
"sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[3]",
acc = inout(vreg) acc,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
),
_ => unreachable!(),
}
acc
}
#[target_feature(enable = "neon,dotprod")]
unsafe fn sdot(mut acc: int32x4_t, a: int8x16_t, b: int8x16_t) -> int32x4_t {
std::arch::asm!(
"sdot {acc:v}.4s, {a:v}.16b, {b:v}.16b",
acc = inout(vreg) acc,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
);
acc
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q4_kx8_q8_k_neon_sdot(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
for x in 0..n_row_groups {
let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q4_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] =
f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let q8_d = act.d[b];
let sb_scale_0123 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d);
let sb_scale_4567 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d);
let sb_min_0123 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d);
let sb_min_4567 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d);
let mut bias_acc = [vdupq_n_s32(0), vdupq_n_s32(0)];
let q8_base = act.q.as_ptr().add(b * Q4_K_BLOCK_ELEMS);
let bsums_ptr = act.bsums.as_ptr().add(b * 16);
let mut bsums_arr = [0i16; 8];
for (i, slot) in bsums_arr.iter_mut().enumerate() {
*slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
}
let scales_base = blk.add(32);
let qs_base = blk.add(128);
for sb in 0..4 {
let mut acc_lo = [vdupq_n_s32(0), vdupq_n_s32(0)];
let mut acc_hi = [vdupq_n_s32(0), vdupq_n_s32(0)];
let mut q4sb_mins = [vdupq_n_s16(0); 2];
let mut q4sb_scales = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut sc_i8 = [0i8; 8];
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
sc_i8[t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q4sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
let mut q8_qs = [vdupq_n_s8(0); 4];
for (i, slot) in q8_qs.iter_mut().enumerate() {
*slot = vld1q_s8(q8_base.add(sb * 64 + i * 16));
}
for c in 0..2 {
let mut q4_cols = [vdupq_n_u8(0); 8];
for (i, slot) in q4_cols.iter_mut().enumerate() {
*slot = vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + i * 32 + 16 * c));
}
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[0], m4b)),
q8_qs[0],
0,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[1], m4b)),
q8_qs[0],
1,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[2], m4b)),
q8_qs[0],
2,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[3], m4b)),
q8_qs[0],
3,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[4], m4b)),
q8_qs[1],
0,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[5], m4b)),
q8_qs[1],
1,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[6], m4b)),
q8_qs[1],
2,
);
acc_lo[c] = sdot_lane(
acc_lo[c],
vreinterpretq_s8_u8(vandq_u8(q4_cols[7], m4b)),
q8_qs[1],
3,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[0], 4)),
q8_qs[2],
0,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[1], 4)),
q8_qs[2],
1,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[2], 4)),
q8_qs[2],
2,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[3], 4)),
q8_qs[2],
3,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[4], 4)),
q8_qs[3],
0,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[5], 4)),
q8_qs[3],
1,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[6], 4)),
q8_qs[3],
2,
);
acc_hi[c] = sdot_lane(
acc_hi[c],
vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[7], 4)),
q8_qs[3],
3,
);
}
let sc_0123_lo = vget_low_s16(q4sb_scales[0]);
let sc_0123_hi = vget_low_s16(q4sb_scales[1]);
let sumf_0123 = vcvtq_f32_s32(vaddq_s32(
vmulq_s32(vmovl_s16(sc_0123_lo), acc_lo[0]),
vmulq_s32(vmovl_s16(sc_0123_hi), acc_hi[0]),
));
acc_f32[0] = vfmaq_f32(acc_f32[0], sb_scale_0123, sumf_0123);
let sc_4567_lo = vget_high_s16(q4sb_scales[0]);
let sc_4567_hi = vget_high_s16(q4sb_scales[1]);
let sumf_4567 = vcvtq_f32_s32(vaddq_s32(
vmulq_s32(vmovl_s16(sc_4567_lo), acc_lo[1]),
vmulq_s32(vmovl_s16(sc_4567_hi), acc_hi[1]),
));
acc_f32[1] = vfmaq_f32(acc_f32[1], sb_scale_4567, sumf_4567);
let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_lo, vget_low_s16(q4sb_mins[0]));
bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_hi, vget_low_s16(q4sb_mins[1]));
bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_lo, vget_high_s16(q4sb_mins[0]));
bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_hi, vget_high_s16(q4sb_mins[1]));
}
acc_f32[0] = vmlsq_f32(acc_f32[0], vcvtq_f32_s32(bias_acc[0]), sb_min_0123);
acc_f32[1] = vmlsq_f32(acc_f32[1], vcvtq_f32_s32(bias_acc[1]), sb_min_4567);
}
let base = x * Q4_KX8_NROWS;
vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q5_kx8_q8_k_neon_sdot(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
let mone = vdupq_n_u8(1);
let mtwo = vdupq_n_u8(2);
for x in 0..n_row_groups {
let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q5_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] =
f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let q8_d = act.d[b];
let sb_scale_0123 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d);
let sb_scale_4567 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d);
let sb_min_0123 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d);
let sb_min_4567 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d);
let mut bias_acc = [vdupq_n_s32(0), vdupq_n_s32(0)];
let q8_base = act.q.as_ptr().add(b * Q5_K_BLOCK_ELEMS);
let bsums_ptr = act.bsums.as_ptr().add(b * 16);
let mut bsums_arr = [0i16; 8];
for (i, slot) in bsums_arr.iter_mut().enumerate() {
*slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
}
let scales_base = blk.add(32);
let qh_base = blk.add(128);
let qs_base = blk.add(384);
let mut qh = [[vdupq_n_u8(0); 8]; 2];
for (c, qh_c) in qh.iter_mut().enumerate() {
for (i, slot) in qh_c.iter_mut().enumerate() {
*slot = vld1q_u8(qh_base.add(i * 32 + 16 * c));
}
}
for sb in 0..4 {
let mut acc_lo = [vdupq_n_s32(0), vdupq_n_s32(0)];
let mut acc_hi = [vdupq_n_s32(0), vdupq_n_s32(0)];
let mut q5sb_mins = [vdupq_n_s16(0); 2];
let mut q5sb_scales = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut sc_i8 = [0i8; 8];
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
sc_i8[t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q5sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
let mut q8_qs = [vdupq_n_s8(0); 4];
for (i, slot) in q8_qs.iter_mut().enumerate() {
*slot = vld1q_s8(q8_base.add(sb * 64 + i * 16));
}
for c in 0..2 {
let mut q5_lo = [vdupq_n_s8(0); 8];
let mut q5_hi = [vdupq_n_s8(0); 8];
for i in 0..8 {
let q5_cols =
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + i * 32 + 16 * c));
let hbit_lo = vandq_u8(qh[c][i], mone);
let hbit_hi = vshlq_n_u8(vandq_u8(qh[c][i], mtwo), 3);
qh[c][i] = vshrq_n_u8(qh[c][i], 2);
q5_lo[i] =
vreinterpretq_s8_u8(vsliq_n_u8(vandq_u8(q5_cols, m4b), hbit_lo, 4));
q5_hi[i] =
vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_cols, 4), hbit_hi));
}
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[0], q8_qs[0], 0);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[1], q8_qs[0], 1);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[2], q8_qs[0], 2);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[3], q8_qs[0], 3);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[4], q8_qs[1], 0);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[5], q8_qs[1], 1);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[6], q8_qs[1], 2);
acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[7], q8_qs[1], 3);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[0], q8_qs[2], 0);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[1], q8_qs[2], 1);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[2], q8_qs[2], 2);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[3], q8_qs[2], 3);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[4], q8_qs[3], 0);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[5], q8_qs[3], 1);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[6], q8_qs[3], 2);
acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[7], q8_qs[3], 3);
}
let sc_0123_lo = vget_low_s16(q5sb_scales[0]);
let sc_0123_hi = vget_low_s16(q5sb_scales[1]);
let sumf_0123 = vcvtq_f32_s32(vaddq_s32(
vmulq_s32(vmovl_s16(sc_0123_lo), acc_lo[0]),
vmulq_s32(vmovl_s16(sc_0123_hi), acc_hi[0]),
));
acc_f32[0] = vfmaq_f32(acc_f32[0], sb_scale_0123, sumf_0123);
let sc_4567_lo = vget_high_s16(q5sb_scales[0]);
let sc_4567_hi = vget_high_s16(q5sb_scales[1]);
let sumf_4567 = vcvtq_f32_s32(vaddq_s32(
vmulq_s32(vmovl_s16(sc_4567_lo), acc_lo[1]),
vmulq_s32(vmovl_s16(sc_4567_hi), acc_hi[1]),
));
acc_f32[1] = vfmaq_f32(acc_f32[1], sb_scale_4567, sumf_4567);
let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_lo, vget_low_s16(q5sb_mins[0]));
bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_hi, vget_low_s16(q5sb_mins[1]));
bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_lo, vget_high_s16(q5sb_mins[0]));
bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_hi, vget_high_s16(q5sb_mins[1]));
}
acc_f32[0] = vmlsq_f32(acc_f32[0], vcvtq_f32_s32(bias_acc[0]), sb_min_0123);
acc_f32[1] = vmlsq_f32(acc_f32[1], vcvtq_f32_s32(bias_acc[1]), sb_min_4567);
}
let base = x * Q5_KX8_NROWS;
vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q4_kx8_q8_k_neon_8x8(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
for x in 0..n_row_groups {
let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q4_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] =
f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let q8_d = act.d[b];
let sb_scale = [
vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d),
vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d),
];
let sb_min = [
vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d),
vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d),
];
let q8_base = act.q.as_ptr().add(b * Q4_K_BLOCK_ELEMS);
let bsums_ptr = act.bsums.as_ptr().add(b * 16);
let mut bsums_arr = [0i16; 8];
for (i, slot) in bsums_arr.iter_mut().enumerate() {
*slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
}
let scales_base = blk.add(32);
let qs_base = blk.add(128);
let mut bias_acc = [vdupq_n_s32(0), vdupq_n_s32(0)];
for sb in 0..4 {
let mut acc_lo = [vdupq_n_s32(0); 4];
let mut acc_hi = [vdupq_n_s32(0); 4];
let mut q4sb_scales = [vdupq_n_s16(0); 2];
let mut q4sb_mins = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut sc_i8 = [0i8; 8];
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
sc_i8[t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q4sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
let q8_sb = q8_base.add(sb * 64);
let mut q8_qs = [vdupq_n_s8(0); 8];
for (i, slot) in q8_qs.iter_mut().enumerate() {
*slot = vreinterpretq_s8_s64(vld1q_dup_s64(q8_sb.add(i * 8) as *const i64));
}
for cp in 0..4 {
let q4_qs = [
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp)),
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 64)),
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 128)),
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 192)),
];
for m in 0..4 {
let q4_lo = vreinterpretq_s8_u8(vandq_u8(q4_qs[m], m4b));
let q4_hi = vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[m], 4));
acc_lo[cp] = sdot(acc_lo[cp], q4_lo, q8_qs[m]);
acc_hi[cp] = sdot(acc_hi[cp], q4_hi, q8_qs[m + 4]);
}
}
for i in 0..2 {
let p = i * 2;
let (scales_lo, scales_hi) = if i == 0 {
(vget_low_s16(q4sb_scales[0]), vget_low_s16(q4sb_scales[1]))
} else {
(vget_high_s16(q4sb_scales[0]), vget_high_s16(q4sb_scales[1]))
};
let sumf_0 = vcvtq_f32_s32(vmulq_s32(
vmovl_s16(scales_lo),
vpaddq_s32(acc_lo[p], acc_lo[p + 1]),
));
acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_0);
let sumf_1 = vcvtq_f32_s32(vmulq_s32(
vmovl_s16(scales_hi),
vpaddq_s32(acc_hi[p], acc_hi[p + 1]),
));
acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_1);
}
let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_lo, vget_low_s16(q4sb_mins[0]));
bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_hi, vget_low_s16(q4sb_mins[1]));
bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_lo, vget_high_s16(q4sb_mins[0]));
bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_hi, vget_high_s16(q4sb_mins[1]));
}
acc_f32[0] = vmlsq_f32(acc_f32[0], vcvtq_f32_s32(bias_acc[0]), sb_min[0]);
acc_f32[1] = vmlsq_f32(acc_f32[1], vcvtq_f32_s32(bias_acc[1]), sb_min[1]);
}
let base = x * Q4_KX8_NROWS;
vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q5_kx8_q8_k_neon_8x8(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
let mone = vdupq_n_u8(1);
let mtwo = vdupq_n_u8(2);
for x in 0..n_row_groups {
let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q5_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] =
f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let q8_d = act.d[b];
let sb_scale = [
vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d),
vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d),
];
let sb_min = [
vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d),
vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d),
];
let q8_base = act.q.as_ptr().add(b * Q5_K_BLOCK_ELEMS);
let bsums_ptr = act.bsums.as_ptr().add(b * 16);
let mut bsums_arr = [0i16; 8];
for (i, slot) in bsums_arr.iter_mut().enumerate() {
*slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
}
let scales_base = blk.add(32);
let qh_base = blk.add(128);
let qs_base = blk.add(384);
let mut qh = [[vdupq_n_u8(0); 4]; 4];
for (cp, qh_cp) in qh.iter_mut().enumerate() {
for (m, slot) in qh_cp.iter_mut().enumerate() {
*slot = vld1q_u8(qh_base.add(16 * cp + 64 * m));
}
}
for sb in 0..4 {
let mut acc_lo = [vdupq_n_s32(0); 4];
let mut acc_hi = [vdupq_n_s32(0); 4];
let mut q5sb_scales = [vdupq_n_s16(0); 2];
let mut q5sb_mins = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut sc_i8 = [0i8; 8];
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
sc_i8[t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q5sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
let q8_sb = q8_base.add(sb * 64);
let mut q8_qs = [vdupq_n_s8(0); 8];
for (i, slot) in q8_qs.iter_mut().enumerate() {
*slot = vreinterpretq_s8_s64(vld1q_dup_s64(q8_sb.add(i * 8) as *const i64));
}
for cp in 0..4 {
let q5_qs = [
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp)),
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 64)),
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 128)),
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 192)),
];
for m in 0..4 {
let hbit_lo = vandq_u8(qh[cp][m], mone);
let hbit_hi = vshlq_n_u8(vandq_u8(qh[cp][m], mtwo), 3);
qh[cp][m] = vshrq_n_u8(qh[cp][m], 2);
let q5_lo = vreinterpretq_s8_u8(vsliq_n_u8(
vandq_u8(q5_qs[m], m4b),
hbit_lo,
4,
));
let q5_hi =
vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_qs[m], 4), hbit_hi));
acc_lo[cp] = sdot(acc_lo[cp], q5_lo, q8_qs[m]);
acc_hi[cp] = sdot(acc_hi[cp], q5_hi, q8_qs[m + 4]);
}
}
let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
for i in 0..2 {
let p = i * 2;
let (scales_lo, scales_hi, mins_lo, mins_hi) = if i == 0 {
(
vget_low_s16(q5sb_scales[0]),
vget_low_s16(q5sb_scales[1]),
vget_low_s16(q5sb_mins[0]),
vget_low_s16(q5sb_mins[1]),
)
} else {
(
vget_high_s16(q5sb_scales[0]),
vget_high_s16(q5sb_scales[1]),
vget_high_s16(q5sb_mins[0]),
vget_high_s16(q5sb_mins[1]),
)
};
let sumf_0 = vcvtq_f32_s32(vmulq_s32(
vmovl_s16(scales_lo),
vpaddq_s32(acc_lo[p], acc_lo[p + 1]),
));
acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_0);
let sumf_1 = vcvtq_f32_s32(vmulq_s32(
vmovl_s16(scales_hi),
vpaddq_s32(acc_hi[p], acc_hi[p + 1]),
));
acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_1);
let mut bias = vmull_s16(bsums_vec_lo, mins_lo);
bias = vmlal_s16(bias, bsums_vec_hi, mins_hi);
acc_f32[i] = vmlsq_f32(acc_f32[i], sb_min[i], vcvtq_f32_s32(bias));
}
}
}
let base = x * Q5_KX8_NROWS;
vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q6_kx8_q8_k_neon_8x8(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q6_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
let mask_lo = vdupq_n_u8(0x03);
let mask_hi = vdupq_n_u8(0x30);
for x in 0..n_row_groups {
let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
let group_off = x * nb * Q6_KX8_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q6_KX8_BLOCK_BYTES);
let scales_base = blk.add(16) as *const i8;
let ql_blk = blk.add(144);
let qh_blk = blk.add(1168);
let mut d_arr = [0f32; 8];
for (j, slot) in d_arr.iter_mut().enumerate() {
*slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
}
let q8_d = act.d[b];
let sb_scale = [
vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d),
vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d),
];
let mut acc = [vdup_n_s32(0); 4];
let mut q6_scales = [0i16; 16 * 8];
for i in 0..16 {
let s16 = vmovl_s8(vld1_s8(scales_base.add(i * 8)));
vst1q_s16(q6_scales.as_mut_ptr().add(i * 8), s16);
}
let mut bias_lo = vdupq_n_s32(0);
let mut bias_hi = vdupq_n_s32(0);
for i in (0..16).step_by(4) {
let bsums_vec = vld1_s16(act.bsums.as_ptr().add(b * 16 + i));
let sc = q6_scales.as_ptr();
bias_lo = vmlal_lane_s16::<0>(bias_lo, vld1_s16(sc.add(i * 8)), bsums_vec);
bias_hi = vmlal_lane_s16::<0>(bias_hi, vld1_s16(sc.add(i * 8 + 4)), bsums_vec);
bias_lo =
vmlal_lane_s16::<1>(bias_lo, vld1_s16(sc.add((i + 1) * 8)), bsums_vec);
bias_hi =
vmlal_lane_s16::<1>(bias_hi, vld1_s16(sc.add((i + 1) * 8 + 4)), bsums_vec);
bias_lo =
vmlal_lane_s16::<2>(bias_lo, vld1_s16(sc.add((i + 2) * 8)), bsums_vec);
bias_hi =
vmlal_lane_s16::<2>(bias_hi, vld1_s16(sc.add((i + 2) * 8 + 4)), bsums_vec);
bias_lo =
vmlal_lane_s16::<3>(bias_lo, vld1_s16(sc.add((i + 3) * 8)), bsums_vec);
bias_hi =
vmlal_lane_s16::<3>(bias_hi, vld1_s16(sc.add((i + 3) * 8 + 4)), bsums_vec);
}
bias_lo = vshlq_n_s32(bias_lo, 5);
bias_hi = vshlq_n_s32(bias_hi, 5);
for half in 0..2 {
let ql_base = ql_blk.add(half * 512);
let qh_base = qh_blk.add(half * 256);
let q8_half = act.q.as_ptr().add(b * Q6_K_BLOCK_ELEMS + half * 128);
for sb in 0..4 {
let q8_base_l = q8_half.add(sb * 16);
let q8_base_h = q8_base_l.add(64);
let mut q8_l = [vdupq_n_s8(0); 2];
let mut q8_h = [vdupq_n_s8(0); 2];
for i in 0..2 {
q8_l[i] = vreinterpretq_s8_s64(vld1q_dup_s64(
q8_base_l.add(i * 8) as *const i64
));
q8_h[i] = vreinterpretq_s8_s64(vld1q_dup_s64(
q8_base_h.add(i * 8) as *const i64
));
}
let ql_off = sb * (Q6_K_BLOCK_ELEMS / 2);
let qh_off = ql_off & 255; let mut q6_ql_0 = [vdupq_n_u8(0); 4];
let mut q6_ql_1 = [vdupq_n_u8(0); 4];
let mut q6_qh_0 = [vdupq_n_u8(0); 4];
let mut q6_qh_1 = [vdupq_n_u8(0); 4];
for k in 0..4 {
q6_ql_0[k] = vld1q_u8(ql_base.add(ql_off + 16 * k));
q6_ql_1[k] = vld1q_u8(ql_base.add(ql_off + 64 + 16 * k));
q6_qh_0[k] = vld1q_u8(qh_base.add(qh_off + 16 * k));
q6_qh_1[k] = vld1q_u8(qh_base.add(qh_off + 64 + 16 * k));
}
if sb > 1 {
for k in 0..4 {
q6_qh_0[k] = vshrq_n_u8(q6_qh_0[k], 2);
q6_qh_1[k] = vshrq_n_u8(q6_qh_1[k], 2);
}
}
for cp in 0..4 {
let hh_0 = vandq_u8(q6_qh_0[cp], mask_hi);
let hh_1 = vandq_u8(q6_qh_1[cp], mask_hi);
let q6_l0 = vreinterpretq_s8_u8(vsliq_n_u8(
vandq_u8(q6_ql_0[cp], m4b),
vandq_u8(q6_qh_0[cp], mask_lo),
4,
));
let q6_l1 = vreinterpretq_s8_u8(vsliq_n_u8(
vandq_u8(q6_ql_1[cp], m4b),
vandq_u8(q6_qh_1[cp], mask_lo),
4,
));
let q6_h0 =
vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_0[cp], 4), hh_0));
let q6_h1 =
vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_1[cp], 4), hh_1));
let mut sb_acc_l = vdupq_n_s32(0);
sb_acc_l = sdot(sb_acc_l, q6_l0, q8_l[0]);
sb_acc_l = sdot(sb_acc_l, q6_l1, q8_l[1]);
let mut sb_acc_h = vdupq_n_s32(0);
sb_acc_h = sdot(sb_acc_h, q6_h0, q8_h[0]);
sb_acc_h = sdot(sb_acc_h, q6_h1, q8_h[1]);
let sum_l = vpadd_s32(vget_low_s32(sb_acc_l), vget_high_s32(sb_acc_l));
let sum_h = vpadd_s32(vget_low_s32(sb_acc_h), vget_high_s32(sb_acc_h));
let scale_idx_l = half * 8 + sb;
let scale_idx_h = half * 8 + sb + 4;
let scale_vec_l = vset_lane_s32::<1>(
i32::from(q6_scales[scale_idx_l * 8 + cp * 2 + 1]),
vdup_n_s32(i32::from(q6_scales[scale_idx_l * 8 + cp * 2])),
);
let scale_vec_h = vset_lane_s32::<1>(
i32::from(q6_scales[scale_idx_h * 8 + cp * 2 + 1]),
vdup_n_s32(i32::from(q6_scales[scale_idx_h * 8 + cp * 2])),
);
acc[cp] = vmla_s32(acc[cp], sum_l, scale_vec_l);
acc[cp] = vmla_s32(acc[cp], sum_h, scale_vec_h);
}
}
}
acc[0] = vsub_s32(acc[0], vget_low_s32(bias_lo));
acc[1] = vsub_s32(acc[1], vget_high_s32(bias_lo));
acc[2] = vsub_s32(acc[2], vget_low_s32(bias_hi));
acc[3] = vsub_s32(acc[3], vget_high_s32(bias_hi));
let w_01 = vmul_f32(vcvt_f32_s32(acc[0]), vget_low_f32(sb_scale[0]));
let w_23 = vmul_f32(vcvt_f32_s32(acc[1]), vget_high_f32(sb_scale[0]));
let w_45 = vmul_f32(vcvt_f32_s32(acc[2]), vget_low_f32(sb_scale[1]));
let w_67 = vmul_f32(vcvt_f32_s32(acc[3]), vget_high_f32(sb_scale[1]));
acc_f32[0] = vaddq_f32(acc_f32[0], vcombine_f32(w_01, w_23));
acc_f32[1] = vaddq_f32(acc_f32[1], vcombine_f32(w_45, w_67));
}
let base = x * Q6_KX8_NROWS;
vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemm_q4_kx8_q8_k_neon_sdot(
packed: &[u8],
acts: &[Q8KActivations],
n_cols: usize,
out: &mut [f32],
) {
let na = acts.len();
debug_assert!(na <= Q4_KX8_GEMM_NC);
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
let mut acc_f32 = [[vdupq_n_f32(0.0); 2]; Q4_KX8_GEMM_NC];
let mut bias_acc = [[vdupq_n_s32(0); 2]; Q4_KX8_GEMM_NC];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q4_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let d_lo = vld1q_f32(d_arr.as_ptr());
let d_hi = vld1q_f32(d_arr.as_ptr().add(4));
let dmin_lo = vld1q_f32(dmin_arr.as_ptr());
let dmin_hi = vld1q_f32(dmin_arr.as_ptr().add(4));
let mut sb_scale = [[vdupq_n_f32(0.0); 2]; Q4_KX8_GEMM_NC];
let mut sb_min = [[vdupq_n_f32(0.0); 2]; Q4_KX8_GEMM_NC];
let mut bsums_arr = [[0i16; 8]; Q4_KX8_GEMM_NC];
for (a, act) in acts.iter().enumerate() {
let q8_d = act.d[b];
sb_scale[a] = [vmulq_n_f32(d_lo, q8_d), vmulq_n_f32(d_hi, q8_d)];
sb_min[a] = [vmulq_n_f32(dmin_lo, q8_d), vmulq_n_f32(dmin_hi, q8_d)];
let bsums_ptr = act.bsums.as_ptr().add(b * 16);
for (i, slot) in bsums_arr[a].iter_mut().enumerate() {
*slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
}
}
let scales_base = blk.add(32);
let qs_base = blk.add(128);
for sb in 0..4 {
let mut q4sb_mins = [vdupq_n_s16(0); 2];
let mut q4sb_scales = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut sc_i8 = [0i8; 8];
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
sc_i8[t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q4sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
for c in 0..2 {
let mut q4_cols = [vdupq_n_u8(0); 8];
for (i, slot) in q4_cols.iter_mut().enumerate() {
*slot = vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + i * 32 + 16 * c));
}
let (sc_lo, sc_hi) = if c == 0 {
(vget_low_s16(q4sb_scales[0]), vget_low_s16(q4sb_scales[1]))
} else {
(vget_high_s16(q4sb_scales[0]), vget_high_s16(q4sb_scales[1]))
};
let lo0 = vreinterpretq_s8_u8(vandq_u8(q4_cols[0], m4b));
let lo1 = vreinterpretq_s8_u8(vandq_u8(q4_cols[1], m4b));
let lo2 = vreinterpretq_s8_u8(vandq_u8(q4_cols[2], m4b));
let lo3 = vreinterpretq_s8_u8(vandq_u8(q4_cols[3], m4b));
let lo4 = vreinterpretq_s8_u8(vandq_u8(q4_cols[4], m4b));
let lo5 = vreinterpretq_s8_u8(vandq_u8(q4_cols[5], m4b));
let lo6 = vreinterpretq_s8_u8(vandq_u8(q4_cols[6], m4b));
let lo7 = vreinterpretq_s8_u8(vandq_u8(q4_cols[7], m4b));
let hi0 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[0], 4));
let hi1 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[1], 4));
let hi2 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[2], 4));
let hi3 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[3], 4));
let hi4 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[4], 4));
let hi5 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[5], 4));
let hi6 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[6], 4));
let hi7 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[7], 4));
let sc_lo_w = vmovl_s16(sc_lo);
let sc_hi_w = vmovl_s16(sc_hi);
for a in 0..na {
let q8_base = acts[a].q.as_ptr().add(b * Q4_K_BLOCK_ELEMS);
let y0 = vld1q_s8(q8_base.add(sb * 64));
let y1 = vld1q_s8(q8_base.add(sb * 64 + 16));
let y2 = vld1q_s8(q8_base.add(sb * 64 + 32));
let y3 = vld1q_s8(q8_base.add(sb * 64 + 48));
let mut acc_lo = vdupq_n_s32(0);
let mut acc_hi = vdupq_n_s32(0);
acc_lo = sdot_lane(acc_lo, lo0, y0, 0);
acc_lo = sdot_lane(acc_lo, lo1, y0, 1);
acc_lo = sdot_lane(acc_lo, lo2, y0, 2);
acc_lo = sdot_lane(acc_lo, lo3, y0, 3);
acc_lo = sdot_lane(acc_lo, lo4, y1, 0);
acc_lo = sdot_lane(acc_lo, lo5, y1, 1);
acc_lo = sdot_lane(acc_lo, lo6, y1, 2);
acc_lo = sdot_lane(acc_lo, lo7, y1, 3);
acc_hi = sdot_lane(acc_hi, hi0, y2, 0);
acc_hi = sdot_lane(acc_hi, hi1, y2, 1);
acc_hi = sdot_lane(acc_hi, hi2, y2, 2);
acc_hi = sdot_lane(acc_hi, hi3, y2, 3);
acc_hi = sdot_lane(acc_hi, hi4, y3, 0);
acc_hi = sdot_lane(acc_hi, hi5, y3, 1);
acc_hi = sdot_lane(acc_hi, hi6, y3, 2);
acc_hi = sdot_lane(acc_hi, hi7, y3, 3);
let sumf = vcvtq_f32_s32(vaddq_s32(
vmulq_s32(sc_lo_w, acc_lo),
vmulq_s32(sc_hi_w, acc_hi),
));
acc_f32[a][c] = vfmaq_f32(acc_f32[a][c], sb_scale[a][c], sumf);
}
}
for a in 0..na {
let bs_lo = vdup_n_s16(bsums_arr[a][2 * sb]);
let bs_hi = vdup_n_s16(bsums_arr[a][2 * sb + 1]);
bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_lo, vget_low_s16(q4sb_mins[0]));
bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_hi, vget_low_s16(q4sb_mins[1]));
bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_lo, vget_high_s16(q4sb_mins[0]));
bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_hi, vget_high_s16(q4sb_mins[1]));
}
}
for a in 0..na {
for c in 0..2 {
acc_f32[a][c] =
vmlsq_f32(acc_f32[a][c], vcvtq_f32_s32(bias_acc[a][c]), sb_min[a][c]);
bias_acc[a][c] = vdupq_n_s32(0);
}
}
}
for a in 0..na {
let mut row = [0f32; Q4_KX8_NROWS];
vst1q_f32(row.as_mut_ptr(), acc_f32[a][0]);
vst1q_f32(row.as_mut_ptr().add(4), acc_f32[a][1]);
for (r, v) in row.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemm_q5_kx8_q8_k_neon_sdot(
packed: &[u8],
acts: &[Q8KActivations],
n_cols: usize,
out: &mut [f32],
) {
let na = acts.len();
debug_assert!(na <= Q5_KX8_GEMM_NC);
let nb = n_cols / Q5_K_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0x0f);
let mone = vdupq_n_u8(1);
let mtwo = vdupq_n_u8(2);
let mut acc_f32 = [[vdupq_n_f32(0.0); 2]; Q5_KX8_GEMM_NC];
let mut bias_acc = [[vdupq_n_s32(0); 2]; Q5_KX8_GEMM_NC];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q5_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let d_lo = vld1q_f32(d_arr.as_ptr());
let d_hi = vld1q_f32(d_arr.as_ptr().add(4));
let dmin_lo = vld1q_f32(dmin_arr.as_ptr());
let dmin_hi = vld1q_f32(dmin_arr.as_ptr().add(4));
let mut sb_scale = [[vdupq_n_f32(0.0); 2]; Q5_KX8_GEMM_NC];
let mut sb_min = [[vdupq_n_f32(0.0); 2]; Q5_KX8_GEMM_NC];
let mut bsums_arr = [[0i16; 8]; Q5_KX8_GEMM_NC];
for (a, act) in acts.iter().enumerate() {
let q8_d = act.d[b];
sb_scale[a] = [vmulq_n_f32(d_lo, q8_d), vmulq_n_f32(d_hi, q8_d)];
sb_min[a] = [vmulq_n_f32(dmin_lo, q8_d), vmulq_n_f32(dmin_hi, q8_d)];
let bsums_ptr = act.bsums.as_ptr().add(b * 16);
for (i, slot) in bsums_arr[a].iter_mut().enumerate() {
*slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
}
}
let scales_base = blk.add(32);
let qh_base = blk.add(128);
let qs_base = blk.add(384);
let mut qh = [[vdupq_n_u8(0); 8]; 2];
for (c, qh_c) in qh.iter_mut().enumerate() {
for (i, slot) in qh_c.iter_mut().enumerate() {
*slot = vld1q_u8(qh_base.add(i * 32 + 16 * c));
}
}
for sb in 0..4 {
let mut q5sb_mins = [vdupq_n_s16(0); 2];
let mut q5sb_scales = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut sc_i8 = [0i8; 8];
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
sc_i8[t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q5sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
for c in 0..2 {
let mut q5_lo = [vdupq_n_s8(0); 8];
let mut q5_hi = [vdupq_n_s8(0); 8];
for i in 0..8 {
let q5_cols =
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + i * 32 + 16 * c));
let hbit_lo = vandq_u8(qh[c][i], mone);
let hbit_hi = vshlq_n_u8(vandq_u8(qh[c][i], mtwo), 3);
qh[c][i] = vshrq_n_u8(qh[c][i], 2);
q5_lo[i] =
vreinterpretq_s8_u8(vsliq_n_u8(vandq_u8(q5_cols, m4b), hbit_lo, 4));
q5_hi[i] = vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_cols, 4), hbit_hi));
}
let (sc_lo, sc_hi) = if c == 0 {
(vget_low_s16(q5sb_scales[0]), vget_low_s16(q5sb_scales[1]))
} else {
(vget_high_s16(q5sb_scales[0]), vget_high_s16(q5sb_scales[1]))
};
let sc_lo_w = vmovl_s16(sc_lo);
let sc_hi_w = vmovl_s16(sc_hi);
for a in 0..na {
let q8_base = acts[a].q.as_ptr().add(b * Q5_K_BLOCK_ELEMS);
let y0 = vld1q_s8(q8_base.add(sb * 64));
let y1 = vld1q_s8(q8_base.add(sb * 64 + 16));
let y2 = vld1q_s8(q8_base.add(sb * 64 + 32));
let y3 = vld1q_s8(q8_base.add(sb * 64 + 48));
let mut acc_lo = vdupq_n_s32(0);
let mut acc_hi = vdupq_n_s32(0);
acc_lo = sdot_lane(acc_lo, q5_lo[0], y0, 0);
acc_lo = sdot_lane(acc_lo, q5_lo[1], y0, 1);
acc_lo = sdot_lane(acc_lo, q5_lo[2], y0, 2);
acc_lo = sdot_lane(acc_lo, q5_lo[3], y0, 3);
acc_lo = sdot_lane(acc_lo, q5_lo[4], y1, 0);
acc_lo = sdot_lane(acc_lo, q5_lo[5], y1, 1);
acc_lo = sdot_lane(acc_lo, q5_lo[6], y1, 2);
acc_lo = sdot_lane(acc_lo, q5_lo[7], y1, 3);
acc_hi = sdot_lane(acc_hi, q5_hi[0], y2, 0);
acc_hi = sdot_lane(acc_hi, q5_hi[1], y2, 1);
acc_hi = sdot_lane(acc_hi, q5_hi[2], y2, 2);
acc_hi = sdot_lane(acc_hi, q5_hi[3], y2, 3);
acc_hi = sdot_lane(acc_hi, q5_hi[4], y3, 0);
acc_hi = sdot_lane(acc_hi, q5_hi[5], y3, 1);
acc_hi = sdot_lane(acc_hi, q5_hi[6], y3, 2);
acc_hi = sdot_lane(acc_hi, q5_hi[7], y3, 3);
let sumf = vcvtq_f32_s32(vaddq_s32(
vmulq_s32(sc_lo_w, acc_lo),
vmulq_s32(sc_hi_w, acc_hi),
));
acc_f32[a][c] = vfmaq_f32(acc_f32[a][c], sb_scale[a][c], sumf);
}
}
for a in 0..na {
let bs_lo = vdup_n_s16(bsums_arr[a][2 * sb]);
let bs_hi = vdup_n_s16(bsums_arr[a][2 * sb + 1]);
bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_lo, vget_low_s16(q5sb_mins[0]));
bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_hi, vget_low_s16(q5sb_mins[1]));
bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_lo, vget_high_s16(q5sb_mins[0]));
bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_hi, vget_high_s16(q5sb_mins[1]));
}
}
for a in 0..na {
for c in 0..2 {
acc_f32[a][c] =
vmlsq_f32(acc_f32[a][c], vcvtq_f32_s32(bias_acc[a][c]), sb_min[a][c]);
bias_acc[a][c] = vdupq_n_s32(0);
}
}
}
for a in 0..na {
let mut row = [0f32; Q5_KX8_NROWS];
vst1q_f32(row.as_mut_ptr(), acc_f32[a][0]);
vst1q_f32(row.as_mut_ptr().add(4), acc_f32[a][1]);
for (r, v) in row.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn gemm_q4_kx8_q8_k_neon_i8mm(
packed: &[u8],
tile: &Q8KActsX4,
n_cols: usize,
out: &mut [f32],
) {
let na = tile.na;
debug_assert!(na <= Q4_KX8_GEMM_NC);
let nb = n_cols / Q4_K_BLOCK_ELEMS;
debug_assert_eq!(tile.n_blocks, nb);
let m4b = vdupq_n_u8(0x0f);
const Q8_K_BLOCKLEN: usize = 4;
let mut acc_f32 = [vdupq_n_f32(0.0); Q4_KX8_GEMM_NC * 2];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q4_KX8_BLOCK_BYTES);
let bsums_base = tile.bsums.as_ptr().add(b * Q8_K_BLOCKLEN * 8);
let mut acc = [vdupq_n_s32(0); 8];
let mut bias_acc = [vdupq_n_s32(0); 8];
for i in 0..8 {
acc[i] = vdupq_n_s32(0);
bias_acc[i] = vdupq_n_s32(0);
}
let scales_base = blk.add(32);
let qs_base = blk.add(128);
let q8_base = tile.qs.as_ptr().add(b * Q4_K_BLOCK_ELEMS * 4);
for sb in 0..4 {
let mut q4sb_scales = [[0i8; 8]; 2];
let mut q4sb_mins = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
q4sb_scales[i][t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
let q8_sb = q8_base.add(sb * 256);
let mut q8_qs_01 = [vdupq_n_s8(0); 8];
let mut q8_qs_23 = [vdupq_n_s8(0); 8];
for i in 0..8 {
q8_qs_01[i] = vld1q_s8(q8_sb.add(i * 32));
q8_qs_23[i] = vld1q_s8(q8_sb.add(i * 32 + 16));
}
let q8s = [q8_qs_01, q8_qs_23];
for cp in 0..4 {
let mut sb_acc = [vdupq_n_s32(0); 4];
let q4_qs = [
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp)),
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 64)),
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 128)),
vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 192)),
];
let q4_nibbles = [
[
vreinterpretq_s8_u8(vandq_u8(q4_qs[0], m4b)),
vreinterpretq_s8_u8(vandq_u8(q4_qs[1], m4b)),
vreinterpretq_s8_u8(vandq_u8(q4_qs[2], m4b)),
vreinterpretq_s8_u8(vandq_u8(q4_qs[3], m4b)),
],
[
vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[0], 4)),
vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[1], 4)),
vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[2], 4)),
vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[3], 4)),
],
];
for rp in 0..2 {
for blk in 0..2 {
let q8 = &q8s[rp][4 * blk..4 * blk + 4];
let q4 = &q4_nibbles[blk];
let mut tile_acc = sb_acc[2 * rp + blk];
for qs_offset in 0..4 {
tile_acc = vmmla_s32(tile_acc, q4[qs_offset], q8[qs_offset]);
}
sb_acc[2 * rp + blk] = tile_acc;
}
}
let scale_offset = cp * 2;
let block_scale_0 = vcombine_s32(
vdup_n_s32(i32::from(q4sb_scales[0][scale_offset])),
vdup_n_s32(i32::from(q4sb_scales[0][scale_offset + 1])),
);
let block_scale_1 = vcombine_s32(
vdup_n_s32(i32::from(q4sb_scales[1][scale_offset])),
vdup_n_s32(i32::from(q4sb_scales[1][scale_offset + 1])),
);
acc[cp] = vmlaq_s32(acc[cp], sb_acc[0], block_scale_0);
acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[2], block_scale_0);
acc[cp] = vmlaq_s32(acc[cp], sb_acc[1], block_scale_1);
acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[3], block_scale_1);
}
for q8_row in 0..Q8_K_BLOCKLEN {
let bs_lo = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb));
let bs_hi = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb + 1));
bias_acc[2 * q8_row] =
vmlal_s16(bias_acc[2 * q8_row], bs_lo, vget_low_s16(q4sb_mins[0]));
bias_acc[2 * q8_row] =
vmlal_s16(bias_acc[2 * q8_row], bs_hi, vget_low_s16(q4sb_mins[1]));
bias_acc[2 * q8_row + 1] =
vmlal_s16(bias_acc[2 * q8_row + 1], bs_lo, vget_high_s16(q4sb_mins[0]));
bias_acc[2 * q8_row + 1] =
vmlal_s16(bias_acc[2 * q8_row + 1], bs_hi, vget_high_s16(q4sb_mins[1]));
}
}
for lane in acc.iter_mut() {
let aux = vzip_s32(vget_low_s32(*lane), vget_high_s32(*lane));
*lane = vcombine_s32(aux.0, aux.1);
}
let reorder_acc = [
vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
vcombine_s32(vget_low_s32(acc[4]), vget_low_s32(acc[5])),
vcombine_s32(vget_low_s32(acc[6]), vget_low_s32(acc[7])),
vcombine_s32(vget_high_s32(acc[4]), vget_high_s32(acc[5])),
vcombine_s32(vget_high_s32(acc[6]), vget_high_s32(acc[7])),
];
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
for i in 0..na {
for j in 0..2 {
let q8_d = vdupq_n_f32(*tile.d.as_ptr().add(b * Q8_K_BLOCKLEN + i));
let dmins = vmulq_f32(vld1q_f32(dmin_arr.as_ptr().add(j * 4)), q8_d);
let scale = vmulq_f32(vld1q_f32(d_arr.as_ptr().add(j * 4)), q8_d);
let idx = 2 * i + j;
acc_f32[idx] = vmlsq_f32(acc_f32[idx], vcvtq_f32_s32(bias_acc[idx]), dmins);
acc_f32[idx] = vmlaq_f32(acc_f32[idx], vcvtq_f32_s32(reorder_acc[idx]), scale);
}
}
}
for a in 0..na {
let mut row = [0f32; Q4_KX8_NROWS];
vst1q_f32(row.as_mut_ptr(), acc_f32[2 * a]);
vst1q_f32(row.as_mut_ptr().add(4), acc_f32[2 * a + 1]);
for (r, v) in row.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn gemm_q5_kx8_q8_k_neon_i8mm(
packed: &[u8],
tile: &Q8KActsX4,
n_cols: usize,
out: &mut [f32],
) {
let na = tile.na;
debug_assert!(na <= Q5_KX8_GEMM_NC);
let nb = n_cols / Q5_K_BLOCK_ELEMS;
debug_assert_eq!(tile.n_blocks, nb);
let m4b = vdupq_n_u8(0x0f);
let mone = vdupq_n_u8(1);
let mtwo = vdupq_n_u8(2);
const Q8_K_BLOCKLEN: usize = 4;
let mut acc_f32 = [vdupq_n_f32(0.0); Q5_KX8_GEMM_NC * 2];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q5_KX8_BLOCK_BYTES);
let bsums_base = tile.bsums.as_ptr().add(b * Q8_K_BLOCKLEN * 8);
let mut acc = [vdupq_n_s32(0); 8];
let mut bias_acc = [vdupq_n_s32(0); 8];
let scales_base = blk.add(32);
let qh_base = blk.add(128);
let qs_base = blk.add(384);
let q8_base = tile.qs.as_ptr().add(b * Q5_K_BLOCK_ELEMS * 4);
let mut qh = [[vdupq_n_u8(0); 4]; 4];
for (cp, qh_cp) in qh.iter_mut().enumerate() {
for (m, slot) in qh_cp.iter_mut().enumerate() {
*slot = vld1q_u8(qh_base.add(16 * cp + 64 * m));
}
}
for sb in 0..4 {
let mut q5sb_scales = [[0i8; 8]; 2];
let mut q5sb_mins = [vdupq_n_s16(0); 2];
for i in 0..2 {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
let offset = sb * 24 + i * 12;
decode_scales_mins(
std::slice::from_raw_parts(scales_base.add(offset), 12),
&mut sc,
&mut mn,
);
let mut mn_i8 = [0i8; 8];
for t in 0..8 {
q5sb_scales[i][t] = sc[t] as i8;
mn_i8[t] = mn[t] as i8;
}
q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
}
let q8_sb = q8_base.add(sb * 256);
let mut q8_qs_01 = [vdupq_n_s8(0); 8];
let mut q8_qs_23 = [vdupq_n_s8(0); 8];
for i in 0..8 {
q8_qs_01[i] = vld1q_s8(q8_sb.add(i * 32));
q8_qs_23[i] = vld1q_s8(q8_sb.add(i * 32 + 16));
}
let q8s = [q8_qs_01, q8_qs_23];
for cp in 0..4 {
let mut sb_acc = [vdupq_n_s32(0); 4];
let q5_qs = [
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp)),
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 64)),
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 128)),
vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 192)),
];
let mut q5_lo = [vdupq_n_s8(0); 4];
let mut q5_hi = [vdupq_n_s8(0); 4];
for m in 0..4 {
let hbit_lo = vandq_u8(qh[cp][m], mone);
let hbit_hi = vshlq_n_u8(vandq_u8(qh[cp][m], mtwo), 3);
qh[cp][m] = vshrq_n_u8(qh[cp][m], 2);
q5_lo[m] =
vreinterpretq_s8_u8(vsliq_n_u8(vandq_u8(q5_qs[m], m4b), hbit_lo, 4));
q5_hi[m] = vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_qs[m], 4), hbit_hi));
}
let q5_vals = [q5_lo, q5_hi];
for rp in 0..2 {
for half in 0..2 {
let q8 = &q8s[rp][4 * half..4 * half + 4];
let q5 = &q5_vals[half];
let mut tile_acc = sb_acc[2 * rp + half];
for m in 0..4 {
tile_acc = vmmla_s32(tile_acc, q5[m], q8[m]);
}
sb_acc[2 * rp + half] = tile_acc;
}
}
let scale_offset = cp * 2;
let block_scale_0 = vcombine_s32(
vdup_n_s32(i32::from(q5sb_scales[0][scale_offset])),
vdup_n_s32(i32::from(q5sb_scales[0][scale_offset + 1])),
);
let block_scale_1 = vcombine_s32(
vdup_n_s32(i32::from(q5sb_scales[1][scale_offset])),
vdup_n_s32(i32::from(q5sb_scales[1][scale_offset + 1])),
);
acc[cp] = vmlaq_s32(acc[cp], sb_acc[0], block_scale_0);
acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[2], block_scale_0);
acc[cp] = vmlaq_s32(acc[cp], sb_acc[1], block_scale_1);
acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[3], block_scale_1);
}
for q8_row in 0..Q8_K_BLOCKLEN {
let bs_lo = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb));
let bs_hi = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb + 1));
bias_acc[2 * q8_row] =
vmlal_s16(bias_acc[2 * q8_row], bs_lo, vget_low_s16(q5sb_mins[0]));
bias_acc[2 * q8_row] =
vmlal_s16(bias_acc[2 * q8_row], bs_hi, vget_low_s16(q5sb_mins[1]));
bias_acc[2 * q8_row + 1] =
vmlal_s16(bias_acc[2 * q8_row + 1], bs_lo, vget_high_s16(q5sb_mins[0]));
bias_acc[2 * q8_row + 1] =
vmlal_s16(bias_acc[2 * q8_row + 1], bs_hi, vget_high_s16(q5sb_mins[1]));
}
}
for lane in acc.iter_mut() {
let aux = vzip_s32(vget_low_s32(*lane), vget_high_s32(*lane));
*lane = vcombine_s32(aux.0, aux.1);
}
let reorder_acc = [
vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
vcombine_s32(vget_low_s32(acc[4]), vget_low_s32(acc[5])),
vcombine_s32(vget_low_s32(acc[6]), vget_low_s32(acc[7])),
vcombine_s32(vget_high_s32(acc[4]), vget_high_s32(acc[5])),
vcombine_s32(vget_high_s32(acc[6]), vget_high_s32(acc[7])),
];
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
for i in 0..na {
for j in 0..2 {
let q8_d = vdupq_n_f32(*tile.d.as_ptr().add(b * Q8_K_BLOCKLEN + i));
let dmins = vmulq_f32(vld1q_f32(dmin_arr.as_ptr().add(j * 4)), q8_d);
let scale = vmulq_f32(vld1q_f32(d_arr.as_ptr().add(j * 4)), q8_d);
let idx = 2 * i + j;
acc_f32[idx] = vmlsq_f32(acc_f32[idx], vcvtq_f32_s32(bias_acc[idx]), dmins);
acc_f32[idx] = vmlaq_f32(acc_f32[idx], vcvtq_f32_s32(reorder_acc[idx]), scale);
}
}
}
for a in 0..na {
let mut row = [0f32; Q5_KX8_NROWS];
vst1q_f32(row.as_mut_ptr(), acc_f32[2 * a]);
vst1q_f32(row.as_mut_ptr().add(4), acc_f32[2 * a + 1]);
for (r, v) in row.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn gemm_q6_kx8_q8_k_neon_i8mm(
packed: &[u8],
tile: &Q8KActsX4,
n_cols: usize,
out: &mut [f32],
) {
let na = tile.na;
debug_assert!(na <= Q8K_ACTS_X4_NC);
let nb = n_cols / Q6_K_BLOCK_ELEMS;
debug_assert_eq!(tile.n_blocks, nb);
let m4b = vdupq_n_u8(0x0f);
let mask_lo = vdupq_n_u8(0x03);
let mask_hi = vdupq_n_u8(0x30);
let m32s = vdupq_n_s8(32);
const Q8_K_BLOCKLEN: usize = 4;
let mut acc_f32 = [vdupq_n_f32(0.0); Q8K_ACTS_X4_NC * 2];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q6_KX8_BLOCK_BYTES);
let scales_base = blk.add(16) as *const i8;
let ql_blk = blk.add(144);
let qh_blk = blk.add(1168);
let q8_blk = tile.qs.as_ptr().add(b * Q6_K_BLOCK_ELEMS * 4);
let mut acc = [vdupq_n_s32(0); 8];
let mut q6_scales = [0i16; 16 * 8];
for i in 0..16 {
let s16 = vmovl_s8(vld1_s8(scales_base.add(i * 8)));
vst1q_s16(q6_scales.as_mut_ptr().add(i * 8), s16);
}
for half in 0..2 {
let ql_base = ql_blk.add(half * 512);
let qh_base = qh_blk.add(half * 256);
for sb in 0..4 {
let q8_base_l = q8_blk.add(half * 512 + sb * 64);
let q8_base_h = q8_blk.add(half * 512 + 256 + sb * 64);
let mut q8_l_01 = [vdupq_n_s8(0); 2];
let mut q8_l_23 = [vdupq_n_s8(0); 2];
let mut q8_h_01 = [vdupq_n_s8(0); 2];
let mut q8_h_23 = [vdupq_n_s8(0); 2];
for i in 0..2 {
q8_l_01[i] = vld1q_s8(q8_base_l.add(i * 32));
q8_l_23[i] = vld1q_s8(q8_base_l.add(i * 32 + 16));
q8_h_01[i] = vld1q_s8(q8_base_h.add(i * 32));
q8_h_23[i] = vld1q_s8(q8_base_h.add(i * 32 + 16));
}
let ql_off = sb * (Q6_K_BLOCK_ELEMS / 2);
let qh_off = ql_off & 255; let mut q6_ql_0 = [vdupq_n_u8(0); 4];
let mut q6_ql_1 = [vdupq_n_u8(0); 4];
let mut q6_qh_0 = [vdupq_n_u8(0); 4];
let mut q6_qh_1 = [vdupq_n_u8(0); 4];
for k in 0..4 {
q6_ql_0[k] = vld1q_u8(ql_base.add(ql_off + 16 * k));
q6_ql_1[k] = vld1q_u8(ql_base.add(ql_off + 64 + 16 * k));
q6_qh_0[k] = vld1q_u8(qh_base.add(qh_off + 16 * k));
q6_qh_1[k] = vld1q_u8(qh_base.add(qh_off + 64 + 16 * k));
}
if sb > 1 {
for k in 0..4 {
q6_qh_0[k] = vshrq_n_u8(q6_qh_0[k], 2);
q6_qh_1[k] = vshrq_n_u8(q6_qh_1[k], 2);
}
}
for cp in 0..4 {
let hh_0 = vandq_u8(q6_qh_0[cp], mask_hi);
let hh_1 = vandq_u8(q6_qh_1[cp], mask_hi);
let q6_l0 = vsubq_s8(
vreinterpretq_s8_u8(vsliq_n_u8(
vandq_u8(q6_ql_0[cp], m4b),
vandq_u8(q6_qh_0[cp], mask_lo),
4,
)),
m32s,
);
let q6_l1 = vsubq_s8(
vreinterpretq_s8_u8(vsliq_n_u8(
vandq_u8(q6_ql_1[cp], m4b),
vandq_u8(q6_qh_1[cp], mask_lo),
4,
)),
m32s,
);
let q6_h0 = vsubq_s8(
vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_0[cp], 4), hh_0)),
m32s,
);
let q6_h1 = vsubq_s8(
vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_1[cp], 4), hh_1)),
m32s,
);
let mut sb_acc_0l = vmmla_s32(vdupq_n_s32(0), q6_l0, q8_l_01[0]);
sb_acc_0l = vmmla_s32(sb_acc_0l, q6_l1, q8_l_01[1]);
let mut sb_acc_0h = vmmla_s32(vdupq_n_s32(0), q6_h0, q8_h_01[0]);
sb_acc_0h = vmmla_s32(sb_acc_0h, q6_h1, q8_h_01[1]);
let mut sb_acc_1l = vmmla_s32(vdupq_n_s32(0), q6_l0, q8_l_23[0]);
sb_acc_1l = vmmla_s32(sb_acc_1l, q6_l1, q8_l_23[1]);
let mut sb_acc_1h = vmmla_s32(vdupq_n_s32(0), q6_h0, q8_h_23[0]);
sb_acc_1h = vmmla_s32(sb_acc_1h, q6_h1, q8_h_23[1]);
let scale_idx_l = half * 8 + sb;
let scale_idx_h = half * 8 + sb + 4;
let scale_l = vcombine_s32(
vdup_n_s32(i32::from(q6_scales[scale_idx_l * 8 + cp * 2])),
vdup_n_s32(i32::from(q6_scales[scale_idx_l * 8 + cp * 2 + 1])),
);
let scale_h = vcombine_s32(
vdup_n_s32(i32::from(q6_scales[scale_idx_h * 8 + cp * 2])),
vdup_n_s32(i32::from(q6_scales[scale_idx_h * 8 + cp * 2 + 1])),
);
acc[cp] = vmlaq_s32(acc[cp], sb_acc_0l, scale_l);
acc[cp] = vmlaq_s32(acc[cp], sb_acc_0h, scale_h);
acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc_1l, scale_l);
acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc_1h, scale_h);
}
}
}
for lane in acc.iter_mut() {
let aux = vzip_s32(vget_low_s32(*lane), vget_high_s32(*lane));
*lane = vcombine_s32(aux.0, aux.1);
}
let reorder_acc = [
vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
vcombine_s32(vget_low_s32(acc[4]), vget_low_s32(acc[5])),
vcombine_s32(vget_low_s32(acc[6]), vget_low_s32(acc[7])),
vcombine_s32(vget_high_s32(acc[4]), vget_high_s32(acc[5])),
vcombine_s32(vget_high_s32(acc[6]), vget_high_s32(acc[7])),
];
let mut d_arr = [0f32; 8];
for (j, slot) in d_arr.iter_mut().enumerate() {
*slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
}
for i in 0..na {
for j in 0..2 {
let q8_d = vdupq_n_f32(*tile.d.as_ptr().add(b * Q8_K_BLOCKLEN + i));
let scale = vmulq_f32(vld1q_f32(d_arr.as_ptr().add(j * 4)), q8_d);
let idx = 2 * i + j;
acc_f32[idx] = vmlaq_f32(acc_f32[idx], vcvtq_f32_s32(reorder_acc[idx]), scale);
}
}
}
for a in 0..na {
let mut row = [0f32; Q6_KX8_NROWS];
vst1q_f32(row.as_mut_ptr(), acc_f32[2 * a]);
vst1q_f32(row.as_mut_ptr().add(4), acc_f32[2 * a + 1]);
for (r, v) in row.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q8_0x4_q8_0_neon_sdot(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q8_0_BLOCK_ELEMS;
for x in 0..n_row_groups {
let mut acc = vdupq_n_f32(0.0);
let group_off = x * nb * Q8_0X4_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q8_0X4_BLOCK_BYTES);
let qs = blk.add(8);
let b0 = vld1q_s8(qs as *const i8);
let b1 = vld1q_s8(qs.add(16) as *const i8);
let b2 = vld1q_s8(qs.add(32) as *const i8);
let b3 = vld1q_s8(qs.add(48) as *const i8);
let b4 = vld1q_s8(qs.add(64) as *const i8);
let b5 = vld1q_s8(qs.add(80) as *const i8);
let b6 = vld1q_s8(qs.add(96) as *const i8);
let b7 = vld1q_s8(qs.add(112) as *const i8);
let a_ptr = act.q.as_ptr().add(b * Q8_0_BLOCK_ELEMS);
let a0 = vld1q_s8(a_ptr);
let a1 = vld1q_s8(a_ptr.add(16));
let mut ret = vdupq_n_s32(0);
ret = sdot_lane(ret, b0, a0, 0);
ret = sdot_lane(ret, b1, a0, 1);
ret = sdot_lane(ret, b2, a0, 2);
ret = sdot_lane(ret, b3, a0, 3);
ret = sdot_lane(ret, b4, a1, 0);
ret = sdot_lane(ret, b5, a1, 1);
ret = sdot_lane(ret, b6, a1, 2);
ret = sdot_lane(ret, b7, a1, 3);
let d_bits = vld1_u16(blk as *const u16);
let mut dw = [0f32; 4];
dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
let scale = vmulq_n_f32(vld1q_f32(dw.as_ptr()), act.d[b]);
acc = vfmaq_f32(acc, vcvtq_f32_s32(ret), scale);
}
vst1q_f32(out.as_mut_ptr().add(x * Q8_0X4_NROWS), acc);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemm_q8_0x4_q8_0_neon_sdot(
group: &[u8],
acts: &[Q8Activations],
n_cols: usize,
out: &mut [f32],
) {
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let n_acts = acts.len();
let mut j0 = 0;
while j0 < n_acts {
let tile = Q8_0X4_GEMM_NC.min(n_acts - j0);
let mut acc = [vdupq_n_f32(0.0); Q8_0X4_GEMM_NC];
for b in 0..nb {
let blk = group.as_ptr().add(b * Q8_0X4_BLOCK_BYTES);
let qs = blk.add(8);
let w = [
vld1q_s8(qs as *const i8),
vld1q_s8(qs.add(16) as *const i8),
vld1q_s8(qs.add(32) as *const i8),
vld1q_s8(qs.add(48) as *const i8),
vld1q_s8(qs.add(64) as *const i8),
vld1q_s8(qs.add(80) as *const i8),
vld1q_s8(qs.add(96) as *const i8),
vld1q_s8(qs.add(112) as *const i8),
];
let d_bits = vld1_u16(blk as *const u16);
let mut dw = [0f32; 4];
dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
let dw_v = vld1q_f32(dw.as_ptr());
for t in 0..tile {
let act = &acts[j0 + t];
let a_ptr = act.q.as_ptr().add(b * Q8_0_BLOCK_ELEMS);
let a0 = vld1q_s8(a_ptr);
let a1 = vld1q_s8(a_ptr.add(16));
let mut ret = vdupq_n_s32(0);
ret = sdot_lane(ret, w[0], a0, 0);
ret = sdot_lane(ret, w[1], a0, 1);
ret = sdot_lane(ret, w[2], a0, 2);
ret = sdot_lane(ret, w[3], a0, 3);
ret = sdot_lane(ret, w[4], a1, 0);
ret = sdot_lane(ret, w[5], a1, 1);
ret = sdot_lane(ret, w[6], a1, 2);
ret = sdot_lane(ret, w[7], a1, 3);
let scale = vmulq_n_f32(dw_v, act.d[b]);
acc[t] = vfmaq_f32(acc[t], vcvtq_f32_s32(ret), scale);
}
}
for t in 0..tile {
let mut lanes = [0f32; Q8_0X4_NROWS];
vst1q_f32(lanes.as_mut_ptr(), acc[t]);
for (r, v) in lanes.iter().enumerate() {
out[r * n_acts + j0 + t] = *v;
}
}
j0 += tile;
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q4_0x4_q8_0_neon_sdot(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let maskf0 = vdupq_n_u8(0xF0);
for x in 0..n_row_groups {
let mut acc = vdupq_n_f32(0.0);
let group_off = x * nb * Q4_0X4_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q4_0X4_BLOCK_BYTES);
let qs = blk.add(8);
let a_ptr = act.q.as_ptr().add(b * Q4_0_BLOCK_ELEMS);
let a0 = vld1q_s8(a_ptr);
let a1 = vld1q_s8(a_ptr.add(16));
let mut ret = vdupq_n_s32(0);
for wi in 0..4u32 {
let w = vld1q_u8(qs.add(wi as usize * 16));
let hi = vreinterpretq_s8_u8(vshlq_n_u8(w, 4));
let lo = vreinterpretq_s8_u8(vandq_u8(w, maskf0));
ret = sdot_lane(ret, hi, a0, wi);
ret = sdot_lane(ret, lo, a1, wi);
}
let d_bits = vld1_u16(blk as *const u16);
let mut dw = [0f32; 4];
dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
let scale = vmulq_n_f32(vld1q_f32(dw.as_ptr()), act.d[b]);
acc = vfmaq_f32(acc, vcvtq_f32_s32(vshrq_n_s32(ret, 4)), scale);
}
vst1q_f32(out.as_mut_ptr().add(x * Q4_0X4_NROWS), acc);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemm_q4_0x4_q8_0_neon_sdot(
group: &[u8],
acts: &[Q8Activations],
n_cols: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let n_acts = acts.len();
let maskf0 = vdupq_n_u8(0xF0);
let mut j0 = 0;
while j0 < n_acts {
let tile = Q4_0X4_GEMM_NC.min(n_acts - j0);
let mut acc = [vdupq_n_f32(0.0); Q4_0X4_GEMM_NC];
for b in 0..nb {
let blk = group.as_ptr().add(b * Q4_0X4_BLOCK_BYTES);
let qs = blk.add(8);
let w = [
vld1q_u8(qs),
vld1q_u8(qs.add(16)),
vld1q_u8(qs.add(32)),
vld1q_u8(qs.add(48)),
];
let d_bits = vld1_u16(blk as *const u16);
let mut dw = [0f32; 4];
dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
let dw_v = vld1q_f32(dw.as_ptr());
for t in 0..tile {
let act = &acts[j0 + t];
let a_ptr = act.q.as_ptr().add(b * Q4_0_BLOCK_ELEMS);
let a0 = vld1q_s8(a_ptr);
let a1 = vld1q_s8(a_ptr.add(16));
let mut ret = vdupq_n_s32(0);
for (wi, wchunk) in w.iter().enumerate() {
let hi = vreinterpretq_s8_u8(vshlq_n_u8(*wchunk, 4));
let lo = vreinterpretq_s8_u8(vandq_u8(*wchunk, maskf0));
ret = sdot_lane(ret, hi, a0, wi as u32);
ret = sdot_lane(ret, lo, a1, wi as u32);
}
let scale = vmulq_n_f32(dw_v, act.d[b]);
acc[t] = vfmaq_f32(acc[t], vcvtq_f32_s32(vshrq_n_s32(ret, 4)), scale);
}
}
for t in 0..tile {
let mut lanes = [0f32; Q4_0X4_NROWS];
vst1q_f32(lanes.as_mut_ptr(), acc[t]);
for (r, v) in lanes.iter().enumerate() {
out[r * n_acts + j0 + t] = *v;
}
}
j0 += tile;
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q8_0x4_q8_0_neon_4x8(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q8_0_BLOCK_ELEMS;
for x in 0..n_row_groups {
let mut acc = vdupq_n_f32(0.0);
let group_off = x * nb * Q8_0X4_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q8_0X4_BLOCK_BYTES);
let qs = blk.add(8) as *const i8;
let mut d_arr = [0f32; 4];
for (j, slot) in d_arr.iter_mut().enumerate() {
*slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
}
let b_d = vld1q_f32(d_arr.as_ptr());
let a_base = act.q.as_ptr().add(b * Q8_0_BLOCK_ELEMS);
let mut ret0 = vdupq_n_s32(0);
let mut ret1 = vdupq_n_s32(0);
for c in 0..4 {
let a = vreinterpretq_s8_s64(vld1q_dup_s64(a_base.add(c * 8) as *const i64));
ret0 = sdot(ret0, vld1q_s8(qs.add(c * 32)), a);
ret1 = sdot(ret1, vld1q_s8(qs.add(c * 32 + 16)), a);
}
let ret = vpaddq_s32(ret0, ret1);
acc = vfmaq_f32(acc, vcvtq_f32_s32(ret), vmulq_n_f32(b_d, act.d[b]));
}
vst1q_f32(out.as_mut_ptr().add(x * Q8_0X4_NROWS), acc);
}
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn gemv_q4_0x4_q8_0_neon_4x8(
packed: &[u8],
act: &Q8Activations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let m4b = vdupq_n_u8(0xf0);
for x in 0..n_row_groups {
let mut acc = vdupq_n_f32(0.0);
let group_off = x * nb * Q4_0X4_BLOCK_BYTES;
for b in 0..nb {
let blk = packed.as_ptr().add(group_off + b * Q4_0X4_BLOCK_BYTES);
let qs = blk.add(8);
let mut d_arr = [0f32; 4];
for (j, slot) in d_arr.iter_mut().enumerate() {
*slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
}
let b_d = vld1q_f32(d_arr.as_ptr());
let a_base = act.q.as_ptr().add(b * Q4_0_BLOCK_ELEMS);
let b0 = vld1q_u8(qs);
let b1 = vld1q_u8(qs.add(16));
let b2 = vld1q_u8(qs.add(32));
let b3 = vld1q_u8(qs.add(48));
let mut a = [vdupq_n_s8(0); 4];
for (c, slot) in a.iter_mut().enumerate() {
*slot = vreinterpretq_s8_s64(vld1q_dup_s64(a_base.add(c * 8) as *const i64));
}
let mut ret0 = vdupq_n_s32(0);
let mut ret1 = vdupq_n_s32(0);
ret0 = sdot(ret0, vreinterpretq_s8_u8(vshlq_n_u8(b0, 4)), a[0]);
ret1 = sdot(ret1, vreinterpretq_s8_u8(vshlq_n_u8(b1, 4)), a[0]);
ret0 = sdot(ret0, vreinterpretq_s8_u8(vshlq_n_u8(b2, 4)), a[1]);
ret1 = sdot(ret1, vreinterpretq_s8_u8(vshlq_n_u8(b3, 4)), a[1]);
ret0 = sdot(ret0, vreinterpretq_s8_u8(vandq_u8(b0, m4b)), a[2]);
ret1 = sdot(ret1, vreinterpretq_s8_u8(vandq_u8(b1, m4b)), a[2]);
ret0 = sdot(ret0, vreinterpretq_s8_u8(vandq_u8(b2, m4b)), a[3]);
ret1 = sdot(ret1, vreinterpretq_s8_u8(vandq_u8(b3, m4b)), a[3]);
let ret = vpaddq_s32(ret0, ret1);
acc = vfmaq_f32(acc, vcvtq_n_f32_s32::<4>(ret), vmulq_n_f32(b_d, act.d[b]));
}
vst1q_f32(out.as_mut_ptr().add(x * Q4_0X4_NROWS), acc);
}
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn gemm_q8_0x4_q8_0_neon_i8mm(
packed: &[u8],
tile: &Q8ActsX4,
n_cols: usize,
out: &mut [f32],
) {
let na = tile.na;
debug_assert!(na <= Q8K_ACTS_X4_NC);
let nb = n_cols / Q8_0_BLOCK_ELEMS;
debug_assert_eq!(tile.n_blocks, nb);
let mut acc_f32 = [vdupq_n_f32(0.0); Q8K_ACTS_X4_NC];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q8_0X4_BLOCK_BYTES);
let qs = blk.add(8) as *const i8;
let a_base = tile.qs.as_ptr().add(b * Q8_0_BLOCK_ELEMS * 4);
let mut acc = [vdupq_n_s32(0); 4];
for chunk in 0..4 {
let a01 = vld1q_s8(a_base.add(chunk * 32));
let a23 = vld1q_s8(a_base.add(chunk * 32 + 16));
let b01 = vld1q_s8(qs.add(chunk * 32));
let b23 = vld1q_s8(qs.add(chunk * 32 + 16));
acc[0] = vmmla_s32(acc[0], a01, b01);
acc[1] = vmmla_s32(acc[1], a01, b23);
acc[2] = vmmla_s32(acc[2], a23, b01);
acc[3] = vmmla_s32(acc[3], a23, b23);
}
let rows = [
vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
];
let mut d_arr = [0f32; 4];
for (j, slot) in d_arr.iter_mut().enumerate() {
*slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
}
let b_d = vld1q_f32(d_arr.as_ptr());
for a in 0..na {
acc_f32[a] = vfmaq_f32(
acc_f32[a],
vcvtq_f32_s32(rows[a]),
vmulq_n_f32(b_d, *tile.d.as_ptr().add(b * 4 + a)),
);
}
}
for a in 0..na {
let mut lanes = [0f32; Q8_0X4_NROWS];
vst1q_f32(lanes.as_mut_ptr(), acc_f32[a]);
for (r, v) in lanes.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn gemm_q4_0x4_q8_0_neon_i8mm(
packed: &[u8],
tile: &Q8ActsX4,
n_cols: usize,
out: &mut [f32],
) {
let na = tile.na;
debug_assert!(na <= Q8K_ACTS_X4_NC);
let nb = n_cols / Q4_0_BLOCK_ELEMS;
debug_assert_eq!(tile.n_blocks, nb);
let m4b = vdupq_n_u8(0xf0);
let mut acc_f32 = [vdupq_n_f32(0.0); Q8K_ACTS_X4_NC];
for b in 0..nb {
let blk = packed.as_ptr().add(b * Q4_0X4_BLOCK_BYTES);
let qs = blk.add(8);
let a_base = tile.qs.as_ptr().add(b * Q4_0_BLOCK_ELEMS * 4);
let bv = [
vld1q_u8(qs),
vld1q_u8(qs.add(16)),
vld1q_u8(qs.add(32)),
vld1q_u8(qs.add(48)),
];
let w = [
[
vreinterpretq_s8_u8(vshlq_n_u8(bv[0], 4)),
vreinterpretq_s8_u8(vshlq_n_u8(bv[1], 4)),
],
[
vreinterpretq_s8_u8(vshlq_n_u8(bv[2], 4)),
vreinterpretq_s8_u8(vshlq_n_u8(bv[3], 4)),
],
[
vreinterpretq_s8_u8(vandq_u8(bv[0], m4b)),
vreinterpretq_s8_u8(vandq_u8(bv[1], m4b)),
],
[
vreinterpretq_s8_u8(vandq_u8(bv[2], m4b)),
vreinterpretq_s8_u8(vandq_u8(bv[3], m4b)),
],
];
let mut acc = [vdupq_n_s32(0); 4];
for (chunk, w_pair) in w.iter().enumerate() {
let a01 = vld1q_s8(a_base.add(chunk * 32));
let a23 = vld1q_s8(a_base.add(chunk * 32 + 16));
acc[0] = vmmla_s32(acc[0], a01, w_pair[0]);
acc[1] = vmmla_s32(acc[1], a01, w_pair[1]);
acc[2] = vmmla_s32(acc[2], a23, w_pair[0]);
acc[3] = vmmla_s32(acc[3], a23, w_pair[1]);
}
let rows = [
vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
];
let mut d_arr = [0f32; 4];
for (j, slot) in d_arr.iter_mut().enumerate() {
*slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
}
let b_d = vld1q_f32(d_arr.as_ptr());
for a in 0..na {
acc_f32[a] = vfmaq_f32(
acc_f32[a],
vcvtq_n_f32_s32::<4>(rows[a]),
vmulq_n_f32(b_d, *tile.d.as_ptr().add(b * 4 + a)),
);
}
}
for a in 0..na {
let mut lanes = [0f32; Q4_0X4_NROWS];
vst1q_f32(lanes.as_mut_ptr(), acc_f32[a]);
for (r, v) in lanes.iter().enumerate() {
out[r * na + a] = *v;
}
}
}
}
#[cfg(target_arch = "x86_64")]
mod avx2 {
use super::*;
use std::arch::x86_64::*;
#[target_feature(enable = "avx2,fma")]
pub unsafe fn gemv_q4_kx8_q8_k_avx2(
packed: &[u8],
act: &Q8KActivations,
n_cols: usize,
n_row_groups: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_K_BLOCK_ELEMS;
let blocklen = 8;
let ncols = Q4_KX8_NROWS;
for x in 0..n_row_groups {
let mut acc = _mm256_setzero_ps();
let mut acc_min = _mm256_setzero_ps();
let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
for l in 0..nb {
let blk = packed.as_ptr().add(group_off + l * Q4_KX8_BLOCK_BYTES);
let mut d_arr = [0f32; 8];
let mut dmin_arr = [0f32; 8];
for j in 0..8 {
d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
dmin_arr[j] =
f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
}
let da = act.d[l];
let d_vec = _mm256_mul_ps(_mm256_loadu_ps(d_arr.as_ptr()), _mm256_set1_ps(da));
let dmin_vec =
_mm256_mul_ps(_mm256_loadu_ps(dmin_arr.as_ptr()), _mm256_set1_ps(da));
let scales = std::slice::from_raw_parts(blk.add(32), 96);
let qs = std::slice::from_raw_parts(blk.add(128), 1024);
let q8 = &act.q[l * Q4_K_BLOCK_ELEMS..(l + 1) * Q4_K_BLOCK_ELEMS];
let bsums = &act.bsums[l * 16..(l + 1) * 16];
let mut all_scales = [[0u8; 8]; 8];
let mut all_mins = [[0u8; 8]; 8];
for sb in 0..8 {
decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
}
let mut isum = [0i32; 8];
let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen);
for k in 0..n_k {
let sb_pair = k / 4;
let sc0 = &all_scales[sb_pair * 2];
let sc1 = &all_scales[sb_pair * 2 + 1];
for j in 0..ncols {
let mut s = 0i32;
for i in 0..blocklen {
let qbyte = qs[k * ncols * blocklen + j * blocklen + i];
let v0 = (qbyte & 0x0F) as i32;
let v1 = (qbyte >> 4) as i32;
let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
s += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
}
isum[j] += s;
}
}
let isum_ps =
_mm256_cvtepi32_ps(_mm256_loadu_si256(isum.as_ptr() as *const __m256i));
acc = _mm256_fmadd_ps(isum_ps, d_vec, acc);
let mut minsum = [0i32; 8];
for sb in 0..8 {
let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
for j in 0..ncols {
minsum[j] += all_mins[sb][j] as i32 * bsum;
}
}
let minsum_ps =
_mm256_cvtepi32_ps(_mm256_loadu_si256(minsum.as_ptr() as *const __m256i));
acc_min = _mm256_fmadd_ps(minsum_ps, dmin_vec, acc_min);
}
_mm256_storeu_ps(out.as_mut_ptr().add(x * ncols), _mm256_sub_ps(acc, acc_min));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
dot_q4_0_q8_scalar, dot_q4_k_q8_scalar, dot_q5_k_q8_scalar, dot_q6_k_q8_scalar,
dot_q8_0_q8_scalar, quantize_activations_q8, quantize_activations_q8_k, Q4_0_BLOCK_BYTES,
Q5_K_BLOCK_BYTES, Q6_K_BLOCK_BYTES,
};
fn synth_q5_k_row(n_blocks: usize, seed: u8) -> Vec<u8> {
let mut weights = Vec::with_capacity(n_blocks * Q5_K_BLOCK_BYTES);
for b in 0..n_blocks {
weights.extend_from_slice(
&f16::from_f32(0.05 + (b as f32 + seed as f32) * 0.01).to_le_bytes(),
);
weights.extend_from_slice(
&f16::from_f32(0.01 + (b as f32 + seed as f32) * 0.002).to_le_bytes(),
);
for i in 0..12u8 {
weights.push(20 + i.wrapping_mul(3).wrapping_add(seed));
}
for i in 0..32u8 {
weights.push(i.wrapping_mul(11).wrapping_add(b as u8).wrapping_add(seed));
}
for i in 0..128u8 {
weights.push(i.wrapping_mul(19).wrapping_add(b as u8).wrapping_add(seed));
}
}
weights
}
fn synth_q6_k_row(n_blocks: usize, seed: u8) -> Vec<u8> {
let mut weights = Vec::with_capacity(n_blocks * Q6_K_BLOCK_BYTES);
for b in 0..n_blocks {
for i in 0..128u8 {
weights.push(i.wrapping_mul(17).wrapping_add(b as u8).wrapping_add(seed));
}
for i in 0..64u8 {
weights.push(i.wrapping_mul(13).wrapping_add(seed).wrapping_add(b as u8));
}
for i in 0..16u8 {
weights.push((20i8).wrapping_add(i as i8).wrapping_add(seed as i8) as u8);
}
weights.extend_from_slice(
&f16::from_f32(0.04 + (b as f32 + seed as f32) * 0.008).to_le_bytes(),
);
}
weights
}
fn synth_q4_k_row(n_blocks: usize, seed: u8) -> Vec<u8> {
let mut weights = Vec::with_capacity(n_blocks * Q4_K_BLOCK_BYTES);
for b in 0..n_blocks {
weights.extend_from_slice(
&f16::from_f32(0.05 + (b as f32 + seed as f32) * 0.01).to_le_bytes(),
);
weights.extend_from_slice(
&f16::from_f32(0.01 + (b as f32 + seed as f32) * 0.002).to_le_bytes(),
);
for i in 0..12u8 {
weights.push(20 + i.wrapping_mul(3).wrapping_add(seed));
}
for i in 0..128u8 {
weights.push(i.wrapping_mul(17).wrapping_add(b as u8).wrapping_add(seed));
}
}
weights
}
fn synth_q4_0_row(n_blocks: usize, seed: u8) -> Vec<u8> {
let mut weights = Vec::with_capacity(n_blocks * Q4_0_BLOCK_BYTES);
for b in 0..n_blocks {
weights.extend_from_slice(
&f16::from_f32(0.05 + (b as f32 + seed as f32) * 0.012).to_le_bytes(),
);
for i in 0..16u8 {
weights.push(i.wrapping_mul(23).wrapping_add(b as u8).wrapping_add(seed));
}
}
weights
}
fn synth_q8_0_row(n_blocks: usize, seed: u8) -> Vec<u8> {
let mut weights = Vec::with_capacity(n_blocks * Q8_0_BLOCK_BYTES);
for b in 0..n_blocks {
weights.extend_from_slice(
&f16::from_f32(0.04 + (b as f32 + seed as f32) * 0.008).to_le_bytes(),
);
for i in 0..32u8 {
let q = ((i as i8)
.wrapping_mul(3)
.wrapping_add(seed as i8)
.wrapping_add(b as i8)) as u8;
weights.push(q);
}
}
weights
}
#[test]
fn pack_and_gemv_matches_scalar_row_dots() {
let n_blocks = 2;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let rows = 16; let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, r as u8));
}
let x: Vec<f32> = (0..cols)
.map(|i| ((i as f32) * 0.017 - 2.1).sin() * 1.8)
.collect();
let act = quantize_activations_q8_k(&x);
let mut reference = vec![0f32; rows];
let row_bytes = n_blocks * Q4_K_BLOCK_BYTES;
for r in 0..rows {
reference[r] = dot_q4_k_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
}
for &interleave in &[4usize, 8] {
let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
let n_groups = rows / Q4_KX8_NROWS;
let mut out = vec![0f32; rows];
gemv_q4_kx8_q8_k(&packed, &act, cols, n_groups, interleave, &mut out);
for r in 0..rows {
let err = (out[r] - reference[r]).abs();
let scale = reference[r].abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"interleave={interleave} row {r}: got {} want {} err={err}",
out[r],
reference[r]
);
}
}
}
#[test]
fn q4_0x4_pack_and_gemv_matches_scalar_row_dots() {
let n_blocks = 3;
let cols = n_blocks * Q4_0_BLOCK_ELEMS;
let rows = 12;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_0_row(n_blocks, r as u8));
}
let x: Vec<f32> = (0..cols)
.map(|i| ((i as f32) * 0.019 - 1.2).sin() * 2.1)
.collect();
let act = quantize_activations_q8(&x);
let row_bytes = n_blocks * Q4_0_BLOCK_BYTES;
let mut reference = vec![0f32; rows];
for r in 0..rows {
reference[r] = dot_q4_0_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
}
for &interleave in &[4usize, 8] {
let packed = pack_q4_0_matrix_x4(&matrix, rows, cols, interleave);
let n_groups = rows / Q4_0X4_NROWS;
let mut out = vec![0f32; rows];
gemv_q4_0x4_q8_0(&packed, &act, cols, n_groups, interleave, &mut out);
for r in 0..rows {
let err = (out[r] - reference[r]).abs();
let scale = reference[r].abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"Q4_0x4 interleave={interleave} row {r}: got {} want {} err={err}",
out[r],
reference[r]
);
}
}
}
#[test]
fn q4_0x4_gemm_matches_the_gemv_run_once_per_activation() {
let n_blocks = 4;
let cols = n_blocks * Q4_0_BLOCK_ELEMS;
let rows = 8;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 2 + 5) as u8));
}
let packed = pack_q4_0_matrix_x4(&matrix, rows, cols, Q4_0X4_INTERLEAVE);
let n_acts = 7;
let acts: Vec<Q8Activations> = (0..n_acts)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 11) as f32) * 0.021 - 0.7).cos() * 1.9)
.collect();
quantize_activations_q8(&x)
})
.collect();
for group in 0..rows / Q4_0X4_NROWS {
let mut gemm_out = vec![0f32; Q4_0X4_NROWS * n_acts];
gemm_q4_0x4_group(
&packed,
group,
&acts,
cols,
Q4_0X4_INTERLEAVE,
&mut gemm_out,
);
for (j, act) in acts.iter().enumerate() {
let mut gemv_out = [0f32; Q4_0X4_NROWS];
gemv_q4_0x4_group(&packed, group, act, cols, Q4_0X4_INTERLEAVE, &mut gemv_out);
for r in 0..Q4_0X4_NROWS {
assert_eq!(
gemm_out[r * n_acts + j],
gemv_out[r],
"group {group} row {r} act {j}: Q4_0 GEMM and GEMV disagree"
);
}
}
}
}
#[test]
fn q4_0x4_gemm_with_no_activations_is_a_no_op() {
let n_blocks = 2;
let cols = n_blocks * Q4_0_BLOCK_ELEMS;
let mut matrix = Vec::new();
for r in 0..Q4_0X4_NROWS {
matrix.extend_from_slice(&synth_q4_0_row(n_blocks, r as u8));
}
let packed = pack_q4_0_matrix_x4(&matrix, Q4_0X4_NROWS, cols, Q4_0X4_INTERLEAVE);
let mut out: Vec<f32> = Vec::new();
gemm_q4_0x4_group(&packed, 0, &[], cols, Q4_0X4_INTERLEAVE, &mut out);
assert!(out.is_empty());
}
#[test]
fn q8_0x4_pack_and_gemv_matches_scalar_row_dots() {
let n_blocks = 3;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
let rows = 12; let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q8_0_row(n_blocks, r as u8));
}
let x: Vec<f32> = (0..cols)
.map(|i| ((i as f32) * 0.023 - 1.4).cos() * 2.2)
.collect();
let act = quantize_activations_q8(&x);
let row_bytes = n_blocks * Q8_0_BLOCK_BYTES;
let mut reference = vec![0f32; rows];
for r in 0..rows {
reference[r] = dot_q8_0_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
}
for &interleave in &[4usize, 8] {
let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, interleave);
let n_groups = rows / Q8_0X4_NROWS;
let mut out = vec![0f32; rows];
gemv_q8_0x4_q8_0(&packed, &act, cols, n_groups, interleave, &mut out);
for r in 0..rows {
let err = (out[r] - reference[r]).abs();
let scale = reference[r].abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"Q8_0x4 interleave={interleave} row {r}: got {} want {} err={err}",
out[r],
reference[r]
);
}
}
}
#[test]
fn q8_0x4_gemm_matches_the_gemv_run_once_per_activation() {
let n_blocks = 4;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
let rows = 8;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 3 + 1) as u8));
}
let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, Q8_0X4_INTERLEAVE);
let n_acts = 7;
let acts: Vec<Q8Activations> = (0..n_acts)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 13) as f32) * 0.017 - 0.9).sin() * 1.7)
.collect();
quantize_activations_q8(&x)
})
.collect();
for group in 0..rows / Q8_0X4_NROWS {
let mut gemm_out = vec![0f32; Q8_0X4_NROWS * n_acts];
gemm_q8_0x4_group(
&packed,
group,
&acts,
cols,
Q8_0X4_INTERLEAVE,
&mut gemm_out,
);
for (j, act) in acts.iter().enumerate() {
let mut gemv_out = [0f32; Q8_0X4_NROWS];
gemv_q8_0x4_group(&packed, group, act, cols, Q8_0X4_INTERLEAVE, &mut gemv_out);
for r in 0..Q8_0X4_NROWS {
assert_eq!(
gemm_out[r * n_acts + j],
gemv_out[r],
"group {group} row {r} act {j}: GEMM and GEMV disagree"
);
}
}
}
}
#[test]
fn q4_kx8_gemm_matches_the_gemv_run_once_per_activation() {
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let rows = 2 * Q4_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 5 + 3) as u8));
}
let interleave = q4_kx8_interleave();
let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
let n_acts = 6;
let acts: Vec<Q8KActivations> = (0..n_acts)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 29) as f32) * 0.011 - 0.4).cos() * 2.3)
.collect();
quantize_activations_q8_k(&x)
})
.collect();
for group in 0..rows / Q4_KX8_NROWS {
for chunk in acts.chunks(Q4_KX8_GEMM_NC) {
let mut gemm_out = vec![0f32; Q4_KX8_NROWS * chunk.len()];
gemm_q4_kx8_group(&packed, group, chunk, cols, interleave, &mut gemm_out);
for (j, act) in chunk.iter().enumerate() {
let mut gemv_out = [0f32; Q4_KX8_NROWS];
gemv_q4_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
for r in 0..Q4_KX8_NROWS {
let got = gemm_out[r * chunk.len() + j];
let want = gemv_out[r];
if interleave == 4 {
assert_eq!(
got, want,
"group {group} row {r} act {j}: Q4_K GEMM and GEMV disagree"
);
} else {
let err = (got - want).abs();
let scale = want.abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-2,
"group {group} row {r} act {j}: GEMM {got} vs GEMV {want} (err={err})"
);
}
}
}
}
}
}
#[test]
#[cfg(target_arch = "aarch64")]
fn q4_kx8_gemm_i8mm_matches_scalar_when_available() {
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return;
}
let interleave = q4_kx8_interleave();
assert_eq!(interleave, 8, "i8mm host should pack with interleave 8");
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let rows = Q4_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 7 + 2) as u8));
}
let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
let n_acts = 4;
let acts: Vec<Q8KActivations> = (0..n_acts)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 17) as f32) * 0.013 - 0.6).sin() * 1.9)
.collect();
quantize_activations_q8_k(&x)
})
.collect();
let mut gemm_out = vec![0f32; Q4_KX8_NROWS * n_acts];
gemm_q4_kx8_group(&packed, 0, &acts, cols, interleave, &mut gemm_out);
for (j, act) in acts.iter().enumerate() {
let mut scalar_out = [0f32; Q4_KX8_NROWS];
gemv_q4_kx8_group(&packed, 0, act, cols, interleave, &mut scalar_out);
for r in 0..Q4_KX8_NROWS {
let got = gemm_out[r * n_acts + j];
let want = scalar_out[r];
let err = (got - want).abs();
let scale = want.abs().max(1.0);
assert!(
err / scale < 1e-5 || err < 1e-3,
"row {r} act {j}: i8mm GEMM {got} vs scalar {want} (err={err})"
);
}
}
}
fn synth_q8_k_acts(n: usize, cols: usize) -> Vec<Q8KActivations> {
(0..n)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 17) as f32) * 0.013 - 0.6).sin() * 1.9)
.collect();
quantize_activations_q8_k(&x)
})
.collect()
}
fn reference_q8_kx4_block_qs(
acts: &[Q8KActivations],
block: usize,
) -> [i8; Q4_K_BLOCK_ELEMS * 4] {
const BLCK: usize = 8;
let na = acts.len();
let mut out = [0i8; Q4_K_BLOCK_ELEMS * 4];
for (j, slot) in out.iter_mut().enumerate() {
let src_offset = (j / (4 * BLCK)) * BLCK + (j % BLCK);
let src_id = (j % (4 * BLCK)) / BLCK;
*slot = if src_id < na {
acts[src_id].q[block * Q4_K_BLOCK_ELEMS + src_offset]
} else {
0
};
}
out
}
#[test]
fn prepare_q8_k_acts_x4_matches_block_interleave_reference() {
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
for na in 1..=Q4_KX8_GEMM_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
assert_eq!(tile.na, na);
assert_eq!(tile.n_blocks, n_blocks);
for b in 0..n_blocks {
let want_qs = reference_q8_kx4_block_qs(&acts, b);
assert_eq!(
&tile.qs[b * Q4_K_BLOCK_ELEMS * 4..][..Q4_K_BLOCK_ELEMS * 4],
&want_qs[..],
"qs mismatch, block {b} na {na}"
);
for a in 0..4 {
let act = acts.get(a);
for i in 0..8 {
let want = act.map_or(0, |act| {
act.bsums[b * 16 + 2 * i] + act.bsums[b * 16 + 2 * i + 1]
});
assert_eq!(
tile.bsums[(b * 4 + a) * 8 + i],
want,
"bsums mismatch, block {b} row {a} pair {i} na {na}"
);
}
let want_d = act.map_or(0.0, |act| act.d[b]);
assert_eq!(tile.d[b * 4 + a], want_d, "d mismatch, block {b} row {a}");
}
}
}
}
#[test]
fn q4_kx8_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let rows = Q4_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 5 + 3) as u8));
}
let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, 8);
for na in 1..=Q4_KX8_GEMM_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
let mut got = vec![0f32; Q4_KX8_NROWS * na];
gemm_q4_kx8_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
for (j, act) in acts.iter().enumerate() {
let mut want = [0f32; Q4_KX8_NROWS];
gemv_q4_kx8_q8_k_scalar_8(&packed, act, cols, 1, &mut want);
for r in 0..Q4_KX8_NROWS {
assert_eq!(
got[r * na + j].to_bits(),
want[r].to_bits(),
"row {r} act {j} na {na}: x4 {} vs GEMV {}",
got[r * na + j],
want[r]
);
}
}
}
}
#[test]
#[cfg(target_arch = "aarch64")]
fn q4_kx8_gemm_x4_i8mm_matches_group_and_scalar() {
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return;
}
let interleave = q4_kx8_interleave();
assert_eq!(interleave, 8, "i8mm host should pack with interleave 8");
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let rows = Q4_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 7 + 2) as u8));
}
let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
assert!(q4_kx8_gemm_uses_acts_x4(interleave));
for na in 1..=Q4_KX8_GEMM_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
let mut x4_out = vec![0f32; Q4_KX8_NROWS * na];
gemm_q4_kx8_group_x4(&packed, 0, &tile, cols, interleave, &mut x4_out);
let mut group_out = vec![0f32; Q4_KX8_NROWS * na];
gemm_q4_kx8_group(&packed, 0, &acts, cols, interleave, &mut group_out);
assert_eq!(
x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"x4 entry diverged from the compat entry, na {na}"
);
for (j, act) in acts.iter().enumerate() {
let mut scalar_out = [0f32; Q4_KX8_NROWS];
gemv_q4_kx8_group(&packed, 0, act, cols, interleave, &mut scalar_out);
for r in 0..Q4_KX8_NROWS {
let got = x4_out[r * na + j];
let want = scalar_out[r];
let err = (got - want).abs();
let scale = want.abs().max(1.0);
assert!(
err / scale < 1e-5 || err < 1e-3,
"row {r} act {j} na {na}: i8mm x4 GEMM {got} vs scalar {want} (err={err})"
);
}
}
}
}
#[test]
fn accel_x4_only_picks_a_kernel_it_never_changes_the_answer() {
let na = 3;
let here = AccelX4::detect();
{
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
let mut q4 = Vec::new();
let mut q5 = Vec::new();
let mut q6 = Vec::new();
for r in 0..Q4_KX8_NROWS {
q4.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 7 + 2) as u8));
q5.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 3 + 5) as u8));
q6.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 11 + 1) as u8));
}
let q4 = pack_q4_k_matrix_x8(&q4, Q4_KX8_NROWS, cols, 8);
let q5 = pack_q5_k_matrix_x8(&q5, Q5_KX8_NROWS, cols, 8);
let q6 = pack_q6_k_matrix_x8(&q6, Q6_KX8_NROWS, cols, 8);
let mut wrapper = vec![0f32; Q4_KX8_NROWS * na];
let mut hoisted = vec![0f32; Q4_KX8_NROWS * na];
let mut portable = vec![0f32; Q4_KX8_NROWS * na];
let mut reference = vec![0f32; Q4_KX8_NROWS * na];
gemm_q4_kx8_group_x4(&q4, 0, &tile, cols, 8, &mut wrapper);
gemm_q4_kx8_group_x4_on(&q4, 0, &tile, cols, 8, here, &mut hoisted);
gemm_q4_kx8_group_x4_on(&q4, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
gemm_q4_kx8_acts_x4_scalar_8(&q4, &tile, cols, &mut reference);
assert_bits_eq("q4_k hoisted", &hoisted, &wrapper);
assert_bits_eq("q4_k portable", &portable, &reference);
gemm_q5_kx8_group_x4(&q5, 0, &tile, cols, 8, &mut wrapper);
gemm_q5_kx8_group_x4_on(&q5, 0, &tile, cols, 8, here, &mut hoisted);
gemm_q5_kx8_group_x4_on(&q5, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
gemm_q5_kx8_acts_x4_scalar_8(&q5, &tile, cols, &mut reference);
assert_bits_eq("q5_k hoisted", &hoisted, &wrapper);
assert_bits_eq("q5_k portable", &portable, &reference);
gemm_q6_kx8_group_x4(&q6, 0, &tile, cols, 8, &mut wrapper);
gemm_q6_kx8_group_x4_on(&q6, 0, &tile, cols, 8, here, &mut hoisted);
gemm_q6_kx8_group_x4_on(&q6, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
gemm_q6_kx8_acts_x4_scalar_8(&q6, &tile, cols, &mut reference);
assert_bits_eq("q6_k hoisted", &hoisted, &wrapper);
assert_bits_eq("q6_k portable", &portable, &reference);
}
{
let n_blocks = 4;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
let acts = synth_q8_0_acts(na, cols);
let tile = prepare_q8_acts_x4(&acts, cols);
let mut q8 = Vec::new();
let mut q4 = Vec::new();
for r in 0..Q8_0X4_NROWS {
q8.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 9 + 4) as u8));
q4.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 13 + 6) as u8));
}
let q8 = pack_q8_0_matrix_x4(&q8, Q8_0X4_NROWS, cols, 8);
let q4 = pack_q4_0_matrix_x4(&q4, Q4_0X4_NROWS, cols, 8);
let mut wrapper = vec![0f32; Q8_0X4_NROWS * na];
let mut hoisted = vec![0f32; Q8_0X4_NROWS * na];
let mut portable = vec![0f32; Q8_0X4_NROWS * na];
let mut reference = vec![0f32; Q8_0X4_NROWS * na];
gemm_q8_0x4_group_x4(&q8, 0, &tile, cols, 8, &mut wrapper);
gemm_q8_0x4_group_x4_on(&q8, 0, &tile, cols, 8, here, &mut hoisted);
gemm_q8_0x4_group_x4_on(&q8, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
gemm_q8_0x4_acts_x4_scalar_8(&q8, &tile, cols, &mut reference);
assert_bits_eq("q8_0 hoisted", &hoisted, &wrapper);
assert_bits_eq("q8_0 portable", &portable, &reference);
gemm_q4_0x4_group_x4(&q4, 0, &tile, cols, 8, &mut wrapper);
gemm_q4_0x4_group_x4_on(&q4, 0, &tile, cols, 8, here, &mut hoisted);
gemm_q4_0x4_group_x4_on(&q4, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
gemm_q4_0x4_acts_x4_scalar_8(&q4, &tile, cols, &mut reference);
assert_bits_eq("q4_0 hoisted", &hoisted, &wrapper);
assert_bits_eq("q4_0 portable", &portable, &reference);
}
}
fn assert_bits_eq(what: &str, got: &[f32], want: &[f32]) {
assert_eq!(
got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"{what}: {got:?} vs {want:?}"
);
}
#[test]
fn q5_kx8_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
let n_blocks = 3;
let cols = n_blocks * Q5_K_BLOCK_ELEMS;
let rows = Q5_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 5 + 3) as u8));
}
let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, 8);
for na in 1..=Q5_KX8_GEMM_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
let mut got = vec![0f32; Q5_KX8_NROWS * na];
gemm_q5_kx8_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
for (j, act) in acts.iter().enumerate() {
let mut want = [0f32; Q5_KX8_NROWS];
gemv_q5_kx8_q8_k_scalar_8(&packed, act, cols, 1, &mut want);
for r in 0..Q5_KX8_NROWS {
assert_eq!(
got[r * na + j].to_bits(),
want[r].to_bits(),
"row {r} act {j} na {na}: x4 {} vs GEMV {}",
got[r * na + j],
want[r]
);
}
}
}
}
#[test]
fn q6_kx8_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
let n_blocks = 2;
let cols = n_blocks * Q6_K_BLOCK_ELEMS;
let rows = Q6_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 3 + 2) as u8));
}
let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, 8);
for na in 1..=Q8K_ACTS_X4_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
let mut got = vec![0f32; Q6_KX8_NROWS * na];
gemm_q6_kx8_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
for (j, act) in acts.iter().enumerate() {
let mut want = [0f32; Q6_KX8_NROWS];
gemv_q6_kx8_q8_k_scalar(&packed, act, cols, 1, 8, &mut want);
for r in 0..Q6_KX8_NROWS {
assert_eq!(
got[r * na + j].to_bits(),
want[r].to_bits(),
"row {r} act {j} na {na}: x4 {} vs GEMV {}",
got[r * na + j],
want[r]
);
}
}
}
}
#[test]
#[cfg(target_arch = "aarch64")]
fn q4_kx8_interleave8_neon_gemv_matches_scalar_when_available() {
if !std::arch::is_aarch64_feature_detected!("dotprod") {
return;
}
let n_blocks = 3;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let n_groups = 2;
let rows = n_groups * Q4_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 11 + 5) as u8));
}
let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, 8);
let acts = synth_q8_k_acts(4, cols);
for act in &acts {
let mut got = vec![0f32; rows];
gemv_q4_kx8_q8_k(&packed, act, cols, n_groups, 8, &mut got);
let mut want = vec![0f32; rows];
gemv_q4_kx8_q8_k_scalar_8(&packed, act, cols, n_groups, &mut want);
for r in 0..rows {
let err = (got[r] - want[r]).abs();
assert!(
err / want[r].abs().max(1.0) < 5e-5 || err < 1e-3,
"gemv row {r}: NEON 8x8 {} vs scalar {}",
got[r],
want[r]
);
}
}
}
#[test]
#[cfg(target_arch = "aarch64")]
fn q5_kx8_interleave8_neon_matches_references_when_available() {
if !std::arch::is_aarch64_feature_detected!("dotprod") {
return;
}
let n_blocks = 3;
let cols = n_blocks * Q5_K_BLOCK_ELEMS;
let n_groups = 2;
let rows = n_groups * Q5_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 7 + 1) as u8));
}
let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, 8);
let acts = synth_q8_k_acts(4, cols);
for act in &acts {
let mut got = vec![0f32; rows];
gemv_q5_kx8_q8_k(&packed, act, cols, n_groups, 8, &mut got);
let mut want = vec![0f32; rows];
gemv_q5_kx8_q8_k_scalar_8(&packed, act, cols, n_groups, &mut want);
for r in 0..rows {
let err = (got[r] - want[r]).abs();
assert!(
err / want[r].abs().max(1.0) < 5e-5 || err < 1e-3,
"gemv row {r}: NEON 8x8 {} vs scalar {}",
got[r],
want[r]
);
}
}
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return;
}
for na in 1..=Q5_KX8_GEMM_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
for group in 0..n_groups {
let mut x4_out = vec![0f32; Q5_KX8_NROWS * na];
gemm_q5_kx8_group_x4(&packed, group, &tile, cols, 8, &mut x4_out);
let mut group_out = vec![0f32; Q5_KX8_NROWS * na];
gemm_q5_kx8_group(&packed, group, &acts, cols, 8, &mut group_out);
assert_eq!(
x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"x4 entry diverged from the compat entry, group {group} na {na}"
);
let nb = cols / Q5_K_BLOCK_ELEMS;
let slice = &packed[group * nb * Q5_KX8_BLOCK_BYTES..][..nb * Q5_KX8_BLOCK_BYTES];
let mut want = vec![0f32; Q5_KX8_NROWS * na];
gemm_q5_kx8_acts_x4_scalar_8(slice, &tile, cols, &mut want);
for (got, want) in x4_out.iter().zip(want.iter()) {
let err = (got - want).abs();
assert!(
err / want.abs().max(1.0) < 5e-5 || err < 1e-3,
"group {group} na {na}: i8mm GEMM {got} vs portable {want}"
);
}
}
}
}
#[test]
#[cfg(target_arch = "aarch64")]
fn q6_kx8_interleave8_neon_matches_references_when_available() {
if !std::arch::is_aarch64_feature_detected!("dotprod") {
return;
}
let n_blocks = 2;
let cols = n_blocks * Q6_K_BLOCK_ELEMS;
let n_groups = 2;
let rows = n_groups * Q6_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 9 + 4) as u8));
}
let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, 8);
let acts = synth_q8_k_acts(4, cols);
for act in &acts {
let mut got = vec![0f32; rows];
gemv_q6_kx8_q8_k(&packed, act, cols, n_groups, 8, &mut got);
let mut want = vec![0f32; rows];
gemv_q6_kx8_q8_k_scalar(&packed, act, cols, n_groups, 8, &mut want);
for r in 0..rows {
let err = (got[r] - want[r]).abs();
assert!(
err / want[r].abs().max(1.0) < 5e-5 || err < 1e-3,
"gemv row {r}: NEON 8x8 {} vs scalar {}",
got[r],
want[r]
);
}
}
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return;
}
for na in 1..=Q8K_ACTS_X4_NC {
let acts = synth_q8_k_acts(na, cols);
let tile = prepare_q8_k_acts_x4(&acts, cols);
for group in 0..n_groups {
let mut x4_out = vec![0f32; Q6_KX8_NROWS * na];
gemm_q6_kx8_group_x4(&packed, group, &tile, cols, 8, &mut x4_out);
let mut group_out = vec![0f32; Q6_KX8_NROWS * na];
gemm_q6_kx8_group(&packed, group, &acts, cols, 8, &mut group_out);
assert_eq!(
x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"x4 entry diverged from the compat entry, group {group} na {na}"
);
let nb = cols / Q6_K_BLOCK_ELEMS;
let slice = &packed[group * nb * Q6_KX8_BLOCK_BYTES..][..nb * Q6_KX8_BLOCK_BYTES];
let mut want = vec![0f32; Q6_KX8_NROWS * na];
gemm_q6_kx8_acts_x4_scalar_8(slice, &tile, cols, &mut want);
for (got, want) in x4_out.iter().zip(want.iter()) {
let err = (got - want).abs();
assert!(
err / want.abs().max(1.0) < 5e-5 || err < 1e-3,
"group {group} na {na}: i8mm GEMM {got} vs portable {want}"
);
}
}
}
}
fn synth_q8_0_acts(n: usize, cols: usize) -> Vec<Q8Activations> {
(0..n)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 13) as f32) * 0.021 - 0.9).sin() * 1.7)
.collect();
quantize_activations_q8(&x)
})
.collect()
}
#[test]
fn prepare_q8_acts_x4_matches_interleave_reference() {
let n_blocks = 3;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
for na in 1..=Q8K_ACTS_X4_NC {
let acts = synth_q8_0_acts(na, cols);
let tile = prepare_q8_acts_x4(&acts, cols);
assert_eq!(tile.na, na);
assert_eq!(tile.n_blocks, n_blocks);
for b in 0..n_blocks {
for (j, got) in tile.qs[b * 128..(b + 1) * 128].iter().enumerate() {
let src_offset = (j / 32) * 8 + (j % 8);
let src_id = (j % 32) / 8;
let want = if src_id < na {
acts[src_id].q[b * Q8_0_BLOCK_ELEMS + src_offset]
} else {
0
};
assert_eq!(*got, want, "qs mismatch, block {b} pos {j} na {na}");
}
for a in 0..4 {
let want_d = acts.get(a).map_or(0.0, |act| act.d[b]);
assert_eq!(tile.d[b * 4 + a], want_d, "d mismatch, block {b} row {a}");
}
}
}
}
#[test]
fn q8_0x4_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
let n_blocks = 3;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
let rows = Q8_0X4_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 7 + 3) as u8));
}
let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, 8);
for na in 1..=Q8K_ACTS_X4_NC {
let acts = synth_q8_0_acts(na, cols);
let tile = prepare_q8_acts_x4(&acts, cols);
let mut got = vec![0f32; Q8_0X4_NROWS * na];
gemm_q8_0x4_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
for (j, act) in acts.iter().enumerate() {
let mut want = [0f32; Q8_0X4_NROWS];
gemv_q8_0x4_q8_0_scalar(&packed, act, cols, 1, 8, &mut want);
for r in 0..Q8_0X4_NROWS {
assert_eq!(
got[r * na + j].to_bits(),
want[r].to_bits(),
"row {r} act {j} na {na}: x4 {} vs GEMV {}",
got[r * na + j],
want[r]
);
}
}
}
}
#[test]
fn q4_0x4_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
let n_blocks = 3;
let cols = n_blocks * Q4_0_BLOCK_ELEMS;
let rows = Q4_0X4_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 5 + 1) as u8));
}
let packed = pack_q4_0_matrix_x4(&matrix, rows, cols, 8);
for na in 1..=Q8K_ACTS_X4_NC {
let acts = synth_q8_0_acts(na, cols);
let tile = prepare_q8_acts_x4(&acts, cols);
let mut got = vec![0f32; Q4_0X4_NROWS * na];
gemm_q4_0x4_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
for (j, act) in acts.iter().enumerate() {
let mut want = [0f32; Q4_0X4_NROWS];
gemv_q4_0x4_q8_0_scalar(&packed, act, cols, 1, 8, &mut want);
for r in 0..Q4_0X4_NROWS {
assert_eq!(
got[r * na + j].to_bits(),
want[r].to_bits(),
"row {r} act {j} na {na}: x4 {} vs GEMV {}",
got[r * na + j],
want[r]
);
}
}
}
}
#[test]
#[cfg(target_arch = "aarch64")]
fn q8_0_q4_0_interleave8_neon_matches_references_when_available() {
if !std::arch::is_aarch64_feature_detected!("dotprod") {
return;
}
let n_blocks = 3;
let n_groups = 2;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
let rows = n_groups * Q8_0X4_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 3 + 2) as u8));
}
let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, 8);
let acts = synth_q8_0_acts(4, cols);
for act in &acts {
let mut got = vec![0f32; rows];
gemv_q8_0x4_q8_0(&packed, act, cols, n_groups, 8, &mut got);
let mut want = vec![0f32; rows];
gemv_q8_0x4_q8_0_scalar(&packed, act, cols, n_groups, 8, &mut want);
for r in 0..rows {
let err = (got[r] - want[r]).abs();
assert!(
err / want[r].abs().max(1.0) < 1e-5 || err < 1e-3,
"q8_0 gemv row {r}: NEON 4x8 {} vs scalar {}",
got[r],
want[r]
);
}
}
let q4_cols = n_blocks * Q4_0_BLOCK_ELEMS;
let q4_rows = n_groups * Q4_0X4_NROWS;
let mut q4_matrix = Vec::new();
for r in 0..q4_rows {
q4_matrix.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 9 + 1) as u8));
}
let q4_packed = pack_q4_0_matrix_x4(&q4_matrix, q4_rows, q4_cols, 8);
let q4_acts = synth_q8_0_acts(4, q4_cols);
for act in &q4_acts {
let mut got = vec![0f32; q4_rows];
gemv_q4_0x4_q8_0(&q4_packed, act, q4_cols, n_groups, 8, &mut got);
let mut want = vec![0f32; q4_rows];
gemv_q4_0x4_q8_0_scalar(&q4_packed, act, q4_cols, n_groups, 8, &mut want);
for r in 0..q4_rows {
let err = (got[r] - want[r]).abs();
assert!(
err / want[r].abs().max(1.0) < 1e-5 || err < 1e-3,
"q4_0 gemv row {r}: NEON 4x8 {} vs scalar {}",
got[r],
want[r]
);
}
}
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return;
}
for na in 1..=Q8K_ACTS_X4_NC {
let acts = synth_q8_0_acts(na, cols);
let tile = prepare_q8_acts_x4(&acts, cols);
for group in 0..n_groups {
let mut x4_out = vec![0f32; Q8_0X4_NROWS * na];
gemm_q8_0x4_group_x4(&packed, group, &tile, cols, 8, &mut x4_out);
let mut group_out = vec![0f32; Q8_0X4_NROWS * na];
gemm_q8_0x4_group(&packed, group, &acts, cols, 8, &mut group_out);
assert_eq!(
x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"q8_0 x4 entry diverged from the compat entry, group {group} na {na}"
);
let slice = &packed[group * n_blocks * Q8_0X4_BLOCK_BYTES..]
[..n_blocks * Q8_0X4_BLOCK_BYTES];
let mut want = vec![0f32; Q8_0X4_NROWS * na];
gemm_q8_0x4_acts_x4_scalar_8(slice, &tile, cols, &mut want);
for (got, want) in x4_out.iter().zip(want.iter()) {
let err = (got - want).abs();
assert!(
err / want.abs().max(1.0) < 1e-5 || err < 1e-3,
"q8_0 group {group} na {na}: i8mm GEMM {got} vs portable {want}"
);
}
}
let q4_acts = synth_q8_0_acts(na, q4_cols);
let q4_tile = prepare_q8_acts_x4(&q4_acts, q4_cols);
for group in 0..n_groups {
let mut x4_out = vec![0f32; Q4_0X4_NROWS * na];
gemm_q4_0x4_group_x4(&q4_packed, group, &q4_tile, q4_cols, 8, &mut x4_out);
let mut group_out = vec![0f32; Q4_0X4_NROWS * na];
gemm_q4_0x4_group(&q4_packed, group, &q4_acts, q4_cols, 8, &mut group_out);
assert_eq!(
x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"q4_0 x4 entry diverged from the compat entry, group {group} na {na}"
);
let slice = &q4_packed[group * n_blocks * Q4_0X4_BLOCK_BYTES..]
[..n_blocks * Q4_0X4_BLOCK_BYTES];
let mut want = vec![0f32; Q4_0X4_NROWS * na];
gemm_q4_0x4_acts_x4_scalar_8(slice, &q4_tile, q4_cols, &mut want);
for (got, want) in x4_out.iter().zip(want.iter()) {
let err = (got - want).abs();
assert!(
err / want.abs().max(1.0) < 1e-5 || err < 1e-3,
"q4_0 group {group} na {na}: i8mm GEMM {got} vs portable {want}"
);
}
}
}
}
#[test]
fn q4_kx8_gemm_with_no_activations_is_a_no_op() {
let n_blocks = 2;
let cols = n_blocks * Q4_K_BLOCK_ELEMS;
let mut matrix = Vec::new();
for r in 0..Q4_KX8_NROWS {
matrix.extend_from_slice(&synth_q4_k_row(n_blocks, r as u8));
}
let packed = pack_q4_k_matrix_x8(&matrix, Q4_KX8_NROWS, cols, 4);
let mut out: Vec<f32> = Vec::new();
gemm_q4_kx8_group(&packed, 0, &[], cols, 4, &mut out);
assert!(out.is_empty());
}
#[test]
fn q8_0x4_gemm_with_no_activations_is_a_no_op() {
let n_blocks = 2;
let cols = n_blocks * Q8_0_BLOCK_ELEMS;
let mut matrix = Vec::new();
for r in 0..Q8_0X4_NROWS {
matrix.extend_from_slice(&synth_q8_0_row(n_blocks, r as u8));
}
let packed = pack_q8_0_matrix_x4(&matrix, Q8_0X4_NROWS, cols, Q8_0X4_INTERLEAVE);
let mut out: Vec<f32> = Vec::new();
gemm_q8_0x4_group(&packed, 0, &[], cols, Q8_0X4_INTERLEAVE, &mut out);
assert!(out.is_empty());
}
#[test]
fn q5_kx8_pack_and_gemv_matches_scalar_row_dots() {
let n_blocks = 2;
let cols = n_blocks * Q5_K_BLOCK_ELEMS;
let rows = 16;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q5_k_row(n_blocks, r as u8));
}
let x: Vec<f32> = (0..cols)
.map(|i| ((i as f32) * 0.019 - 1.8).sin() * 1.6)
.collect();
let act = quantize_activations_q8_k(&x);
let row_bytes = n_blocks * Q5_K_BLOCK_BYTES;
let mut reference = vec![0f32; rows];
for r in 0..rows {
reference[r] = dot_q5_k_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
}
for &interleave in &[4usize, 8] {
let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, interleave);
let n_groups = rows / Q5_KX8_NROWS;
let mut out = vec![0f32; rows];
gemv_q5_kx8_q8_k(&packed, &act, cols, n_groups, interleave, &mut out);
for r in 0..rows {
let err = (out[r] - reference[r]).abs();
let scale = reference[r].abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"interleave={interleave} row {r}: got {} want {} err={err}",
out[r],
reference[r]
);
}
}
}
#[test]
fn q5_kx8_gemm_matches_the_gemv_run_once_per_activation() {
let n_blocks = 3;
let cols = n_blocks * Q5_K_BLOCK_ELEMS;
let rows = 2 * Q5_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 5 + 3) as u8));
}
let interleave = q5_kx8_interleave();
let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, interleave);
let n_acts = 6;
let acts: Vec<Q8KActivations> = (0..n_acts)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 29) as f32) * 0.011 - 0.4).cos() * 2.3)
.collect();
quantize_activations_q8_k(&x)
})
.collect();
let row_bytes = n_blocks * Q5_K_BLOCK_BYTES;
for group in 0..rows / Q5_KX8_NROWS {
for chunk in acts.chunks(Q5_KX8_GEMM_NC) {
let mut gemm_out = vec![0f32; Q5_KX8_NROWS * chunk.len()];
gemm_q5_kx8_group(&packed, group, chunk, cols, interleave, &mut gemm_out);
for (j, act) in chunk.iter().enumerate() {
let mut gemv_out = [0f32; Q5_KX8_NROWS];
gemv_q5_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
for r in 0..Q5_KX8_NROWS {
let got = gemm_out[r * chunk.len() + j];
let want = gemv_out[r];
let err = (got - want).abs();
let scale = want.abs().max(1.0);
assert!(
err / scale < 5e-5 || err < 1e-3,
"group {group} row {r} act {j}: Q5_K GEMM {got} vs GEMV {want}"
);
}
}
}
for (j, act) in acts.iter().enumerate() {
for r in 0..Q5_KX8_NROWS {
let row_idx = group * Q5_KX8_NROWS + r;
let row = &matrix[row_idx * row_bytes..(row_idx + 1) * row_bytes];
let want = dot_q5_k_q8_scalar(row, act);
let mut gemv_out = [0f32; Q5_KX8_NROWS];
gemv_q5_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
let err = (gemv_out[r] - want).abs();
let scale = want.abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"group {group} row {r} act {j}: packed gemv {} vs dot {want}",
gemv_out[r]
);
}
}
}
}
#[test]
fn q5_kx8_gemm_with_no_activations_is_a_no_op() {
let n_blocks = 2;
let cols = n_blocks * Q5_K_BLOCK_ELEMS;
let mut matrix = Vec::new();
for r in 0..Q5_KX8_NROWS {
matrix.extend_from_slice(&synth_q5_k_row(n_blocks, r as u8));
}
let packed = pack_q5_k_matrix_x8(&matrix, Q5_KX8_NROWS, cols, 4);
let mut out: Vec<f32> = Vec::new();
gemm_q5_kx8_group(&packed, 0, &[], cols, 4, &mut out);
assert!(out.is_empty());
}
#[test]
fn q6_kx8_pack_and_gemv_matches_scalar_row_dots() {
let n_blocks = 3;
let cols = n_blocks * Q6_K_BLOCK_ELEMS;
let rows = 2 * Q6_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 5 + 1) as u8));
}
let x: Vec<f32> = (0..cols)
.map(|i| ((i as f32) * 0.017 - 0.8).cos() * 1.8)
.collect();
let act = quantize_activations_q8_k(&x);
let row_bytes = n_blocks * Q6_K_BLOCK_BYTES;
let mut reference = vec![0f32; rows];
for r in 0..rows {
reference[r] = dot_q6_k_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
}
for interleave in [4usize, 8] {
let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, interleave);
let n_groups = rows / Q6_KX8_NROWS;
let mut out = vec![0f32; rows];
gemv_q6_kx8_q8_k(&packed, &act, cols, n_groups, interleave, &mut out);
for r in 0..rows {
let err = (out[r] - reference[r]).abs();
let scale = reference[r].abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"interleave={interleave} row {r}: got {} want {} err={err}",
out[r],
reference[r]
);
}
}
}
#[test]
fn q6_kx8_gemm_matches_the_gemv_run_once_per_activation() {
let n_blocks = 2;
let cols = n_blocks * Q6_K_BLOCK_ELEMS;
let rows = 2 * Q6_KX8_NROWS;
let mut matrix = Vec::new();
for r in 0..rows {
matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 3 + 2) as u8));
}
let interleave = q6_kx8_interleave();
let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, interleave);
let acts: Vec<_> = (0..Q6_KX8_GEMM_NC)
.map(|j| {
let x: Vec<f32> = (0..cols)
.map(|i| (((i + j * 11) as f32) * 0.015 - 0.7).sin() * 2.0)
.collect();
quantize_activations_q8_k(&x)
})
.collect();
let row_bytes = n_blocks * Q6_K_BLOCK_BYTES;
for group in 0..rows / Q6_KX8_NROWS {
for chunk in acts.chunks(Q6_KX8_GEMM_NC) {
let mut gemm_out = vec![0f32; Q6_KX8_NROWS * chunk.len()];
gemm_q6_kx8_group(&packed, group, chunk, cols, interleave, &mut gemm_out);
for (j, act) in chunk.iter().enumerate() {
let mut gemv_out = [0f32; Q6_KX8_NROWS];
gemv_q6_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
for r in 0..Q6_KX8_NROWS {
let got = gemm_out[r * chunk.len() + j];
let want = gemv_out[r];
let err = (got - want).abs();
let scale = want.abs().max(1.0);
assert!(
err / scale < 1e-4 || err < 1e-3,
"group {group} row {r} act {j}: gemm {got} vs gemv {want}"
);
let row_idx = group * Q6_KX8_NROWS + r;
let row = &matrix[row_idx * row_bytes..(row_idx + 1) * row_bytes];
let dot = dot_q6_k_q8_scalar(row, act);
let err2 = (got - dot).abs();
let scale2 = dot.abs().max(1.0);
assert!(
err2 / scale2 < 1e-4 || err2 < 1e-3,
"group {group} row {r} act {j}: gemm {got} vs dot {dot}"
);
}
}
}
}
}
#[test]
fn block_size_matches_ggml() {
assert_eq!(Q4_KX8_BLOCK_BYTES, 16 + 16 + 96 + 1024);
assert_eq!(Q5_KX8_BLOCK_BYTES, 16 + 16 + 96 + 256 + 1024);
assert_eq!(Q6_KX8_BLOCK_BYTES, 16 + 128 + 1024 + 512);
assert_eq!(Q8_0X4_BLOCK_BYTES, 4 * 2 + Q8_0_BLOCK_ELEMS * Q8_0X4_NROWS);
assert_eq!(Q4_0X4_BLOCK_BYTES, 4 * 2 + Q4_0_BLOCK_ELEMS * 2);
}
}