use crate::quant::int4::VALID_GROUP_SIZES;
#[inline]
#[must_use]
pub fn sign_extend_nibble(nib: u8) -> i8 {
((nib << 4) as i8) >> 4
}
#[inline]
#[must_use]
pub fn unpack_byte(b: u8) -> (i8, i8) {
(sign_extend_nibble(b & 0x0F), sign_extend_nibble(b >> 4))
}
#[must_use]
pub fn unpack_to_i8(b_packed: &[u8], n: usize, k: usize) -> Vec<i8> {
assert!(
k.is_multiple_of(2),
"unpack_to_i8: k {k} must be even (two nibbles/byte)"
);
let packed_len = super::scalar::checked_len("unpack_to_i8", n, k / 2, "n*k/2");
let out_len = super::scalar::checked_len("unpack_to_i8", n, k, "n*k");
assert_eq!(
b_packed.len(),
packed_len,
"unpack_to_i8: packed len {} != n*k/2 {}",
b_packed.len(),
packed_len
);
let mut out = vec![0i8; out_len];
accel::unpack_nibbles(b_packed, &mut out);
out
}
#[allow(clippy::too_many_arguments)]
pub fn igemm_s4s8(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
assert!(k.is_multiple_of(2), "igemm_s4s8: k {k} must be even");
assert!(
VALID_GROUP_SIZES.contains(&group),
"igemm_s4s8: group {group} must be 16 or 32"
);
assert!(
k.is_multiple_of(group),
"igemm_s4s8: group {group} must divide k {k}"
);
let groups = k / group;
let a_len = super::scalar::checked_len("igemm_s4s8", m, k, "m*k");
let packed_len = super::scalar::checked_len("igemm_s4s8", n, k / 2, "n*k/2");
let scales_len = super::scalar::checked_len("igemm_s4s8", n, groups, "n*(k/group)");
let out_len = super::scalar::checked_len("igemm_s4s8", m, n, "m*n");
assert_eq!(
a.len(),
a_len,
"igemm_s4s8: a len {} != m*k {}",
a.len(),
a_len
);
assert_eq!(
b_packed.len(),
packed_len,
"igemm_s4s8: b_packed len {} != n*k/2 {}",
b_packed.len(),
packed_len
);
assert_eq!(
scales.len(),
scales_len,
"igemm_s4s8: scales len {} != n*(k/group) {}",
scales.len(),
scales_len
);
assert_eq!(
out.len(),
out_len,
"igemm_s4s8: out len {} != m*n {}",
out.len(),
out_len
);
let b_i8 = unpack_to_i8(b_packed, n, k);
igemm_s4s8_unpacked(a, &b_i8, scales, group, m, k, n, out);
}
#[allow(clippy::too_many_arguments)]
fn igemm_s4s8_unpacked(
a: &[i8],
b_i8: &[i8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
let groups = k / group;
for mi in 0..m {
let a_row = &a[mi * k..(mi + 1) * k];
for ni in 0..n {
let b_row = &b_i8[ni * k..(ni + 1) * k];
let scale_row = &scales[ni * groups..(ni + 1) * groups];
let mut acc_f = 0.0f32;
for (g, &s) in scale_row.iter().enumerate() {
let lo = g * group;
let hi = lo + group;
let mut acc_i: i32 = 0;
for kk in lo..hi {
acc_i += i32::from(a_row[kk]) * i32::from(b_row[kk]);
}
acc_f += s * acc_i as f32;
}
out[mi * n + ni] = acc_f;
}
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn s4s8_packed_kernel_scalar(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
let groups = k / group;
let kbytes = k / 2;
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 g in 0..groups {
let lo = g * group;
let hi = lo + group;
let mut acc_i: i32 = 0;
for kk in lo..hi {
let byte = b_packed[wbase + kk / 2];
let w = if kk & 1 == 0 {
sign_extend_nibble(byte & 0x0F)
} else {
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;
}
}
}
#[allow(clippy::too_many_arguments)]
fn assert_packed_shapes(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &[f32],
) {
assert!(k.is_multiple_of(2), "igemm_s4s8_packed: k {k} must be even");
assert!(
VALID_GROUP_SIZES.contains(&group),
"igemm_s4s8_packed: group {group} must be 16 or 32"
);
assert!(
k.is_multiple_of(group),
"igemm_s4s8_packed: group {group} must divide k {k}"
);
let groups = k / group;
let a_len = super::scalar::checked_len("igemm_s4s8_packed", m, k, "m*k");
let packed_len = super::scalar::checked_len("igemm_s4s8_packed", n, k / 2, "n*k/2");
let scales_len = super::scalar::checked_len("igemm_s4s8_packed", n, groups, "n*(k/group)");
let out_len = super::scalar::checked_len("igemm_s4s8_packed", m, n, "m*n");
assert_eq!(
a.len(),
a_len,
"igemm_s4s8_packed: a len {} != m*k {}",
a.len(),
a_len
);
assert_eq!(
b_packed.len(),
packed_len,
"igemm_s4s8_packed: b_packed len {} != n*k/2 {}",
b_packed.len(),
packed_len
);
assert_eq!(
scales.len(),
scales_len,
"igemm_s4s8_packed: scales len {} != n*(k/group) {}",
scales.len(),
scales_len
);
assert_eq!(
out.len(),
out_len,
"igemm_s4s8_packed: out len {} != m*n {}",
out.len(),
out_len
);
}
#[allow(clippy::too_many_arguments)]
pub fn igemm_s4s8_packed_scalar(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
assert_packed_shapes(a, b_packed, scales, group, m, k, n, out);
s4s8_packed_kernel_scalar(a, b_packed, scales, group, m, k, n, out);
}
#[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],
) {
assert_packed_shapes(a, b_packed, scales, group, m, k, n, out);
#[cfg(target_arch = "aarch64")]
{
super::arm::igemm_s4s8_packed(a, b_packed, scales, group, m, k, n, out);
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
super::wasm128::igemm_s4s8_packed(a, b_packed, scales, group, m, k, n, out);
}
#[cfg(not(any(
target_arch = "aarch64",
all(target_arch = "wasm32", target_feature = "simd128")
)))]
{
s4s8_packed_kernel_scalar(a, b_packed, scales, group, m, k, n, out);
}
}
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
mod accel {
#[inline]
pub(super) fn unpack_nibbles_scalar(src: &[u8], out: &mut [i8]) {
debug_assert_eq!(out.len(), src.len() * 2, "unpack: out must be 2x src");
for (j, &b) in src.iter().enumerate() {
out[2 * j] = ((b << 4) as i8) >> 4; out[2 * j + 1] = (b as i8) >> 4; }
}
#[inline]
pub(super) fn unpack_nibbles(src: &[u8], out: &mut [i8]) {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe {
return unpack_nibbles_neon(src, out);
}
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") {
unsafe {
return unpack_nibbles_avx2(src, out);
}
}
}
unpack_nibbles_scalar(src, out);
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn unpack_nibbles_neon(src: &[u8], out: &mut [i8]) {
use std::arch::aarch64::*;
debug_assert_eq!(out.len(), src.len() * 2);
let blocks = src.len() / 16;
let lo_mask = vdupq_n_u8(0x0F);
for blk in 0..blocks {
let s_off = blk * 16;
let o_off = blk * 32;
let v = vld1q_u8(src.as_ptr().add(s_off));
let lo_u = vandq_u8(v, lo_mask);
let lo_s = vshrq_n_s8(vshlq_n_s8(vreinterpretq_s8_u8(lo_u), 4), 4);
let hi_u = vshrq_n_u8(v, 4);
let hi_s = vshrq_n_s8(vshlq_n_s8(vreinterpretq_s8_u8(hi_u), 4), 4);
let pair = int8x16x2_t(lo_s, hi_s);
vst2q_s8(out.as_mut_ptr().add(o_off), pair);
}
let done = blocks * 16;
if done < src.len() {
unpack_nibbles_scalar(&src[done..], &mut out[done * 2..]);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn unpack_nibbles_avx2(src: &[u8], out: &mut [i8]) {
use std::arch::x86_64::*;
debug_assert_eq!(out.len(), src.len() * 2);
let blocks = src.len() / 32;
let lo_mask = _mm256_set1_epi8(0x0F);
let bias = _mm256_set1_epi8(0x08);
for blk in 0..blocks {
let s_off = blk * 32;
let o_off = blk * 64;
let v = _mm256_loadu_si256(src.as_ptr().add(s_off) as *const __m256i);
let lo_n = _mm256_and_si256(v, lo_mask);
let lo_s = _mm256_sub_epi8(_mm256_xor_si256(lo_n, bias), bias);
let hi_n = _mm256_and_si256(_mm256_srli_epi16(v, 4), lo_mask);
let hi_s = _mm256_sub_epi8(_mm256_xor_si256(hi_n, bias), bias);
let il = _mm256_unpacklo_epi8(lo_s, hi_s); let ih = _mm256_unpackhi_epi8(lo_s, hi_s); let out0 = _mm256_permute2x128_si256(il, ih, 0x20);
let out1 = _mm256_permute2x128_si256(il, ih, 0x31);
_mm256_storeu_si256(out.as_mut_ptr().add(o_off) as *mut __m256i, out0);
_mm256_storeu_si256(out.as_mut_ptr().add(o_off + 32) as *mut __m256i, out1);
}
let done = blocks * 32;
if done < src.len() {
unpack_nibbles_scalar(&src[done..], &mut out[done * 2..]);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sign_extend_covers_full_int4_domain() {
let expected: [i8; 16] = [0, 1, 2, 3, 4, 5, 6, 7, -8, -7, -6, -5, -4, -3, -2, -1];
for nib in 0u8..16 {
assert_eq!(
sign_extend_nibble(nib),
expected[nib as usize],
"nibble 0x{nib:X} sign-extends wrong"
);
}
}
#[test]
fn unpack_byte_matches_pinned_low_then_high() {
assert_eq!(unpack_byte(0x21), (1, 2));
assert_eq!(unpack_byte(0x43), (3, 4));
assert_eq!(unpack_byte(0x87), (7, -8));
assert_eq!(unpack_byte(0xF0), (0, -1));
assert_eq!(unpack_byte(0x00), (0, 0));
assert_eq!(unpack_byte(0x88), (-8, -8)); }
#[test]
fn unpack_to_i8_reproduces_quant_core_fixture() {
let unpacked = unpack_to_i8(&[0x21, 0x43], 1, 4);
assert_eq!(unpacked, vec![1i8, 2, 3, 4]);
let two_rows = unpack_to_i8(&[0x21, 0x43, 0x65, 0x87], 2, 4);
assert_eq!(two_rows, vec![1i8, 2, 3, 4, 5, 6, 7, -8]);
}
#[test]
#[should_panic(expected = "unpack_to_i8: n*k/2 overflow")]
fn unpack_to_i8_rejects_packed_shape_overflow_before_allocating() {
let _ = unpack_to_i8(&[], usize::MAX, 4);
}
#[test]
#[should_panic(expected = "igemm_s4s8: m*k overflow")]
fn igemm_s4s8_rejects_activation_shape_overflow_before_len_checks() {
let mut out = [];
igemm_s4s8(&[], &[], &[], 16, usize::MAX, 16, 1, &mut out);
}
struct Rng(u64);
impl Rng {
fn next_u64(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn byte(&mut self) -> u8 {
(self.next_u64() & 0xFF) as u8
}
fn i8q(&mut self) -> i8 {
((self.next_u64() % 255) as i64 - 127) as i8
}
}
#[test]
fn accel_unpack_equals_scalar_oracle_randomized() {
let mut rng = Rng(0x1234_5678_9abc_def0);
for &len in &[
0usize, 1, 7, 15, 16, 17, 31, 32, 33, 48, 63, 64, 100, 257, 1000,
] {
let src: Vec<u8> = (0..len).map(|_| rng.byte()).collect();
let mut got = vec![0i8; len * 2];
let mut want = vec![0i8; len * 2];
super::accel::unpack_nibbles(&src, &mut got);
super::accel::unpack_nibbles_scalar(&src, &mut want);
assert_eq!(got, want, "dispatched unpack != scalar oracle at len {len}");
}
}
#[test]
fn accel_unpack_equals_scalar_oracle_adversarial() {
for fill in [0x88u8, 0x77, 0xFF, 0x00, 0x8F, 0xF8] {
for &len in &[16usize, 32, 48, 100] {
let src = vec![fill; len];
let mut got = vec![0i8; len * 2];
let mut want = vec![0i8; len * 2];
super::accel::unpack_nibbles(&src, &mut got);
super::accel::unpack_nibbles_scalar(&src, &mut want);
assert_eq!(got, want, "adversarial fill 0x{fill:X} len {len}");
}
}
}
fn reference_s4s8(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
) -> Vec<f32> {
let groups = k / group;
let mut b = vec![0i8; n * k];
for o in 0..n {
for j in 0..k / 2 {
let byte = b_packed[o * (k / 2) + j];
b[o * k + 2 * j] = sign_extend_nibble(byte & 0x0F);
b[o * k + 2 * j + 1] = sign_extend_nibble(byte >> 4);
}
}
let mut out = vec![0f32; m * n];
for mi in 0..m {
for ni in 0..n {
let mut acc = 0.0f32;
for g in 0..groups {
let s = scales[ni * groups + g];
let mut gi: i32 = 0;
for kk in g * group..(g + 1) * group {
gi += i32::from(a[mi * k + kk]) * i32::from(b[ni * k + kk]);
}
acc += s * gi as f32;
}
out[mi * n + ni] = acc;
}
}
out
}
#[test]
fn igemm_s4s8_matches_reference_randomized() {
let mut rng = Rng(0xdead_beef_cafe_babe);
let cases = [
(1usize, 16usize, 1usize, 16usize),
(2, 32, 3, 16),
(4, 64, 5, 32),
(3, 96, 2, 16),
(2, 128, 7, 32),
];
for (m, k, n, group) in cases {
let groups = k / group;
let a: Vec<i8> = (0..m * k).map(|_| rng.i8q()).collect();
let b_packed: Vec<u8> = (0..n * (k / 2)).map(|_| rng.byte()).collect();
let scales: Vec<f32> = (0..n * groups)
.map(|i| {
let v = ((rng.next_u64() % 17) as f32 - 8.0) / 8.0; if v == 0.0 { 0.125 } else { v + 0.0 * i as f32 }
})
.collect();
let mut out = vec![0f32; m * n];
igemm_s4s8(&a, &b_packed, &scales, group, m, k, n, &mut out);
let want = reference_s4s8(&a, &b_packed, &scales, group, m, k, n);
assert_eq!(
out, want,
"igemm_s4s8 != reference for (m={m},k={k},n={n},g={group})"
);
}
}
#[test]
fn igemm_s4s8_adversarial_max_operands() {
let (m, k, n, group) = (2usize, 32usize, 3usize, 16usize);
let groups = k / group;
let b_packed = vec![0x88u8; n * (k / 2)];
let a: Vec<i8> = (0..m * k)
.map(|i| if i % 2 == 0 { 127 } else { -127 })
.collect();
let scales = vec![1.0f32; n * groups];
let mut out = vec![0f32; m * n];
igemm_s4s8(&a, &b_packed, &scales, group, m, k, n, &mut out);
let want = reference_s4s8(&a, &b_packed, &scales, group, m, k, n);
assert_eq!(out, want, "adversarial igemm_s4s8 != reference");
assert!(
out.iter().all(|&v| v == 0.0),
"expected 0.0 cells for ±127 cancel"
);
let a_pos = vec![-127i8; m * k];
let mut out2 = vec![0f32; m * n];
igemm_s4s8(&a_pos, &b_packed, &scales, group, m, k, n, &mut out2);
let want2 = reference_s4s8(&a_pos, &b_packed, &scales, group, m, k, n);
assert_eq!(out2, want2);
assert!(
out2.iter().all(|&v| (v - 32_512.0).abs() < 1e-3),
"group-sum value wrong"
);
}
#[test]
fn igemm_s4s8_single_group_equals_int8_per_channel() {
let mut rng = Rng(0x0f0f_0f0f_1234_5678);
let (m, k, n) = (3usize, 16usize, 4usize);
let group = k; let a: Vec<i8> = (0..m * k).map(|_| rng.i8q()).collect();
let b_packed: Vec<u8> = (0..n * (k / 2)).map(|_| rng.byte()).collect();
let scales: Vec<f32> = (0..n).map(|i| 0.0625 + i as f32 * 0.03125).collect();
let mut out = vec![0f32; m * n];
igemm_s4s8(&a, &b_packed, &scales, group, m, k, n, &mut out);
let b = unpack_to_i8(&b_packed, n, k);
let mut want = vec![0f32; m * n];
for mi in 0..m {
for ni in 0..n {
let mut acc: i32 = 0;
for kk in 0..k {
acc += i32::from(a[mi * k + kk]) * i32::from(b[ni * k + kk]);
}
want[mi * n + ni] = scales[ni] * acc as f32;
}
}
assert_eq!(out, want);
}
#[test]
#[should_panic(expected = "must be 16 or 32")]
fn igemm_s4s8_rejects_noncanonical_group_even_when_it_divides_k() {
let (m, k, n, group) = (1usize, 32usize, 1usize, 8usize);
let a = vec![1i8; m * k];
let b_packed = vec![0u8; n * (k / 2)];
let scales = vec![1.0f32; n * (k / group)];
let mut out = vec![0f32; m * n];
igemm_s4s8(&a, &b_packed, &scales, group, m, k, n, &mut out);
}
#[test]
#[should_panic(expected = "must divide k")]
fn igemm_s4s8_rejects_non_dividing_group() {
let (m, k, n, group) = (1usize, 24usize, 1usize, 16usize); let a = vec![1i8; m * k];
let b_packed = vec![0u8; n * (k / 2)];
let scales = vec![1.0f32; n * (k / group).max(1)];
let mut out = vec![0f32; m * n];
igemm_s4s8(&a, &b_packed, &scales, group, m, k, n, &mut out);
}
}