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, Debug, Eq, PartialEq)]
enum QuantizedFormat {
Q8_0,
Mxfp4,
}
impl QuantizedFormat {
fn from_tensor_type(kind: TensorType) -> Option<Self> {
match kind {
TensorType::Q8_0 => Some(Self::Q8_0),
TensorType::Mxfp4 => Some(Self::Mxfp4),
TensorType::F32 | TensorType::Bf16 => None,
}
}
fn row<'a>(self, matrix: &Tensor<'a>, row: usize) -> Result<&'a [u8]> {
ensure!(
Self::from_tensor_type(matrix.tensor_type()) == Some(self),
"quantized format differs from tensor"
);
matrix.encoded_row(row)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum WeightFormat {
F32,
Bf16,
Quantized(QuantizedFormat),
}
impl WeightFormat {
fn from_tensor_type(kind: TensorType) -> Self {
match kind {
TensorType::F32 => Self::F32,
TensorType::Bf16 => Self::Bf16,
TensorType::Q8_0 | TensorType::Mxfp4 => {
let Some(format) = QuantizedFormat::from_tensor_type(kind) else {
unreachable!("quantized tensor type")
};
Self::Quantized(format)
}
}
}
fn quantized(self) -> Option<QuantizedFormat> {
match self {
Self::Quantized(format) => Some(format),
Self::F32 | Self::Bf16 => None,
}
}
fn project_vector(
self,
matrix: &Tensor<'_>,
vector: &[f32],
activation: Option<&Q8Activation>,
result: &mut [f32],
) -> Result<()> {
match self {
Self::F32 => {
result
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = dot_f32(matrix.f32_row(row)?, vector);
Ok(())
})
}
Self::Bf16 => {
result
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = dot_bf16(matrix.bf16_row(row)?, vector);
Ok(())
})
}
Self::Quantized(format) => {
if let Some(act) = activation {
project_quantized(matrix, format, act, result)
} else {
let act = Q8Activation::new(vector)?;
project_quantized(matrix, format, &act, result)
}
}
}
}
fn project_batch(
self,
matrix: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
input_size: usize,
output: &mut [f32],
) -> Result<()> {
match self {
Self::F32 => project_f32_batch(matrix, vectors, output),
Self::Bf16 => project_bf16_batch(matrix, vectors, output),
Self::Quantized(format) => {
let activations = matrix_activations(vectors, row_count, input_size)?;
project_quantized_batch(matrix, format, &activations, output)
}
}
}
fn row_dot(
self,
matrix: &Tensor<'_>,
row: usize,
vector: &[f32],
activation: Option<&Q8Activation>,
kernel: IsaKernel,
) -> Result<f32> {
Ok(match self {
Self::F32 => dot_f32(matrix.f32_row(row)?, vector),
Self::Bf16 => dot_bf16(matrix.bf16_row(row)?, vector),
Self::Quantized(format) => kernel.dot(
format,
format.row(matrix, row)?,
activation.expect("quantized activation"),
),
})
}
}
#[derive(Clone, Copy)]
enum IsaKernel {
Scalar,
#[cfg(target_arch = "x86_64")]
Avx2,
#[cfg(target_arch = "x86_64")]
Avx512,
#[cfg(target_arch = "aarch64")]
Neon,
}
impl IsaKernel {
fn detect() -> Self {
static KERNEL: OnceLock<IsaKernel> = 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, format: QuantizedFormat, row: &[u8], activation: &Q8Activation) -> f32 {
match (self, format) {
(Self::Scalar, QuantizedFormat::Q8_0) => dot_q8_scalar(row, activation),
(Self::Scalar, QuantizedFormat::Mxfp4) => dot_mxfp4_scalar(row, activation),
#[cfg(target_arch = "x86_64")]
(Self::Avx2, QuantizedFormat::Q8_0) => unsafe { dot_q8_avx2(row, activation) },
#[cfg(target_arch = "x86_64")]
(Self::Avx2, QuantizedFormat::Mxfp4) => unsafe { dot_mxfp4_avx2(row, activation) },
#[cfg(target_arch = "x86_64")]
(Self::Avx512, QuantizedFormat::Q8_0) => unsafe { dot_q8_avx512(row, activation) },
#[cfg(target_arch = "x86_64")]
(Self::Avx512, QuantizedFormat::Mxfp4) => unsafe { dot_mxfp4_avx512(row, activation) },
#[cfg(target_arch = "aarch64")]
(Self::Neon, QuantizedFormat::Q8_0) => unsafe { dot_q8_neon(row, activation) },
#[cfg(target_arch = "aarch64")]
(Self::Neon, QuantizedFormat::Mxfp4) => unsafe { dot_mxfp4_neon(row, activation) },
}
}
fn dot_batch(
self,
format: QuantizedFormat,
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
match (self, format) {
(Self::Scalar, QuantizedFormat::Q8_0) => {
std::array::from_fn(|index| dot_q8_scalar(row, &activations[index]))
}
(Self::Scalar, QuantizedFormat::Mxfp4) => {
std::array::from_fn(|index| dot_mxfp4_scalar(row, &activations[index]))
}
#[cfg(target_arch = "x86_64")]
(Self::Avx2, QuantizedFormat::Q8_0) => unsafe { dot_q8_avx2_batch(row, activations) },
#[cfg(target_arch = "x86_64")]
(Self::Avx2, QuantizedFormat::Mxfp4) => unsafe {
dot_mxfp4_avx2_batch(row, activations)
},
#[cfg(target_arch = "x86_64")]
(Self::Avx512, QuantizedFormat::Q8_0) => unsafe {
dot_q8_avx512_batch(row, activations)
},
#[cfg(target_arch = "x86_64")]
(Self::Avx512, QuantizedFormat::Mxfp4) => unsafe {
dot_mxfp4_avx512_batch(row, activations)
},
#[cfg(target_arch = "aarch64")]
(Self::Neon, QuantizedFormat::Q8_0) => {
std::array::from_fn(|index| unsafe { dot_q8_neon(row, &activations[index]) })
}
#[cfg(target_arch = "aarch64")]
(Self::Neon, QuantizedFormat::Mxfp4) => {
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 multiply_i8_block_avx512(
weights: std::arch::x86_64::__m256i,
activation: std::arch::x86_64::__m256i,
) -> std::arch::x86_64::__m512i {
use std::arch::x86_64::{_mm512_cvtepi8_epi16, _mm512_madd_epi16};
let weights = _mm512_cvtepi8_epi16(weights);
let activation = _mm512_cvtepi8_epi16(activation);
_mm512_madd_epi16(weights, activation)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn horizontal_sum_i32_avx512(values: std::arch::x86_64::__m512i) -> f32 {
use std::arch::x86_64::_mm512_reduce_add_epi32;
#[allow(clippy::cast_precision_loss)]
{
_mm512_reduce_add_epi32(values) as f32
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn horizontal_sum_ps_avx512(values: std::arch::x86_64::__m512) -> f32 {
use std::arch::x86_64::_mm512_reduce_add_ps;
_mm512_reduce_add_ps(values)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn dot_q8_avx512(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::__m256i;
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 prod0 = unsafe { multiply_i8_block_avx512(w0, a0) };
let prod1 = unsafe { multiply_i8_block_avx512(w1, a1) };
let ws0 = Fp16::decode_le(&b0[0..2]).to_f32();
let ws1 = Fp16::decode_le(&b1[0..2]).to_f32();
sum += ws0 * s0 * unsafe { horizontal_sum_i32_avx512(prod0) };
sum += ws1 * s1 * unsafe { horizontal_sum_i32_avx512(prod1) };
}
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 prod0 = unsafe { multiply_i8_block_avx512(w0, a0) };
let ws0 = Fp16::decode_le(&b0[0..2]).to_f32();
sum += ws0 * s0 * unsafe { horizontal_sum_i32_avx512(prod0) };
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn decode_mxfp4_weights_avx512(block: &[u8]) -> std::arch::x86_64::__m256i {
use std::arch::x86_64::{
__m128i, _mm_and_si128, _mm_set1_epi8, _mm_shuffle_epi8, _mm_srli_epi16, _mm256_set_m128i,
};
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 packed: __m128i = bytemuck::pod_read_unaligned(&block[1..17]);
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);
_mm256_set_m128i(high, low)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn dot_mxfp4_avx512(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::__m256i;
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 w0 = unsafe { decode_mxfp4_weights_avx512(b0) };
let a0: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v0));
let w1 = unsafe { decode_mxfp4_weights_avx512(b1) };
let a1: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v1));
let prod0 = unsafe { multiply_i8_block_avx512(w0, a0) };
let prod1 = unsafe { multiply_i8_block_avx512(w1, a1) };
let sc0 = E8m0::decode_le(&b0[0..1]).to_f32();
let sc1 = E8m0::decode_le(&b1[0..1]).to_f32();
sum += sc0 * s0 * unsafe { horizontal_sum_i32_avx512(prod0) };
sum += sc1 * s1 * unsafe { horizontal_sum_i32_avx512(prod1) };
}
while let (Some(b0), Some(v0), Some(s0)) = (blocks.next(), values_chunks.next(), scales.next())
{
let w0 = unsafe { decode_mxfp4_weights_avx512(b0) };
let a0: __m256i = bytemuck::pod_read_unaligned(bytemuck::cast_slice(v0));
let prod0 = unsafe { multiply_i8_block_avx512(w0, a0) };
let sc0 = E8m0::decode_le(&b0[0..1]).to_f32();
sum += sc0 * s0 * unsafe { horizontal_sum_i32_avx512(prod0) };
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn dot_q8_avx512_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m256i, __m512, _mm512_add_ps, _mm512_cvtepi32_ps, _mm512_mul_ps, _mm512_set1_ps,
_mm512_setzero_ps,
};
let mut sums: [__m512; MATRIX_ROW_TILE] = [_mm512_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 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_avx512(weights, values) };
let scale = weight_scale * activations[index].scales[block_index];
sums[index] = _mm512_add_ps(
sums[index],
_mm512_mul_ps(_mm512_cvtepi32_ps(products), _mm512_set1_ps(scale)),
);
}
}
std::array::from_fn(|index| unsafe { horizontal_sum_ps_avx512(sums[index]) })
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn dot_mxfp4_avx512_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m256i, __m512, _mm512_add_ps, _mm512_cvtepi32_ps, _mm512_mul_ps, _mm512_set1_ps,
_mm512_setzero_ps,
};
let mut sums: [__m512; MATRIX_ROW_TILE] = [_mm512_setzero_ps(); MATRIX_ROW_TILE];
for (block_index, block) in row.chunks_exact(MXFP4_BLOCK_BYTES).enumerate() {
let weights = unsafe { decode_mxfp4_weights_avx512(block) };
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_avx512(weights, values) };
let scale = weight_scale * activations[index].scales[block_index];
sums[index] = _mm512_add_ps(
sums[index],
_mm512_mul_ps(_mm512_cvtepi32_ps(products), _mm512_set1_ps(scale)),
);
}
}
std::array::from_fn(|index| unsafe { horizontal_sum_ps_avx512(sums[index]) })
}
#[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];
WeightFormat::from_tensor_type(matrix.tensor_type()).project_vector(
matrix,
vector,
activation,
&mut result,
)?;
Ok(result)
}
pub(super) fn matrix_vector_pair(
first: &Tensor<'_>,
second: &Tensor<'_>,
vector: &[f32],
activation: Option<&Q8Activation>,
) -> 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, output] = first_dimensions;
ensure!(vector.len() == input, "matrix input width differs");
let first_format = WeightFormat::from_tensor_type(first.tensor_type()).quantized();
let second_format = WeightFormat::from_tensor_type(second.tensor_type()).quantized();
if let (Some(format), Some(other_format)) = (first_format, second_format)
&& format == other_format
{
let owned_activation;
let activation = if let Some(activation) = activation {
activation
} else {
owned_activation = Q8Activation::new(vector)?;
&owned_activation
};
let (first_output, second_output) = rayon::join(
|| {
let mut result = vec![0.0; output];
project_quantized(first, format, activation, &mut result).map(|()| result)
},
|| {
let mut result = vec![0.0; output];
project_quantized(second, format, activation, &mut result).map(|()| result)
},
);
Ok((first_output?, second_output?))
} else {
let (first_output, second_output) = rayon::join(
|| matrix_vector_with_activation(first, vector, activation),
|| matrix_vector_with_activation(second, vector, activation),
);
Ok((first_output?, second_output?))
}
}
pub(super) fn matrix_vector_triple(
first: &Tensor<'_>,
second: &Tensor<'_>,
third: &Tensor<'_>,
vector: &[f32],
) -> Result<(Vec<f32>, Vec<f32>, Vec<f32>)> {
let formats = (
WeightFormat::from_tensor_type(first.tensor_type()),
WeightFormat::from_tensor_type(second.tensor_type()),
WeightFormat::from_tensor_type(third.tensor_type()),
);
if let (
WeightFormat::Quantized(format),
WeightFormat::Quantized(second_format),
WeightFormat::Quantized(third_format),
) = formats
&& format == second_format
&& format == third_format
&& format == QuantizedFormat::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 = IsaKernel::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(format, first.q8_row(row)?, &activation);
} else if row < first_rows + second_rows {
*value = kernel.dot(format, second.q8_row(row - first_rows)?, &activation);
} else {
*value = kernel.dot(
format,
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];
WeightFormat::from_tensor_type(matrix.tensor_type()).project_batch(
matrix,
vectors,
row_count,
input_size,
&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)?;
let first_format = WeightFormat::from_tensor_type(first.tensor_type()).quantized();
let second_format = WeightFormat::from_tensor_type(second.tensor_type()).quantized();
if first_format.is_some() && first_format == second_format {
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?))
} else {
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);
}
let formats = (
WeightFormat::from_tensor_type(first.tensor_type()),
WeightFormat::from_tensor_type(second.tensor_type()),
WeightFormat::from_tensor_type(third.tensor_type()),
);
if matches!(
formats,
(
WeightFormat::Quantized(QuantizedFormat::Q8_0),
WeightFormat::Quantized(QuantizedFormat::Q8_0),
WeightFormat::Quantized(QuantizedFormat::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];
let Some(format) = WeightFormat::from_tensor_type(matrix.tensor_type()).quantized() else {
unreachable!("validated quantized matrix")
};
project_quantized_batch(matrix, format, activations, &mut output)?;
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_quantized_batch(
matrix: &Tensor<'_>,
format: QuantizedFormat,
activations: &[Q8Activation],
output: &mut [f32],
) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
ensure!(
QuantizedFormat::from_tensor_type(matrix.tensor_type()) == Some(format),
"projection format differs from tensor"
);
validate_quantized_batch(activations, output, input_size, output_size)?;
let kernel = IsaKernel::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(format, format.row(matrix, 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 = format.row(matrix, output_channel)?;
for (output_row, activation) in output_rows
.chunks_exact_mut(output_size)
.zip(activation_rows)
{
output_row[output_channel] = kernel.dot(format, 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 format = WeightFormat::from_tensor_type(matrix.tensor_type());
let q8_activation = if format.quantized().is_some() {
Some(Q8Activation::new(vector)?)
} else {
None
};
let kernel = IsaKernel::detect();
let best = (0..output)
.into_par_iter()
.map(|index| -> Result<(usize, f32)> {
let value = format.row_dot(matrix, index, vector, q8_activation.as_ref(), kernel)?;
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_quantized(
matrix: &Tensor<'_>,
format: QuantizedFormat,
activation: &Q8Activation,
output: &mut [f32],
) -> Result<()> {
let [input, rows] = matrix_dimensions(matrix)?;
ensure!(
QuantizedFormat::from_tensor_type(matrix.tensor_type()) == Some(format),
"projection format differs from tensor"
);
ensure!(
activation.values.len() == input,
"matrix input width differs"
);
ensure!(output.len() == rows, "matrix output height differs");
let kernel = IsaKernel::detect();
output
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = kernel.dot(format, format.row(matrix, 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 dot_f32_fp16(query: &[f32], key: &[Fp16]) -> f32 {
debug_assert_eq!(query.len(), key.len());
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2")
&& std::is_x86_feature_detected!("fma")
&& std::is_x86_feature_detected!("f16c")
{
return unsafe { dot_f32_fp16_avx2(query, key) };
}
query
.iter()
.zip(key)
.map(|(query, key)| query * f32::from(*key))
.sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn dot_f32_fp16_avx2(query: &[f32], key: &[Fp16]) -> f32 {
use std::arch::x86_64::{
__m128i, _mm256_cvtph_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps,
_mm256_storeu_ps,
};
let vectorized = query.len() / 8 * 8;
let mut sums = _mm256_setzero_ps();
for index in (0..vectorized).step_by(8) {
let query = unsafe { _mm256_loadu_ps(query.as_ptr().add(index)) };
let key_half: __m128i =
bytemuck::pod_read_unaligned(bytemuck::cast_slice(&key[index..index + 8]));
let key = _mm256_cvtph_ps(key_half);
sums = _mm256_fmadd_ps(query, key, sums);
}
let mut lanes = [0.0; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sums) };
lanes.into_iter().sum::<f32>()
+ query[vectorized..]
.iter()
.zip(&key[vectorized..])
.map(|(query, key)| query * f32::from(*key))
.sum::<f32>()
}
pub(super) fn accumulate_fp16(output: &mut [f32], values: &[Fp16], weight: f32) {
debug_assert_eq!(output.len(), values.len());
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2")
&& std::is_x86_feature_detected!("fma")
&& std::is_x86_feature_detected!("f16c")
{
unsafe { accumulate_fp16_avx2(output, values, weight) };
return;
}
for (output, value) in output.iter_mut().zip(values) {
*output += weight * f32::from(*value);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn accumulate_fp16_avx2(output: &mut [f32], values: &[Fp16], weight: f32) {
use std::arch::x86_64::{
__m128i, _mm256_cvtph_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_set1_ps,
_mm256_storeu_ps,
};
let vectorized = output.len() / 8 * 8;
let weight_vec = _mm256_set1_ps(weight);
for index in (0..vectorized).step_by(8) {
let current = unsafe { _mm256_loadu_ps(output.as_ptr().add(index)) };
let value_half: __m128i =
bytemuck::pod_read_unaligned(bytemuck::cast_slice(&values[index..index + 8]));
let value = _mm256_cvtph_ps(value_half);
let sum = _mm256_fmadd_ps(value, weight_vec, current);
unsafe { _mm256_storeu_ps(output.as_mut_ptr().add(index), sum) };
}
for index in vectorized..output.len() {
output[index] += weight * f32::from(values[index]);
}
}
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];
}
}
pub(super) fn swiglu_inplace(gate: &mut [f32], up: &[f32], alpha: f32, limit: f32) -> Result<()> {
ensure!(gate.len() == up.len(), "swiglu widths differ");
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") {
unsafe { swiglu_inplace_avx2(gate, up, alpha, limit) };
return Ok(());
}
for (gate, up) in gate.iter_mut().zip(up) {
*gate = (*gate).min(limit);
let up = (*up).clamp(-limit, limit);
let glu = *gate / (1.0 + (-alpha * *gate).exp());
*gate = (up + 1.0) * glu;
}
Ok(())
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn swiglu_inplace_avx2(gate: &mut [f32], up: &[f32], alpha: f32, limit: f32) {
use std::arch::x86_64::{
_mm256_add_ps, _mm256_div_ps, _mm256_loadu_ps, _mm256_max_ps, _mm256_min_ps, _mm256_mul_ps,
_mm256_set1_ps, _mm256_storeu_ps,
};
let limit_pos = _mm256_set1_ps(limit);
let limit_neg = _mm256_set1_ps(-limit);
let one = _mm256_set1_ps(1.0);
let vectorized = gate.len() / 8 * 8;
for index in (0..vectorized).step_by(8) {
let mut gate_vec = unsafe { _mm256_loadu_ps(gate.as_ptr().add(index)) };
let mut up_vec = unsafe { _mm256_loadu_ps(up.as_ptr().add(index)) };
gate_vec = _mm256_min_ps(gate_vec, limit_pos);
up_vec = _mm256_min_ps(_mm256_max_ps(up_vec, limit_neg), limit_pos);
let mut gate_lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(gate_lanes.as_mut_ptr(), gate_vec) };
let mut exp_lanes = [0.0_f32; 8];
for lane in 0..8 {
exp_lanes[lane] = (-alpha * gate_lanes[lane]).exp();
}
let exp_vec = unsafe { _mm256_loadu_ps(exp_lanes.as_ptr()) };
let denom = _mm256_add_ps(one, exp_vec);
let glu = _mm256_div_ps(gate_vec, denom);
let result = _mm256_mul_ps(_mm256_add_ps(up_vec, one), glu);
unsafe { _mm256_storeu_ps(gate.as_mut_ptr().add(index), result) };
}
for index in vectorized..gate.len() {
let clamped_gate = gate[index].min(limit);
let clamped_up = up[index].clamp(-limit, limit);
let glu = clamped_gate / (1.0 + (-alpha * clamped_gate).exp());
gate[index] = (clamped_up + 1.0) * glu;
}
}