use std::sync::OnceLock;
use anyhow::{Context as _, Result, bail, ensure};
use rayon::prelude::*;
use super::gguf::{Tensor, TensorType};
use super::types::{Bf16, E8m0, Fp16, QuantFloat};
const BLOCK_VALUES: usize = 32;
const BLOCK_BYTES: usize = 34;
const MXFP4_BLOCK_BYTES: usize = 17;
const MATRIX_ROW_TILE: usize = 4;
const MXFP4_VALUES: [i8; 16] = [0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12];
pub(super) fn dim_to_f32(value: usize) -> f32 {
f32::from(u16::try_from(value).expect("model dimension/index exceeds u16"))
}
pub(super) fn dequantize_row(row: &[u8], output: &mut [f32]) -> Result<()> {
ensure!(
output.len().is_multiple_of(BLOCK_VALUES),
"Q8_0 output width is not divisible by 32"
);
ensure!(
row.len() == output.len() / BLOCK_VALUES * BLOCK_BYTES,
"invalid Q8_0 row size"
);
for (block, values) in row
.chunks_exact(BLOCK_BYTES)
.zip(output.chunks_exact_mut(BLOCK_VALUES))
{
let scale = Fp16::decode_le(&block[0..2]).to_f32();
for (value, quantized) in values.iter_mut().zip(&block[2..]) {
*value = scale * f32::from(i8::from_ne_bytes([*quantized]));
}
}
Ok(())
}
pub(super) struct Q8Activation {
scales: Vec<f32>,
values: Vec<i8>,
}
impl Q8Activation {
pub(super) fn new(vector: &[f32]) -> Result<Self> {
ensure!(
vector.len().is_multiple_of(BLOCK_VALUES),
"Q8 activation width is not divisible by 32"
);
let mut scales = Vec::with_capacity(vector.len() / BLOCK_VALUES);
let mut values = Vec::with_capacity(vector.len());
for block in vector.chunks_exact(BLOCK_VALUES) {
let mut maximum = 0.0_f32;
for &value in block {
ensure!(value.is_finite(), "Q8 activation is not finite");
maximum = maximum.max(value.abs());
}
let scale = maximum / 127.0;
let inverse = if scale == 0.0 { 0.0 } else { scale.recip() };
scales.push(scale);
for &value in block {
let quantized = (value * inverse).round().clamp(-127.0, 127.0);
#[allow(clippy::cast_possible_truncation)]
values.push(quantized as i8);
}
}
Ok(Self { scales, values })
}
}
#[derive(Clone, Copy)]
enum Q8Kernel {
Scalar,
#[cfg(target_arch = "x86_64")]
Avx2,
#[cfg(target_arch = "x86_64")]
Avx512,
#[cfg(target_arch = "aarch64")]
Neon,
}
impl Q8Kernel {
fn detect() -> Self {
static KERNEL: OnceLock<Q8Kernel> = OnceLock::new();
*KERNEL.get_or_init(|| {
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bw")
{
return Self::Avx512;
}
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") {
return Self::Avx2;
}
#[cfg(target_arch = "aarch64")]
{
return Self::Neon;
}
#[allow(unreachable_code)]
Self::Scalar
})
}
fn dot(self, row: &[u8], activation: &Q8Activation) -> f32 {
match self {
Self::Scalar => dot_q8_scalar(row, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => unsafe { dot_q8_avx2(row, activation) },
#[cfg(target_arch = "x86_64")]
Self::Avx512 => unsafe { dot_q8_avx512(row, activation) },
#[cfg(target_arch = "aarch64")]
Self::Neon => unsafe { dot_q8_neon(row, activation) },
}
}
fn dot_batch(
self,
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
match self {
Self::Scalar => std::array::from_fn(|index| dot_q8_scalar(row, &activations[index])),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => unsafe { dot_q8_avx2_batch(row, activations) },
#[cfg(target_arch = "x86_64")]
Self::Avx512 => unsafe { dot_q8_avx2_batch(row, activations) },
#[cfg(target_arch = "aarch64")]
Self::Neon => {
std::array::from_fn(|index| unsafe { dot_q8_neon(row, &activations[index]) })
}
}
}
}
#[derive(Clone, Copy)]
enum Mxfp4Kernel {
Scalar,
#[cfg(target_arch = "x86_64")]
Avx2,
#[cfg(target_arch = "x86_64")]
Avx512,
#[cfg(target_arch = "aarch64")]
Neon,
}
impl Mxfp4Kernel {
fn detect() -> Self {
static KERNEL: OnceLock<Mxfp4Kernel> = OnceLock::new();
*KERNEL.get_or_init(|| {
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bw")
{
return Self::Avx512;
}
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") {
return Self::Avx2;
}
#[cfg(target_arch = "aarch64")]
{
return Self::Neon;
}
#[allow(unreachable_code)]
Self::Scalar
})
}
fn dot(self, row: &[u8], activation: &Q8Activation) -> f32 {
match self {
Self::Scalar => dot_mxfp4_scalar(row, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => unsafe { dot_mxfp4_avx2(row, activation) },
#[cfg(target_arch = "x86_64")]
Self::Avx512 => unsafe { dot_mxfp4_avx512(row, activation) },
#[cfg(target_arch = "aarch64")]
Self::Neon => unsafe { dot_mxfp4_neon(row, activation) },
}
}
fn dot_batch(
self,
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
match self {
Self::Scalar => std::array::from_fn(|index| dot_mxfp4_scalar(row, &activations[index])),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => unsafe { dot_mxfp4_avx2_batch(row, activations) },
#[cfg(target_arch = "x86_64")]
Self::Avx512 => unsafe { dot_mxfp4_avx2_batch(row, activations) },
#[cfg(target_arch = "aarch64")]
Self::Neon => {
std::array::from_fn(|index| unsafe { dot_mxfp4_neon(row, &activations[index]) })
}
}
}
}
fn dot_q8_scalar(row: &[u8], activation: &Q8Activation) -> f32 {
row.chunks_exact(BLOCK_BYTES)
.zip(activation.values.chunks_exact(BLOCK_VALUES))
.zip(&activation.scales)
.map(|((block, values), activation_scale)| {
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
let sum_i32: i32 = block[2..]
.iter()
.zip(values)
.map(|(&weight, &value)| i32::from(i8::from_ne_bytes([weight])) * i32::from(value))
.sum();
#[allow(clippy::cast_precision_loss)]
let sum = sum_i32 as f32;
weight_scale * activation_scale * sum
})
.sum()
}
fn dot_mxfp4_scalar(row: &[u8], activation: &Q8Activation) -> f32 {
row.chunks_exact(MXFP4_BLOCK_BYTES)
.zip(activation.values.chunks_exact(BLOCK_VALUES))
.zip(&activation.scales)
.map(|((block, values), activation_scale)| {
let mut sum_i32 = 0_i32;
for (index, packed) in block[1..].iter().copied().enumerate() {
let low = i32::from(MXFP4_VALUES[usize::from(packed & 0x0f)]);
let high = i32::from(MXFP4_VALUES[usize::from(packed >> 4)]);
sum_i32 += i32::from(values[index]) * low;
sum_i32 += i32::from(values[index + BLOCK_VALUES / 2]) * high;
}
#[allow(clippy::cast_precision_loss)]
let sum = sum_i32 as f32;
E8m0::decode_le(&block[0..1]).to_f32() * activation_scale * sum
})
.sum()
}
fn dot_bf16(row: &[u8], vector: &[f32]) -> f32 {
row.chunks_exact(2)
.zip(vector)
.map(|(bytes, value)| {
let weight = Bf16::decode_le(&bytes[0..2]).to_f32();
weight * value
})
.sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_q8_avx2(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::{
__m256i, _mm_add_epi32, _mm_cvtepi32_ps, _mm_cvtss_f32, _mm_shuffle_epi32,
_mm_unpackhi_epi64, _mm256_abs_epi8, _mm256_castsi256_si128, _mm256_extracti128_si256,
_mm256_madd_epi16, _mm256_maddubs_epi16, _mm256_set1_epi16, _mm256_sign_epi8,
};
let ones = _mm256_set1_epi16(1);
let mut sum = 0.0;
let mut blocks = row.chunks_exact(BLOCK_BYTES);
let mut values_chunks = activation.values.chunks_exact(BLOCK_VALUES);
let mut scales = activation.scales.iter().copied();
while let (Some(b0), Some(v0), Some(s0), Some(b1), Some(v1), Some(s1)) = (
blocks.next(),
values_chunks.next(),
scales.next(),
blocks.next(),
values_chunks.next(),
scales.next(),
) {
let w0: __m256i = bytemuck::pod_read_unaligned(&b0[2..]);
let a0: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v0));
let w1: __m256i = bytemuck::pod_read_unaligned(&b1[2..]);
let a1: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v1));
let signed0 = _mm256_sign_epi8(a0, w0);
let signed1 = _mm256_sign_epi8(a1, w1);
let mag0 = _mm256_abs_epi8(w0);
let mag1 = _mm256_abs_epi8(w1);
let pairs0 = _mm256_maddubs_epi16(mag0, signed0);
let pairs1 = _mm256_maddubs_epi16(mag1, signed1);
let prod0 = _mm256_madd_epi16(pairs0, ones);
let prod1 = _mm256_madd_epi16(pairs1, ones);
let low0 = _mm256_castsi256_si128(prod0);
let high0 = _mm256_extracti128_si256::<1>(prod0);
let lanes0 = _mm_add_epi32(low0, high0);
let p0 = _mm_add_epi32(lanes0, _mm_unpackhi_epi64(lanes0, lanes0));
let tot0 = _mm_add_epi32(p0, _mm_shuffle_epi32::<0x55>(p0));
let low1 = _mm256_castsi256_si128(prod1);
let high1 = _mm256_extracti128_si256::<1>(prod1);
let lanes1 = _mm_add_epi32(low1, high1);
let p1 = _mm_add_epi32(lanes1, _mm_unpackhi_epi64(lanes1, lanes1));
let tot1 = _mm_add_epi32(p1, _mm_shuffle_epi32::<0x55>(p1));
let ws0 = Fp16::decode_le(&b0[0..2]).to_f32();
let ws1 = Fp16::decode_le(&b1[0..2]).to_f32();
sum += ws0 * s0 * _mm_cvtss_f32(_mm_cvtepi32_ps(tot0));
sum += ws1 * s1 * _mm_cvtss_f32(_mm_cvtepi32_ps(tot1));
}
while let (Some(b0), Some(v0), Some(s0)) = (blocks.next(), values_chunks.next(), scales.next())
{
let w0: __m256i = bytemuck::pod_read_unaligned(&b0[2..]);
let a0: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v0));
let signed0 = _mm256_sign_epi8(a0, w0);
let mag0 = _mm256_abs_epi8(w0);
let pairs0 = _mm256_maddubs_epi16(mag0, signed0);
let prod0 = _mm256_madd_epi16(pairs0, ones);
let low0 = _mm256_castsi256_si128(prod0);
let high0 = _mm256_extracti128_si256::<1>(prod0);
let lanes0 = _mm_add_epi32(low0, high0);
let p0 = _mm_add_epi32(lanes0, _mm_unpackhi_epi64(lanes0, lanes0));
let tot0 = _mm_add_epi32(p0, _mm_shuffle_epi32::<0x55>(p0));
let ws0 = Fp16::decode_le(&b0[0..2]).to_f32();
sum += ws0 * s0 * _mm_cvtss_f32(_mm_cvtepi32_ps(tot0));
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_mxfp4_avx2(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::{
__m128i, __m256i, _mm_add_epi32, _mm_and_si128, _mm_cvtepi32_ps, _mm_cvtss_f32,
_mm_set1_epi8, _mm_shuffle_epi8, _mm_shuffle_epi32, _mm_srli_epi16, _mm_unpackhi_epi64,
_mm256_abs_epi8, _mm256_castsi256_si128, _mm256_extracti128_si256, _mm256_madd_epi16,
_mm256_maddubs_epi16, _mm256_set_m128i, _mm256_set1_epi16, _mm256_sign_epi8,
};
const LUT: [i8; 16] = [0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12];
let lut_128: __m128i = bytemuck::cast(LUT);
let mask_0f = _mm_set1_epi8(0x0F);
let ones = _mm256_set1_epi16(1);
let mut sum = 0.0;
let mut blocks = row.chunks_exact(MXFP4_BLOCK_BYTES);
let mut values_chunks = activation.values.chunks_exact(BLOCK_VALUES);
let mut scales = activation.scales.iter().copied();
while let (Some(b0), Some(v0), Some(s0), Some(b1), Some(v1), Some(s1)) = (
blocks.next(),
values_chunks.next(),
scales.next(),
blocks.next(),
values_chunks.next(),
scales.next(),
) {
let p128_0: __m128i = bytemuck::pod_read_unaligned(&b0[1..17]);
let low0 = _mm_and_si128(p128_0, mask_0f);
let high0 = _mm_and_si128(_mm_srli_epi16(p128_0, 4), mask_0f);
let w_low0 = _mm_shuffle_epi8(lut_128, low0);
let w_high0 = _mm_shuffle_epi8(lut_128, high0);
let w0 = _mm256_set_m128i(w_high0, w_low0);
let a0: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v0));
let p128_1: __m128i = bytemuck::pod_read_unaligned(&b1[1..17]);
let low1 = _mm_and_si128(p128_1, mask_0f);
let high1 = _mm_and_si128(_mm_srli_epi16(p128_1, 4), mask_0f);
let w_low1 = _mm_shuffle_epi8(lut_128, low1);
let w_high1 = _mm_shuffle_epi8(lut_128, high1);
let w1 = _mm256_set_m128i(w_high1, w_low1);
let a1: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v1));
let signed0 = _mm256_sign_epi8(a0, w0);
let signed1 = _mm256_sign_epi8(a1, w1);
let mag0 = _mm256_abs_epi8(w0);
let mag1 = _mm256_abs_epi8(w1);
let pairs0 = _mm256_maddubs_epi16(mag0, signed0);
let pairs1 = _mm256_maddubs_epi16(mag1, signed1);
let prod0 = _mm256_madd_epi16(pairs0, ones);
let prod1 = _mm256_madd_epi16(pairs1, ones);
let l0 = _mm256_castsi256_si128(prod0);
let h0 = _mm256_extracti128_si256::<1>(prod0);
let lanes0 = _mm_add_epi32(l0, h0);
let p0 = _mm_add_epi32(lanes0, _mm_unpackhi_epi64(lanes0, lanes0));
let tot0 = _mm_add_epi32(p0, _mm_shuffle_epi32::<0x55>(p0));
let l1 = _mm256_castsi256_si128(prod1);
let h1 = _mm256_extracti128_si256::<1>(prod1);
let lanes1 = _mm_add_epi32(l1, h1);
let p1 = _mm_add_epi32(lanes1, _mm_unpackhi_epi64(lanes1, lanes1));
let tot1 = _mm_add_epi32(p1, _mm_shuffle_epi32::<0x55>(p1));
let sc0 = E8m0::decode_le(&b0[0..1]).to_f32();
let sc1 = E8m0::decode_le(&b1[0..1]).to_f32();
sum += sc0 * s0 * _mm_cvtss_f32(_mm_cvtepi32_ps(tot0));
sum += sc1 * s1 * _mm_cvtss_f32(_mm_cvtepi32_ps(tot1));
}
while let (Some(b0), Some(v0), Some(s0)) = (blocks.next(), values_chunks.next(), scales.next())
{
let p128_0: __m128i = bytemuck::pod_read_unaligned(&b0[1..17]);
let low0 = _mm_and_si128(p128_0, mask_0f);
let high0 = _mm_and_si128(_mm_srli_epi16(p128_0, 4), mask_0f);
let w_low0 = _mm_shuffle_epi8(lut_128, low0);
let w_high0 = _mm_shuffle_epi8(lut_128, high0);
let w0 = _mm256_set_m128i(w_high0, w_low0);
let a0: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v0));
let signed0 = _mm256_sign_epi8(a0, w0);
let mag0 = _mm256_abs_epi8(w0);
let pairs0 = _mm256_maddubs_epi16(mag0, signed0);
let prod0 = _mm256_madd_epi16(pairs0, ones);
let l0 = _mm256_castsi256_si128(prod0);
let h0 = _mm256_extracti128_si256::<1>(prod0);
let lanes0 = _mm_add_epi32(l0, h0);
let p0 = _mm_add_epi32(lanes0, _mm_unpackhi_epi64(lanes0, lanes0));
let tot0 = _mm_add_epi32(p0, _mm_shuffle_epi32::<0x55>(p0));
let sc0 = E8m0::decode_le(&b0[0..1]).to_f32();
sum += sc0 * s0 * _mm_cvtss_f32(_mm_cvtepi32_ps(tot0));
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn multiply_i8_block_avx2(
weights: std::arch::x86_64::__m256i,
magnitudes: std::arch::x86_64::__m256i,
activation: std::arch::x86_64::__m256i,
ones: std::arch::x86_64::__m256i,
) -> std::arch::x86_64::__m256i {
use std::arch::x86_64::{_mm256_madd_epi16, _mm256_maddubs_epi16, _mm256_sign_epi8};
let signed = _mm256_sign_epi8(activation, weights);
let pairs = _mm256_maddubs_epi16(magnitudes, signed);
_mm256_madd_epi16(pairs, ones)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_q8_avx2_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m256, __m256i, _mm256_abs_epi8, _mm256_add_ps, _mm256_cvtepi32_ps, _mm256_mul_ps,
_mm256_set1_epi16, _mm256_set1_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
let ones = _mm256_set1_epi16(1);
let mut sums: [__m256; MATRIX_ROW_TILE] = [_mm256_setzero_ps(); MATRIX_ROW_TILE];
for (block_index, block) in row.chunks_exact(BLOCK_BYTES).enumerate() {
let weights: __m256i = bytemuck::pod_read_unaligned(&block[2..]);
let magnitudes = _mm256_abs_epi8(weights);
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
let start = block_index * BLOCK_VALUES;
for index in 0..MATRIX_ROW_TILE {
let values: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(
&activations[index].values[start..start + BLOCK_VALUES],
));
let products = unsafe { multiply_i8_block_avx2(weights, magnitudes, values, ones) };
let scale = weight_scale * activations[index].scales[block_index];
sums[index] = _mm256_add_ps(
sums[index],
_mm256_mul_ps(_mm256_cvtepi32_ps(products), _mm256_set1_ps(scale)),
);
}
}
std::array::from_fn(|index| {
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sums[index]) };
lanes.into_iter().sum()
})
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_mxfp4_avx2_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m128i, __m256, __m256i, _mm_and_si128, _mm_set1_epi8, _mm_shuffle_epi8, _mm_srli_epi16,
_mm256_abs_epi8, _mm256_add_ps, _mm256_cvtepi32_ps, _mm256_mul_ps, _mm256_set_m128i,
_mm256_set1_epi16, _mm256_set1_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
const LUT: [i8; 16] = [0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12];
let lut: __m128i = bytemuck::cast(LUT);
let mask = _mm_set1_epi8(0x0f);
let ones = _mm256_set1_epi16(1);
let mut sums: [__m256; MATRIX_ROW_TILE] = [_mm256_setzero_ps(); MATRIX_ROW_TILE];
for (block_index, block) in row.chunks_exact(MXFP4_BLOCK_BYTES).enumerate() {
let packed: __m128i = bytemuck::pod_read_unaligned(&block[1..]);
let low = _mm_and_si128(packed, mask);
let high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
let low = _mm_shuffle_epi8(lut, low);
let high = _mm_shuffle_epi8(lut, high);
let weights = _mm256_set_m128i(high, low);
let magnitudes = _mm256_abs_epi8(weights);
let weight_scale = E8m0::decode_le(&block[0..1]).to_f32();
let start = block_index * BLOCK_VALUES;
for index in 0..MATRIX_ROW_TILE {
let values: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(
&activations[index].values[start..start + BLOCK_VALUES],
));
let products = unsafe { multiply_i8_block_avx2(weights, magnitudes, values, ones) };
let scale = weight_scale * activations[index].scales[block_index];
sums[index] = _mm256_add_ps(
sums[index],
_mm256_mul_ps(_mm256_cvtepi32_ps(products), _mm256_set1_ps(scale)),
);
}
}
std::array::from_fn(|index| {
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sums[index]) };
lanes.into_iter().sum()
})
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn dot_q8_avx512(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_avx2(row, activation) }
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn dot_mxfp4_avx512(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_mxfp4_avx2(row, activation) }
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn dot_q8_neon(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::aarch64::{
vaddq_s32, vaddvq_s32, vget_low_s8, vget_low_s16, vld1q_s8, vmovl_high_s8, vmovl_s8,
vmull_high_s16, vmull_s16,
};
let mut sum = 0.0_f32;
let blocks = row.chunks_exact(BLOCK_BYTES);
let values_chunks = activation.values.chunks_exact(BLOCK_VALUES);
let scales = activation.scales.iter().copied();
for ((block, values), activation_scale) in blocks.zip(values_chunks).zip(scales) {
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
let w_ptr = block[2..].as_ptr().cast::<i8>();
let v_ptr = values.as_ptr();
let block_sum = unsafe {
let w0 = vld1q_s8(w_ptr);
let w1 = vld1q_s8(w_ptr.add(16));
let v0 = vld1q_s8(v_ptr);
let v1 = vld1q_s8(v_ptr.add(16));
let w0_low = vmovl_s8(vget_low_s8(w0));
let w0_high = vmovl_high_s8(w0);
let v0_low = vmovl_s8(vget_low_s8(v0));
let v0_high = vmovl_high_s8(v0);
let w1_low = vmovl_s8(vget_low_s8(w1));
let w1_high = vmovl_high_s8(w1);
let v1_low = vmovl_s8(vget_low_s8(v1));
let v1_high = vmovl_high_s8(v1);
let p0 = vmull_s16(vget_low_s16(w0_low), vget_low_s16(v0_low));
let p1 = vmull_high_s16(w0_low, v0_low);
let p2 = vmull_s16(vget_low_s16(w0_high), vget_low_s16(v0_high));
let p3 = vmull_high_s16(w0_high, v0_high);
let p4 = vmull_s16(vget_low_s16(w1_low), vget_low_s16(v1_low));
let p5 = vmull_high_s16(w1_low, v1_low);
let p6 = vmull_s16(vget_low_s16(w1_high), vget_low_s16(v1_high));
let p7 = vmull_high_s16(w1_high, v1_high);
let acc0 = vaddq_s32(vaddq_s32(p0, p1), vaddq_s32(p2, p3));
let acc1 = vaddq_s32(vaddq_s32(p4, p5), vaddq_s32(p6, p7));
let acc = vaddq_s32(acc0, acc1);
vaddvq_s32(acc)
};
#[allow(clippy::cast_precision_loss)]
let block_sum_f32 = block_sum as f32;
sum += weight_scale * activation_scale * block_sum_f32;
}
sum
}
#[cfg(target_arch = "aarch64")]
unsafe fn dot_mxfp4_neon(row: &[u8], activation: &Q8Activation) -> f32 {
let mut sum = 0.0_f32;
let blocks = row.chunks_exact(MXFP4_BLOCK_BYTES);
let values_chunks = activation.values.chunks_exact(BLOCK_VALUES);
let scales = activation.scales.iter().copied();
for ((block, values), activation_scale) in blocks.zip(values_chunks).zip(scales) {
let mut block_sum_i32 = 0_i32;
for (index, packed) in block[1..].iter().copied().enumerate() {
let low = i32::from(MXFP4_VALUES[usize::from(packed & 0x0f)]);
let high = i32::from(MXFP4_VALUES[usize::from(packed >> 4)]);
block_sum_i32 += i32::from(values[index]) * low;
block_sum_i32 += i32::from(values[index + BLOCK_VALUES / 2]) * high;
}
#[allow(clippy::cast_precision_loss)]
let block_sum = block_sum_i32 as f32;
sum += E8m0::decode_le(&block[0..1]).to_f32() * activation_scale * block_sum;
}
sum
}
pub(super) fn matrix_vector(matrix: &Tensor<'_>, vector: &[f32]) -> Result<Vec<f32>> {
matrix_vector_with_activation(matrix, vector, None)
}
pub(super) fn matrix_vector_with_activation(
matrix: &Tensor<'_>,
vector: &[f32],
activation: Option<&Q8Activation>,
) -> Result<Vec<f32>> {
let [input, output] = matrix_dimensions(matrix)?;
ensure!(vector.len() == input, "matrix input width differs");
let mut result = vec![0.0; output];
match matrix.tensor_type() {
TensorType::F32 => {
result
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = dot_f32(matrix.f32_row(row)?, vector);
Ok(())
})?;
}
TensorType::Bf16 => {
result
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = dot_bf16(matrix.bf16_row(row)?, vector);
Ok(())
})?;
}
TensorType::Q8_0 => {
if let Some(act) = activation {
project_q8(matrix, act, &mut result)?;
} else {
let act = Q8Activation::new(vector)?;
project_q8(matrix, &act, &mut result)?;
}
}
TensorType::Mxfp4 => {
if let Some(act) = activation {
project_mxfp4(matrix, act, &mut result)?;
} else {
let act = Q8Activation::new(vector)?;
project_mxfp4(matrix, &act, &mut result)?;
}
}
}
Ok(result)
}
pub(super) fn matrix_vector_triple(
first: &Tensor<'_>,
second: &Tensor<'_>,
third: &Tensor<'_>,
vector: &[f32],
) -> Result<(Vec<f32>, Vec<f32>, Vec<f32>)> {
if first.tensor_type() == TensorType::Q8_0
&& second.tensor_type() == TensorType::Q8_0
&& third.tensor_type() == TensorType::Q8_0
{
let activation = Q8Activation::new(vector)?;
let first_rows = matrix_dimensions(first)?[1];
let second_rows = matrix_dimensions(second)?[1];
let third_rows = matrix_dimensions(third)?[1];
let mut first_output = vec![0.0; first_rows];
let mut second_output = vec![0.0; second_rows];
let mut third_output = vec![0.0; third_rows];
let kernel = Q8Kernel::detect();
first_output
.par_iter_mut()
.chain(second_output.par_iter_mut())
.chain(third_output.par_iter_mut())
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
if row < first_rows {
*value = kernel.dot(first.q8_row(row)?, &activation);
} else if row < first_rows + second_rows {
*value = kernel.dot(second.q8_row(row - first_rows)?, &activation);
} else {
*value = kernel.dot(third.q8_row(row - first_rows - second_rows)?, &activation);
}
Ok(())
})?;
return Ok((first_output, second_output, third_output));
}
Ok((
matrix_vector(first, vector)?,
matrix_vector(second, vector)?,
matrix_vector(third, vector)?,
))
}
pub(super) fn matrix_matrix(
matrix: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
) -> Result<Vec<f32>> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
validate_matrix_input(vectors, row_count, input_size)?;
if row_count == 1 {
return matrix_vector(matrix, vectors);
}
let output_len = row_count
.checked_mul(output_size)
.context("matrix-matrix output size overflow")?;
let mut output = vec![0.0; output_len];
match matrix.tensor_type() {
TensorType::F32 => project_f32_batch(matrix, vectors, &mut output)?,
TensorType::Bf16 => project_bf16_batch(matrix, vectors, &mut output)?,
TensorType::Q8_0 => {
let activations = matrix_activations(vectors, row_count, input_size)?;
project_q8_batch(matrix, &activations, &mut output)?;
}
TensorType::Mxfp4 => {
let activations = matrix_activations(vectors, row_count, input_size)?;
project_mxfp4_batch(matrix, &activations, &mut output)?;
}
}
Ok(output)
}
pub(super) fn matrix_matrix_pair(
first: &Tensor<'_>,
second: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
) -> Result<(Vec<f32>, Vec<f32>)> {
let first_dimensions = matrix_dimensions(first)?;
let second_dimensions = matrix_dimensions(second)?;
ensure!(
first_dimensions == second_dimensions,
"paired matrix dimensions differ"
);
let [input_size, output_size] = first_dimensions;
validate_matrix_input(vectors, row_count, input_size)?;
match (first.tensor_type(), second.tensor_type()) {
(TensorType::Q8_0, TensorType::Q8_0) => {
let activations = matrix_activations(vectors, row_count, input_size)?;
let (first_output, second_output) = rayon::join(
|| quantized_matrix_matrix(first, output_size, &activations),
|| quantized_matrix_matrix(second, output_size, &activations),
);
Ok((first_output?, second_output?))
}
(TensorType::Mxfp4, TensorType::Mxfp4) => {
let activations = matrix_activations(vectors, row_count, input_size)?;
let (first_output, second_output) = rayon::join(
|| quantized_matrix_matrix(first, output_size, &activations),
|| quantized_matrix_matrix(second, output_size, &activations),
);
Ok((first_output?, second_output?))
}
_ => {
let (first_output, second_output) = rayon::join(
|| matrix_matrix(first, vectors, row_count),
|| matrix_matrix(second, vectors, row_count),
);
Ok((first_output?, second_output?))
}
}
}
pub(super) fn matrix_matrix_triple(
first: &Tensor<'_>,
second: &Tensor<'_>,
third: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
) -> Result<(Vec<f32>, Vec<f32>, Vec<f32>)> {
let [input_size, first_size] = matrix_dimensions(first)?;
let [second_input, second_size] = matrix_dimensions(second)?;
let [third_input, third_size] = matrix_dimensions(third)?;
ensure!(
second_input == input_size && third_input == input_size,
"triple matrix input dimensions differ"
);
validate_matrix_input(vectors, row_count, input_size)?;
if row_count == 1 {
return matrix_vector_triple(first, second, third, vectors);
}
if first.tensor_type() == TensorType::Q8_0
&& second.tensor_type() == TensorType::Q8_0
&& third.tensor_type() == TensorType::Q8_0
{
let activations = matrix_activations(vectors, row_count, input_size)?;
let (first_output, (second_output, third_output)) = rayon::join(
|| quantized_matrix_matrix(first, first_size, &activations),
|| {
rayon::join(
|| quantized_matrix_matrix(second, second_size, &activations),
|| quantized_matrix_matrix(third, third_size, &activations),
)
},
);
return Ok((first_output?, second_output?, third_output?));
}
let (first_output, (second_output, third_output)) = rayon::join(
|| matrix_matrix(first, vectors, row_count),
|| {
rayon::join(
|| matrix_matrix(second, vectors, row_count),
|| matrix_matrix(third, vectors, row_count),
)
},
);
Ok((first_output?, second_output?, third_output?))
}
fn validate_matrix_input(vectors: &[f32], row_count: usize, input_size: usize) -> Result<()> {
ensure!(row_count != 0, "matrix-matrix row count is zero");
let expected = row_count
.checked_mul(input_size)
.context("matrix-matrix input size overflow")?;
ensure!(
vectors.len() == expected,
"matrix-matrix input has {} values, expected {expected}",
vectors.len()
);
Ok(())
}
fn matrix_activations(
vectors: &[f32],
row_count: usize,
input_size: usize,
) -> Result<Vec<Q8Activation>> {
validate_matrix_input(vectors, row_count, input_size)?;
vectors
.chunks_exact(input_size)
.map(Q8Activation::new)
.collect()
}
fn quantized_matrix_matrix(
matrix: &Tensor<'_>,
output_size: usize,
activations: &[Q8Activation],
) -> Result<Vec<f32>> {
let output_len = activations
.len()
.checked_mul(output_size)
.context("matrix-matrix output size overflow")?;
let mut output = vec![0.0; output_len];
match matrix.tensor_type() {
TensorType::Q8_0 => project_q8_batch(matrix, activations, &mut output)?,
TensorType::Mxfp4 => project_mxfp4_batch(matrix, activations, &mut output)?,
_ => unreachable!("validated quantized matrix"),
}
Ok(output)
}
fn project_f32_batch(matrix: &Tensor<'_>, vectors: &[f32], output: &mut [f32]) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
ensure!(
matrix.tensor_type() == TensorType::F32,
"projection is not F32"
);
validate_dense_batch(vectors, output, input_size, output_size)?;
output
.par_chunks_mut(MATRIX_ROW_TILE * output_size)
.zip(vectors.par_chunks(MATRIX_ROW_TILE * input_size))
.try_for_each(|(output_rows, input_rows)| -> Result<()> {
for output_channel in 0..output_size {
let weights = matrix.f32_row(output_channel)?;
for (output_row, input_row) in output_rows
.chunks_exact_mut(output_size)
.zip(input_rows.chunks_exact(input_size))
{
output_row[output_channel] = dot_f32(weights, input_row);
}
}
Ok(())
})
}
fn project_bf16_batch(matrix: &Tensor<'_>, vectors: &[f32], output: &mut [f32]) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
ensure!(
matrix.tensor_type() == TensorType::Bf16,
"projection is not BF16"
);
validate_dense_batch(vectors, output, input_size, output_size)?;
output
.par_chunks_mut(MATRIX_ROW_TILE * output_size)
.zip(vectors.par_chunks(MATRIX_ROW_TILE * input_size))
.try_for_each(|(output_rows, input_rows)| -> Result<()> {
for output_channel in 0..output_size {
let weights = matrix.bf16_row(output_channel)?;
for (output_row, input_row) in output_rows
.chunks_exact_mut(output_size)
.zip(input_rows.chunks_exact(input_size))
{
output_row[output_channel] = dot_bf16(weights, input_row);
}
}
Ok(())
})
}
fn validate_dense_batch(
vectors: &[f32],
output: &[f32],
input_size: usize,
output_size: usize,
) -> Result<()> {
ensure!(
!vectors.is_empty() && vectors.len().is_multiple_of(input_size),
"invalid matrix-matrix input shape"
);
let row_count = vectors.len() / input_size;
ensure!(
output.len() == row_count * output_size,
"matrix output shape differs"
);
Ok(())
}
fn project_q8_batch(
matrix: &Tensor<'_>,
activations: &[Q8Activation],
output: &mut [f32],
) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
ensure!(
matrix.tensor_type() == TensorType::Q8_0,
"projection is not Q8_0"
);
validate_quantized_batch(activations, output, input_size, output_size)?;
let kernel = Q8Kernel::detect();
output
.par_chunks_mut(MATRIX_ROW_TILE * output_size)
.zip(activations.par_chunks(MATRIX_ROW_TILE))
.try_for_each(|(output_rows, activation_rows)| -> Result<()> {
if let Ok(activations) = <&[Q8Activation; MATRIX_ROW_TILE]>::try_from(activation_rows) {
for output_channel in 0..output_size {
let values = kernel.dot_batch(matrix.q8_row(output_channel)?, activations);
for (output_row, value) in output_rows.chunks_exact_mut(output_size).zip(values)
{
output_row[output_channel] = value;
}
}
} else {
for output_channel in 0..output_size {
let weights = matrix.q8_row(output_channel)?;
for (output_row, activation) in output_rows
.chunks_exact_mut(output_size)
.zip(activation_rows)
{
output_row[output_channel] = kernel.dot(weights, activation);
}
}
}
Ok(())
})
}
fn project_mxfp4_batch(
matrix: &Tensor<'_>,
activations: &[Q8Activation],
output: &mut [f32],
) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
ensure!(
matrix.tensor_type() == TensorType::Mxfp4,
"projection is not MXFP4"
);
validate_quantized_batch(activations, output, input_size, output_size)?;
let kernel = Mxfp4Kernel::detect();
output
.par_chunks_mut(MATRIX_ROW_TILE * output_size)
.zip(activations.par_chunks(MATRIX_ROW_TILE))
.try_for_each(|(output_rows, activation_rows)| -> Result<()> {
if let Ok(activations) = <&[Q8Activation; MATRIX_ROW_TILE]>::try_from(activation_rows) {
for output_channel in 0..output_size {
let values = kernel.dot_batch(matrix.mxfp4_row(output_channel)?, activations);
for (output_row, value) in output_rows.chunks_exact_mut(output_size).zip(values)
{
output_row[output_channel] = value;
}
}
} else {
for output_channel in 0..output_size {
let weights = matrix.mxfp4_row(output_channel)?;
for (output_row, activation) in output_rows
.chunks_exact_mut(output_size)
.zip(activation_rows)
{
output_row[output_channel] = kernel.dot(weights, activation);
}
}
}
Ok(())
})
}
fn validate_quantized_batch(
activations: &[Q8Activation],
output: &[f32],
input_size: usize,
output_size: usize,
) -> Result<()> {
ensure!(!activations.is_empty(), "matrix-matrix row count is zero");
ensure!(
activations
.iter()
.all(|activation| activation.values.len() == input_size),
"matrix input width differs"
);
ensure!(
output.len() == activations.len() * output_size,
"matrix output shape differs"
);
Ok(())
}
pub(super) fn matrix_argmax(matrix: &Tensor<'_>, vector: &[f32]) -> Result<usize> {
let [input, output] = matrix_dimensions(matrix)?;
ensure!(vector.len() == input, "matrix input width differs");
let q8_activation = match matrix.tensor_type() {
TensorType::F32 | TensorType::Bf16 => None,
TensorType::Q8_0 | TensorType::Mxfp4 => Some(Q8Activation::new(vector)?),
};
let kernel = Q8Kernel::detect();
let mxfp4_kernel = Mxfp4Kernel::detect();
let best = (0..output)
.into_par_iter()
.map(|index| -> Result<(usize, f32)> {
let value = match matrix.tensor_type() {
TensorType::F32 => dot_f32(matrix.f32_row(index)?, vector),
TensorType::Bf16 => dot_bf16(matrix.bf16_row(index)?, vector),
TensorType::Q8_0 => kernel.dot(
matrix.q8_row(index)?,
q8_activation.as_ref().expect("quantized activation"),
),
TensorType::Mxfp4 => mxfp4_kernel.dot(
matrix.mxfp4_row(index)?,
q8_activation.as_ref().expect("quantized activation"),
),
};
if !value.is_finite() {
bail!("matrix output {index} is not finite");
}
Ok((index, value))
})
.try_reduce_with(|left, right| {
Ok(match right.1.total_cmp(&left.1) {
std::cmp::Ordering::Greater => right,
std::cmp::Ordering::Equal if right.0 < left.0 => right,
_ => left,
})
})
.transpose()?
.expect("validated matrix has output rows");
Ok(best.0)
}
fn project_q8(matrix: &Tensor<'_>, activation: &Q8Activation, output: &mut [f32]) -> Result<()> {
let [input, rows] = matrix_dimensions(matrix)?;
ensure!(
matrix.tensor_type() == TensorType::Q8_0,
"projection is not Q8_0"
);
ensure!(
activation.values.len() == input,
"matrix input width differs"
);
ensure!(output.len() == rows, "matrix output height differs");
let kernel = Q8Kernel::detect();
output
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = kernel.dot(matrix.q8_row(row)?, activation);
Ok(())
})?;
Ok(())
}
fn project_mxfp4(matrix: &Tensor<'_>, activation: &Q8Activation, output: &mut [f32]) -> Result<()> {
let [input, rows] = matrix_dimensions(matrix)?;
ensure!(
matrix.tensor_type() == TensorType::Mxfp4,
"projection is not MXFP4"
);
ensure!(
activation.values.len() == input,
"matrix input width differs"
);
ensure!(output.len() == rows, "matrix output height differs");
let kernel = Mxfp4Kernel::detect();
output
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = kernel.dot(matrix.mxfp4_row(row)?, activation);
Ok(())
})?;
Ok(())
}
fn matrix_dimensions(matrix: &Tensor<'_>) -> Result<[usize; 2]> {
match matrix.dimensions() {
[input, output] => Ok([*input, *output]),
dimensions => bail!("matrix has dimensions {dimensions:?}, expected two"),
}
}
fn dot_f32(left: &[f32], right: &[f32]) -> f32 {
debug_assert_eq!(left.len(), right.len());
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
return unsafe { dot_f32_avx2(left, right) };
}
left.iter()
.zip(right)
.map(|(left, right)| left * right)
.sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_f32_avx2(left: &[f32], right: &[f32]) -> f32 {
use std::arch::x86_64::{
_mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
let vectorized = left.len() / 8 * 8;
let mut sums = _mm256_setzero_ps();
for index in (0..vectorized).step_by(8) {
let left = unsafe { _mm256_loadu_ps(left.as_ptr().add(index)) };
let right = unsafe { _mm256_loadu_ps(right.as_ptr().add(index)) };
sums = _mm256_fmadd_ps(left, right, sums);
}
let mut lanes = [0.0; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sums) };
lanes.into_iter().sum::<f32>()
+ left[vectorized..]
.iter()
.zip(&right[vectorized..])
.map(|(left, right)| left * right)
.sum::<f32>()
}
pub(super) fn rms_norm(
values: &[f32],
width: usize,
weight: &[f32],
epsilon: f32,
) -> Result<Vec<f32>> {
ensure!(
width != 0 && values.len().is_multiple_of(width),
"invalid RMS norm shape"
);
ensure!(weight.len() == width, "invalid RMS norm weight");
let mut output = vec![0.0; values.len()];
for (input, output) in values
.chunks_exact(width)
.zip(output.chunks_exact_mut(width))
{
let width_f32 = dim_to_f32(width);
let mean_square = dot_f32(input, input) / width_f32;
let scale = (mean_square + epsilon).sqrt().recip();
for index in 0..width {
output[index] = input[index] * scale * weight[index];
}
}
Ok(output)
}
pub(super) fn softmax(values: &mut [f32]) {
let maximum = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let sum = values
.iter_mut()
.map(|value| {
*value = (*value - maximum).exp();
*value
})
.sum::<f32>();
let inv_sum = sum.recip();
for value in values {
*value *= inv_sum;
}
}
pub(super) fn vector_add(left: &mut [f32], right: &[f32]) -> Result<()> {
ensure!(left.len() == right.len(), "vector lengths differ");
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") {
unsafe { vector_add_avx2(left, right) };
return Ok(());
}
left.iter_mut()
.zip(right)
.for_each(|(left, right)| *left += right);
Ok(())
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn vector_add_avx2(left: &mut [f32], right: &[f32]) {
use std::arch::x86_64::{_mm256_add_ps, _mm256_loadu_ps, _mm256_storeu_ps};
let vectorized = left.len() / 8 * 8;
for index in (0..vectorized).step_by(8) {
let l = unsafe { _mm256_loadu_ps(left.as_ptr().add(index)) };
let r = unsafe { _mm256_loadu_ps(right.as_ptr().add(index)) };
let sum = _mm256_add_ps(l, r);
unsafe { _mm256_storeu_ps(left.as_mut_ptr().add(index), sum) };
}
for index in vectorized..left.len() {
left[index] += right[index];
}
}