#![allow(unsafe_code, unsafe_op_in_unsafe_fn)]
#[cfg(target_feature = "simd128")]
use core::arch::wasm32::*;
#[must_use]
pub const fn simd128_enabled() -> bool {
cfg!(target_feature = "simd128")
}
#[cfg(target_feature = "simd128")]
#[inline(always)]
fn hsum_i32x4(v: v128) -> i32 {
let s = i32x4_add(v, i32x4_shuffle::<2, 3, 0, 1>(v, v));
let s = i32x4_add(s, i32x4_shuffle::<1, 0, 3, 2>(s, s));
i32x4_extract_lane::<0>(s)
}
#[cfg(target_feature = "simd128")]
#[inline]
unsafe fn dot_s8s8(a: &[i8], b: &[i8], k: usize) -> i32 {
let ap = a.as_ptr();
let bp = b.as_ptr();
let mut acc0 = i32x4_splat(0);
let mut acc1 = i32x4_splat(0);
let mut acc2 = i32x4_splat(0);
let mut acc3 = i32x4_splat(0);
let mut p = 0usize;
while p + 64 <= k {
for (i, acc) in [&mut acc0, &mut acc1, &mut acc2, &mut acc3]
.into_iter()
.enumerate()
{
let off = p + i * 16;
let va = v128_load(ap.add(off).cast());
let vb = v128_load(bp.add(off).cast());
let lo = i32x4_dot_i16x8(i16x8_extend_low_i8x16(va), i16x8_extend_low_i8x16(vb));
let hi = i32x4_dot_i16x8(i16x8_extend_high_i8x16(va), i16x8_extend_high_i8x16(vb));
*acc = i32x4_add(*acc, i32x4_add(lo, hi));
}
p += 64;
}
while p + 16 <= k {
let va = v128_load(ap.add(p).cast());
let vb = v128_load(bp.add(p).cast());
let lo = i32x4_dot_i16x8(i16x8_extend_low_i8x16(va), i16x8_extend_low_i8x16(vb));
let hi = i32x4_dot_i16x8(i16x8_extend_high_i8x16(va), i16x8_extend_high_i8x16(vb));
acc0 = i32x4_add(acc0, i32x4_add(lo, hi));
p += 16;
}
let acc = i32x4_add(i32x4_add(acc0, acc1), i32x4_add(acc2, acc3));
let mut sum = hsum_i32x4(acc);
while p < k {
sum += i32::from(a[p]) * i32::from(b[p]);
p += 1;
}
sum
}
#[cfg(target_feature = "simd128")]
#[inline]
unsafe fn dot_u8s8(a: &[u8], b: &[i8], k: usize) -> i32 {
let ap = a.as_ptr();
let bp = b.as_ptr();
let mut acc0 = i32x4_splat(0);
let mut acc1 = i32x4_splat(0);
let mut acc2 = i32x4_splat(0);
let mut acc3 = i32x4_splat(0);
let mut p = 0usize;
while p + 64 <= k {
for (i, acc) in [&mut acc0, &mut acc1, &mut acc2, &mut acc3]
.into_iter()
.enumerate()
{
let off = p + i * 16;
let va = v128_load(ap.add(off).cast());
let vb = v128_load(bp.add(off).cast());
let lo = i32x4_dot_i16x8(u16x8_extend_low_u8x16(va), i16x8_extend_low_i8x16(vb));
let hi = i32x4_dot_i16x8(u16x8_extend_high_u8x16(va), i16x8_extend_high_i8x16(vb));
*acc = i32x4_add(*acc, i32x4_add(lo, hi));
}
p += 64;
}
while p + 16 <= k {
let va = v128_load(ap.add(p).cast());
let vb = v128_load(bp.add(p).cast());
let lo = i32x4_dot_i16x8(u16x8_extend_low_u8x16(va), i16x8_extend_low_i8x16(vb));
let hi = i32x4_dot_i16x8(u16x8_extend_high_u8x16(va), i16x8_extend_high_i8x16(vb));
acc0 = i32x4_add(acc0, i32x4_add(lo, hi));
p += 16;
}
let acc = i32x4_add(i32x4_add(acc0, acc1), i32x4_add(acc2, acc3));
let mut sum = hsum_i32x4(acc);
while p < k {
sum += i32::from(a[p]) * i32::from(b[p]);
p += 1;
}
sum
}
pub fn igemm_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
#[cfg(not(target_feature = "simd128"))]
{
super::scalar::igemm_s8s8(a, b, m, k, n, out);
}
#[cfg(target_feature = "simd128")]
{
super::scalar::assert_gemm_shapes("igemm_s8s8", a.len(), b.len(), out.len(), m, k, n);
for i in 0..m {
let a_row = &a[i * k..i * k + k];
let out_row = &mut out[i * n..i * n + n];
for o in 0..n {
let b_row = &b[o * k..o * k + k];
out_row[o] += unsafe { dot_s8s8(a_row, b_row, k) };
}
}
}
}
pub fn igemm_u8s8(a: &[u8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
#[cfg(not(target_feature = "simd128"))]
{
super::scalar::igemm_u8s8(a, b, m, k, n, out);
}
#[cfg(target_feature = "simd128")]
{
super::scalar::assert_gemm_shapes("igemm_u8s8", a.len(), b.len(), out.len(), m, k, n);
for i in 0..m {
let a_row = &a[i * k..i * k + k];
let out_row = &mut out[i * n..i * n + n];
for o in 0..n {
let b_row = &b[o * k..o * k + k];
out_row[o] += unsafe { dot_u8s8(a_row, b_row, k) };
}
}
}
}
#[cfg(target_feature = "simd128")]
#[inline]
unsafe fn dot_s4s8_block32(a: *const i8, w: *const u8) -> (i32, i32) {
let wv = v128_load(w.cast());
let lo_nib = i8x16_shr(i8x16_shl(wv, 4), 4); let hi_nib = i8x16_shr(wv, 4); let a0 = v128_load(a.cast());
let a1 = v128_load(a.add(16).cast());
let a_even = i8x16_shuffle::<0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 28, 30>(a0, a1);
let a_odd = i8x16_shuffle::<1, 3, 5, 7, 9, 11, 13, 15, 17, 19, 21, 23, 25, 27, 29, 31>(a0, a1);
let first = i32x4_add(
i32x4_dot_i16x8(
i16x8_extend_low_i8x16(a_even),
i16x8_extend_low_i8x16(lo_nib),
),
i32x4_dot_i16x8(
i16x8_extend_low_i8x16(a_odd),
i16x8_extend_low_i8x16(hi_nib),
),
);
let second = i32x4_add(
i32x4_dot_i16x8(
i16x8_extend_high_i8x16(a_even),
i16x8_extend_high_i8x16(lo_nib),
),
i32x4_dot_i16x8(
i16x8_extend_high_i8x16(a_odd),
i16x8_extend_high_i8x16(hi_nib),
),
);
(hsum_i32x4(first), hsum_i32x4(second))
}
#[allow(clippy::too_many_arguments)]
pub fn igemm_s4s8_packed(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
#[cfg(not(target_feature = "simd128"))]
{
super::int4::s4s8_packed_kernel_scalar(a, b_packed, scales, group, m, k, n, out);
}
#[cfg(target_feature = "simd128")]
{
let groups = k / group;
let kbytes = k / 2;
let blocks = k / 32;
let vec_groups = blocks * (32 / group);
for mi in 0..m {
let a_row = &a[mi * k..mi * k + k];
for ni in 0..n {
let wbase = ni * kbytes;
let sbase = ni * groups;
let mut acc_f = 0.0f32;
for blk in 0..blocks {
let kk = blk * 32;
let (d0, d1) = unsafe {
dot_s4s8_block32(
a_row.as_ptr().add(kk),
b_packed.as_ptr().add(wbase + kk / 2),
)
};
if group == 32 {
acc_f += scales[sbase + blk] * (d0 + d1) as f32;
} else {
let g = blk * 2;
acc_f += scales[sbase + g] * d0 as f32;
acc_f += scales[sbase + g + 1] * d1 as f32;
}
}
for g in vec_groups..groups {
let lo = g * group;
let mut acc_i: i32 = 0;
for kk in lo..lo + group {
let byte = b_packed[wbase + kk / 2];
let w = if kk & 1 == 0 {
super::int4::sign_extend_nibble(byte & 0x0F)
} else {
super::int4::sign_extend_nibble(byte >> 4)
};
acc_i += i32::from(a_row[kk]) * i32::from(w);
}
acc_f += scales[sbase + g] * acc_i as f32;
}
out[mi * n + ni] = acc_f;
}
}
}
}