use std::sync::OnceLock;
use anyhow::{Context as _, Result, bail, ensure};
use rayon::prelude::*;
use super::gguf::{Tensor, TensorType};
use super::types::Fp16;
const BLOCK_VALUES: usize = 32;
const BLOCK_BYTES: usize = 34;
const Q1_BLOCK_VALUES: usize = 128;
const Q1_BLOCK_BYTES: usize = 18;
const Q1_VALUES_PER_SIGN_BYTE: usize = 8;
const Q1_SIGN_BYTES_PER_Q8_BLOCK: usize = BLOCK_VALUES / Q1_VALUES_PER_SIGN_BYTE;
const Q1_Q8_BLOCKS: usize = Q1_BLOCK_VALUES / BLOCK_VALUES;
const MATRIX_ROW_TILE: usize = 4;
const WIDE_MATRIX_ROW_TILE: usize = 8;
const DOUBLE_WIDE_MATRIX_ROW_TILE: usize = 16;
#[cfg(target_arch = "x86_64")]
const Q1_SIGN_NIBBLE_LUT: [i32; 16] = q1_sign_nibble_lut();
#[cfg(target_arch = "x86_64")]
const fn q1_sign_nibble_lut() -> [i32; 16] {
let mut lut = [0_i32; 16];
let mut pattern = 0;
while pattern < lut.len() {
let mut packed = 0_u32;
let mut bit = 0;
while bit < 4 {
let sign = if pattern & (1 << bit) == 0 {
u8::MAX
} else {
1
};
packed |= (sign as u32) << (bit * 8);
bit += 1;
}
lut[pattern] = packed.cast_signed();
pattern += 1;
}
lut
}
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 with_capacity(vector_len: usize) -> Result<Self> {
ensure!(
vector_len.is_multiple_of(BLOCK_VALUES),
"Q8 activation width is not divisible by 32"
);
Ok(Self {
scales: Vec::with_capacity(vector_len / BLOCK_VALUES),
values: Vec::with_capacity(vector_len),
})
}
pub(super) fn new(vector: &[f32]) -> Result<Self> {
let mut activation = Self::with_capacity(vector.len())?;
activation.quantize_into(vector)?;
Ok(activation)
}
pub(super) fn quantize_into(&mut self, vector: &[f32]) -> Result<()> {
ensure!(
vector.len().is_multiple_of(BLOCK_VALUES),
"Q8 activation width is not divisible by 32"
);
self.scales.resize(vector.len() / BLOCK_VALUES, 0.0);
self.values.resize(vector.len(), 0);
IsaKernel::detect().quantize_into(vector, self)
}
pub(super) fn scale_repeating(
&mut self,
values_per_group: usize,
scales: &[f32],
) -> Result<()> {
ensure!(
values_per_group != 0 && values_per_group.is_multiple_of(BLOCK_VALUES),
"Q8 scale group width is not divisible by 32"
);
let blocks_per_group = values_per_group / BLOCK_VALUES;
ensure!(
self.scales.len() == scales.len() * blocks_per_group,
"Q8 scale group count differs"
);
for (block_scales, &scale) in self.scales.chunks_exact_mut(blocks_per_group).zip(scales) {
for block_scale in block_scales {
*block_scale *= scale;
}
}
Ok(())
}
fn validate_len(&self, vector_len: usize) -> Result<()> {
ensure!(
self.values.len() == vector_len && self.scales.len() == vector_len / BLOCK_VALUES,
"Q8 activation width differs"
);
Ok(())
}
}
#[derive(Default)]
struct Q8Activations {
rows: Vec<Q8Activation>,
segment_lengths: Vec<usize>,
packed: Option<PackedQ8Activations>,
}
struct PackedQ8Activations {
block_count: usize,
values: Vec<i8>,
scales: Vec<f32>,
tiles: Vec<PackedQ8ActivationTile>,
}
#[derive(Default)]
pub(super) struct MatrixWorkspace {
activations: Q8Activations,
columns: [Vec<f32>; 3],
}
#[derive(Clone, Copy)]
struct PackedQ8ActivationTile {
row_start: usize,
row_count: usize,
value_start: usize,
scale_start: usize,
}
impl Q8Activations {
fn new(
rows: Vec<Q8Activation>,
segment_lengths: &[usize],
packed_tile_rows: Option<usize>,
) -> Result<Self> {
let packed =
if let Some(tile_rows) = packed_tile_rows.filter(|_| rows.len() >= MATRIX_ROW_TILE) {
Some(PackedQ8Activations::new(&rows, segment_lengths, tile_rows)?)
} else {
None
};
Ok(Self {
rows,
segment_lengths: segment_lengths.to_vec(),
packed,
})
}
fn prepare(
&mut self,
vectors: &[f32],
row_count: usize,
input_size: usize,
segment_lengths: &[usize],
packed_tile_rows: Option<usize>,
) -> Result<()> {
validate_matrix_input(vectors, row_count, input_size)?;
ensure!(
segmented_row_count(segment_lengths)? == row_count,
"Q8 activation segment rows differ"
);
while self.rows.len() < row_count {
self.rows.push(Q8Activation::with_capacity(input_size)?);
}
self.rows.truncate(row_count);
for (row, values) in self.rows.iter_mut().zip(vectors.chunks_exact(input_size)) {
row.quantize_into(values)?;
}
self.segment_lengths.clear();
self.segment_lengths.extend_from_slice(segment_lengths);
if let Some(tile_rows) = packed_tile_rows {
match self.packed.as_mut() {
Some(packed) => packed.prepare(&self.rows, segment_lengths, tile_rows)?,
None => {
self.packed = Some(PackedQ8Activations::new(
&self.rows,
segment_lengths,
tile_rows,
)?);
}
}
} else {
self.packed = None;
}
Ok(())
}
fn len(&self) -> usize {
self.rows.len()
}
}
impl PackedQ8Activations {
fn new(rows: &[Q8Activation], segment_lengths: &[usize], max_tile_rows: usize) -> Result<Self> {
let mut packed = Self {
block_count: 0,
values: Vec::new(),
scales: Vec::new(),
tiles: Vec::new(),
};
packed.prepare(rows, segment_lengths, max_tile_rows)?;
Ok(packed)
}
fn prepare(
&mut self,
rows: &[Q8Activation],
segment_lengths: &[usize],
max_tile_rows: usize,
) -> Result<()> {
let first = rows.first().context("Q8 activation batch is empty")?;
self.block_count = first.scales.len();
ensure!(
rows.iter().all(|row| {
row.scales.len() == self.block_count
&& row.values.len() == self.block_count * BLOCK_VALUES
}),
"Q8 activation batch shape differs"
);
ensure!(
segment_lengths.iter().sum::<usize>() == rows.len(),
"Q8 activation segment rows differ"
);
ensure!(
max_tile_rows == WIDE_MATRIX_ROW_TILE || max_tile_rows == DOUBLE_WIDE_MATRIX_ROW_TILE,
"invalid packed Q8 maximum tile width"
);
self.values.clear();
self.scales.clear();
self.tiles.clear();
let mut row_start = 0;
while rows.len() - row_start >= MATRIX_ROW_TILE {
let remaining = rows.len() - row_start;
let row_count = if remaining >= max_tile_rows {
max_tile_rows
} else if remaining >= WIDE_MATRIX_ROW_TILE {
WIDE_MATRIX_ROW_TILE
} else {
MATRIX_ROW_TILE
};
self.push_tile(rows, row_start, row_count);
row_start += row_count;
}
Ok(())
}
fn push_tile(&mut self, rows: &[Q8Activation], row_start: usize, row_count: usize) {
let value_start = self.values.len();
let scale_start = self.scales.len();
let tile_rows = &rows[row_start..row_start + row_count];
let group_values = if row_count <= WIDE_MATRIX_ROW_TILE {
8
} else {
4
};
for block_index in 0..self.block_count {
let block_start = block_index * BLOCK_VALUES;
for group in 0..BLOCK_VALUES / group_values {
let group_start = block_start + group * group_values;
for row in tile_rows {
self.values
.extend_from_slice(&row.values[group_start..group_start + group_values]);
}
}
self.scales
.extend(tile_rows.iter().map(|row| row.scales[block_index]));
}
self.tiles.push(PackedQ8ActivationTile {
row_start,
row_count,
value_start,
scale_start,
});
}
fn values(&self, tile: PackedQ8ActivationTile) -> Result<&[i8]> {
let len = self
.block_count
.checked_mul(tile.row_count)
.and_then(|values| values.checked_mul(BLOCK_VALUES))
.context("packed Q8 tile value count overflow")?;
let end = tile
.value_start
.checked_add(len)
.context("packed Q8 tile value range overflow")?;
self.values
.get(tile.value_start..end)
.context("packed Q8 tile values are out of range")
}
fn scales(&self, tile: PackedQ8ActivationTile) -> Result<&[f32]> {
let len = self
.block_count
.checked_mul(tile.row_count)
.context("packed Q8 tile scale count overflow")?;
let end = tile
.scale_start
.checked_add(len)
.context("packed Q8 tile scale range overflow")?;
self.scales
.get(tile.scale_start..end)
.context("packed Q8 tile scales are out of range")
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum QuantizedFormat {
Q8_0,
Q1_0,
}
impl QuantizedFormat {
fn from_tensor_type(kind: TensorType) -> Option<Self> {
match kind {
TensorType::Q8_0 => Some(Self::Q8_0),
TensorType::Q1_0 => Some(Self::Q1_0),
TensorType::F32 => None,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum WeightFormat {
F32,
Quantized(QuantizedFormat),
}
impl WeightFormat {
fn from_tensor_type(kind: TensorType) -> Self {
match kind {
TensorType::F32 => Self::F32,
TensorType::Q8_0 => Self::Quantized(QuantizedFormat::Q8_0),
TensorType::Q1_0 => Self::Quantized(QuantizedFormat::Q1_0),
}
}
fn quantized(self) -> Option<QuantizedFormat> {
match self {
Self::Quantized(format) => Some(format),
Self::F32 => None,
}
}
fn project_vector(
self,
matrix: &Tensor<'_>,
vector: &[f32],
activation: Option<&Q8Activation>,
result: &mut [f32],
) -> Result<()> {
match self {
Self::F32 => project_f32(matrix, vector, result),
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::Quantized(format) => {
let kernel = IsaKernel::detect();
let activations = matrix_activations(
vectors,
row_count,
input_size,
&[row_count],
kernel.packed_tile_rows(),
)?;
project_quantized_batch(kernel, matrix, format, &activations, output)
}
}
}
}
#[derive(Clone, Copy)]
enum DenseIsaKernel {
Scalar,
#[cfg(target_arch = "x86_64")]
Avx2Fma,
}
#[derive(Clone, Copy)]
enum IsaKernel {
#[cfg(not(target_arch = "aarch64"))]
Scalar,
#[cfg(target_arch = "x86_64")]
Avx2,
#[cfg(target_arch = "x86_64")]
Avx2Fma,
#[cfg(target_arch = "x86_64")]
Avx512,
#[cfg(target_arch = "aarch64")]
Neon,
#[cfg(target_arch = "aarch64")]
NeonDotprod,
#[cfg(target_arch = "aarch64")]
NeonI8mm,
}
trait Sealed {}
trait DenseKernel: Sealed + Sync {
fn dot_f32(left: &[f32], right: &[f32]) -> f32;
}
trait ActivationKernel: Sealed {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32>;
}
trait QuantizedFormatMarker: Sealed {}
struct Q8_0Format;
struct Q1_0Format;
impl Sealed for Q8_0Format {}
impl Sealed for Q1_0Format {}
impl QuantizedFormatMarker for Q8_0Format {}
impl QuantizedFormatMarker for Q1_0Format {}
trait QuantizedKernel<F: QuantizedFormatMarker>: Sealed + Sync {
fn dot(row: &[u8], activation: &Q8Activation) -> f32;
fn dot_rows(
rows: [&[u8]; MATRIX_ROW_TILE],
activation: &Q8Activation,
) -> [f32; MATRIX_ROW_TILE] {
std::array::from_fn(|index| Self::dot(rows[index], activation))
}
fn dot_wide_rows(
rows: [&[u8]; WIDE_MATRIX_ROW_TILE],
activation: &Q8Activation,
) -> [f32; WIDE_MATRIX_ROW_TILE] {
std::array::from_fn(|index| Self::dot(rows[index], activation))
}
fn dot_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
std::array::from_fn(|index| Self::dot(row, &activations[index]))
}
fn dot_wide_batch(
_row: &[u8],
_activations: &[Q8Activation; WIDE_MATRIX_ROW_TILE],
) -> Option<[f32; WIDE_MATRIX_ROW_TILE]> {
None
}
fn dot_packed_batch(
_row: &[u8],
_values: &[i8],
_scales: &[f32],
) -> Option<[f32; MATRIX_ROW_TILE]> {
None
}
fn dot_packed_wide_batch(
_row: &[u8],
_values: &[i8],
_scales: &[f32],
) -> Option<[f32; WIDE_MATRIX_ROW_TILE]> {
None
}
fn dot_packed_double_wide_batch(
_row: &[u8],
_values: &[i8],
_scales: &[f32],
) -> Option<[f32; DOUBLE_WIDE_MATRIX_ROW_TILE]> {
None
}
fn supports_weight_row_pairs() -> bool {
false
}
fn dot_packed_batch_pair(
_rows: [&[u8]; 2],
_values: &[i8],
_scales: &[f32],
) -> Option<[[f32; MATRIX_ROW_TILE]; 2]> {
None
}
fn dot_packed_wide_batch_pair(
_rows: [&[u8]; 2],
_values: &[i8],
_scales: &[f32],
) -> Option<[[f32; WIDE_MATRIX_ROW_TILE]; 2]> {
None
}
}
struct Scalar;
impl Sealed for Scalar {}
impl DenseKernel for Scalar {
fn dot_f32(left: &[f32], right: &[f32]) -> f32 {
dot_f32_scalar(left, right)
}
}
impl ActivationKernel for Scalar {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
quantize_block_scalar(input, output)
}
}
impl QuantizedKernel<Q8_0Format> for Scalar {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
dot_q8_scalar(row, activation)
}
}
impl QuantizedKernel<Q1_0Format> for Scalar {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
dot_q1_q8_scalar(row, activation)
}
}
#[cfg(target_arch = "x86_64")]
struct Avx2;
#[cfg(target_arch = "x86_64")]
impl Sealed for Avx2 {}
#[cfg(target_arch = "x86_64")]
impl ActivationKernel for Avx2 {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
unsafe { quantize_block_avx2(input, output) }
}
}
#[cfg(target_arch = "x86_64")]
impl QuantizedKernel<Q8_0Format> for Avx2 {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_avx2(row, activation) }
}
fn dot_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
unsafe { dot_q8_avx2_batch(row, activations) }
}
}
#[cfg(target_arch = "x86_64")]
impl QuantizedKernel<Q1_0Format> for Avx2 {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q1_q8_avx2(row, activation) }
}
}
#[cfg(target_arch = "x86_64")]
struct Avx2Fma;
#[cfg(target_arch = "x86_64")]
impl Sealed for Avx2Fma {}
#[cfg(target_arch = "x86_64")]
impl ActivationKernel for Avx2Fma {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
unsafe { quantize_block_avx2(input, output) }
}
}
#[cfg(target_arch = "x86_64")]
impl DenseKernel for Avx2Fma {
fn dot_f32(left: &[f32], right: &[f32]) -> f32 {
unsafe { dot_f32_avx2(left, right) }
}
}
#[cfg(target_arch = "x86_64")]
impl QuantizedKernel<Q8_0Format> for Avx2Fma {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_avx2_fma(row, activation) }
}
fn dot_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
unsafe { dot_q8_avx2_fma_batch(row, activations) }
}
fn dot_wide_batch(
row: &[u8],
activations: &[Q8Activation; WIDE_MATRIX_ROW_TILE],
) -> Option<[f32; WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_avx2_fma_batch(row, activations) })
}
fn dot_packed_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_avx2_fma_packed_4(row, values, scales) })
}
fn dot_packed_wide_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_avx2_fma_packed_8(row, values, scales) })
}
fn dot_packed_double_wide_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; DOUBLE_WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_avx2_fma_packed_16(row, values, scales) })
}
}
#[cfg(target_arch = "x86_64")]
impl QuantizedKernel<Q1_0Format> for Avx2Fma {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q1_q8_avx2_fma(row, activation) }
}
fn dot_rows(
rows: [&[u8]; MATRIX_ROW_TILE],
activation: &Q8Activation,
) -> [f32; MATRIX_ROW_TILE] {
unsafe { dot_q1_q8_avx2_fma_rows(rows, activation) }
}
fn dot_wide_rows(
rows: [&[u8]; WIDE_MATRIX_ROW_TILE],
activation: &Q8Activation,
) -> [f32; WIDE_MATRIX_ROW_TILE] {
unsafe { dot_q1_q8_avx2_fma_rows(rows, activation) }
}
fn dot_packed_double_wide_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; DOUBLE_WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q1_q8_avx2_fma_packed_16(row, values, scales) })
}
}
#[cfg(target_arch = "x86_64")]
struct Avx512;
#[cfg(target_arch = "x86_64")]
impl Sealed for Avx512 {}
#[cfg(target_arch = "x86_64")]
impl ActivationKernel for Avx512 {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
unsafe { quantize_block_avx2(input, output) }
}
}
#[cfg(target_arch = "x86_64")]
impl QuantizedKernel<Q8_0Format> for Avx512 {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_avx512(row, activation) }
}
fn dot_batch(
row: &[u8],
activations: &[Q8Activation; MATRIX_ROW_TILE],
) -> [f32; MATRIX_ROW_TILE] {
unsafe { dot_q8_avx512_batch(row, activations) }
}
}
#[cfg(target_arch = "x86_64")]
impl QuantizedKernel<Q1_0Format> for Avx512 {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q1_q8_avx2(row, activation) }
}
fn dot_rows(
rows: [&[u8]; MATRIX_ROW_TILE],
activation: &Q8Activation,
) -> [f32; MATRIX_ROW_TILE] {
unsafe { dot_q1_q8_avx2_fma_rows(rows, activation) }
}
fn dot_wide_rows(
rows: [&[u8]; WIDE_MATRIX_ROW_TILE],
activation: &Q8Activation,
) -> [f32; WIDE_MATRIX_ROW_TILE] {
unsafe { dot_q1_q8_avx2_fma_rows(rows, activation) }
}
}
#[cfg(target_arch = "aarch64")]
struct Neon;
#[cfg(target_arch = "aarch64")]
struct NeonDotprod;
#[cfg(target_arch = "aarch64")]
struct NeonI8mm;
#[cfg(target_arch = "aarch64")]
impl Sealed for Neon {}
#[cfg(target_arch = "aarch64")]
impl Sealed for NeonDotprod {}
#[cfg(target_arch = "aarch64")]
impl Sealed for NeonI8mm {}
#[cfg(target_arch = "aarch64")]
impl ActivationKernel for Neon {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
unsafe { quantize_block_neon(input, output) }
}
}
#[cfg(target_arch = "aarch64")]
impl ActivationKernel for NeonDotprod {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
unsafe { quantize_block_neon(input, output) }
}
}
#[cfg(target_arch = "aarch64")]
impl ActivationKernel for NeonI8mm {
fn quantize_block(input: &[f32; BLOCK_VALUES], output: &mut [i8; BLOCK_VALUES]) -> Option<f32> {
unsafe { quantize_block_neon(input, output) }
}
}
#[cfg(target_arch = "aarch64")]
impl QuantizedKernel<Q8_0Format> for Neon {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_neon(row, activation) }
}
}
#[cfg(target_arch = "aarch64")]
impl QuantizedKernel<Q1_0Format> for Neon {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
dot_q1_q8_scalar(row, activation)
}
}
#[cfg(target_arch = "aarch64")]
impl QuantizedKernel<Q8_0Format> for NeonDotprod {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_neon_dotprod(row, activation) }
}
fn dot_packed_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_neon_dotprod_packed_8::<MATRIX_ROW_TILE>(row, values, scales) })
}
fn dot_packed_wide_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_neon_dotprod_packed_8::<WIDE_MATRIX_ROW_TILE>(row, values, scales) })
}
fn dot_packed_double_wide_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; DOUBLE_WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_neon_dotprod_packed_16(row, values, scales) })
}
}
#[cfg(target_arch = "aarch64")]
impl QuantizedKernel<Q1_0Format> for NeonDotprod {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
dot_q1_q8_scalar(row, activation)
}
}
#[cfg(target_arch = "aarch64")]
impl QuantizedKernel<Q8_0Format> for NeonI8mm {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
unsafe { dot_q8_neon_dotprod(row, activation) }
}
fn dot_packed_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_neon_dotprod_packed_8::<MATRIX_ROW_TILE>(row, values, scales) })
}
fn dot_packed_wide_batch(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> Option<[f32; WIDE_MATRIX_ROW_TILE]> {
Some(unsafe { dot_q8_neon_dotprod_packed_8::<WIDE_MATRIX_ROW_TILE>(row, values, scales) })
}
fn supports_weight_row_pairs() -> bool {
true
}
fn dot_packed_batch_pair(
rows: [&[u8]; 2],
values: &[i8],
scales: &[f32],
) -> Option<[[f32; MATRIX_ROW_TILE]; 2]> {
Some(unsafe { dot_q8_neon_i8mm_packed_pair::<MATRIX_ROW_TILE>(rows, values, scales) })
}
fn dot_packed_wide_batch_pair(
rows: [&[u8]; 2],
values: &[i8],
scales: &[f32],
) -> Option<[[f32; WIDE_MATRIX_ROW_TILE]; 2]> {
Some(unsafe { dot_q8_neon_i8mm_packed_pair::<WIDE_MATRIX_ROW_TILE>(rows, values, scales) })
}
}
#[cfg(target_arch = "aarch64")]
impl QuantizedKernel<Q1_0Format> for NeonI8mm {
fn dot(row: &[u8], activation: &Q8Activation) -> f32 {
dot_q1_q8_scalar(row, activation)
}
}
impl DenseIsaKernel {
fn detect() -> Self {
static KERNEL: OnceLock<DenseIsaKernel> = OnceLock::new();
*KERNEL.get_or_init(|| {
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
return Self::Avx2Fma;
}
Self::Scalar
})
}
fn project(self, matrix: &Tensor<'_>, vector: &[f32], output: &mut [f32]) -> Result<()> {
match self {
Self::Scalar => project_f32_with::<Scalar>(matrix, vector, output),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => project_f32_with::<Avx2Fma>(matrix, vector, output),
}
}
fn project_batch(
self,
matrix: &Tensor<'_>,
vectors: &[f32],
input_size: usize,
output_size: usize,
output: &mut [f32],
) -> Result<()> {
match self {
Self::Scalar => {
project_f32_batch_with::<Scalar>(matrix, vectors, input_size, output_size, output)
}
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => {
project_f32_batch_with::<Avx2Fma>(matrix, vectors, input_size, output_size, output)
}
}
}
fn argmax(self, matrix: &Tensor<'_>, vector: &[f32], output_size: usize) -> Result<usize> {
match self {
Self::Scalar => matrix_argmax_f32_with::<Scalar>(matrix, vector, output_size),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => matrix_argmax_f32_with::<Avx2Fma>(matrix, vector, output_size),
}
}
fn dot_f32(self, left: &[f32], right: &[f32]) -> f32 {
match self {
Self::Scalar => <Scalar as DenseKernel>::dot_f32(left, right),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => <Avx2Fma as DenseKernel>::dot_f32(left, right),
}
}
}
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!("avx2")
&& 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") && std::is_x86_feature_detected!("fma") {
return Self::Avx2Fma;
}
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx2") {
return Self::Avx2;
}
#[cfg(target_arch = "aarch64")]
if std::arch::is_aarch64_feature_detected!("i8mm")
&& std::arch::is_aarch64_feature_detected!("dotprod")
{
return Self::NeonI8mm;
}
#[cfg(target_arch = "aarch64")]
if std::arch::is_aarch64_feature_detected!("dotprod") {
return Self::NeonDotprod;
}
#[cfg(target_arch = "aarch64")]
{
Self::Neon
}
#[cfg(not(target_arch = "aarch64"))]
{
Self::Scalar
}
})
}
fn packed_tile_rows(self) -> Option<usize> {
match self {
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => Some(DOUBLE_WIDE_MATRIX_ROW_TILE),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => Some(DOUBLE_WIDE_MATRIX_ROW_TILE),
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => Some(WIDE_MATRIX_ROW_TILE),
_ => None,
}
}
fn quantize_into(self, vector: &[f32], activation: &mut Q8Activation) -> Result<()> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => quantize_activation_with::<Scalar>(vector, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => quantize_activation_with::<Avx2>(vector, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => quantize_activation_with::<Avx2Fma>(vector, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx512 => quantize_activation_with::<Avx512>(vector, activation),
#[cfg(target_arch = "aarch64")]
Self::Neon => quantize_activation_with::<Neon>(vector, activation),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => quantize_activation_with::<NeonDotprod>(vector, activation),
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => quantize_activation_with::<NeonI8mm>(vector, activation),
}
}
fn rms_norm_quantize_into(
self,
values: &[f32],
width: usize,
weight: &[f32],
epsilon: f32,
activation: &mut Q8Activation,
) -> Result<()> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => rms_norm_quantize_activation_with::<Scalar>(
values, width, weight, epsilon, activation,
),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => rms_norm_quantize_activation_with::<Avx2>(
values, width, weight, epsilon, activation,
),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => rms_norm_quantize_activation_with::<Avx2Fma>(
values, width, weight, epsilon, activation,
),
#[cfg(target_arch = "x86_64")]
Self::Avx512 => rms_norm_quantize_activation_with::<Avx512>(
values, width, weight, epsilon, activation,
),
#[cfg(target_arch = "aarch64")]
Self::Neon => rms_norm_quantize_activation_with::<Neon>(
values, width, weight, epsilon, activation,
),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => rms_norm_quantize_activation_with::<NeonDotprod>(
values, width, weight, epsilon, activation,
),
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => rms_norm_quantize_activation_with::<NeonI8mm>(
values, width, weight, epsilon, activation,
),
}
}
fn project(
self,
matrix: &Tensor<'_>,
format: QuantizedFormat,
activation: &Q8Activation,
output: &mut [f32],
) -> Result<()> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => project_quantized_with::<Scalar>(matrix, format, activation, output),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => project_quantized_with::<Avx2>(matrix, format, activation, output),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => project_quantized_with::<Avx2Fma>(matrix, format, activation, output),
#[cfg(target_arch = "x86_64")]
Self::Avx512 => project_quantized_with::<Avx512>(matrix, format, activation, output),
#[cfg(target_arch = "aarch64")]
Self::Neon => project_quantized_with::<Neon>(matrix, format, activation, output),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => {
project_quantized_with::<NeonDotprod>(matrix, format, activation, output)
}
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => {
project_quantized_with::<NeonI8mm>(matrix, format, activation, output)
}
}
}
fn project_q1_add(
self,
matrix: &Tensor<'_>,
activation: &Q8Activation,
dest: &mut [f32],
) -> Result<()> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => {
project_quantized_q1_rows_with::<Scalar>(matrix, activation, dest, true)
}
#[cfg(target_arch = "x86_64")]
Self::Avx2 => project_quantized_q1_rows_with::<Avx2>(matrix, activation, dest, true),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => {
project_quantized_q1_rows_with::<Avx2Fma>(matrix, activation, dest, true)
}
#[cfg(target_arch = "x86_64")]
Self::Avx512 => {
project_quantized_q1_rows_with::<Avx512>(matrix, activation, dest, true)
}
#[cfg(target_arch = "aarch64")]
Self::Neon => project_quantized_q1_rows_with::<Neon>(matrix, activation, dest, true),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => {
project_quantized_q1_rows_with::<NeonDotprod>(matrix, activation, dest, true)
}
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => {
project_quantized_q1_rows_with::<NeonI8mm>(matrix, activation, dest, true)
}
}
}
fn project_q1_swiglu(
self,
gate: &Tensor<'_>,
up: &Tensor<'_>,
input: &Q8Activation,
output: &mut Q8Activation,
) -> Result<()> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => project_quantized_q1_swiglu_with::<Scalar>(gate, up, input, output),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => project_quantized_q1_swiglu_with::<Avx2>(gate, up, input, output),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => project_quantized_q1_swiglu_with::<Avx2Fma>(gate, up, input, output),
#[cfg(target_arch = "x86_64")]
Self::Avx512 => project_quantized_q1_swiglu_with::<Avx512>(gate, up, input, output),
#[cfg(target_arch = "aarch64")]
Self::Neon => project_quantized_q1_swiglu_with::<Neon>(gate, up, input, output),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => {
project_quantized_q1_swiglu_with::<NeonDotprod>(gate, up, input, output)
}
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => project_quantized_q1_swiglu_with::<NeonI8mm>(gate, up, input, output),
}
}
fn project_batch(
self,
matrix: &Tensor<'_>,
format: QuantizedFormat,
activations: &Q8Activations,
output: &mut [f32],
columns: Option<&mut Vec<f32>>,
) -> Result<()> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => {
project_quantized_batch_with::<Scalar>(matrix, format, activations, output, columns)
}
#[cfg(target_arch = "x86_64")]
Self::Avx2 => {
project_quantized_batch_with::<Avx2>(matrix, format, activations, output, columns)
}
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => project_quantized_batch_with::<Avx2Fma>(
matrix,
format,
activations,
output,
columns,
),
#[cfg(target_arch = "x86_64")]
Self::Avx512 => {
project_quantized_batch_with::<Avx512>(matrix, format, activations, output, columns)
}
#[cfg(target_arch = "aarch64")]
Self::Neon => {
project_quantized_batch_with::<Neon>(matrix, format, activations, output, columns)
}
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => project_quantized_batch_with::<NeonDotprod>(
matrix,
format,
activations,
output,
columns,
),
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => project_quantized_batch_with::<NeonI8mm>(
matrix,
format,
activations,
output,
columns,
),
}
}
fn argmax(
self,
matrix: &Tensor<'_>,
format: QuantizedFormat,
activation: &Q8Activation,
) -> Result<usize> {
match self {
#[cfg(not(target_arch = "aarch64"))]
Self::Scalar => matrix_argmax_quantized_with::<Scalar>(matrix, format, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx2 => matrix_argmax_quantized_with::<Avx2>(matrix, format, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx2Fma => matrix_argmax_quantized_with::<Avx2Fma>(matrix, format, activation),
#[cfg(target_arch = "x86_64")]
Self::Avx512 => matrix_argmax_quantized_with::<Avx512>(matrix, format, activation),
#[cfg(target_arch = "aarch64")]
Self::Neon => matrix_argmax_quantized_with::<Neon>(matrix, format, activation),
#[cfg(target_arch = "aarch64")]
Self::NeonDotprod => {
matrix_argmax_quantized_with::<NeonDotprod>(matrix, format, activation)
}
#[cfg(target_arch = "aarch64")]
Self::NeonI8mm => matrix_argmax_quantized_with::<NeonI8mm>(matrix, format, activation),
}
}
}
fn quantize_activation_with<K: ActivationKernel>(
vector: &[f32],
activation: &mut Q8Activation,
) -> Result<()> {
ensure!(
activation.scales.len() == vector.len() / BLOCK_VALUES
&& activation.values.len() == vector.len(),
"Q8 activation output shape differs"
);
for (index, (input, output)) in vector
.chunks_exact(BLOCK_VALUES)
.zip(activation.values.chunks_exact_mut(BLOCK_VALUES))
.enumerate()
{
let input =
<&[f32; BLOCK_VALUES]>::try_from(input).context("Q8 input block width differs")?;
let output =
<&mut [i8; BLOCK_VALUES]>::try_from(output).context("Q8 output block width differs")?;
let Some(scale) = K::quantize_block(input, output) else {
bail!("Q8 activation is not finite");
};
activation.scales[index] = scale;
}
Ok(())
}
fn rms_norm_quantize_activation_with<K: ActivationKernel>(
values: &[f32],
width: usize,
weight: &[f32],
epsilon: f32,
activation: &mut Q8Activation,
) -> Result<()> {
ensure!(
activation.scales.len() == values.len() / BLOCK_VALUES
&& activation.values.len() == values.len(),
"Q8 activation output shape differs"
);
let width_f32 = f32::from(u16::try_from(width).context("RMS norm width exceeds u16")?);
let mut block = [0.0_f32; BLOCK_VALUES];
let mut scales = activation.scales.iter_mut();
let mut quantized = activation.values.chunks_exact_mut(BLOCK_VALUES);
for input in values.chunks_exact(width) {
let mean_square = dot_f32(input, input)? / width_f32;
let scale = (mean_square + epsilon).sqrt().recip();
for (input_block, weight_block) in input
.chunks_exact(BLOCK_VALUES)
.zip(weight.chunks_exact(BLOCK_VALUES))
{
for ((normalized, &value), &gain) in block.iter_mut().zip(input_block).zip(weight_block)
{
*normalized = value * scale * gain;
}
let output = quantized.next().context("Q8 output block width differs")?;
let output = <&mut [i8; BLOCK_VALUES]>::try_from(output)
.context("Q8 output block width differs")?;
let scale = K::quantize_block(&block, output).context("Q8 activation is not finite")?;
*scales.next().context("Q8 scale count differs")? = scale;
}
}
Ok(())
}
fn quantize_block_scalar(
input: &[f32; BLOCK_VALUES],
output: &mut [i8; BLOCK_VALUES],
) -> Option<f32> {
let mut maximum = 0.0_f32;
for &value in input {
if !value.is_finite() {
return None;
}
maximum = maximum.max(value.abs());
}
let scale = maximum / 127.0;
let inverse = if scale == 0.0 { 0.0 } else { scale.recip() };
for (quantized, &value) in output.iter_mut().zip(input) {
let value = (value * inverse).round().clamp(-127.0, 127.0);
*quantized = unsafe { value.to_int_unchecked::<i8>() };
}
Some(scale)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn quantize_block_avx2(
input: &[f32; BLOCK_VALUES],
output: &mut [i8; BLOCK_VALUES],
) -> Option<f32> {
use std::arch::x86_64::{
__m128i, __m256, __m256i, _CMP_EQ_OQ, _CMP_LE_OQ, _CMP_LT_OQ, _MM_FROUND_NO_EXC,
_MM_FROUND_TO_NEAREST_INT, _mm_packs_epi16, _mm_packs_epi32, _mm256_add_ps, _mm256_and_ps,
_mm256_andnot_ps, _mm256_castsi256_si128, _mm256_cmp_ps, _mm256_cvttps_epi32,
_mm256_extracti128_si256, _mm256_loadu_ps, _mm256_max_epi32, _mm256_max_ps,
_mm256_min_epi32, _mm256_movemask_ps, _mm256_mul_ps, _mm256_or_ps, _mm256_round_ps,
_mm256_set_m128i, _mm256_set1_epi32, _mm256_set1_ps, _mm256_sub_ps,
};
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn round_away(value: __m256) -> __m256i {
let nearest = _mm256_round_ps::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(value);
let sign_mask = _mm256_set1_ps(-0.0);
let absolute_value = _mm256_andnot_ps(sign_mask, value);
let absolute_nearest = _mm256_andnot_ps(sign_mask, nearest);
let difference = _mm256_andnot_ps(sign_mask, _mm256_sub_ps(value, nearest));
let ties = _mm256_cmp_ps::<_CMP_EQ_OQ>(difference, _mm256_set1_ps(0.5));
let rounded_toward_zero = _mm256_cmp_ps::<_CMP_LT_OQ>(absolute_nearest, absolute_value);
let correction_mask = _mm256_and_ps(ties, rounded_toward_zero);
let signed_one = _mm256_or_ps(_mm256_and_ps(value, sign_mask), _mm256_set1_ps(1.0));
let rounded = _mm256_add_ps(nearest, _mm256_and_ps(correction_mask, signed_one));
let quantized = _mm256_cvttps_epi32(rounded);
_mm256_min_epi32(
_mm256_max_epi32(quantized, _mm256_set1_epi32(-127)),
_mm256_set1_epi32(127),
)
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn pack_pair(first: __m256i, second: __m256i) -> __m128i {
let first = _mm_packs_epi32(
_mm256_castsi256_si128(first),
_mm256_extracti128_si256::<1>(first),
);
let second = _mm_packs_epi32(
_mm256_castsi256_si128(second),
_mm256_extracti128_si256::<1>(second),
);
_mm_packs_epi16(first, second)
}
let sign_mask = _mm256_set1_ps(-0.0);
let maximum_finite = _mm256_set1_ps(f32::MAX);
let mut maximum = _mm256_set1_ps(0.0);
for offset in (0..BLOCK_VALUES).step_by(8) {
let value = unsafe { _mm256_loadu_ps(input.as_ptr().add(offset)) };
let absolute = _mm256_andnot_ps(sign_mask, value);
let finite = _mm256_cmp_ps::<_CMP_LE_OQ>(absolute, maximum_finite);
if _mm256_movemask_ps(finite) != 0xff {
return None;
}
maximum = _mm256_max_ps(maximum, absolute);
}
let maximum_lanes: [f32; 8] = bytemuck::cast(maximum);
let maximum = maximum_lanes.into_iter().fold(0.0_f32, f32::max);
let scale = maximum / 127.0;
if scale == 0.0 {
output.fill(0);
return Some(scale);
}
let inverse = _mm256_set1_ps(scale.recip());
let quantized: [__m256i; 4] = std::array::from_fn(|index| {
let value = unsafe { _mm256_loadu_ps(input.as_ptr().add(index * 8)) };
unsafe { round_away(_mm256_mul_ps(value, inverse)) }
});
let low = unsafe { pack_pair(quantized[0], quantized[1]) };
let high = unsafe { pack_pair(quantized[2], quantized[3]) };
*output = bytemuck::cast(_mm256_set_m128i(high, low));
Some(scale)
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn quantize_block_neon(
input: &[f32; BLOCK_VALUES],
output: &mut [i8; BLOCK_VALUES],
) -> Option<f32> {
use std::arch::aarch64::{
vabsq_f32, vcleq_f32, vcombine_s16, vcvtaq_s32_f32, vdupq_n_f32, vdupq_n_s32, vld1q_f32,
vmaxq_f32, vmaxq_s32, vmaxvq_f32, vminq_s32, vminvq_u32, vmulq_n_f32, vqmovn_s16,
vqmovn_s32, vst1_s8,
};
let mut maximum = vdupq_n_f32(0.0);
for offset in (0..BLOCK_VALUES).step_by(4) {
let absolute = vabsq_f32(unsafe { vld1q_f32(input.as_ptr().add(offset)) });
if vminvq_u32(vcleq_f32(absolute, vdupq_n_f32(f32::MAX))) == 0 {
return None;
}
maximum = vmaxq_f32(maximum, absolute);
}
let maximum = vmaxvq_f32(maximum);
let scale = maximum / 127.0;
if scale == 0.0 {
output.fill(0);
return Some(scale);
}
let inverse = scale.recip();
let minimum = vdupq_n_s32(-127);
let maximum = vdupq_n_s32(127);
for offset in (0..BLOCK_VALUES).step_by(8) {
let first = vcvtaq_s32_f32(vmulq_n_f32(
unsafe { vld1q_f32(input.as_ptr().add(offset)) },
inverse,
));
let second = vcvtaq_s32_f32(vmulq_n_f32(
unsafe { vld1q_f32(input.as_ptr().add(offset + 4)) },
inverse,
));
let first = vminq_s32(vmaxq_s32(first, minimum), maximum);
let second = vminq_s32(vmaxq_s32(second, minimum), maximum);
let packed = vqmovn_s16(vcombine_s16(vqmovn_s32(first), vqmovn_s32(second)));
unsafe { vst1_s8(output.as_mut_ptr().add(offset), packed) };
}
Some(scale)
}
fn i32_to_f32(value: i32) -> f32 {
let bytes = value.to_le_bytes();
let low = u16::from_le_bytes([bytes[0], bytes[1]]);
let high = i16::from_le_bytes([bytes[2], bytes[3]]);
f32::from(high).mul_add(65_536.0, f32::from(low))
}
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();
let sum = i32_to_f32(sum_i32);
weight_scale * activation_scale * sum
})
.sum()
}
fn dot_q1_q8_scalar(row: &[u8], activation: &Q8Activation) -> f32 {
let mut total = 0.0_f32;
let mut value_blocks = activation.values.chunks_exact(BLOCK_VALUES);
let mut scales = activation.scales.iter().copied();
for block in row.chunks_exact(Q1_BLOCK_BYTES) {
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
let signs = &block[2..];
let mut weighted_sum = 0.0_f32;
for activation_block in 0..Q1_Q8_BLOCKS {
let Some(values) = value_blocks.next() else {
return total;
};
let Some(activation_scale) = scales.next() else {
return total;
};
let mut sum_i32 = 0_i32;
let sign_start = activation_block * Q1_SIGN_BYTES_PER_Q8_BLOCK;
for sign_byte in 0..Q1_SIGN_BYTES_PER_Q8_BLOCK {
let bits = signs[sign_start + sign_byte];
let value_start = sign_byte * Q1_VALUES_PER_SIGN_BYTE;
for bit in 0..Q1_VALUES_PER_SIGN_BYTE {
let value = i32::from(values[value_start + bit]);
sum_i32 += if bits & (1_u8 << bit) == 0 {
-value
} else {
value
};
}
}
weighted_sum += activation_scale * i32_to_f32(sum_i32);
}
total += weight_scale * weighted_sum;
}
total
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_q1_q8_avx2(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::{
_mm256_add_ps, _mm256_and_si256, _mm256_cmpeq_epi8, _mm256_cvtepi32_ps, _mm256_loadu_si256,
_mm256_madd_epi16, _mm256_maddubs_epi16, _mm256_mul_ps, _mm256_set1_epi8,
_mm256_set1_epi16, _mm256_set1_epi32, _mm256_set1_ps, _mm256_setr_epi8, _mm256_setzero_ps,
_mm256_setzero_si256, _mm256_shuffle_epi8, _mm256_storeu_ps, _mm256_sub_epi8,
_mm256_xor_si256,
};
let ones_8 = _mm256_set1_epi8(1);
let ones_16 = _mm256_set1_epi16(1);
let byte_shuf = _mm256_setr_epi8(
0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3,
3, 3,
);
let bit_masks = _mm256_setr_epi8(
1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128,
1, 2, 4, 8, 16, 32, 64, -128,
);
let zero = _mm256_setzero_si256();
let mut acc = _mm256_setzero_ps();
let values = activation.values.as_slice();
let scales = activation.scales.as_slice();
let block_count = row.len() / Q1_BLOCK_BYTES;
if values.len() < block_count * Q1_BLOCK_VALUES || scales.len() < block_count * Q1_Q8_BLOCKS {
return 0.0;
}
let values_ptr = values.as_ptr();
let scales_ptr = scales.as_ptr();
for block_index in 0..block_count {
let block = unsafe { row.as_ptr().add(block_index * Q1_BLOCK_BYTES) };
let scale_bits = u16::from_le(unsafe { std::ptr::read_unaligned(block.cast()) });
let weight_scale = _mm256_set1_ps(Fp16(scale_bits).to_f32());
let signs = unsafe { block.add(2) };
let value_base = block_index * Q1_BLOCK_VALUES;
let scale_base = block_index * Q1_Q8_BLOCKS;
let mut acc_block = {
let qy = unsafe { _mm256_loadu_si256(values_ptr.add(value_base).cast()) };
let qs32 = i32::from_le(unsafe { std::ptr::read_unaligned(signs.cast()) });
let sm = _mm256_cmpeq_epi8(
_mm256_and_si256(
_mm256_shuffle_epi8(_mm256_set1_epi32(qs32), byte_shuf),
bit_masks,
),
zero,
);
let sy = _mm256_sub_epi8(_mm256_xor_si256(qy, sm), sm);
let s32 = _mm256_madd_epi16(_mm256_maddubs_epi16(ones_8, sy), ones_16);
let act_scale = unsafe { *scales_ptr.add(scale_base) };
_mm256_mul_ps(_mm256_set1_ps(act_scale), _mm256_cvtepi32_ps(s32))
};
for q8 in 1..Q1_Q8_BLOCKS {
let qy = unsafe {
_mm256_loadu_si256(values_ptr.add(value_base + q8 * BLOCK_VALUES).cast())
};
let qs32 = i32::from_le(unsafe {
std::ptr::read_unaligned(signs.add(q8 * Q1_SIGN_BYTES_PER_Q8_BLOCK).cast())
});
let sm = _mm256_cmpeq_epi8(
_mm256_and_si256(
_mm256_shuffle_epi8(_mm256_set1_epi32(qs32), byte_shuf),
bit_masks,
),
zero,
);
let sy = _mm256_sub_epi8(_mm256_xor_si256(qy, sm), sm);
let s32 = _mm256_madd_epi16(_mm256_maddubs_epi16(ones_8, sy), ones_16);
let act_scale = unsafe { *scales_ptr.add(scale_base + q8) };
let sub = _mm256_mul_ps(_mm256_set1_ps(act_scale), _mm256_cvtepi32_ps(s32));
acc_block = _mm256_add_ps(acc_block, sub);
}
acc = _mm256_add_ps(acc, _mm256_mul_ps(weight_scale, acc_block));
}
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), acc) };
lanes.into_iter().sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_q1_q8_avx2_fma(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::{
_mm256_and_si256, _mm256_cmpeq_epi8, _mm256_cvtepi32_ps, _mm256_fmadd_ps,
_mm256_loadu_si256, _mm256_madd_epi16, _mm256_maddubs_epi16, _mm256_mul_ps,
_mm256_set1_epi8, _mm256_set1_epi16, _mm256_set1_epi32, _mm256_set1_ps, _mm256_setr_epi8,
_mm256_setzero_ps, _mm256_setzero_si256, _mm256_shuffle_epi8, _mm256_storeu_ps,
_mm256_sub_epi8, _mm256_xor_si256,
};
let ones_8 = _mm256_set1_epi8(1);
let ones_16 = _mm256_set1_epi16(1);
let byte_shuf = _mm256_setr_epi8(
0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3,
3, 3,
);
let bit_masks = _mm256_setr_epi8(
1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128,
1, 2, 4, 8, 16, 32, 64, -128,
);
let zero = _mm256_setzero_si256();
let mut acc = _mm256_setzero_ps();
let values = activation.values.as_slice();
let scales = activation.scales.as_slice();
let block_count = row.len() / Q1_BLOCK_BYTES;
if values.len() < block_count * Q1_BLOCK_VALUES || scales.len() < block_count * Q1_Q8_BLOCKS {
return 0.0;
}
let values_ptr = values.as_ptr();
let scales_ptr = scales.as_ptr();
for block_index in 0..block_count {
let block = unsafe { row.as_ptr().add(block_index * Q1_BLOCK_BYTES) };
let scale_bits = u16::from_le(unsafe { std::ptr::read_unaligned(block.cast()) });
let weight_scale = _mm256_set1_ps(Fp16(scale_bits).to_f32());
let signs = unsafe { block.add(2) };
let value_base = block_index * Q1_BLOCK_VALUES;
let scale_base = block_index * Q1_Q8_BLOCKS;
let mut acc_block = {
let qy = unsafe { _mm256_loadu_si256(values_ptr.add(value_base).cast()) };
let qs32 = i32::from_le(unsafe { std::ptr::read_unaligned(signs.cast()) });
let sm = _mm256_cmpeq_epi8(
_mm256_and_si256(
_mm256_shuffle_epi8(_mm256_set1_epi32(qs32), byte_shuf),
bit_masks,
),
zero,
);
let sy = _mm256_sub_epi8(_mm256_xor_si256(qy, sm), sm);
let s32 = _mm256_madd_epi16(_mm256_maddubs_epi16(ones_8, sy), ones_16);
let act_scale = unsafe { *scales_ptr.add(scale_base) };
_mm256_mul_ps(_mm256_set1_ps(act_scale), _mm256_cvtepi32_ps(s32))
};
for q8 in 1..Q1_Q8_BLOCKS {
let qy = unsafe {
_mm256_loadu_si256(values_ptr.add(value_base + q8 * BLOCK_VALUES).cast())
};
let qs32 = i32::from_le(unsafe {
std::ptr::read_unaligned(signs.add(q8 * Q1_SIGN_BYTES_PER_Q8_BLOCK).cast())
});
let sm = _mm256_cmpeq_epi8(
_mm256_and_si256(
_mm256_shuffle_epi8(_mm256_set1_epi32(qs32), byte_shuf),
bit_masks,
),
zero,
);
let sy = _mm256_sub_epi8(_mm256_xor_si256(qy, sm), sm);
let s32 = _mm256_madd_epi16(_mm256_maddubs_epi16(ones_8, sy), ones_16);
let act_scale = unsafe { *scales_ptr.add(scale_base + q8) };
acc_block = _mm256_fmadd_ps(
_mm256_set1_ps(act_scale),
_mm256_cvtepi32_ps(s32),
acc_block,
);
}
acc = _mm256_fmadd_ps(weight_scale, acc_block, acc);
}
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), acc) };
lanes.into_iter().sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_q1_q8_avx2_fma_rows<const N: usize>(
rows: [&[u8]; N],
activation: &Q8Activation,
) -> [f32; N] {
use std::arch::x86_64::{
_mm256_and_si256, _mm256_cmpeq_epi8, _mm256_cvtepi32_ps, _mm256_fmadd_ps,
_mm256_loadu_si256, _mm256_madd_epi16, _mm256_maddubs_epi16, _mm256_set1_epi8,
_mm256_set1_epi16, _mm256_set1_epi32, _mm256_set1_ps, _mm256_setr_epi8, _mm256_setzero_ps,
_mm256_setzero_si256, _mm256_shuffle_epi8, _mm256_storeu_ps, _mm256_sub_epi8,
_mm256_xor_si256,
};
let ones_8 = _mm256_set1_epi8(1);
let ones_16 = _mm256_set1_epi16(1);
let byte_shuf = _mm256_setr_epi8(
0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3,
3, 3,
);
let bit_masks = _mm256_setr_epi8(
1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128,
1, 2, 4, 8, 16, 32, 64, -128,
);
let zero = _mm256_setzero_si256();
let mut acc = [_mm256_setzero_ps(); N];
let values = activation.values.as_slice();
let scales = activation.scales.as_slice();
let block_count = rows[0].len() / Q1_BLOCK_BYTES;
if rows.iter().any(|row| row.len() != rows[0].len())
|| values.len() < block_count * Q1_BLOCK_VALUES
|| scales.len() < block_count * Q1_Q8_BLOCKS
{
return [0.0; N];
}
let values_ptr = values.as_ptr();
let scales_ptr = scales.as_ptr();
let row_ptrs = rows.map(<[u8]>::as_ptr);
for block_index in 0..block_count {
let value_base = block_index * Q1_BLOCK_VALUES;
let scale_base = block_index * Q1_Q8_BLOCKS;
let signs = row_ptrs.map(|row| unsafe { row.add(block_index * Q1_BLOCK_BYTES + 2) });
let weight_scales = row_ptrs.map(|row| {
let scale_bits = u16::from_le(unsafe {
std::ptr::read_unaligned(row.add(block_index * Q1_BLOCK_BYTES).cast())
});
Fp16(scale_bits).to_f32()
});
for q8 in 0..Q1_Q8_BLOCKS {
let qy = unsafe {
_mm256_loadu_si256(values_ptr.add(value_base + q8 * BLOCK_VALUES).cast())
};
let act_scale = unsafe { *scales_ptr.add(scale_base + q8) };
for row in 0..N {
let qs32 = i32::from_le(unsafe {
std::ptr::read_unaligned(signs[row].add(q8 * Q1_SIGN_BYTES_PER_Q8_BLOCK).cast())
});
let sm = _mm256_cmpeq_epi8(
_mm256_and_si256(
_mm256_shuffle_epi8(_mm256_set1_epi32(qs32), byte_shuf),
bit_masks,
),
zero,
);
let sy = _mm256_sub_epi8(_mm256_xor_si256(qy, sm), sm);
let s32 = _mm256_madd_epi16(_mm256_maddubs_epi16(ones_8, sy), ones_16);
let scale = _mm256_set1_ps(weight_scales[row] * act_scale);
acc[row] = _mm256_fmadd_ps(scale, _mm256_cvtepi32_ps(s32), acc[row]);
}
}
}
let mut output = [0.0_f32; N];
for (value, acc) in output.iter_mut().zip(acc) {
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), acc) };
*value = lanes.into_iter().sum();
}
output
}
#[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 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,fma")]
unsafe fn dot_q8_avx2_fma(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::x86_64::{
_mm256_abs_epi8, _mm256_add_ps, _mm256_cvtepi32_ps, _mm256_fmadd_ps, _mm256_loadu_si256,
_mm256_set1_epi16, _mm256_set1_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
let block_count = row.len() / BLOCK_BYTES;
let ones = _mm256_set1_epi16(1);
let mut sum0 = _mm256_setzero_ps();
let mut sum1 = _mm256_setzero_ps();
let row_ptr = row.as_ptr();
let values_ptr = activation.values.as_ptr();
let scales_ptr = activation.scales.as_ptr();
for pair_index in 0..block_count / 2 {
let block_index = pair_index * 2;
let block0 = unsafe { row_ptr.add(block_index * BLOCK_BYTES) };
let weights0 = unsafe { _mm256_loadu_si256(block0.add(2).cast()) };
let magnitudes0 = _mm256_abs_epi8(weights0);
let values0 =
unsafe { _mm256_loadu_si256(values_ptr.add(block_index * BLOCK_VALUES).cast()) };
let products0 = unsafe { multiply_i8_block_avx2(weights0, magnitudes0, values0, ones) };
let scale0 = u16::from_le(unsafe { std::ptr::read_unaligned(block0.cast()) });
let factor0 = Fp16(scale0).to_f32() * unsafe { *scales_ptr.add(block_index) };
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(products0), _mm256_set1_ps(factor0), sum0);
let block_index = block_index + 1;
let block1 = unsafe { row_ptr.add(block_index * BLOCK_BYTES) };
let weights1 = unsafe { _mm256_loadu_si256(block1.add(2).cast()) };
let magnitudes1 = _mm256_abs_epi8(weights1);
let values1 =
unsafe { _mm256_loadu_si256(values_ptr.add(block_index * BLOCK_VALUES).cast()) };
let products1 = unsafe { multiply_i8_block_avx2(weights1, magnitudes1, values1, ones) };
let scale1 = u16::from_le(unsafe { std::ptr::read_unaligned(block1.cast()) });
let factor1 = Fp16(scale1).to_f32() * unsafe { *scales_ptr.add(block_index) };
sum1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(products1), _mm256_set1_ps(factor1), sum1);
}
if !block_count.is_multiple_of(2) {
let block_index = block_count - 1;
let block = unsafe { row_ptr.add(block_index * BLOCK_BYTES) };
let weights = unsafe { _mm256_loadu_si256(block.add(2).cast()) };
let magnitudes = _mm256_abs_epi8(weights);
let values =
unsafe { _mm256_loadu_si256(values_ptr.add(block_index * BLOCK_VALUES).cast()) };
let products = unsafe { multiply_i8_block_avx2(weights, magnitudes, values, ones) };
let scale = u16::from_le(unsafe { std::ptr::read_unaligned(block.cast()) });
let factor = Fp16(scale).to_f32() * unsafe { *scales_ptr.add(block_index) };
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(products), _mm256_set1_ps(factor), sum0);
}
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), _mm256_add_ps(sum0, sum1)) };
lanes.into_iter().sum()
}
#[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,fma")]
unsafe fn dot_q8_avx2_fma_batch<const ROWS: usize>(
row: &[u8],
activations: &[Q8Activation; ROWS],
) -> [f32; ROWS] {
use std::arch::x86_64::{
__m256, __m256i, _mm256_abs_epi8, _mm256_cvtepi32_ps, _mm256_fmadd_ps, _mm256_set1_epi16,
_mm256_set1_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
let ones = _mm256_set1_epi16(1);
let mut sums: [__m256; ROWS] = [_mm256_setzero_ps(); ROWS];
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, sum) in sums.iter_mut().enumerate() {
let activation = unsafe { activations.get_unchecked(index) };
let value_bytes = unsafe {
activation
.values
.as_ptr()
.add(start)
.cast::<[i8; BLOCK_VALUES]>()
.read_unaligned()
};
let values: __m256i = bytemuck::cast(value_bytes);
let products = unsafe { multiply_i8_block_avx2(weights, magnitudes, values, ones) };
let scale = weight_scale * unsafe { *activation.scales.get_unchecked(block_index) };
*sum = _mm256_fmadd_ps(_mm256_cvtepi32_ps(products), _mm256_set1_ps(scale), *sum);
}
}
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")]
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn accumulate_i8_packed_rows_16(
weight_group: i32,
values: *const i8,
sums: [std::arch::x86_64::__m256i; 2],
ones: std::arch::x86_64::__m256i,
) -> [std::arch::x86_64::__m256i; 2] {
use std::arch::x86_64::{_mm256_abs_epi8, _mm256_add_epi32, _mm256_set1_epi32};
let weights = _mm256_set1_epi32(weight_group);
let magnitudes = _mm256_abs_epi8(weights);
let first_bytes = unsafe { values.cast::<[i8; 32]>().read_unaligned() };
let second_bytes = unsafe {
values
.add(WIDE_MATRIX_ROW_TILE * 4)
.cast::<[i8; 32]>()
.read_unaligned()
};
let first = bytemuck::cast(first_bytes);
let second = bytemuck::cast(second_bytes);
let first = unsafe { multiply_i8_block_avx2(weights, magnitudes, first, ones) };
let second = unsafe { multiply_i8_block_avx2(weights, magnitudes, second, ones) };
[
_mm256_add_epi32(sums[0], first),
_mm256_add_epi32(sums[1], second),
]
}
#[cfg(target_arch = "x86_64")]
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn accumulate_i8_packed_rows_8_wide(
weight_group: i64,
values: *const i8,
sum: std::arch::x86_64::__m256i,
ones: std::arch::x86_64::__m256i,
) -> std::arch::x86_64::__m256i {
use std::arch::x86_64::{
_mm_unpacklo_epi64, _mm256_abs_epi8, _mm256_add_epi32, _mm256_castsi256_si128,
_mm256_extracti128_si256, _mm256_hadd_epi32, _mm256_set_m128i, _mm256_set1_epi64x,
};
let weights = _mm256_set1_epi64x(weight_group);
let magnitudes = _mm256_abs_epi8(weights);
let first_bytes = unsafe { values.cast::<[i8; 32]>().read_unaligned() };
let second_bytes = unsafe {
values
.add(MATRIX_ROW_TILE * 8)
.cast::<[i8; 32]>()
.read_unaligned()
};
let first = bytemuck::cast(first_bytes);
let second = bytemuck::cast(second_bytes);
let first = unsafe { multiply_i8_block_avx2(weights, magnitudes, first, ones) };
let second = unsafe { multiply_i8_block_avx2(weights, magnitudes, second, ones) };
let first = _mm256_hadd_epi32(first, first);
let second = _mm256_hadd_epi32(second, second);
let first = _mm_unpacklo_epi64(
_mm256_castsi256_si128(first),
_mm256_extracti128_si256::<1>(first),
);
let second = _mm_unpacklo_epi64(
_mm256_castsi256_si128(second),
_mm256_extracti128_si256::<1>(second),
);
_mm256_add_epi32(sum, _mm256_set_m128i(second, first))
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn accumulate_i8_packed_rows_4(
weight_group: i64,
values: *const i8,
sum: std::arch::x86_64::__m128i,
ones: std::arch::x86_64::__m256i,
) -> std::arch::x86_64::__m128i {
use std::arch::x86_64::{
_mm_add_epi32, _mm_unpacklo_epi64, _mm256_abs_epi8, _mm256_castsi256_si128,
_mm256_extracti128_si256, _mm256_hadd_epi32, _mm256_set1_epi64x,
};
let weights = _mm256_set1_epi64x(weight_group);
let magnitudes = _mm256_abs_epi8(weights);
let bytes = unsafe { values.cast::<[i8; 32]>().read_unaligned() };
let activation = bytemuck::cast(bytes);
let products = unsafe { multiply_i8_block_avx2(weights, magnitudes, activation, ones) };
let pairs = _mm256_hadd_epi32(products, products);
let rows = _mm_unpacklo_epi64(
_mm256_castsi256_si128(pairs),
_mm256_extracti128_si256::<1>(pairs),
);
_mm_add_epi32(sum, rows)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_q8_avx2_fma_packed_16(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> [f32; DOUBLE_WIDE_MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m256, _mm256_cvtepi32_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_mul_ps,
_mm256_set1_epi16, _mm256_set1_ps, _mm256_setzero_ps, _mm256_setzero_si256,
_mm256_storeu_ps,
};
let block_count = row.len() / BLOCK_BYTES;
let mut sums: [__m256; 2] = [_mm256_setzero_ps(); 2];
let ones = _mm256_set1_epi16(1);
for block_index in 0..block_count {
let block = unsafe { row.as_ptr().add(block_index * BLOCK_BYTES) };
let activation_values = unsafe {
values
.as_ptr()
.add(block_index * DOUBLE_WIDE_MATRIX_ROW_TILE * BLOCK_VALUES)
};
let mut products = [_mm256_setzero_si256(); 2];
for group in 0..BLOCK_VALUES / 4 {
let weight_bytes =
unsafe { block.add(2 + group * 4).cast::<[u8; 4]>().read_unaligned() };
products = unsafe {
accumulate_i8_packed_rows_16(
i32::from_ne_bytes(weight_bytes),
activation_values.add(group * DOUBLE_WIDE_MATRIX_ROW_TILE * 4),
products,
ones,
)
};
}
let scale_bytes = unsafe { block.cast::<[u8; 2]>().read_unaligned() };
let weight_scale = Fp16(u16::from_le_bytes(scale_bytes)).to_f32();
for row_half in 0..2 {
let activation_scales = unsafe {
_mm256_loadu_ps(scales.as_ptr().add(
block_index * DOUBLE_WIDE_MATRIX_ROW_TILE + row_half * WIDE_MATRIX_ROW_TILE,
))
};
let factor = _mm256_mul_ps(_mm256_set1_ps(weight_scale), activation_scales);
sums[row_half] = _mm256_fmadd_ps(
_mm256_cvtepi32_ps(products[row_half]),
factor,
sums[row_half],
);
}
}
let mut output = [0.0; DOUBLE_WIDE_MATRIX_ROW_TILE];
unsafe { _mm256_storeu_ps(output.as_mut_ptr(), sums[0]) };
unsafe { _mm256_storeu_ps(output.as_mut_ptr().add(WIDE_MATRIX_ROW_TILE), sums[1]) };
output
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_q8_avx2_fma_packed_8(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> [f32; WIDE_MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m256, _mm256_cvtepi32_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_mul_ps,
_mm256_set1_epi16, _mm256_set1_ps, _mm256_setzero_ps, _mm256_setzero_si256,
_mm256_storeu_ps,
};
let block_count = row.len() / BLOCK_BYTES;
let mut sum: __m256 = _mm256_setzero_ps();
let ones = _mm256_set1_epi16(1);
for block_index in 0..block_count {
let block = unsafe { row.as_ptr().add(block_index * BLOCK_BYTES) };
let weight_bytes = unsafe { block.add(2).cast::<[u8; BLOCK_VALUES]>().read_unaligned() };
let weight_groups: [i64; 4] = bytemuck::cast(weight_bytes);
let activation_values = unsafe {
values
.as_ptr()
.add(block_index * WIDE_MATRIX_ROW_TILE * BLOCK_VALUES)
};
let mut products = _mm256_setzero_si256();
for (group, weight_group) in weight_groups.iter().copied().enumerate() {
products = unsafe {
accumulate_i8_packed_rows_8_wide(
weight_group,
activation_values.add(group * WIDE_MATRIX_ROW_TILE * 8),
products,
ones,
)
};
}
let scale_bytes = unsafe { block.cast::<[u8; 2]>().read_unaligned() };
let weight_scale = Fp16(u16::from_le_bytes(scale_bytes)).to_f32();
let activation_scales =
unsafe { _mm256_loadu_ps(scales.as_ptr().add(block_index * WIDE_MATRIX_ROW_TILE)) };
let factor = _mm256_mul_ps(_mm256_set1_ps(weight_scale), activation_scales);
sum = _mm256_fmadd_ps(_mm256_cvtepi32_ps(products), factor, sum);
}
let mut output = [0.0; WIDE_MATRIX_ROW_TILE];
unsafe { _mm256_storeu_ps(output.as_mut_ptr(), sum) };
output
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_q8_avx2_fma_packed_4(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> [f32; MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m128, _mm_cvtepi32_ps, _mm_fmadd_ps, _mm_loadu_ps, _mm_mul_ps, _mm_set1_ps,
_mm_setzero_ps, _mm_setzero_si128, _mm_storeu_ps, _mm256_set1_epi16,
};
let block_count = row.len() / BLOCK_BYTES;
let mut sum: __m128 = _mm_setzero_ps();
let ones = _mm256_set1_epi16(1);
for block_index in 0..block_count {
let block = unsafe { row.as_ptr().add(block_index * BLOCK_BYTES) };
let weight_bytes = unsafe { block.add(2).cast::<[u8; BLOCK_VALUES]>().read_unaligned() };
let weight_groups: [i64; 4] = bytemuck::cast(weight_bytes);
let activation_values = unsafe {
values
.as_ptr()
.add(block_index * MATRIX_ROW_TILE * BLOCK_VALUES)
};
let mut products = _mm_setzero_si128();
for (group, weight_group) in weight_groups.iter().copied().enumerate() {
products = unsafe {
accumulate_i8_packed_rows_4(
weight_group,
activation_values.add(group * MATRIX_ROW_TILE * 8),
products,
ones,
)
};
}
let scale_bytes = unsafe { block.cast::<[u8; 2]>().read_unaligned() };
let weight_scale = Fp16(u16::from_le_bytes(scale_bytes)).to_f32();
let activation_scales =
unsafe { _mm_loadu_ps(scales.as_ptr().add(block_index * MATRIX_ROW_TILE)) };
let factor = _mm_mul_ps(_mm_set1_ps(weight_scale), activation_scales);
sum = _mm_fmadd_ps(_mm_cvtepi32_ps(products), factor, sum);
}
let mut output = [0.0; MATRIX_ROW_TILE];
unsafe { _mm_storeu_ps(output.as_mut_ptr(), sum) };
output
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_q1_q8_avx2_fma_packed_16(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> [f32; DOUBLE_WIDE_MATRIX_ROW_TILE] {
use std::arch::x86_64::{
__m256, _mm256_cvtepi32_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_mul_ps,
_mm256_set1_epi16, _mm256_set1_ps, _mm256_setzero_ps, _mm256_setzero_si256,
_mm256_storeu_ps,
};
let block_count = row.len() / Q1_BLOCK_BYTES;
let mut sums: [__m256; 2] = [_mm256_setzero_ps(); 2];
let ones = _mm256_set1_epi16(1);
for block_index in 0..block_count {
let block = unsafe { row.as_ptr().add(block_index * Q1_BLOCK_BYTES) };
let scale_bytes = unsafe { block.cast::<[u8; 2]>().read_unaligned() };
let weight_scale = Fp16(u16::from_le_bytes(scale_bytes)).to_f32();
let signs = unsafe { block.add(2) };
let activation_values = unsafe {
values
.as_ptr()
.add(block_index * DOUBLE_WIDE_MATRIX_ROW_TILE * Q1_BLOCK_VALUES)
};
for q8_block in 0..Q1_Q8_BLOCKS {
let mut products = [_mm256_setzero_si256(); 2];
for group in 0..BLOCK_VALUES / 4 {
let bits = unsafe { *signs.add(q8_block * Q1_SIGN_BYTES_PER_Q8_BLOCK + group / 2) };
let nibble = if group.is_multiple_of(2) {
bits & 0x0f
} else {
bits >> 4
};
products = unsafe {
accumulate_i8_packed_rows_16(
Q1_SIGN_NIBBLE_LUT[usize::from(nibble)],
activation_values.add(
q8_block * DOUBLE_WIDE_MATRIX_ROW_TILE * BLOCK_VALUES
+ group * DOUBLE_WIDE_MATRIX_ROW_TILE * 4,
),
products,
ones,
)
};
}
for row_half in 0..2 {
let activation_scales = unsafe {
_mm256_loadu_ps(scales.as_ptr().add(
(block_index * Q1_Q8_BLOCKS + q8_block) * DOUBLE_WIDE_MATRIX_ROW_TILE
+ row_half * WIDE_MATRIX_ROW_TILE,
))
};
let factor = _mm256_mul_ps(_mm256_set1_ps(weight_scale), activation_scales);
sums[row_half] = _mm256_fmadd_ps(
_mm256_cvtepi32_ps(products[row_half]),
factor,
sums[row_half],
);
}
}
}
let mut output = [0.0; DOUBLE_WIDE_MATRIX_ROW_TILE];
unsafe { _mm256_storeu_ps(output.as_mut_ptr(), sums[0]) };
unsafe { _mm256_storeu_ps(output.as_mut_ptr().add(WIDE_MATRIX_ROW_TILE), sums[1]) };
output
}
#[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;
i32_to_f32(_mm512_reduce_add_epi32(values))
}
#[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 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 = "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)
};
let block_sum_f32 = i32_to_f32(block_sum);
sum += weight_scale * activation_scale * block_sum_f32;
}
sum
}
#[cfg(target_arch = "aarch64")]
#[inline]
#[target_feature(enable = "dotprod")]
unsafe fn neon_dotprod_8(
mut accumulator: std::arch::aarch64::int32x2_t,
left: std::arch::aarch64::int8x8_t,
right: std::arch::aarch64::int8x8_t,
) -> std::arch::aarch64::int32x2_t {
unsafe {
std::arch::asm!(
"sdot {accumulator:v}.2s, {left:v}.8b, {right:v}.8b",
accumulator = inout(vreg) accumulator,
left = in(vreg) left,
right = in(vreg) right,
options(pure, nomem, nostack),
);
};
accumulator
}
#[cfg(target_arch = "aarch64")]
#[inline]
#[target_feature(enable = "dotprod")]
unsafe fn neon_dotprod_16(
mut accumulator: std::arch::aarch64::int32x4_t,
left: std::arch::aarch64::int8x16_t,
right: std::arch::aarch64::int8x16_t,
) -> std::arch::aarch64::int32x4_t {
unsafe {
std::arch::asm!(
"sdot {accumulator:v}.4s, {left:v}.16b, {right:v}.16b",
accumulator = inout(vreg) accumulator,
left = in(vreg) left,
right = in(vreg) right,
options(pure, nomem, nostack),
);
};
accumulator
}
#[cfg(target_arch = "aarch64")]
#[inline]
#[target_feature(enable = "i8mm")]
unsafe fn neon_i8mm(
mut accumulator: std::arch::aarch64::int32x4_t,
left: std::arch::aarch64::int8x16_t,
right: std::arch::aarch64::int8x16_t,
) -> std::arch::aarch64::int32x4_t {
unsafe {
std::arch::asm!(
"smmla {accumulator:v}.4s, {left:v}.16b, {right:v}.16b",
accumulator = inout(vreg) accumulator,
left = in(vreg) left,
right = in(vreg) right,
options(pure, nomem, nostack),
);
};
accumulator
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "dotprod")]
unsafe fn dot_q8_neon_dotprod(row: &[u8], activation: &Q8Activation) -> f32 {
use std::arch::aarch64::{vaddvq_s32, vdupq_n_s32, vld1q_s8};
let mut sum = 0.0_f32;
for ((block, values), activation_scale) in row
.chunks_exact(BLOCK_BYTES)
.zip(activation.values.chunks_exact(BLOCK_VALUES))
.zip(activation.scales.iter().copied())
{
let weights = block[2..].as_ptr().cast::<i8>();
let mut products = vdupq_n_s32(0);
products =
unsafe { neon_dotprod_16(products, vld1q_s8(weights), vld1q_s8(values.as_ptr())) };
products = unsafe {
neon_dotprod_16(
products,
vld1q_s8(weights.add(16)),
vld1q_s8(values.as_ptr().add(16)),
)
};
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
sum += weight_scale * activation_scale * i32_to_f32(vaddvq_s32(products));
}
sum
}
#[cfg(target_arch = "aarch64")]
#[inline]
#[target_feature(enable = "dotprod")]
unsafe fn accumulate_neon_dotprod_packed_8<const ROWS: usize>(
weights: std::arch::aarch64::int8x8_t,
values: *const i8,
products: &mut [std::arch::aarch64::int32x2_t; ROWS],
) {
use std::arch::aarch64::vld1_s8;
for (index, product) in products.iter_mut().enumerate() {
*product = unsafe { neon_dotprod_8(*product, weights, vld1_s8(values.add(index * 8))) };
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
#[target_feature(enable = "dotprod")]
unsafe fn accumulate_neon_dotprod_packed_16(
weight_group: i32,
values: *const i8,
products: &mut [std::arch::aarch64::int32x4_t; 4],
) {
use std::arch::aarch64::{vdupq_n_s32, vld1q_s8, vreinterpretq_s8_s32};
let weights = vreinterpretq_s8_s32(vdupq_n_s32(weight_group));
for (index, product) in products.iter_mut().enumerate() {
*product = unsafe { neon_dotprod_16(*product, weights, vld1q_s8(values.add(index * 16))) };
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "dotprod")]
unsafe fn dot_q8_neon_dotprod_packed_8<const ROWS: usize>(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> [f32; ROWS] {
use std::arch::aarch64::{vaddv_s32, vdup_n_s32, vld1_s8};
let mut sums = [0.0_f32; ROWS];
for (block_index, block) in row.chunks_exact(BLOCK_BYTES).enumerate() {
let weights = block[2..].as_ptr().cast::<i8>();
let activation_values = unsafe { values.as_ptr().add(block_index * ROWS * BLOCK_VALUES) };
let mut products = [vdup_n_s32(0); ROWS];
for group in 0..BLOCK_VALUES / 8 {
unsafe {
accumulate_neon_dotprod_packed_8(
vld1_s8(weights.add(group * 8)),
activation_values.add(group * ROWS * 8),
&mut products,
);
};
}
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
for index in 0..ROWS {
let activation_scale = unsafe { *scales.get_unchecked(block_index * ROWS + index) };
sums[index] += weight_scale * activation_scale * i32_to_f32(vaddv_s32(products[index]));
}
}
sums
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "dotprod")]
unsafe fn dot_q8_neon_dotprod_packed_16(
row: &[u8],
values: &[i8],
scales: &[f32],
) -> [f32; DOUBLE_WIDE_MATRIX_ROW_TILE] {
use std::arch::aarch64::{vdupq_n_s32, vst1q_s32};
let mut sums = [0.0_f32; DOUBLE_WIDE_MATRIX_ROW_TILE];
for (block_index, block) in row.chunks_exact(BLOCK_BYTES).enumerate() {
let activation_values = unsafe {
values
.as_ptr()
.add(block_index * DOUBLE_WIDE_MATRIX_ROW_TILE * BLOCK_VALUES)
};
let mut products = [vdupq_n_s32(0); 4];
for group in 0..BLOCK_VALUES / 4 {
let weight_bytes = unsafe {
block[2..]
.as_ptr()
.add(group * 4)
.cast::<[u8; 4]>()
.read_unaligned()
};
unsafe {
accumulate_neon_dotprod_packed_16(
i32::from_ne_bytes(weight_bytes),
activation_values.add(group * DOUBLE_WIDE_MATRIX_ROW_TILE * 4),
&mut products,
);
};
}
let mut block_sums = [0_i32; DOUBLE_WIDE_MATRIX_ROW_TILE];
for (block_sum, product) in block_sums.chunks_exact_mut(4).zip(products) {
unsafe { vst1q_s32(block_sum.as_mut_ptr(), product) };
}
let weight_scale = Fp16::decode_le(&block[0..2]).to_f32();
for index in 0..DOUBLE_WIDE_MATRIX_ROW_TILE {
let activation_scale =
unsafe { *scales.get_unchecked(block_index * DOUBLE_WIDE_MATRIX_ROW_TILE + index) };
sums[index] += weight_scale * activation_scale * i32_to_f32(block_sums[index]);
}
}
sums
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,i8mm")]
unsafe fn dot_q8_neon_i8mm_packed_pair<const ROWS: usize>(
rows: [&[u8]; 2],
values: &[i8],
scales: &[f32],
) -> [[f32; ROWS]; 2] {
use std::arch::aarch64::{vcombine_s8, vdupq_n_s32, vld1_s8, vld1q_s8, vst1q_s32};
let mut sums = [[0.0_f32; ROWS]; 2];
for (block_index, (first, second)) in rows[0]
.chunks_exact(BLOCK_BYTES)
.zip(rows[1].chunks_exact(BLOCK_BYTES))
.enumerate()
{
let mut products = [vdupq_n_s32(0); WIDE_MATRIX_ROW_TILE / 2];
for group in 0..BLOCK_VALUES / 8 {
let weights = vcombine_s8(
unsafe { vld1_s8(first.as_ptr().add(2 + group * 8).cast()) },
unsafe { vld1_s8(second.as_ptr().add(2 + group * 8).cast()) },
);
let group_values = unsafe {
values
.as_ptr()
.add(block_index * ROWS * BLOCK_VALUES + group * ROWS * 8)
};
for (pair, product) in products[..ROWS / 2].iter_mut().enumerate() {
*product =
unsafe { neon_i8mm(*product, weights, vld1q_s8(group_values.add(pair * 16))) };
}
}
let weight_scales = [
Fp16::decode_le(&first[0..2]).to_f32(),
Fp16::decode_le(&second[0..2]).to_f32(),
];
let block_scales = unsafe { scales.as_ptr().add(block_index * ROWS) };
for (pair, product) in products[..ROWS / 2].iter().copied().enumerate() {
let mut lanes = [0_i32; 4];
unsafe { vst1q_s32(lanes.as_mut_ptr(), product) };
for offset in 0..2 {
let row = pair * 2 + offset;
let activation_scale = unsafe { *block_scales.add(row) };
sums[0][row] += weight_scales[0] * activation_scale * i32_to_f32(lanes[offset]);
sums[1][row] += weight_scales[1] * activation_scale * i32_to_f32(lanes[2 + offset]);
}
}
}
sums
}
pub(super) fn matrix_vector_into(
matrix: &Tensor<'_>,
vector: &[f32],
activation: Option<&Q8Activation>,
output: &mut [f32],
) -> Result<()> {
let [input, output_size] = matrix_dimensions(matrix)?;
ensure!(vector.len() == input, "matrix input width differs");
if let Some(activation) = activation {
activation.validate_len(vector.len())?;
}
ensure!(output.len() == output_size, "matrix output width differs");
WeightFormat::from_tensor_type(matrix.tensor_type())
.project_vector(matrix, vector, activation, output)
}
pub(super) fn matrix_vector_q1_add_into(
matrix: &Tensor<'_>,
activation: &Q8Activation,
dest: &mut [f32],
) -> Result<()> {
let [input, output_size] = matrix_dimensions(matrix)?;
activation.validate_len(input)?;
ensure!(dest.len() == output_size, "matrix output width differs");
ensure!(
matrix.tensor_type() == TensorType::Q1_0,
"accumulate matrix is not Q1_0"
);
IsaKernel::detect().project_q1_add(matrix, activation, dest)
}
pub(super) fn matrix_vector_q1_swiglu_quantize_into(
gate: &Tensor<'_>,
up: &Tensor<'_>,
activation: &mut Q8Activation,
) -> Result<()> {
let gate_dimensions = matrix_dimensions(gate)?;
let up_dimensions = matrix_dimensions(up)?;
ensure!(
gate_dimensions == up_dimensions,
"paired matrix dimensions differ"
);
let [input, output] = gate_dimensions;
activation.validate_len(input)?;
ensure!(
gate.tensor_type() == TensorType::Q1_0 && up.tensor_type() == TensorType::Q1_0,
"swiglu pair is not Q1_0"
);
ensure!(
output.is_multiple_of(BLOCK_VALUES),
"Q8 activation width is not divisible by 32"
);
let input_activation = Q8Activation {
scales: activation.scales.clone(),
values: activation.values.clone(),
};
activation.scales.resize(output / BLOCK_VALUES, 0.0);
activation.values.resize(output, 0);
IsaKernel::detect().project_q1_swiglu(gate, up, &input_activation, activation)
}
pub(super) fn matrix_vector_q1_triple_into(
matrices: [&Tensor<'_>; 3],
vector: &[f32],
activation: &Q8Activation,
outputs: [&mut [f32]; 3],
) -> Result<()> {
activation.validate_len(vector.len())?;
for (matrix, output) in matrices.iter().zip(outputs.iter()) {
let [input, output_size] = matrix_dimensions(matrix)?;
ensure!(vector.len() == input, "matrix input width differs");
ensure!(output.len() == output_size, "matrix output width differs");
ensure!(
matrix.tensor_type() == TensorType::Q1_0,
"triple matrix is not Q1_0"
);
}
let [first, second, third] = matrices;
let [first_output, second_output, third_output] = outputs;
let kernel = IsaKernel::detect();
let (first_result, (second_result, third_result)) = rayon::join(
|| kernel.project(first, QuantizedFormat::Q1_0, activation, first_output),
|| {
rayon::join(
|| kernel.project(second, QuantizedFormat::Q1_0, activation, second_output),
|| kernel.project(third, QuantizedFormat::Q1_0, activation, third_output),
)
},
);
first_result?;
second_result?;
third_result
}
fn prepare_matrix_workspace(
matrix: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
workspace: &mut MatrixWorkspace,
) -> Result<QuantizedFormat> {
let [input_size, _] = matrix_dimensions(matrix)?;
let format = WeightFormat::from_tensor_type(matrix.tensor_type())
.quantized()
.context("matrix is not quantized")?;
let kernel = IsaKernel::detect();
workspace.activations.prepare(
vectors,
row_count,
input_size,
&[row_count],
kernel.packed_tile_rows(),
)?;
Ok(format)
}
fn prepare_matrix_workspace_segmented(
matrix: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
segment_lengths: &[usize],
workspace: &mut MatrixWorkspace,
) -> Result<QuantizedFormat> {
let [input_size, _] = matrix_dimensions(matrix)?;
let format = WeightFormat::from_tensor_type(matrix.tensor_type())
.quantized()
.context("matrix is not quantized")?;
let kernel = IsaKernel::detect();
workspace.activations.prepare(
vectors,
row_count,
input_size,
segment_lengths,
kernel.packed_tile_rows(),
)?;
Ok(format)
}
fn project_prepared_matrix_into(
matrix: &Tensor<'_>,
format: QuantizedFormat,
activations: &Q8Activations,
output: &mut [f32],
columns: &mut Vec<f32>,
) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
validate_quantized_batch(&activations.rows, output, input_size, output_size)?;
IsaKernel::detect().project_batch(matrix, format, activations, output, Some(columns))
}
pub(super) fn matrix_matrix_into(
matrix: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
output: &mut [f32],
workspace: &mut MatrixWorkspace,
) -> Result<()> {
let [input_size, output_size] = matrix_dimensions(matrix)?;
validate_matrix_input(vectors, row_count, input_size)?;
let output_len = row_count
.checked_mul(output_size)
.context("matrix-matrix output size overflow")?;
ensure!(
output.len() == output_len,
"matrix-matrix output shape differs"
);
if row_count == 1 {
return matrix_vector_into(matrix, vectors, None, output);
}
match WeightFormat::from_tensor_type(matrix.tensor_type()) {
format @ WeightFormat::F32 => {
format.project_batch(matrix, vectors, row_count, input_size, output)
}
WeightFormat::Quantized(_) => {
let format = prepare_matrix_workspace(matrix, vectors, row_count, workspace)?;
let MatrixWorkspace {
activations,
columns,
} = workspace;
project_prepared_matrix_into(matrix, format, activations, output, &mut columns[0])
}
}
}
pub(super) fn matrix_matrix_pair_into(
first: &Tensor<'_>,
second: &Tensor<'_>,
vectors: &[f32],
row_count: usize,
first_output: &mut [f32],
second_output: &mut [f32],
workspace: &mut MatrixWorkspace,
) -> Result<()> {
let first_dimensions = matrix_dimensions(first)?;
ensure!(
first_dimensions == matrix_dimensions(second)?,
"paired matrix dimensions differ"
);
let [input_size, output_size] = first_dimensions;
validate_matrix_input(vectors, row_count, input_size)?;
let output_len = row_count
.checked_mul(output_size)
.context("paired matrix output size overflow")?;
ensure!(
first_output.len() == output_len && second_output.len() == output_len,
"paired matrix output shape differs"
);
let first_format = WeightFormat::from_tensor_type(first.tensor_type()).quantized();
let second_format = WeightFormat::from_tensor_type(second.tensor_type()).quantized();
let Some(format) = first_format.filter(|format| Some(*format) == second_format) else {
matrix_matrix_into(first, vectors, row_count, first_output, workspace)?;
return matrix_matrix_into(second, vectors, row_count, second_output, workspace);
};
prepare_matrix_workspace(first, vectors, row_count, workspace)?;
let MatrixWorkspace {
activations,
columns,
} = workspace;
let [first_columns, second_columns, _] = columns;
let (first_result, second_result) = rayon::join(
|| project_prepared_matrix_into(first, format, activations, first_output, first_columns),
|| project_prepared_matrix_into(second, format, activations, second_output, second_columns),
);
first_result?;
second_result
}
pub(super) fn matrix_matrix_triple_into(
matrices: [&Tensor<'_>; 3],
vectors: &[f32],
row_count: usize,
outputs: [&mut [f32]; 3],
workspace: &mut MatrixWorkspace,
) -> Result<()> {
let [first, second, third] = matrices;
let [first_output, second_output, third_output] = outputs;
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)?;
let expected_lengths = [first_size, second_size, third_size].map(|size| {
row_count
.checked_mul(size)
.context("triple matrix output size overflow")
});
let [first_len, second_len, third_len] = expected_lengths;
ensure!(
first_output.len() == first_len?
&& second_output.len() == second_len?
&& third_output.len() == third_len?,
"triple matrix output shape differs"
);
let formats = [first, second, third]
.map(|matrix| WeightFormat::from_tensor_type(matrix.tensor_type()).quantized());
let Some(format) =
formats[0].filter(|format| formats[1] == Some(*format) && formats[2] == Some(*format))
else {
matrix_matrix_into(first, vectors, row_count, first_output, workspace)?;
matrix_matrix_into(second, vectors, row_count, second_output, workspace)?;
return matrix_matrix_into(third, vectors, row_count, third_output, workspace);
};
prepare_matrix_workspace(first, vectors, row_count, workspace)?;
let MatrixWorkspace {
activations,
columns,
} = workspace;
let [first_columns, second_columns, third_columns] = columns;
let (first_result, (second_result, third_result)) = rayon::join(
|| project_prepared_matrix_into(first, format, activations, first_output, first_columns),
|| {
rayon::join(
|| {
project_prepared_matrix_into(
second,
format,
activations,
second_output,
second_columns,
)
},
|| {
project_prepared_matrix_into(
third,
format,
activations,
third_output,
third_columns,
)
},
)
},
);
first_result?;
second_result?;
third_result
}
pub(super) fn matrix_matrix_segmented_into(
matrix: &Tensor<'_>,
vectors: &[f32],
segment_lengths: &[usize],
output: &mut [f32],
workspace: &mut MatrixWorkspace,
) -> Result<()> {
let row_count = segmented_row_count(segment_lengths)?;
if segment_lengths.len() == 1 {
return matrix_matrix_into(matrix, vectors, row_count, output, workspace);
}
let [input_size, output_size] = matrix_dimensions(matrix)?;
validate_matrix_input(vectors, row_count, input_size)?;
let output_len = row_count
.checked_mul(output_size)
.context("matrix-matrix output size overflow")?;
ensure!(
output.len() == output_len,
"matrix-matrix output shape differs"
);
match WeightFormat::from_tensor_type(matrix.tensor_type()) {
format @ WeightFormat::F32 => {
format.project_batch(matrix, vectors, row_count, input_size, output)
}
WeightFormat::Quantized(_) => {
let format = prepare_matrix_workspace_segmented(
matrix,
vectors,
row_count,
segment_lengths,
workspace,
)?;
let MatrixWorkspace {
activations,
columns,
} = workspace;
project_prepared_matrix_into(matrix, format, activations, output, &mut columns[0])
}
}
}
pub(super) fn matrix_matrix_triple_segmented_into(
matrices: [&Tensor<'_>; 3],
vectors: &[f32],
segment_lengths: &[usize],
outputs: [&mut [f32]; 3],
workspace: &mut MatrixWorkspace,
) -> Result<()> {
let row_count = segmented_row_count(segment_lengths)?;
if segment_lengths.len() == 1 {
return matrix_matrix_triple_into(matrices, vectors, row_count, outputs, workspace);
}
let [first, second, third] = matrices;
let [first_output, second_output, third_output] = outputs;
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)?;
let expected_lengths = [first_size, second_size, third_size].map(|size| {
row_count
.checked_mul(size)
.context("triple matrix output size overflow")
});
let [first_len, second_len, third_len] = expected_lengths;
ensure!(
first_output.len() == first_len?
&& second_output.len() == second_len?
&& third_output.len() == third_len?,
"triple matrix output shape differs"
);
let formats = [first, second, third]
.map(|matrix| WeightFormat::from_tensor_type(matrix.tensor_type()).quantized());
let Some(format) =
formats[0].filter(|format| formats[1] == Some(*format) && formats[2] == Some(*format))
else {
matrix_matrix_segmented_into(first, vectors, segment_lengths, first_output, workspace)?;
matrix_matrix_segmented_into(second, vectors, segment_lengths, second_output, workspace)?;
return matrix_matrix_segmented_into(
third,
vectors,
segment_lengths,
third_output,
workspace,
);
};
prepare_matrix_workspace_segmented(first, vectors, row_count, segment_lengths, workspace)?;
let MatrixWorkspace {
activations,
columns,
} = workspace;
let [first_columns, second_columns, third_columns] = columns;
let (first_result, (second_result, third_result)) = rayon::join(
|| project_prepared_matrix_into(first, format, activations, first_output, first_columns),
|| {
rayon::join(
|| {
project_prepared_matrix_into(
second,
format,
activations,
second_output,
second_columns,
)
},
|| {
project_prepared_matrix_into(
third,
format,
activations,
third_output,
third_columns,
)
},
)
},
);
first_result?;
second_result?;
third_result
}
fn segmented_row_count(segment_lengths: &[usize]) -> Result<usize> {
ensure!(
!segment_lengths.is_empty() && segment_lengths.iter().all(|rows| *rows != 0),
"matrix-matrix segment is empty"
);
segment_lengths.iter().try_fold(0_usize, |total, rows| {
total
.checked_add(*rows)
.context("matrix-matrix row count overflow")
})
}
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,
segment_lengths: &[usize],
packed_tile_rows: Option<usize>,
) -> Result<Q8Activations> {
validate_matrix_input(vectors, row_count, input_size)?;
ensure!(
segmented_row_count(segment_lengths)? == row_count,
"Q8 activation segment rows differ"
);
let rows = vectors
.chunks_exact(input_size)
.map(Q8Activation::new)
.collect::<Result<Vec<_>>>()?;
Q8Activations::new(rows, segment_lengths, packed_tile_rows)
}
fn project_f32(matrix: &Tensor<'_>, vector: &[f32], output: &mut [f32]) -> Result<()> {
DenseIsaKernel::detect().project(matrix, vector, output)
}
fn project_f32_with<K: DenseKernel>(
matrix: &Tensor<'_>,
vector: &[f32],
output: &mut [f32],
) -> Result<()> {
output
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = K::dot_f32(matrix.f32_row(row)?, vector);
Ok(())
})
}
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)?;
DenseIsaKernel::detect().project_batch(matrix, vectors, input_size, output_size, output)
}
fn project_f32_batch_with<K: DenseKernel>(
matrix: &Tensor<'_>,
vectors: &[f32],
input_size: usize,
output_size: usize,
output: &mut [f32],
) -> Result<()> {
let row_count = vectors.len() / input_size;
let mut columns = vec![0.0; output.len()];
columns.par_chunks_mut(row_count).enumerate().try_for_each(
|(output_channel, column)| -> Result<()> {
let weights = matrix.f32_row(output_channel)?;
for (value, input_row) in column.iter_mut().zip(vectors.chunks_exact(input_size)) {
*value = K::dot_f32(weights, input_row);
}
Ok(())
},
)?;
transpose_batch_output(&columns, row_count, output_size, output)?;
Ok(())
}
fn transpose_batch_output(
columns: &[f32],
row_count: usize,
output_size: usize,
output: &mut [f32],
) -> Result<()> {
let expected = row_count
.checked_mul(output_size)
.context("matrix transpose size overflow")?;
ensure!(
columns.len() == expected,
"matrix transpose input shape differs"
);
ensure!(
output.len() == columns.len(),
"matrix transpose output shape differs"
);
output
.par_chunks_exact_mut(output_size)
.enumerate()
.for_each(|(row, output_row)| {
for (output_channel, value) in output_row.iter_mut().enumerate() {
let index = output_channel * row_count + row;
*value = unsafe { *columns.get_unchecked(index) };
}
});
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(
kernel: IsaKernel,
matrix: &Tensor<'_>,
format: QuantizedFormat,
activations: &Q8Activations,
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.rows, output, input_size, output_size)?;
kernel.project_batch(matrix, format, activations, output, None)
}
fn project_quantized_batch_with<K>(
matrix: &Tensor<'_>,
format: QuantizedFormat,
activations: &Q8Activations,
output: &mut [f32],
columns: Option<&mut Vec<f32>>,
) -> Result<()>
where
K: QuantizedKernel<Q8_0Format> + QuantizedKernel<Q1_0Format>,
{
match format {
QuantizedFormat::Q8_0 => project_quantized_batch_format_with::<K, Q8_0Format>(
matrix,
activations,
output,
columns,
),
QuantizedFormat::Q1_0 => project_quantized_batch_format_with::<K, Q1_0Format>(
matrix,
activations,
output,
columns,
),
}
}
fn project_quantized_batch_format_with<K, F>(
matrix: &Tensor<'_>,
activations: &Q8Activations,
output: &mut [f32],
columns: Option<&mut Vec<f32>>,
) -> Result<()>
where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
project_quantized_batch_format_segmented_with::<K, F>(
matrix,
activations,
&activations.segment_lengths,
output,
columns,
)
}
fn project_quantized_batch_format_segmented_with<K, F>(
matrix: &Tensor<'_>,
activations: &Q8Activations,
segment_lengths: &[usize],
output: &mut [f32],
columns: Option<&mut Vec<f32>>,
) -> Result<()>
where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
ensure!(
activations.segment_lengths == segment_lengths,
"matrix-matrix segment plan differs"
);
let row_count = activations.len();
let output_size = output.len() / row_count;
let mut owned_columns = Vec::new();
let columns = columns.unwrap_or(&mut owned_columns);
columns.resize(output.len(), 0.0);
if <K as QuantizedKernel<F>>::supports_weight_row_pairs()
&& let Some(packed) = &activations.packed
{
let pair_rows = row_count
.checked_mul(2)
.context("paired matrix output size overflow")?;
columns.par_chunks_mut(pair_rows).enumerate().try_for_each(
|(pair, pair_columns)| -> Result<()> {
let output_channel = pair * 2;
let first_weights = matrix.encoded_row(output_channel)?;
let (first_column, second_column) = pair_columns.split_at_mut(row_count);
if second_column.is_empty() {
project_quantized_packed::<K, F>(
first_weights,
activations,
packed,
first_column,
)?;
} else {
ensure!(
second_column.len() == row_count,
"paired matrix output shape differs"
);
let second_weights = matrix.encoded_row(output_channel + 1)?;
project_quantized_packed_pair::<K, F>(
[first_weights, second_weights],
activations,
packed,
[first_column, second_column],
)?;
}
Ok(())
},
)?;
} else {
columns.par_chunks_mut(row_count).enumerate().try_for_each(
|(output_channel, column)| -> Result<()> {
let weights = matrix.encoded_row(output_channel)?;
if let Some(packed) = &activations.packed {
project_quantized_packed::<K, F>(weights, activations, packed, column)?;
} else {
project_quantized_activation_rows::<K, F>(weights, &activations.rows, column);
}
Ok(())
},
)?;
}
transpose_batch_output(columns, row_count, output_size, output)?;
Ok(())
}
fn project_quantized_packed<K, F>(
weights: &[u8],
activations: &Q8Activations,
packed: &PackedQ8Activations,
values: &mut [f32],
) -> Result<()>
where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
let mut written = 0;
for &tile in &packed.tiles {
ensure!(
tile.row_start == written,
"packed Q8 tile rows are not contiguous"
);
let tile_values = packed.values(tile)?;
let tile_scales = packed.scales(tile)?;
let end = written
.checked_add(tile.row_count)
.context("packed Q8 output range overflow")?;
ensure!(
end <= values.len(),
"packed Q8 output range is out of bounds"
);
let projected = match tile.row_count {
DOUBLE_WIDE_MATRIX_ROW_TILE => {
if let Some(result) = <K as QuantizedKernel<F>>::dot_packed_double_wide_batch(
weights,
tile_values,
tile_scales,
) {
values[written..end].copy_from_slice(&result);
true
} else {
false
}
}
WIDE_MATRIX_ROW_TILE => {
if let Some(result) = <K as QuantizedKernel<F>>::dot_packed_wide_batch(
weights,
tile_values,
tile_scales,
) {
values[written..end].copy_from_slice(&result);
true
} else {
false
}
}
MATRIX_ROW_TILE => {
if let Some(result) =
<K as QuantizedKernel<F>>::dot_packed_batch(weights, tile_values, tile_scales)
{
values[written..end].copy_from_slice(&result);
true
} else {
false
}
}
_ => bail!("invalid packed Q8 tile width {}", tile.row_count),
};
if !projected {
let rows = activations
.rows
.get(tile.row_start..end)
.context("packed Q8 activation rows are out of range")?;
let output = values
.get_mut(written..end)
.context("packed Q8 output rows are out of range")?;
project_quantized_activation_rows::<K, F>(weights, rows, output);
}
written = end;
}
let remaining_activations = activations
.rows
.get(written..)
.context("packed Q8 activation tail is out of range")?;
let remaining_values = values
.get_mut(written..)
.context("packed Q8 output tail is out of range")?;
project_quantized_activation_rows::<K, F>(weights, remaining_activations, remaining_values);
Ok(())
}
fn project_quantized_packed_pair<K, F>(
weights: [&[u8]; 2],
activations: &Q8Activations,
packed: &PackedQ8Activations,
values: [&mut [f32]; 2],
) -> Result<()>
where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
let [first_values, second_values] = values;
let mut written = 0;
for &tile in &packed.tiles {
ensure!(
tile.row_start == written,
"packed Q8 tile rows are not contiguous"
);
let tile_values = packed.values(tile)?;
let tile_scales = packed.scales(tile)?;
let end = written
.checked_add(tile.row_count)
.context("packed Q8 output range overflow")?;
ensure!(
end <= first_values.len() && end <= second_values.len(),
"packed Q8 output range is out of bounds"
);
let projected = match tile.row_count {
WIDE_MATRIX_ROW_TILE => {
if let Some([first, second]) = <K as QuantizedKernel<F>>::dot_packed_wide_batch_pair(
weights,
tile_values,
tile_scales,
) {
first_values[written..end].copy_from_slice(&first);
second_values[written..end].copy_from_slice(&second);
true
} else {
false
}
}
MATRIX_ROW_TILE => {
if let Some([first, second]) = <K as QuantizedKernel<F>>::dot_packed_batch_pair(
weights,
tile_values,
tile_scales,
) {
first_values[written..end].copy_from_slice(&first);
second_values[written..end].copy_from_slice(&second);
true
} else {
false
}
}
DOUBLE_WIDE_MATRIX_ROW_TILE => false,
_ => bail!("invalid packed Q8 tile width {}", tile.row_count),
};
if !projected {
let rows = activations
.rows
.get(tile.row_start..end)
.context("packed Q8 activation rows are out of range")?;
let first_output = first_values
.get_mut(written..end)
.context("packed Q8 output rows are out of range")?;
let second_output = second_values
.get_mut(written..end)
.context("packed Q8 output rows are out of range")?;
project_quantized_activation_rows::<K, F>(weights[0], rows, first_output);
project_quantized_activation_rows::<K, F>(weights[1], rows, second_output);
}
written = end;
}
let remaining_activations = activations
.rows
.get(written..)
.context("packed Q8 activation tail is out of range")?;
let first_remaining = first_values
.get_mut(written..)
.context("packed Q8 output tail is out of range")?;
let second_remaining = second_values
.get_mut(written..)
.context("packed Q8 output tail is out of range")?;
project_quantized_activation_rows::<K, F>(weights[0], remaining_activations, first_remaining);
project_quantized_activation_rows::<K, F>(weights[1], remaining_activations, second_remaining);
Ok(())
}
fn project_quantized_activation_rows<K, F>(
weights: &[u8],
activations: &[Q8Activation],
values: &mut [f32],
) where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
for (wide_activation_rows, wide_values) in activations
.chunks(WIDE_MATRIX_ROW_TILE)
.zip(values.chunks_mut(WIDE_MATRIX_ROW_TILE))
{
if let Ok(activation_tile) =
<&[Q8Activation; WIDE_MATRIX_ROW_TILE]>::try_from(wide_activation_rows)
&& let Some(tile) = <K as QuantizedKernel<F>>::dot_wide_batch(weights, activation_tile)
{
wide_values.copy_from_slice(&tile);
continue;
}
for (activation_rows, values) in wide_activation_rows
.chunks(MATRIX_ROW_TILE)
.zip(wide_values.chunks_mut(MATRIX_ROW_TILE))
{
if let Ok(activation_tile) =
<&[Q8Activation; MATRIX_ROW_TILE]>::try_from(activation_rows)
{
let tile = <K as QuantizedKernel<F>>::dot_batch(weights, activation_tile);
values.copy_from_slice(&tile);
} else {
for (value, activation) in values.iter_mut().zip(activation_rows) {
*value = <K as QuantizedKernel<F>>::dot(weights, activation);
}
}
}
}
}
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());
match format {
WeightFormat::Quantized(format) => {
let activation = Q8Activation::new(vector)?;
IsaKernel::detect().argmax(matrix, format, &activation)
}
WeightFormat::F32 => DenseIsaKernel::detect().argmax(matrix, vector, output),
}
}
pub(super) fn matrix_argmax_with_activation(
matrix: &Tensor<'_>,
vector: &[f32],
activation: &Q8Activation,
) -> Result<usize> {
let [input, output] = matrix_dimensions(matrix)?;
ensure!(vector.len() == input, "matrix input width differs");
activation.validate_len(input)?;
match WeightFormat::from_tensor_type(matrix.tensor_type()) {
WeightFormat::Quantized(format) => IsaKernel::detect().argmax(matrix, format, activation),
WeightFormat::F32 => DenseIsaKernel::detect().argmax(matrix, vector, output),
}
}
fn matrix_argmax_f32_with<K: DenseKernel>(
matrix: &Tensor<'_>,
vector: &[f32],
output_size: usize,
) -> Result<usize> {
matrix_argmax_with(output_size, |index| {
Ok(K::dot_f32(matrix.f32_row(index)?, vector))
})
}
fn matrix_argmax_quantized_with<K>(
matrix: &Tensor<'_>,
format: QuantizedFormat,
activation: &Q8Activation,
) -> Result<usize>
where
K: QuantizedKernel<Q8_0Format> + QuantizedKernel<Q1_0Format>,
{
match format {
QuantizedFormat::Q8_0 => {
matrix_argmax_quantized_format_with::<K, Q8_0Format>(matrix, activation)
}
QuantizedFormat::Q1_0 => matrix_argmax_q1_rows_with::<K>(matrix, activation),
}
}
fn matrix_argmax_quantized_format_with<K, F>(
matrix: &Tensor<'_>,
activation: &Q8Activation,
) -> Result<usize>
where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
let output = matrix_dimensions(matrix)?[1];
matrix_argmax_with(output, |index| {
Ok(<K as QuantizedKernel<F>>::dot(
matrix.encoded_row(index)?,
activation,
))
})
}
fn matrix_argmax_q1_rows_with<K>(matrix: &Tensor<'_>, activation: &Q8Activation) -> Result<usize>
where
K: QuantizedKernel<Q1_0Format>,
{
let output = matrix_dimensions(matrix)?[1];
let wide_count = output / WIDE_MATRIX_ROW_TILE;
let wide_end = wide_count * WIDE_MATRIX_ROW_TILE;
let mid_count = (output - wide_end) / MATRIX_ROW_TILE;
let mid_end = wide_end + mid_count * MATRIX_ROW_TILE;
let wide_best = (0..wide_count)
.into_par_iter()
.map(|tile| -> Result<(usize, f32)> {
let start = tile * WIDE_MATRIX_ROW_TILE;
let values = <K as QuantizedKernel<Q1_0Format>>::dot_wide_rows(
encoded_rows(matrix, start)?,
activation,
);
select_argmax_scores(start, values)
})
.try_reduce_with(|left, right| Ok(select_argmax(left, right)))
.transpose()?;
let mid_best = (0..mid_count)
.into_par_iter()
.map(|tile| -> Result<(usize, f32)> {
let start = wide_end + tile * MATRIX_ROW_TILE;
let values = <K as QuantizedKernel<Q1_0Format>>::dot_rows(
encoded_rows(matrix, start)?,
activation,
);
select_argmax_scores(start, values)
})
.try_reduce_with(|left, right| Ok(select_argmax(left, right)))
.transpose()?;
let tail_best = (mid_end..output)
.into_par_iter()
.map(|index| -> Result<(usize, f32)> {
let value =
<K as QuantizedKernel<Q1_0Format>>::dot(matrix.encoded_row(index)?, activation);
if !value.is_finite() {
bail!("matrix output {index} is not finite");
}
Ok((index, value))
})
.try_reduce_with(|left, right| Ok(select_argmax(left, right)))
.transpose()?;
[wide_best, mid_best, tail_best]
.into_iter()
.flatten()
.reduce(select_argmax)
.map(|best| best.0)
.context("matrix has no output rows")
}
fn select_argmax_scores<const N: usize>(start: usize, values: [f32; N]) -> Result<(usize, f32)> {
let Some(&first) = values.first() else {
bail!("matrix tile has no output rows");
};
let mut best = (start, first);
if !best.1.is_finite() {
bail!("matrix output {start} is not finite");
}
for (offset, &value) in values.iter().enumerate().skip(1) {
if !value.is_finite() {
bail!("matrix output {} is not finite", start + offset);
}
best = select_argmax(best, (start + offset, value));
}
Ok(best)
}
fn select_argmax(left: (usize, f32), right: (usize, f32)) -> (usize, f32) {
match right.1.total_cmp(&left.1) {
std::cmp::Ordering::Greater => right,
std::cmp::Ordering::Equal if right.0 < left.0 => right,
_ => left,
}
}
fn matrix_argmax_with(output: usize, dot: impl Fn(usize) -> Result<f32> + Sync) -> Result<usize> {
let best = (0..output)
.into_par_iter()
.map(|index| -> Result<(usize, f32)> {
let value = dot(index)?;
if !value.is_finite() {
bail!("matrix output {index} is not finite");
}
Ok((index, value))
})
.try_reduce_with(|left, right| Ok(select_argmax(left, right)))
.transpose()?
.context("matrix has no 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");
IsaKernel::detect().project(matrix, format, activation, output)
}
fn project_quantized_with<K>(
matrix: &Tensor<'_>,
format: QuantizedFormat,
activation: &Q8Activation,
output: &mut [f32],
) -> Result<()>
where
K: QuantizedKernel<Q8_0Format> + QuantizedKernel<Q1_0Format>,
{
match format {
QuantizedFormat::Q8_0 => {
project_quantized_format_with::<K, Q8_0Format>(matrix, activation, output)
}
QuantizedFormat::Q1_0 => {
project_quantized_q1_rows_with::<K>(matrix, activation, output, false)
}
}
}
fn project_quantized_format_with<K, F>(
matrix: &Tensor<'_>,
activation: &Q8Activation,
output: &mut [f32],
) -> Result<()>
where
F: QuantizedFormatMarker,
K: QuantizedKernel<F>,
{
output
.par_iter_mut()
.enumerate()
.try_for_each(|(row, value)| -> Result<()> {
*value = <K as QuantizedKernel<F>>::dot(matrix.encoded_row(row)?, activation);
Ok(())
})
}
fn project_quantized_q1_rows_with<K>(
matrix: &Tensor<'_>,
activation: &Q8Activation,
output: &mut [f32],
add: bool,
) -> Result<()>
where
K: QuantizedKernel<Q1_0Format>,
{
let wide_count = output.len() / WIDE_MATRIX_ROW_TILE;
let (wide, rest) = output.split_at_mut(wide_count * WIDE_MATRIX_ROW_TILE);
wide.par_chunks_exact_mut(WIDE_MATRIX_ROW_TILE)
.enumerate()
.try_for_each(|(tile, values)| -> Result<()> {
let start = tile * WIDE_MATRIX_ROW_TILE;
apply_q1_dots(
values,
<K as QuantizedKernel<Q1_0Format>>::dot_wide_rows(
encoded_rows(matrix, start)?,
activation,
),
add,
);
Ok(())
})?;
let rest_start = wide_count * WIDE_MATRIX_ROW_TILE;
let tile_count = rest.len() / MATRIX_ROW_TILE;
let (head, tail) = rest.split_at_mut(tile_count * MATRIX_ROW_TILE);
head.par_chunks_exact_mut(MATRIX_ROW_TILE)
.enumerate()
.try_for_each(|(tile, values)| -> Result<()> {
let start = rest_start + tile * MATRIX_ROW_TILE;
apply_q1_dots(
values,
<K as QuantizedKernel<Q1_0Format>>::dot_rows(
encoded_rows(matrix, start)?,
activation,
),
add,
);
Ok(())
})?;
let tail_start = rest_start + tile_count * MATRIX_ROW_TILE;
tail.iter_mut()
.enumerate()
.try_for_each(|(offset, value)| -> Result<()> {
let dot = <K as QuantizedKernel<Q1_0Format>>::dot(
matrix.encoded_row(tail_start + offset)?,
activation,
);
if add {
*value += dot;
} else {
*value = dot;
}
Ok(())
})
}
fn encoded_rows<'a, const N: usize>(matrix: &Tensor<'a>, start: usize) -> Result<[&'a [u8]; N]> {
let mut rows = [&[][..]; N];
for (offset, row) in rows.iter_mut().enumerate() {
*row = matrix.encoded_row(start + offset)?;
}
Ok(rows)
}
fn apply_q1_dots<const N: usize>(values: &mut [f32], dots: [f32; N], add: bool) {
if add {
for (value, dot) in values.iter_mut().zip(dots) {
*value += dot;
}
} else {
values.copy_from_slice(&dots);
}
}
fn project_quantized_q1_swiglu_with<K>(
gate: &Tensor<'_>,
up: &Tensor<'_>,
input: &Q8Activation,
output: &mut Q8Activation,
) -> Result<()>
where
K: QuantizedKernel<Q1_0Format> + ActivationKernel,
{
output
.values
.par_chunks_exact_mut(BLOCK_VALUES)
.zip(output.scales.par_iter_mut())
.enumerate()
.try_for_each(|(tile, (values, scale))| -> Result<()> {
let start = tile * BLOCK_VALUES;
let gate_dots = q1_block_dots::<K>(gate, start, input)?;
let up_dots = q1_block_dots::<K>(up, start, input)?;
let mut block = [0.0_f32; BLOCK_VALUES];
for ((value, gate), up) in block.iter_mut().zip(gate_dots).zip(up_dots) {
*value = (gate / (1.0 + (-gate).exp())) * up;
}
let quantized = <&mut [i8; BLOCK_VALUES]>::try_from(values)
.context("Q8 output block width differs")?;
*scale = K::quantize_block(&block, quantized).context("Q8 activation is not finite")?;
Ok(())
})
}
fn q1_block_dots<K>(
matrix: &Tensor<'_>,
start: usize,
activation: &Q8Activation,
) -> Result<[f32; BLOCK_VALUES]>
where
K: QuantizedKernel<Q1_0Format>,
{
let mut dots = [0.0_f32; BLOCK_VALUES];
for tile in 0..BLOCK_VALUES / WIDE_MATRIX_ROW_TILE {
let offset = tile * WIDE_MATRIX_ROW_TILE;
let values = dots
.get_mut(offset..offset + WIDE_MATRIX_ROW_TILE)
.context("Q1 SwiGLU tile is out of range")?;
apply_q1_dots(
values,
<K as QuantizedKernel<Q1_0Format>>::dot_wide_rows(
encoded_rows(matrix, start + offset)?,
activation,
),
false,
);
}
Ok(dots)
}
fn matrix_dimensions(matrix: &Tensor<'_>) -> Result<[usize; 2]> {
match matrix.dimensions() {
[input, output] => Ok([*input, *output]),
dimensions => bail!("matrix has dimensions {dimensions:?}, expected two"),
}
}
pub(super) fn dot_f32(left: &[f32], right: &[f32]) -> Result<f32> {
ensure!(left.len() == right.len(), "dot product lengths differ");
Ok(DenseIsaKernel::detect().dot_f32(left, right))
}
fn dot_f32_scalar(left: &[f32], right: &[f32]) -> f32 {
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>()
}
const GQA_QUERY_GROUP: usize = 4;
pub(super) fn dot_f32_fp16_group(queries: &[f32], key: &[Fp16]) -> Result<[f32; GQA_QUERY_GROUP]> {
let width = key.len();
let expected = width
.checked_mul(GQA_QUERY_GROUP)
.context("GQA query group size overflow")?;
ensure!(queries.len() == expected, "GQA query group width differs");
ensure!(width != 0, "GQA key width is zero");
#[cfg(target_arch = "x86_64")]
if fp16_fma_avx2() {
return Ok(unsafe { dot_f32_fp16_group_avx2(queries, key) });
}
Ok(dot_f32_fp16_group_scalar(queries, key))
}
pub(super) fn scale_accumulate_fp16_group(
dests: &mut [f32],
value: &[Fp16],
rescales: [f32; GQA_QUERY_GROUP],
weights: [f32; GQA_QUERY_GROUP],
) -> Result<()> {
let width = value.len();
let expected = width
.checked_mul(GQA_QUERY_GROUP)
.context("GQA accumulate group size overflow")?;
ensure!(
dests.len() == expected,
"GQA accumulate group width differs"
);
ensure!(width != 0, "GQA value width is zero");
#[cfg(target_arch = "x86_64")]
if fp16_fma_avx2() {
unsafe { scale_accumulate_fp16_group_avx2(dests, value, rescales, weights) };
return Ok(());
}
scale_accumulate_fp16_group_scalar(dests, value, rescales, weights);
Ok(())
}
fn dot_f32_fp16_group_scalar(queries: &[f32], key: &[Fp16]) -> [f32; GQA_QUERY_GROUP] {
let width = key.len();
let mut scores = [0.0_f32; GQA_QUERY_GROUP];
for (index, &key) in key.iter().enumerate() {
let key = f32::from(key);
for (query_index, score) in scores.iter_mut().enumerate() {
if let Some(&query) = queries.get(query_index * width + index) {
*score += query * key;
}
}
}
scores
}
fn scale_accumulate_fp16_group_scalar(
dests: &mut [f32],
value: &[Fp16],
rescales: [f32; GQA_QUERY_GROUP],
weights: [f32; GQA_QUERY_GROUP],
) {
let width = value.len();
for (index, &value) in value.iter().enumerate() {
let value = f32::from(value);
for query_index in 0..GQA_QUERY_GROUP {
if let Some(dest) = dests.get_mut(query_index * width + index) {
*dest = *dest * rescales[query_index] + weights[query_index] * value;
}
}
}
}
#[cfg(target_arch = "x86_64")]
fn fp16_fma_avx2() -> bool {
static AVAILABLE: OnceLock<bool> = OnceLock::new();
*AVAILABLE.get_or_init(|| {
std::is_x86_feature_detected!("avx2")
&& std::is_x86_feature_detected!("f16c")
&& std::is_x86_feature_detected!("fma")
})
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn load_fp16x8(values: &[Fp16], index: usize) -> std::arch::x86_64::__m256 {
use std::arch::x86_64::{__m128i, _mm256_cvtph_ps};
let packed = unsafe {
values
.as_ptr()
.add(index)
.cast::<[Fp16; 8]>()
.read_unaligned()
};
let half: __m128i = bytemuck::cast(packed);
_mm256_cvtph_ps(half)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn horizontal_sum_avx2(sums: std::arch::x86_64::__m256) -> f32 {
use std::arch::x86_64::_mm256_storeu_ps;
let mut lanes = [0.0; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sums) };
lanes.into_iter().sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn dot_f32_fp16_group_avx2(queries: &[f32], key: &[Fp16]) -> [f32; GQA_QUERY_GROUP] {
use std::arch::x86_64::{_mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps};
let width = key.len();
let vectorized = width / 8 * 8;
let mut sums = [_mm256_setzero_ps(); GQA_QUERY_GROUP];
for index in (0..vectorized).step_by(8) {
let key = unsafe { load_fp16x8(key, index) };
for (query_index, sum) in sums.iter_mut().enumerate() {
let query =
unsafe { _mm256_loadu_ps(queries.as_ptr().add(query_index * width + index)) };
*sum = _mm256_fmadd_ps(query, key, *sum);
}
}
let mut scores = [0.0_f32; GQA_QUERY_GROUP];
for (query_index, (score, sum)) in scores.iter_mut().zip(sums).enumerate() {
let start = query_index * width + vectorized;
let end = query_index * width + width;
*score = unsafe { horizontal_sum_avx2(sum) }
+ queries
.get(start..end)
.into_iter()
.flatten()
.zip(&key[vectorized..])
.map(|(query, key)| query * f32::from(*key))
.sum::<f32>();
}
scores
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn scale_accumulate_fp16_group_avx2(
dests: &mut [f32],
value: &[Fp16],
rescales: [f32; GQA_QUERY_GROUP],
weights: [f32; GQA_QUERY_GROUP],
) {
use std::arch::x86_64::{
_mm256_fmadd_ps, _mm256_loadu_ps, _mm256_mul_ps, _mm256_set1_ps, _mm256_storeu_ps,
};
let width = value.len();
let vectorized = width / 8 * 8;
let rescale_v = [
_mm256_set1_ps(rescales[0]),
_mm256_set1_ps(rescales[1]),
_mm256_set1_ps(rescales[2]),
_mm256_set1_ps(rescales[3]),
];
let weight_v = [
_mm256_set1_ps(weights[0]),
_mm256_set1_ps(weights[1]),
_mm256_set1_ps(weights[2]),
_mm256_set1_ps(weights[3]),
];
for index in (0..vectorized).step_by(8) {
let value = unsafe { load_fp16x8(value, index) };
for query_index in 0..GQA_QUERY_GROUP {
let dest = unsafe { dests.as_mut_ptr().add(query_index * width + index) };
let dest_v = unsafe { _mm256_loadu_ps(dest) };
let updated = _mm256_fmadd_ps(
weight_v[query_index],
value,
_mm256_mul_ps(dest_v, rescale_v[query_index]),
);
unsafe { _mm256_storeu_ps(dest, updated) };
}
}
for (index, &value) in value[vectorized..].iter().enumerate() {
let value = f32::from(value);
let offset = vectorized + index;
for query_index in 0..GQA_QUERY_GROUP {
if let Some(dest) = dests.get_mut(query_index * width + offset) {
*dest = *dest * rescales[query_index] + weights[query_index] * value;
}
}
}
}
pub(super) fn rms_norm_into(
values: &[f32],
width: usize,
weight: &[f32],
epsilon: f32,
output: &mut [f32],
) -> Result<()> {
ensure!(
width != 0 && values.len().is_multiple_of(width),
"invalid RMS norm shape"
);
ensure!(weight.len() == width, "invalid RMS norm weight");
ensure!(output.len() == values.len(), "invalid RMS norm output");
for (input, output) in values
.chunks_exact(width)
.zip(output.chunks_exact_mut(width))
{
let width_f32 = f32::from(u16::try_from(width).context("RMS norm width exceeds u16")?);
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(())
}
pub(super) fn rms_norm_quantize_into(
values: &[f32],
width: usize,
weight: &[f32],
epsilon: f32,
activation: &mut Q8Activation,
) -> Result<()> {
ensure!(
width != 0 && values.len().is_multiple_of(width),
"invalid RMS norm shape"
);
ensure!(weight.len() == width, "invalid RMS norm weight");
ensure!(
width.is_multiple_of(BLOCK_VALUES) && values.len().is_multiple_of(BLOCK_VALUES),
"Q8 activation width is not divisible by 32"
);
activation.scales.resize(values.len() / BLOCK_VALUES, 0.0);
activation.values.resize(values.len(), 0);
IsaKernel::detect().rms_norm_quantize_into(values, width, weight, epsilon, activation)
}
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];
}
}