use crate::{Q8Activations, Q8KActivations, Q4_K_BLOCK_ELEMS, Q8_0_BLOCK_ELEMS};
use half::f16;
pub(crate) const KMASK1: u32 = 0x3f3f_3f3f;
pub(crate) const KMASK2: u32 = 0x0f0f_0f0f;
pub(crate) const KMASK3: u32 = 0x0303_0303;
#[inline]
pub(crate) 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,
Avx2,
Portable,
}
impl AccelX4 {
#[inline]
pub fn detect() -> Self {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("i8mm") {
return AccelX4::NeonI8mm;
}
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return AccelX4::Avx2;
}
}
AccelX4::Portable
}
#[inline]
pub fn is_simd(self) -> bool {
match self {
AccelX4::NeonI8mm | AccelX4::Avx2 => true,
AccelX4::Portable => false,
}
}
}
#[inline]
pub fn interleaved_gemm_is_accelerated(interleave: usize) -> bool {
interleave == 8 && AccelX4::detect().is_simd()
}
#[inline]
pub fn preferred_interleave() -> usize {
if AccelX4::detect().is_simd() {
8
} else {
4
}
}
#[inline]
pub(crate) 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]);
}
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);
debug_assert!(
!act.q.contains(&i8::MIN),
"the AVX2 Q8_0 GEMM negates the activation with `_mm256_sign_epi8`, \
and -128 negates to itself; every ggml quantizer clamps to +-127"
);
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,
}
}
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,
}
}