#![warn(missing_docs, clippy::missing_docs_in_private_items)]
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn ensure_rayon_global_pool() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
if super::cpu_features::env_disabled("CERA_RAYON_GLOBAL") {
return;
}
let n = rayon_pool_width(
super::cpu_features::env_usize("RAYON_NUM_THREADS"),
super::cpu_features::performance_core_count(),
);
let cores: &'static [usize] = super::threadpool::perf_pinned_cores();
if let Err(err) = rayon::ThreadPoolBuilder::new()
.num_threads(n)
.start_handler(move |worker_index| {
#[cfg(target_os = "macos")]
super::threadpool::set_macos_thread_qos_interactive();
if !cores.is_empty() && !super::threadpool::set_current_thread_affinity(cores) {
tracing::debug!(
"cera: rayon worker {worker_index} kept its inherited CPU mask; \
sched_setaffinity({cores:?}) was refused"
);
}
})
.build_global()
{
tracing::warn!(
"cera: rayon global pool already initialized; P-core width ({n}) and affinity not applied to residual rayon sites: {err}"
);
}
});
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
fn rayon_pool_width(env_override: Option<usize>, perf_cores: usize) -> usize {
env_override.unwrap_or(perf_cores).max(1)
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn ensure_rayon_global_pool() {}
#[cfg(not(feature = "parallel"))]
pub fn ensure_rayon_global_pool() {}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn configure_thread_pool() -> usize {
ensure_rayon_global_pool();
super::threadpool::RowPool::prefill().num_threads()
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn configure_thread_pool() -> usize {
1
}
#[cfg(not(feature = "parallel"))]
pub fn configure_thread_pool() -> usize {
1
}
use crate::quant::{
BlockQ4_0, BlockQ4_1, BlockQ4KM, BlockQ5K, BlockQ8_0, f16_to_f32, vec_dot_q4_0_f32,
vec_dot_q4_1_f32, vec_dot_q4_k_m_f32, vec_dot_q5_k_f32, vec_dot_q8_0_f32,
};
#[cfg(not(target_arch = "aarch64"))]
use crate::quant::{BlockQ6K, vec_dot_q6_k_f32};
use crate::tensor::DType;
use std::mem::size_of;
pub fn matmul_f32(a: &[f32], b: &[f32], c: &mut [f32], m: usize, n: usize, k: usize) {
debug_assert_eq!(a.len(), m * k);
debug_assert_eq!(b.len(), k * n);
debug_assert_eq!(c.len(), m * n);
for i in 0..m {
for p in 0..k {
let a_val = a[i * k + p];
for j in 0..n {
c[i * n + j] += a_val * b[p * n + j];
}
}
}
}
pub fn matmul_q4_0_f32(a_quant: &[u8], b: &[f32], c: &mut [f32], m: usize, n: usize, k: usize) {
debug_assert_eq!(k % 32, 0);
let blocks_per_row = k / 32;
let bytes_per_row = blocks_per_row * size_of::<BlockQ4_0>();
debug_assert_eq!(a_quant.len(), m * bytes_per_row);
debug_assert_eq!(b.len(), k * n);
debug_assert_eq!(c.len(), m * n);
for i in 0..m {
let row_start = i * bytes_per_row;
for j in 0..n {
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let block_offset = row_start + bi * size_of::<BlockQ4_0>();
let block = unsafe { &*(a_quant.as_ptr().add(block_offset) as *const BlockQ4_0) };
let col_start = bi * 32;
let b_slice: Vec<f32> = (0..32).map(|l| b[(col_start + l) * n + j]).collect();
sum += vec_dot_q4_0_f32(block, &b_slice);
}
c[i * n + j] = sum;
}
}
}
pub fn matmul_q8_0_f32(a_quant: &[u8], b: &[f32], c: &mut [f32], m: usize, n: usize, k: usize) {
debug_assert_eq!(k % 32, 0);
let blocks_per_row = k / 32;
let bytes_per_row = blocks_per_row * size_of::<BlockQ8_0>();
debug_assert_eq!(a_quant.len(), m * bytes_per_row);
debug_assert_eq!(b.len(), k * n);
debug_assert_eq!(c.len(), m * n);
for i in 0..m {
let row_start = i * bytes_per_row;
for j in 0..n {
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let block_offset = row_start + bi * size_of::<BlockQ8_0>();
let block = unsafe { &*(a_quant.as_ptr().add(block_offset) as *const BlockQ8_0) };
let col_start = bi * 32;
let b_slice: Vec<f32> = (0..32).map(|l| b[(col_start + l) * n + j]).collect();
sum += vec_dot_q8_0_f32(block, &b_slice);
}
c[i * n + j] = sum;
}
}
}
pub fn matmul_q4km_f32(a_quant: &[u8], b: &[f32], c: &mut [f32], m: usize, n: usize, k: usize) {
debug_assert_eq!(k % 256, 0);
let blocks_per_row = k / 256;
let bytes_per_row = blocks_per_row * size_of::<BlockQ4KM>();
debug_assert_eq!(a_quant.len(), m * bytes_per_row);
debug_assert_eq!(b.len(), k * n);
debug_assert_eq!(c.len(), m * n);
for i in 0..m {
let row_start = i * bytes_per_row;
for j in 0..n {
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let block_offset = row_start + bi * size_of::<BlockQ4KM>();
let block = unsafe { &*(a_quant.as_ptr().add(block_offset) as *const BlockQ4KM) };
let col_start = bi * 256;
let b_slice: Vec<f32> = (0..256).map(|l| b[(col_start + l) * n + j]).collect();
sum += vec_dot_q4_k_m_f32(block, &b_slice);
}
c[i * n + j] = sum;
}
}
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_rows(y: &mut [f32], min_rows: usize, f: impl Fn((usize, &mut f32)) + Sync + Send) {
super::threadpool::RowPool::decode().dispatch_rows(y, 1, min_rows, |row, slice| {
f((row, &mut slice[0]));
});
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_rows_slice(
y: &mut [f32],
min_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
if y.is_empty() {
return;
}
let total_rows = y.len();
let y_ptr = y.as_mut_ptr() as usize;
par_range(total_rows, min_rows, move |start, count| {
let slice =
unsafe { core::slice::from_raw_parts_mut((y_ptr as *mut f32).add(start), count) };
f((start, slice));
});
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_range(total_rows: usize, min_rows: usize, f: impl Fn(usize, usize) + Sync + Send) {
if total_rows == 0 {
return;
}
let pool = super::threadpool::RowPool::decode();
let nth = pool.num_threads().max(1);
let chunk = total_rows.div_ceil(nth).max(min_rows).max(1);
let n_chunks = total_rows.div_ceil(chunk);
let mut dummy_chunks = [0.0f32; 128];
if n_chunks <= dummy_chunks.len() {
pool.dispatch_rows_chunked(&mut dummy_chunks[..n_chunks], 1, 1, 1, |t, slice| {
let m_start = t * chunk;
let count = (total_rows.saturating_sub(m_start)).min(chunk * slice.len());
if count > 0 {
f(m_start, count);
}
});
} else {
let mut dummy = vec![0.0f32; n_chunks];
pool.dispatch_rows_chunked(&mut dummy, 1, 1, 1, |t, slice| {
let m_start = t * chunk;
let count = (total_rows.saturating_sub(m_start)).min(chunk * slice.len());
if count > 0 {
f(m_start, count);
}
});
}
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_range(total_rows: usize, _min_rows: usize, f: impl Fn(usize, usize) + Sync + Send) {
f(0, total_rows);
}
#[cfg(not(feature = "parallel"))]
pub fn par_range(total_rows: usize, _min_rows: usize, f: impl Fn(usize, usize) + Sync + Send) {
f(0, total_rows);
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_rows(y: &mut [f32], min_rows: usize, f: impl Fn((usize, &mut f32)) + Sync + Send) {
use crate::par::{IndexedParallelIterator, ParallelIterator, ParallelSliceMut};
const WASM_MIN_CHUNK_ROWS: usize = 512;
let chunk_size = (y.len() / crate::par::current_num_threads())
.max(min_rows)
.max(WASM_MIN_CHUNK_ROWS);
y.par_chunks_mut(chunk_size)
.enumerate()
.for_each(|(ci, chunk)| {
let base = ci * chunk_size;
for (j, yi) in chunk.iter_mut().enumerate() {
f((base + j, yi));
}
});
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_rows_slice(
y: &mut [f32],
min_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
use crate::par::{IndexedParallelIterator, ParallelIterator, ParallelSliceMut};
const WASM_MIN_CHUNK_ROWS: usize = 512;
let chunk_size = (y.len() / crate::par::current_num_threads())
.max(min_rows)
.max(WASM_MIN_CHUNK_ROWS);
y.par_chunks_mut(chunk_size)
.enumerate()
.for_each(|(ci, chunk)| {
f((ci * chunk_size, chunk));
});
}
#[cfg(not(feature = "parallel"))]
pub fn par_rows(y: &mut [f32], _min_rows: usize, f: impl Fn((usize, &mut f32))) {
for (i, yi) in y.iter_mut().enumerate() {
f((i, yi));
}
}
#[cfg(not(feature = "parallel"))]
pub fn par_rows_slice(y: &mut [f32], _min_rows: usize, f: impl Fn((usize, &mut [f32]))) {
f((0, y));
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_rows_n(
y: &mut [f32],
n: usize,
min_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
debug_assert_ne!(n, 0, "par_rows_n: n must be > 0");
if n == 0 || y.is_empty() {
return;
}
super::threadpool::RowPool::prefill().dispatch_rows(y, n, min_rows, |row, row_slice| {
f((row, row_slice));
});
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_rows_n(
y: &mut [f32],
n: usize,
min_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
debug_assert_ne!(n, 0, "par_rows_n: n must be > 0");
if n == 0 || y.is_empty() {
return;
}
use crate::par::{IndexedParallelIterator, ParallelIterator, ParallelSliceMut};
let m = y.len() / n;
let rows_per_chunk = (m / crate::par::current_num_threads()).max(min_rows.max(1));
let elems_per_chunk = rows_per_chunk * n;
y.par_chunks_mut(elems_per_chunk)
.enumerate()
.for_each(|(ci, chunk)| {
let base_row = ci * rows_per_chunk;
for (j, row) in chunk.chunks_mut(n).enumerate() {
f((base_row + j, row));
}
});
}
#[cfg(not(feature = "parallel"))]
pub fn par_rows_n(y: &mut [f32], n: usize, _min_rows: usize, f: impl Fn((usize, &mut [f32]))) {
debug_assert_ne!(n, 0, "par_rows_n: n must be > 0");
if n == 0 || y.is_empty() {
return;
}
for (j, row) in y.chunks_mut(n).enumerate() {
f((j, row));
}
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_rows_n_chunked(
y: &mut [f32],
n: usize,
min_rows: usize,
min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
par_rows_n_chunked_on(
super::threadpool::RowPool::prefill(),
y,
n,
min_rows,
min_chunk_rows,
f,
);
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
fn par_rows_n_chunked_on(
pool: &'static super::threadpool::RowPool,
y: &mut [f32],
n: usize,
min_rows: usize,
min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
debug_assert_ne!(n, 0, "par_rows_n_chunked_on: n must be > 0");
if n == 0 || y.is_empty() {
return;
}
pool.dispatch_rows_chunked(y, n, min_rows, min_chunk_rows, |row, row_slice| {
f((row, row_slice));
});
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_rows_n_chunked(
y: &mut [f32],
n: usize,
min_rows: usize,
_min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
par_rows_n(y, n, min_rows, f);
}
#[cfg(not(feature = "parallel"))]
pub fn par_rows_n_chunked(
y: &mut [f32],
n: usize,
min_rows: usize,
_min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])),
) {
par_rows_n(y, n, min_rows, f);
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_rows_n_chunked_decode(
y: &mut [f32],
n: usize,
min_rows: usize,
min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
par_rows_n_chunked_on(
super::threadpool::RowPool::decode(),
y,
n,
min_rows,
min_chunk_rows,
f,
);
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_rows_n_chunked_decode(
y: &mut [f32],
n: usize,
min_rows: usize,
min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
par_rows_n_chunked(y, n, min_rows, min_chunk_rows, f);
}
#[cfg(not(feature = "parallel"))]
pub fn par_rows_n_chunked_decode(
y: &mut [f32],
n: usize,
min_rows: usize,
min_chunk_rows: usize,
f: impl Fn((usize, &mut [f32])),
) {
par_rows_n_chunked(y, n, min_rows, min_chunk_rows, f);
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn decode_par_threads() -> usize {
super::threadpool::RowPool::decode().num_threads()
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn decode_par_threads() -> usize {
crate::par::current_num_threads()
}
#[cfg(not(feature = "parallel"))]
pub fn decode_par_threads() -> usize {
1
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
pub fn par_rows_n_work(
y: &mut [f32],
n: usize,
min_rows: usize,
depth: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
debug_assert_ne!(n, 0, "par_rows_n_work: n must be > 0");
if n == 0 || y.is_empty() {
return;
}
super::threadpool::RowPool::prefill().dispatch_rows_work(
y,
n,
min_rows,
depth,
|row, row_slice| {
f((row, row_slice));
},
);
}
#[cfg(all(feature = "parallel", target_arch = "wasm32"))]
pub fn par_rows_n_work(
y: &mut [f32],
n: usize,
min_rows: usize,
_depth: usize,
f: impl Fn((usize, &mut [f32])) + Sync + Send,
) {
par_rows_n(y, n, min_rows, f);
}
#[cfg(not(feature = "parallel"))]
pub fn par_rows_n_work(
y: &mut [f32],
n: usize,
min_rows: usize,
_depth: usize,
f: impl Fn((usize, &mut [f32])),
) {
par_rows_n(y, n, min_rows, f);
}
#[allow(clippy::ptr_arg)]
pub fn gemv_q4_0_f32(
a_quant: &[u8],
x: &[f32],
y: &mut [f32],
m: usize,
k: usize,
q8_scales: &mut Vec<f32>,
q8_quants: &mut Vec<i8>,
) {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(y.len(), m);
debug_assert_eq!(k % 32, 0, "Q4_0 GEMV: k must be divisible by 32");
let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * size_of::<BlockQ4_0>();
debug_assert_eq!(a_quant.len(), m * row_bytes);
#[cfg(target_arch = "aarch64")]
{
unsafe {
crate::backend::simd::neon::gemv_q4_0_f32_neon(
a_quant, x, y, m, k, q8_scales, q8_quants,
);
}
}
#[cfg(not(target_arch = "aarch64"))]
{
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
if vnni_int8_available() {
q8_scales.resize(blocks_per_row, 0.0);
q8_quants.resize(k, 0);
unsafe {
quantize_f32_to_q8_0_into(x, q8_scales, q8_quants);
crate::backend::simd::avx512_vnni::gemv_q4_0_q8_0(
a_quant, q8_scales, q8_quants, y, m, k,
);
}
return;
}
#[cfg(target_arch = "x86_64")]
if avx2_int8_available() {
q8_scales.resize(blocks_per_row, 0.0);
q8_quants.resize(k, 0);
unsafe {
quantize_f32_to_q8_0_into(x, q8_scales, q8_quants);
crate::backend::simd::avx2_int8::gemv_q4_0_q8_0(
a_quant, q8_scales, q8_quants, y, m, k,
);
}
return;
}
let _ = (q8_scales, q8_quants);
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_start = i * row_bytes;
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let offset = row_start + bi * size_of::<BlockQ4_0>();
let block = unsafe { &*(a_quant.as_ptr().add(offset) as *const BlockQ4_0) };
sum += vec_dot_q4_0_f32(block, &x[bi * 32..(bi + 1) * 32]);
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
}
pub const GEMV_PAR_THRESHOLD_DEFAULT: usize = 256;
#[cfg(feature = "parallel")]
pub fn gemv_par_threshold() -> usize {
use std::sync::OnceLock;
static THRESHOLD: OnceLock<usize> = OnceLock::new();
*THRESHOLD.get_or_init(|| {
super::cpu_features::env_usize("CERA_PAR_THRESHOLD").unwrap_or(GEMV_PAR_THRESHOLD_DEFAULT)
})
}
#[cfg(not(feature = "parallel"))]
pub fn gemv_par_threshold() -> usize {
usize::MAX
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
pub const PREQUANT_PAR_MIN_COLS_DEFAULT: usize = 256;
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
pub fn prequant_par_min_cols() -> usize {
use std::sync::OnceLock;
static MIN_COLS: OnceLock<usize> = OnceLock::new();
*MIN_COLS.get_or_init(|| {
super::cpu_features::env_usize("CERA_PREQUANT_MIN_COLS")
.unwrap_or(PREQUANT_PAR_MIN_COLS_DEFAULT)
})
}
pub const GEMV_MIN_ROWS_DEFAULT: usize = 128;
#[cfg(feature = "parallel")]
pub fn gemv_min_rows() -> usize {
use std::sync::OnceLock;
static MIN_ROWS: OnceLock<usize> = OnceLock::new();
*MIN_ROWS.get_or_init(|| {
super::cpu_features::env_usize("CERA_MIN_ROWS").unwrap_or(GEMV_MIN_ROWS_DEFAULT)
})
}
#[cfg(not(feature = "parallel"))]
pub fn gemv_min_rows() -> usize {
GEMV_MIN_ROWS_DEFAULT
}
#[cfg(not(target_arch = "aarch64"))]
pub(crate) fn quantize_f32_to_q8_0_scalar(x: &[f32], scales: &mut [f32], quants: &mut [i8]) {
for (bi, blk) in x.chunks(32).enumerate() {
let amax = blk.iter().fold(0.0f32, |a, &v| a.max(v.abs()));
let d = amax / 127.0;
let id = match 1.0 / d {
r if d != 0.0 && r.is_finite() => r,
_ => 0.0,
};
scales[bi] = crate::quant::f16_to_f32(crate::quant::f32_to_f16(d));
for (t, &v) in blk.iter().enumerate() {
quants[bi * 32 + t] = (v * id).round_ties_even().clamp(-128.0, 127.0) as i8;
}
}
}
#[doc(hidden)]
pub fn quantize_f32_to_q8_0_into(x: &[f32], scales: &mut [f32], quants: &mut [i8]) {
assert_eq!(
x.len() % 32,
0,
"quantize_f32_to_q8_0_into: x.len() must be divisible by 32"
);
assert!(
scales.len() >= x.len() / 32 && quants.len() >= x.len(),
"quantize_f32_to_q8_0_into: scales/quants too small for x.len()={} \
(need {} scales, {} quants; got {} and {})",
x.len(),
x.len() / 32,
x.len(),
scales.len(),
quants.len()
);
#[cfg(target_arch = "aarch64")]
unsafe {
crate::backend::simd::neon::quantize_f32_to_q8_0_neon(x, scales, quants);
}
#[cfg(not(target_arch = "aarch64"))]
{
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
if avx512_quantizer_available() {
unsafe {
crate::backend::simd::avx512_vnni::quantize_f32_to_q8_0_avx512(x, scales, quants);
}
return;
}
quantize_f32_to_q8_0_scalar(x, scales, quants);
}
}
#[doc(hidden)]
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
pub fn avx512_quantizer_available() -> bool {
let f = crate::backend::cpu_features::cpu_features();
f.tier >= crate::backend::cpu_features::CpuTier::Avx512 && f.avx512vl && f.avx2
}
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
pub(crate) fn vnni_int8_available() -> bool {
crate::backend::cpu_features::cpu_features().tier
>= crate::backend::cpu_features::CpuTier::Avx512Vnni
}
#[cfg(target_arch = "x86_64")]
pub(crate) fn avx2_int8_available() -> bool {
crate::backend::cpu_features::cpu_features().tier >= crate::backend::cpu_features::CpuTier::Avx2
}
pub fn int8_gemm_available() -> bool {
#[cfg(target_arch = "aarch64")]
{
true
}
#[cfg(target_arch = "x86_64")]
{
avx2_int8_available()
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
{
false
}
}
#[cfg(target_arch = "aarch64")]
#[cfg_attr(has_blas, allow(dead_code))]
pub(crate) fn repack_q4_0_8x8(src: &[u8], m: usize, k: usize) -> (Vec<u8>, Vec<f32>) {
assert!(
m.is_multiple_of(8),
"repack_q4_0_8x8: m must be a multiple of 8"
);
assert!(
k.is_multiple_of(32),
"repack_q4_0_8x8: k must be a multiple of 32"
);
let nb = k / 32;
let bsz = size_of::<crate::quant::BlockQ4_0>();
assert_eq!(
src.len(),
m * nb * bsz,
"repack_q4_0_8x8: src is {} bytes, need {} for {m}x{k}",
src.len(),
m * nb * bsz,
);
let sr_count = m / 8;
let mut packed = vec![0u8; sr_count * nb * 256];
let mut scales = vec![0.0f32; sr_count * nb * 8];
let nibble = |row: usize, b: usize, e: usize| -> i8 {
let qs = (row * nb + b) * bsz + 2;
let raw = if e < 16 {
src[qs + e] & 0x0F
} else {
src[qs + e - 16] >> 4
};
(raw as i8) - 8
};
for sr in 0..sr_count {
for b in 0..nb {
for r in 0..8 {
let off = ((8 * sr + r) * nb + b) * bsz;
scales[(sr * nb + b) * 8 + r] =
f16_to_f32(u16::from_le_bytes([src[off], src[off + 1]]));
}
let base = (sr * nb + b) * 256;
for g in 0..8usize {
for chunk in 0..2usize {
for r in 0..4usize {
for c in 0..4usize {
let e = 4 * g + c;
let s = nibble(8 * sr + chunk * 4 + r, b, e);
packed[base + g * 32 + chunk * 16 + r * 4 + c] = s as u8;
}
}
}
}
}
}
(packed, scales)
}
#[cfg(target_arch = "x86_64")]
#[cfg_attr(has_blas, allow(dead_code))]
pub(crate) fn repack_q4_0_8x8(src: &[u8], m: usize, k: usize) -> (Vec<u8>, Vec<f32>) {
assert!(
m.is_multiple_of(8),
"repack_q4_0_8x8: m must be a multiple of 8"
);
assert!(
k.is_multiple_of(32),
"repack_q4_0_8x8: k must be a multiple of 32"
);
let nb = k / 32;
let bsz = size_of::<crate::quant::BlockQ4_0>();
assert_eq!(
src.len(),
m * nb * bsz,
"repack_q4_0_8x8: src is {} bytes, need {} for {m}x{k}",
src.len(),
m * nb * bsz,
);
let sr_count = m / 8;
let mut packed = vec![0u8; sr_count * nb * 128];
let mut scales = vec![0.0f32; sr_count * nb * 8];
let nibble = |row: usize, b: usize, e: usize| -> u8 {
let qs = (row * nb + b) * bsz + 2;
if e < 16 {
src[qs + e] & 0x0F
} else {
src[qs + e - 16] >> 4
}
};
for sr in 0..sr_count {
for b in 0..nb {
for r in 0..8 {
let off = ((8 * sr + r) * nb + b) * bsz;
scales[(sr * nb + b) * 8 + r] =
f16_to_f32(u16::from_le_bytes([src[off], src[off + 1]]));
}
for g in 0..8usize {
for i in 0..16usize {
let e = 4 * g + (i % 4);
let lo = nibble(8 * sr + i / 4, b, e);
let hi = nibble(8 * sr + 4 + i / 4, b, e);
packed[(sr * nb + b) * 128 + g * 16 + i] = lo | (hi << 4);
}
}
}
}
(packed, scales)
}
#[cfg(all(any(target_arch = "x86_64", target_arch = "aarch64"), not(has_blas)))]
pub(crate) fn q4_0_repack_supported(m: usize, k: usize) -> bool {
#[cfg(target_arch = "aarch64")]
return m.is_multiple_of(16) && k.is_multiple_of(32);
#[cfg(target_arch = "x86_64")]
return int8_gemm_available() && m.is_multiple_of(8) && k.is_multiple_of(32);
}
#[cfg(all(any(target_arch = "x86_64", target_arch = "aarch64"), not(has_blas)))]
#[allow(clippy::too_many_arguments)]
pub(crate) fn gemm_preq_repacked_q4_0_dispatch(
packed: &[u8],
scales: &[f32],
b_scales: &[f32],
b_quants: &[i8],
out: &mut [f32],
m: usize,
n: usize,
k: usize,
) -> bool {
let nb = k / 32;
#[cfg(target_arch = "aarch64")]
{
assert!(
k.is_multiple_of(32) && m.is_multiple_of(16),
"gemm_preq_repacked_q4_0_dispatch: need k%32==0 and m%16==0, got m={m} k={k}"
);
let expected_bytes_per_block = 512;
assert!(
packed.len() >= (m / 16) * nb * expected_bytes_per_block
&& scales.len() >= (m / 16) * nb * 16,
"gemm_preq_repacked_q4_0_dispatch: repacked weights too small for {m}x{k}"
);
}
#[cfg(target_arch = "x86_64")]
{
assert!(
k.is_multiple_of(32) && m.is_multiple_of(8),
"gemm_preq_repacked_q4_0_dispatch: need k%32==0 and m%8==0, got m={m} k={k}"
);
let expected_bytes_per_block = 128;
assert!(
packed.len() >= (m / 8) * nb * expected_bytes_per_block
&& scales.len() >= (m / 8) * nb * 8,
"gemm_preq_repacked_q4_0_dispatch: repacked weights too small for {m}x{k}"
);
}
assert!(
b_quants.len() >= n * k && b_scales.len() >= n * nb && out.len() == m * n,
"gemm_preq_repacked_q4_0_dispatch: activation/output buffers wrong for {m}x{n}x{k}"
);
#[cfg(target_arch = "aarch64")]
if crate::backend::simd::neon::k_quant_gemm_available() {
unsafe {
crate::backend::simd::neon::gemm_q4_0_8x8_q8_0(
packed, scales, b_scales, b_quants, out, m, n, k,
);
}
return true;
}
#[cfg(target_arch = "x86_64")]
{
#[cfg(feature = "avx512")]
if vnni_int8_available() {
unsafe {
crate::backend::simd::avx512_vnni::gemm_q4_0_8x8_q8_0(
packed, scales, b_scales, b_quants, out, m, n, k,
);
}
return true;
}
if avx2_int8_available() {
unsafe {
crate::backend::simd::avx2_int8::gemm_q4_0_8x8_q8_0(
packed, scales, b_scales, b_quants, out, m, n, k,
);
}
return true;
}
}
false
}
#[cfg(all(any(target_arch = "x86_64", target_arch = "aarch64"), not(has_blas)))]
#[allow(clippy::too_many_arguments)]
pub(crate) fn gemm_preq_repacked_q4_0_rowmajor_dispatch(
packed: &[u8],
scales: &[f32],
b_scales: &[f32],
b_quants: &[i8],
out: &mut [f32],
n: usize,
m: usize,
k: usize,
) -> bool {
let nb = k / 32;
#[cfg(target_arch = "aarch64")]
{
assert!(
k.is_multiple_of(32) && m.is_multiple_of(8),
"gemm_preq_repacked_q4_0_rowmajor_dispatch: need k%32==0 and m%8==0, got m={m} k={k}"
);
let expected_bytes_per_block = 256;
assert!(
packed.len() >= (m / 8) * nb * expected_bytes_per_block
&& scales.len() >= (m / 8) * nb * 8,
"gemm_preq_repacked_q4_0_rowmajor_dispatch: repacked weights too small for {m}x{k}"
);
}
#[cfg(target_arch = "x86_64")]
{
assert!(
k.is_multiple_of(32) && m.is_multiple_of(8),
"gemm_preq_repacked_q4_0_rowmajor_dispatch: need k%32==0 and m%8==0, got m={m} k={k}"
);
let expected_bytes_per_block = 128;
assert!(
packed.len() >= (m / 8) * nb * expected_bytes_per_block
&& scales.len() >= (m / 8) * nb * 8,
"gemm_preq_repacked_q4_0_rowmajor_dispatch: repacked weights too small for {m}x{k}"
);
}
assert!(
b_quants.len() >= n * k && b_scales.len() >= n * nb && out.len() == m * n,
"gemm_preq_repacked_q4_0_rowmajor_dispatch: activation/output buffers wrong for {m}x{n}x{k}"
);
#[cfg(target_arch = "aarch64")]
if crate::backend::simd::neon::k_quant_gemm_available() {
unsafe {
crate::backend::simd::neon::gemm_q4_0_8x8_q8_0_rowmajor(
packed, scales, b_scales, b_quants, out, n, m, k,
);
}
return true;
}
false
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_preq_repacked_q4_0_gate_up_silu_dispatch(
gate_packed: &[u8],
gate_scales: &[f32],
up_packed: &[u8],
up_scales: &[f32],
b_scales: &[f32],
b_quants: &[i8],
out: &mut [f32],
m: usize,
n: usize,
k: usize,
) -> bool {
let _ = (gate_packed, gate_scales, up_packed, up_scales);
let nb = k / 32;
#[cfg(target_arch = "aarch64")]
{
assert!(
m.is_multiple_of(8) && k.is_multiple_of(32),
"gemm_preq_repacked_q4_0_gate_up_silu_dispatch: m={m} must be %8 and k={k} must be %32"
);
let sr_count = m / 8;
assert_eq!(
gate_packed.len(),
sr_count * nb * 256,
"gate packed bytes mismatch"
);
assert_eq!(
up_packed.len(),
sr_count * nb * 256,
"up packed bytes mismatch"
);
assert_eq!(gate_scales.len(), sr_count * nb * 8, "gate scales mismatch");
assert_eq!(up_scales.len(), sr_count * nb * 8, "up scales mismatch");
}
assert!(
b_quants.len() >= n * k && b_scales.len() >= n * nb && out.len() == m * n,
"gemm_preq_repacked_q4_0_gate_up_silu_dispatch: activation/output buffers wrong for {m}x{n}x{k}"
);
#[cfg(target_arch = "aarch64")]
if crate::backend::simd::neon::k_quant_gemm_available() {
unsafe {
crate::backend::simd::neon::gemm_q4_0_gate_up_silu_rowmajor(
gate_packed,
gate_scales,
up_packed,
up_scales,
b_scales,
b_quants,
out,
n,
m,
k,
);
}
return true;
}
false
}
#[cfg(target_arch = "x86_64")]
#[cfg_attr(has_blas, allow(dead_code))]
pub(crate) fn repack_q4_k_8x8(src: &[u8], m: usize, k: usize) -> (Vec<u8>, Vec<f32>, Vec<f32>) {
assert!(
m.is_multiple_of(8),
"repack_q4_k_8x8: m must be a multiple of 8"
);
assert!(
k.is_multiple_of(256),
"repack_q4_k_8x8: k must be a multiple of 256"
);
let sb = k / 256; let nb32 = k / 32; let bsz = size_of::<crate::quant::BlockQ4KM>();
assert_eq!(
src.len(),
m * sb * bsz,
"repack_q4_k_8x8: src is {} bytes, need {} for {m}x{k}",
src.len(),
m * sb * bsz,
);
let sr_count = m / 8;
let mut packed = vec![0u8; sr_count * nb32 * 128];
let mut dsc = vec![0.0f32; sr_count * nb32 * 8];
let mut dmn = vec![0.0f32; sr_count * nb32 * 8];
const D_OFF: usize = std::mem::offset_of!(crate::quant::BlockQ4KM, d);
const DMIN_OFF: usize = std::mem::offset_of!(crate::quant::BlockQ4KM, dmin);
const SC_OFF: usize = std::mem::offset_of!(crate::quant::BlockQ4KM, scales);
const QS_OFF: usize = std::mem::offset_of!(crate::quant::BlockQ4KM, qs);
let nibble = |row: usize, block: usize, e: usize| -> u8 {
let bi = block / 8;
let s = block % 8;
let qs = (row * sb + bi) * bsz + QS_OFF;
let byte = src[qs + (s / 2) * 32 + e];
if s.is_multiple_of(2) {
byte & 0x0F
} else {
byte >> 4
}
};
for sr in 0..sr_count {
for bi in 0..sb {
for r in 0..8 {
let off = ((8 * sr + r) * sb + bi) * bsz;
let d = f16_to_f32(u16::from_le_bytes([src[off + D_OFF], src[off + D_OFF + 1]]));
let dmin = f16_to_f32(u16::from_le_bytes([
src[off + DMIN_OFF],
src[off + DMIN_OFF + 1],
]));
let scales_bytes: &[u8; 12] =
src[off + SC_OFF..off + SC_OFF + 12].try_into().unwrap();
let (sc, mn) = crate::quant::decode_q4km_scales(scales_bytes);
for s in 0..8 {
let block = bi * 8 + s;
dsc[(sr * nb32 + block) * 8 + r] = d * sc[s] as f32;
dmn[(sr * nb32 + block) * 8 + r] = dmin * mn[s] as f32;
}
}
for s in 0..8usize {
let block = bi * 8 + s;
for g in 0..8usize {
for i in 0..16usize {
let e = 4 * g + (i % 4);
let lo = nibble(8 * sr + i / 4, block, e);
let hi = nibble(8 * sr + 4 + i / 4, block, e);
packed[(sr * nb32 + block) * 128 + g * 16 + i] = lo | (hi << 4);
}
}
}
}
}
(packed, dsc, dmn)
}
#[cfg(all(target_arch = "x86_64", not(has_blas)))]
pub(crate) fn q4_k_repack_supported(m: usize, k: usize) -> bool {
m.is_multiple_of(8) && k.is_multiple_of(256) && int8_gemm_available()
}
#[cfg(all(target_arch = "x86_64", not(has_blas)))]
#[allow(clippy::too_many_arguments)]
pub(crate) fn gemm_preq_repacked_q4_k_dispatch(
packed: &[u8],
dsc: &[f32],
dmn: &[f32],
b_scales: &[f32],
b_quants: &[i8],
out: &mut [f32],
m: usize,
n: usize,
k: usize,
) -> bool {
let nb32 = k / 32;
assert!(
k.is_multiple_of(256) && m.is_multiple_of(8),
"gemm_preq_repacked_q4_k_dispatch: need k%256==0 and m%8==0, got m={m} k={k}"
);
assert!(
packed.len() >= (m / 8) * nb32 * 128
&& dsc.len() >= (m / 8) * nb32 * 8
&& dmn.len() >= (m / 8) * nb32 * 8,
"gemm_preq_repacked_q4_k_dispatch: repacked weights too small for {m}x{k}"
);
assert!(
b_quants.len() >= n * k && b_scales.len() >= n * nb32 && out.len() == m * n,
"gemm_preq_repacked_q4_k_dispatch: activation/output buffers wrong for {m}x{n}x{k}"
);
#[cfg(feature = "avx512")]
if vnni_int8_available() {
unsafe {
crate::backend::simd::avx512_vnni::gemm_q4_k_8x8_q8_0(
packed, dsc, dmn, b_scales, b_quants, out, m, n, k,
);
}
return true;
}
if avx2_int8_available() {
unsafe {
crate::backend::simd::avx2_int8::gemm_q4_k_8x8_q8_0(
packed, dsc, dmn, b_scales, b_quants, out, m, n, k,
);
}
return true;
}
false
}
#[allow(unused_variables)]
#[allow(clippy::too_many_arguments)]
#[doc(hidden)]
pub fn gemm_preq_dispatch(
dtype: DType,
data: &[u8],
b_scales: &[f32],
b_quants: &[i8],
out: &mut [f32],
m: usize,
n: usize,
k: usize,
) -> bool {
let blocks = k / dtype.block_size();
assert!(
k.is_multiple_of(dtype.block_size()),
"gemm_preq_dispatch: k={k} is not a multiple of {:?}'s block size",
dtype
);
assert!(
data.len() >= m * blocks * dtype.block_bytes(),
"gemm_preq_dispatch: weights are {} bytes, need {} for {m}x{k} {:?}",
data.len(),
m * blocks * dtype.block_bytes(),
dtype
);
assert!(
b_quants.len() >= n * k && b_scales.len() >= n * (k / 32) && out.len() == m * n,
"gemm_preq_dispatch: activation/output buffers wrong for {m}x{n}x{k} \
(quants {}, scales {}, out {} — out must be exactly {})",
b_quants.len(),
b_scales.len(),
out.len(),
m * n
);
#[cfg(target_arch = "aarch64")]
unsafe {
use crate::backend::simd::neon;
match dtype {
DType::Q4_0 => {
neon::gemm_q4_0_q8_0_neon(data, b_scales, b_quants, out, m, n, k);
true
}
DType::Q8_0 => {
neon::gemm_q8_0_q8_0_neon(data, b_scales, b_quants, out, m, n, k);
true
}
DType::Q4_1 => neon::gemm_q4_1_q8_0_neon(data, b_scales, b_quants, out, m, n, k),
DType::Q4KM => neon::gemm_q4_k_q8_0_neon(data, b_scales, b_quants, out, m, n, k),
DType::Q6K => neon::gemm_q6_k_q8_0_neon(data, b_scales, b_quants, out, m, n, k),
_ => false,
}
}
#[cfg(not(target_arch = "aarch64"))]
{
#[cfg(target_arch = "x86_64")]
macro_rules! x86_int8_gemm {
($m:path) => {{
use $m as kern;
match dtype {
DType::Q4_0 => {
kern::gemm_q4_0_q8_0(data, b_scales, b_quants, out, m, n, k);
true
}
DType::Q8_0 => {
kern::gemm_q8_0_q8_0(data, b_scales, b_quants, out, m, n, k);
true
}
DType::Q4_1 => {
kern::gemm_q4_1_q8_0(data, b_scales, b_quants, out, m, n, k);
true
}
DType::Q4KM => {
kern::gemm_q4_k_q8_0(data, b_scales, b_quants, out, m, n, k);
true
}
DType::Q6K => {
kern::gemm_q6_k_q8_0(data, b_scales, b_quants, out, m, n, k);
true
}
_ => false,
}
}};
}
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
if vnni_int8_available() {
return unsafe { x86_int8_gemm!(crate::backend::simd::avx512_vnni) };
}
#[cfg(target_arch = "x86_64")]
if avx2_int8_available() {
return unsafe { x86_int8_gemm!(crate::backend::simd::avx2_int8) };
}
false
}
}
#[cfg(target_arch = "aarch64")]
pub fn quantize_f32_to_q8_0(x: &[f32]) -> (Vec<f32>, Vec<i8>) {
assert_eq!(
x.len() % 32,
0,
"quantize_f32_to_q8_0: x.len() must be divisible by 32"
);
let n_blocks = x.len() / 32;
let mut scales = vec![0.0f32; n_blocks];
let mut quants = vec![0i8; x.len()];
unsafe {
crate::backend::simd::neon::quantize_f32_to_q8_0_neon(x, &mut scales, &mut quants);
}
(scales, quants)
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
pub fn gemv_with_preq(
dtype: DType,
a_quant: &[u8],
x_scales: &[f32],
x_quants: &[i8],
x_f32: &[f32],
y: &mut [f32],
m: usize,
k: usize,
) {
match dtype {
DType::Q4_0 => gemv_q4_0_with_q8(a_quant, x_scales, x_quants, y, m, k),
DType::Q8_0 => unsafe {
crate::backend::simd::neon::gemv_q8_0_q8_0_neon(a_quant, x_scales, x_quants, y, m, k)
},
DType::Q6K => unsafe {
crate::backend::simd::neon::gemv_q6k_q8_0_neon(a_quant, x_scales, x_quants, y, m, k)
},
DType::Q5KM => unsafe {
crate::backend::simd::neon::gemv_q5k_q8_0_neon(a_quant, x_scales, x_quants, y, m, k)
},
DType::Q4_1 if q4_1_gemm_available() => {
if !gemm_preq_dispatch(
DType::Q4_1,
a_quant,
&x_scales[..k / 32],
&x_quants[..k],
y,
m,
1,
k,
) {
gemv_dispatch(dtype, a_quant, x_f32, y, m, k, None);
}
}
_ => gemv_dispatch(dtype, a_quant, x_f32, y, m, k, None),
}
}
#[cfg(target_arch = "aarch64")]
#[allow(dead_code)]
pub fn gemv_with_preq_argmax(
dtype: DType,
a_quant: &[u8],
x_scales: &[f32],
x_quants: &[i8],
x_f32: &[f32],
m: usize,
k: usize,
) -> usize {
match dtype {
DType::Q6K => unsafe {
crate::backend::simd::neon::gemv_q6k_q8_0_argmax_neon(a_quant, x_scales, x_quants, m, k)
},
_ => {
let mut logits = vec![0.0f32; m];
gemv_with_preq(dtype, a_quant, x_scales, x_quants, x_f32, &mut logits, m, k);
crate::sampler::argmax(&logits) as usize
}
}
}
#[cfg(target_arch = "aarch64")]
pub fn gemv_q4_0_with_q8(
a_quant: &[u8],
x_scales: &[f32],
x_quants: &[i8],
y: &mut [f32],
m: usize,
k: usize,
) {
unsafe {
crate::backend::simd::neon::gemv_q4_0_q8_0_neon(a_quant, x_scales, x_quants, y, m, k);
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
pub fn gemv_q4_0_fused2_with_q8(
a1_quant: &[u8],
a2_quant: &[u8],
x_scales: &[f32],
x_quants: &[i8],
y1: &mut [f32],
y2: &mut [f32],
m: usize,
k: usize,
) {
unsafe {
crate::backend::simd::neon::gemv_q4_0_q8_0_fused2_neon(
a1_quant, a2_quant, x_scales, x_quants, y1, y2, m, k,
);
}
}
#[cfg(target_arch = "aarch64")]
pub fn gemv_q4_0_gate_up_swiglu_with_q8(
gate_quant: &[u8],
up_quant: &[u8],
x_scales: &[f32],
x_quants: &[i8],
out: &mut [f32],
m: usize,
k: usize,
) {
unsafe {
crate::backend::simd::neon::gemv_q4_0_gate_up_swiglu_neon(
gate_quant, up_quant, x_scales, x_quants, out, m, k,
);
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
pub fn gemv_q4_0_concat3_with_q8(
a1_quant: &[u8],
a2_quant: &[u8],
a3_quant: &[u8],
x_scales: &[f32],
x_quants: &[i8],
y1: &mut [f32],
y2: &mut [f32],
y3: &mut [f32],
m1: usize,
m2: usize,
m3: usize,
k: usize,
) {
let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * std::mem::size_of::<crate::quant::BlockQ4_0>();
assert!(
k.is_multiple_of(32),
"gemv_q4_0_concat3_with_q8: k must be a multiple of 32"
);
assert!(
a1_quant.len() >= m1 * row_bytes,
"a1_quant buffer underflow"
);
assert!(
a2_quant.len() >= m2 * row_bytes,
"a2_quant buffer underflow"
);
assert!(
a3_quant.len() >= m3 * row_bytes,
"a3_quant buffer underflow"
);
assert!(
x_scales.len() >= blocks_per_row,
"x_scales buffer underflow"
);
assert!(x_quants.len() >= k, "x_quants buffer underflow");
assert!(y1.len() >= m1, "y1 buffer underflow");
assert!(y2.len() >= m2, "y2 buffer underflow");
assert!(y3.len() >= m3, "y3 buffer underflow");
unsafe {
crate::backend::simd::neon::gemv_q4_0_q8_0_concat3_neon(
a1_quant, a2_quant, a3_quant, x_scales, x_quants, y1, y2, y3, m1, m2, m3, k,
);
}
}
#[allow(clippy::ptr_arg)]
pub fn gemv_q8_0_f32(
a_quant: &[u8],
x: &[f32],
y: &mut [f32],
m: usize,
k: usize,
q8_scales: &mut Vec<f32>,
q8_quants: &mut Vec<i8>,
) {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(y.len(), m);
debug_assert_eq!(k % 32, 0, "Q8_0 GEMV: k must be divisible by 32");
let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * size_of::<BlockQ8_0>();
debug_assert_eq!(a_quant.len(), m * row_bytes);
#[cfg(target_arch = "aarch64")]
{
unsafe {
crate::backend::simd::neon::gemv_q8_0_f32_neon(
a_quant, x, y, m, k, q8_scales, q8_quants,
);
}
}
#[cfg(not(target_arch = "aarch64"))]
{
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
if vnni_int8_available() {
q8_scales.resize(blocks_per_row, 0.0);
q8_quants.resize(k, 0);
unsafe {
quantize_f32_to_q8_0_into(x, q8_scales, q8_quants);
crate::backend::simd::avx512_vnni::gemv_q8_0_q8_0(
a_quant, q8_scales, q8_quants, y, m, k,
);
}
return;
}
#[cfg(target_arch = "x86_64")]
if avx2_int8_available() {
q8_scales.resize(blocks_per_row, 0.0);
q8_quants.resize(k, 0);
unsafe {
quantize_f32_to_q8_0_into(x, q8_scales, q8_quants);
crate::backend::simd::avx2_int8::gemv_q8_0_q8_0(
a_quant, q8_scales, q8_quants, y, m, k,
);
}
return;
}
let _ = (q8_scales, q8_quants);
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_start = i * row_bytes;
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let offset = row_start + bi * size_of::<BlockQ8_0>();
let block = unsafe { &*(a_quant.as_ptr().add(offset) as *const BlockQ8_0) };
sum += vec_dot_q8_0_f32(block, &x[bi * 32..(bi + 1) * 32]);
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
}
#[allow(clippy::ptr_arg)]
#[allow(unused_variables)]
pub fn gemv_q6k_f32(
a_quant: &[u8],
x: &[f32],
y: &mut [f32],
m: usize,
k: usize,
q8_scales: &mut Vec<f32>,
q8_quants: &mut Vec<i8>,
) {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(y.len(), m);
debug_assert_eq!(k % 256, 0, "Q6_K GEMV: k must be divisible by 256");
#[cfg(target_arch = "aarch64")]
{
unsafe {
crate::backend::simd::neon::gemv_q6k_f32_neon(
a_quant, x, y, m, k, q8_scales, q8_quants,
);
}
}
#[cfg(not(target_arch = "aarch64"))]
{
let blocks_per_row = k / 256;
let row_bytes = blocks_per_row * size_of::<BlockQ6K>();
debug_assert_eq!(a_quant.len(), m * row_bytes);
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_start = i * row_bytes;
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let offset = row_start + bi * size_of::<BlockQ6K>();
let block = unsafe { &*(a_quant.as_ptr().add(offset) as *const BlockQ6K) };
sum += vec_dot_q6_k_f32(block, &x[bi * 256..(bi + 1) * 256]);
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
}
pub fn gemv_q4km_f32(a_quant: &[u8], x: &[f32], y: &mut [f32], m: usize, k: usize) {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(y.len(), m);
debug_assert_eq!(k % 256, 0, "Q4_K_M GEMV: k must be divisible by 256");
let blocks_per_row = k / 256;
let row_bytes = blocks_per_row * size_of::<BlockQ4KM>();
debug_assert_eq!(a_quant.len(), m * row_bytes);
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_start = i * row_bytes;
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let offset = row_start + bi * size_of::<BlockQ4KM>();
let block = unsafe { &*(a_quant.as_ptr().add(offset) as *const BlockQ4KM) };
sum += vec_dot_q4_k_m_f32(block, &x[bi * 256..(bi + 1) * 256]);
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
pub fn gemv_q5km_f32(a_quant: &[u8], x: &[f32], y: &mut [f32], m: usize, k: usize) {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(y.len(), m);
debug_assert_eq!(k % 256, 0, "Q5_K GEMV: k must be divisible by 256");
let blocks_per_row = k / 256;
let row_bytes = blocks_per_row * size_of::<BlockQ5K>();
debug_assert_eq!(a_quant.len(), m * row_bytes);
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_start = i * row_bytes;
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let offset = row_start + bi * size_of::<BlockQ5K>();
let block = unsafe { &*(a_quant.as_ptr().add(offset) as *const BlockQ5K) };
sum += vec_dot_q5_k_f32(block, &x[bi * 256..(bi + 1) * 256]);
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
fn q4_1_gemm_available() -> bool {
#[cfg(target_arch = "aarch64")]
{
crate::backend::simd::neon::k_quant_gemm_available()
}
#[cfg(target_arch = "x86_64")]
{
int8_gemm_available()
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
{
false
}
}
#[allow(clippy::ptr_arg)]
pub fn gemv_q4_1_f32(
a_quant: &[u8],
x: &[f32],
y: &mut [f32],
m: usize,
k: usize,
q8_scales: &mut Vec<f32>,
q8_quants: &mut Vec<i8>,
) {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(y.len(), m);
debug_assert_eq!(k % 32, 0, "Q4_1 GEMV: k must be divisible by 32");
let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * size_of::<BlockQ4_1>();
debug_assert_eq!(a_quant.len(), m * row_bytes);
if q4_1_gemm_available() {
q8_scales.resize(blocks_per_row, 0.0);
q8_quants.resize(k, 0);
quantize_f32_to_q8_0_into(x, q8_scales, q8_quants);
if gemm_preq_dispatch(DType::Q4_1, a_quant, q8_scales, q8_quants, y, m, 1, k) {
return;
}
}
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_start = i * row_bytes;
let mut sum = 0.0f32;
for bi in 0..blocks_per_row {
let offset = row_start + bi * size_of::<BlockQ4_1>();
let block = unsafe { &*(a_quant.as_ptr().add(offset) as *const BlockQ4_1) };
sum += vec_dot_q4_1_f32(block, &x[bi * 32..(bi + 1) * 32]);
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
#[allow(clippy::chunks_exact_to_as_chunks)]
fn dot_f32_neon(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "dot_f32: input lengths must match");
use std::arch::aarch64::*;
let mut sum_v0 = unsafe { vdupq_n_f32(0.0) };
let mut sum_v1 = unsafe { vdupq_n_f32(0.0) };
let mut sum_v2 = unsafe { vdupq_n_f32(0.0) };
let mut sum_v3 = unsafe { vdupq_n_f32(0.0) };
let mut a_chunks = a.chunks_exact(16);
let mut b_chunks = b.chunks_exact(16);
for (ca, cb) in a_chunks.by_ref().zip(b_chunks.by_ref()) {
unsafe {
sum_v0 = vfmaq_f32(sum_v0, vld1q_f32(ca.as_ptr()), vld1q_f32(cb.as_ptr()));
sum_v1 = vfmaq_f32(
sum_v1,
vld1q_f32(ca.as_ptr().add(4)),
vld1q_f32(cb.as_ptr().add(4)),
);
sum_v2 = vfmaq_f32(
sum_v2,
vld1q_f32(ca.as_ptr().add(8)),
vld1q_f32(cb.as_ptr().add(8)),
);
sum_v3 = vfmaq_f32(
sum_v3,
vld1q_f32(ca.as_ptr().add(12)),
vld1q_f32(cb.as_ptr().add(12)),
);
}
}
let sum_v01 = unsafe { vaddq_f32(sum_v0, sum_v1) };
let sum_v23 = unsafe { vaddq_f32(sum_v2, sum_v3) };
let mut sum = unsafe { vaddvq_f32(vaddq_f32(sum_v01, sum_v23)) };
for (&x, &y) in a_chunks.remainder().iter().zip(b_chunks.remainder().iter()) {
sum += x * y;
}
sum
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[inline(always)]
#[allow(clippy::chunks_exact_to_as_chunks)]
fn dot_f32_wasm_simd128(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "dot_f32: input lengths must match");
use core::arch::wasm32::*;
let mut sum_v0 = f32x4_splat(0.0);
let mut sum_v1 = f32x4_splat(0.0);
let mut a_chunks = a.chunks_exact(8);
let mut b_chunks = b.chunks_exact(8);
for (ca, cb) in a_chunks.by_ref().zip(b_chunks.by_ref()) {
unsafe {
let va0 = v128_load(ca.as_ptr() as *const v128);
let vb0 = v128_load(cb.as_ptr() as *const v128);
sum_v0 = f32x4_add(sum_v0, f32x4_mul(va0, vb0));
let va1 = v128_load(ca.as_ptr().add(4) as *const v128);
let vb1 = v128_load(cb.as_ptr().add(4) as *const v128);
sum_v1 = f32x4_add(sum_v1, f32x4_mul(va1, vb1));
}
}
let sum_v = f32x4_add(sum_v0, sum_v1);
let mut sum = f32x4_extract_lane::<0>(sum_v)
+ f32x4_extract_lane::<1>(sum_v)
+ f32x4_extract_lane::<2>(sum_v)
+ f32x4_extract_lane::<3>(sum_v);
for (&x, &y) in a_chunks.remainder().iter().zip(b_chunks.remainder().iter()) {
sum += x * y;
}
sum
}
#[inline(always)]
fn dot_f32_scalar_fallback(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "dot_f32: input lengths must match");
let (a_chunks, a_rem) = a.as_chunks::<8>();
let (b_chunks, b_rem) = b.as_chunks::<8>();
let mut sum0 = 0.0f32;
let mut sum1 = 0.0f32;
let mut sum2 = 0.0f32;
let mut sum3 = 0.0f32;
let mut sum4 = 0.0f32;
let mut sum5 = 0.0f32;
let mut sum6 = 0.0f32;
let mut sum7 = 0.0f32;
for (ca, cb) in a_chunks.iter().zip(b_chunks.iter()) {
sum0 += ca[0] * cb[0];
sum1 += ca[1] * cb[1];
sum2 += ca[2] * cb[2];
sum3 += ca[3] * cb[3];
sum4 += ca[4] * cb[4];
sum5 += ca[5] * cb[5];
sum6 += ca[6] * cb[6];
sum7 += ca[7] * cb[7];
}
let mut sum = ((sum0 + sum1) + (sum2 + sum3)) + ((sum4 + sum5) + (sum6 + sum7));
for (&x, &y) in a_rem.iter().zip(b_rem.iter()) {
sum += x * y;
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx", enable = "fma")]
#[allow(clippy::chunks_exact_to_as_chunks)]
unsafe fn dot_f32_avx_fma(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "dot_f32: input lengths must match");
use std::arch::x86_64::*;
unsafe {
let mut sum_v0 = _mm256_setzero_ps();
let mut sum_v1 = _mm256_setzero_ps();
let mut a_chunks = a.chunks_exact(16);
let mut b_chunks = b.chunks_exact(16);
for (ca, cb) in a_chunks.by_ref().zip(b_chunks.by_ref()) {
let va0 = _mm256_loadu_ps(ca.as_ptr());
let vb0 = _mm256_loadu_ps(cb.as_ptr());
sum_v0 = _mm256_fmadd_ps(va0, vb0, sum_v0);
let va1 = _mm256_loadu_ps(ca.as_ptr().add(8));
let vb1 = _mm256_loadu_ps(cb.as_ptr().add(8));
sum_v1 = _mm256_fmadd_ps(va1, vb1, sum_v1);
}
let sum256 = _mm256_add_ps(sum_v0, sum_v1);
let sum128 = _mm_add_ps(
_mm256_castps256_ps128(sum256),
_mm256_extractf128_ps(sum256, 1),
);
let sum64 = _mm_add_ps(sum128, _mm_movehl_ps(sum128, sum128));
let sum32 = _mm_add_ss(sum64, _mm_shuffle_ps(sum64, sum64, 1));
let mut sum = _mm_cvtss_f32(sum32);
for (&x, &y) in a_chunks.remainder().iter().zip(b_chunks.remainder().iter()) {
sum += x * y;
}
sum
}
}
#[inline(always)]
pub fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "dot_f32 input lengths must match");
#[cfg(target_arch = "aarch64")]
{
return dot_f32_neon(a, b);
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
return dot_f32_wasm_simd128(a, b);
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("fma") && is_x86_feature_detected!("avx") {
return unsafe { dot_f32_avx_fma(a, b) };
}
}
#[allow(unreachable_code)]
dot_f32_scalar_fallback(a, b)
}
pub fn gemv_f32(a: &[u8], x: &[f32], y: &mut [f32], m: usize, k: usize) {
assert_eq!(x.len(), k, "gemv_f32: x must have k elements");
assert_eq!(y.len(), m, "gemv_f32: y must have m elements");
assert!(
a.len() >= m * k * std::mem::size_of::<f32>(),
"gemv_f32: a buffer too small"
);
if let Ok(a_f32) = bytemuck::try_cast_slice::<u8, f32>(a) {
let compute_row = |(i, yi): (usize, &mut f32)| {
let row = &a_f32[i * k..(i + 1) * k];
*yi = dot_f32(row, x);
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
} else {
let k_bytes = k * std::mem::size_of::<f32>();
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_bytes = &a[i * k_bytes..(i + 1) * k_bytes];
let f32_ptr = row_bytes.as_ptr() as *const f32;
let mut sum = 0.0f32;
for (j, &xj) in x.iter().enumerate() {
let val = unsafe { std::ptr::read_unaligned(f32_ptr.add(j)) };
sum += val * xj;
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
}
pub fn gemv_f16(a: &[u8], x: &[f32], y: &mut [f32], m: usize, k: usize) {
assert_eq!(x.len(), k, "gemv_f16: x must have k elements");
debug_assert_eq!(y.len(), m);
if let Ok(a16) = bytemuck::try_cast_slice::<u8, u16>(a) {
let compute_row = |(i, yi): (usize, &mut f32)| {
let row = &a16[i * k..(i + 1) * k];
*yi = dot_f16_f32(row, x);
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
} else {
let k_bytes = k * std::mem::size_of::<u16>();
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_bytes = &a[i * k_bytes..(i + 1) * k_bytes];
let u16_ptr = row_bytes.as_ptr() as *const u16;
let mut sum = 0.0f32;
for (j, &xj) in x.iter().enumerate() {
let val = unsafe { std::ptr::read_unaligned(u16_ptr.add(j)) };
sum += crate::quant::f16_to_f32(val) * xj;
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
}
#[inline]
fn dot_f16_f32(row: &[u16], x: &[f32]) -> f32 {
debug_assert_eq!(row.len(), x.len());
#[cfg(target_arch = "aarch64")]
{
if crate::backend::cpu_features::cpu_features().fp16 {
return unsafe { dot_f16_f32_neon(row, x) };
}
}
row.iter()
.zip(x)
.map(|(&w, &xv)| crate::quant::f16_to_f32(w) * xv)
.sum()
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16")]
unsafe fn dot_f16_f32_neon(row: &[u16], x: &[f32]) -> f32 {
use std::arch::aarch64::*;
unsafe {
let (wp, xp) = (row.as_ptr(), x.as_ptr());
let n = row.len().min(x.len());
let mut acc = [vdupq_n_f32(0.0); 4];
let mut i = 0;
while i + 16 <= n {
for (c, a) in acc.iter_mut().enumerate() {
let off = i + c * 4;
let w = vcvt_f32_f16(vreinterpret_f16_u16(vld1_u16(wp.add(off))));
*a = vfmaq_f32(*a, w, vld1q_f32(xp.add(off)));
}
i += 16;
}
let mut sum = vaddvq_f32(vaddq_f32(
vaddq_f32(acc[0], acc[1]),
vaddq_f32(acc[2], acc[3]),
));
while i < n {
sum += crate::quant::f16_to_f32(*wp.add(i)) * *xp.add(i);
i += 1;
}
sum
}
}
pub fn gemv_bf16(a: &[u8], x: &[f32], y: &mut [f32], m: usize, k: usize) {
assert_eq!(x.len(), k, "gemv_bf16: x must have k elements");
debug_assert_eq!(y.len(), m);
if let Ok(a16) = bytemuck::try_cast_slice::<u8, u16>(a) {
let compute_row = |(i, yi): (usize, &mut f32)| {
let row = &a16[i * k..(i + 1) * k];
*yi = row
.iter()
.zip(x)
.map(|(&w, &xv)| crate::quant::bf16_to_f32(w) * xv)
.sum();
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
} else {
let k_bytes = k * std::mem::size_of::<u16>();
let compute_row = |(i, yi): (usize, &mut f32)| {
let row_bytes = &a[i * k_bytes..(i + 1) * k_bytes];
let u16_ptr = row_bytes.as_ptr() as *const u16;
let mut sum = 0.0f32;
for (j, &xj) in x.iter().enumerate() {
let val = unsafe { std::ptr::read_unaligned(u16_ptr.add(j)) };
sum += crate::quant::bf16_to_f32(val) * xj;
}
*yi = sum;
};
if m >= gemv_par_threshold() {
par_rows(y, gemv_min_rows(), compute_row);
} else {
y.iter_mut().enumerate().for_each(compute_row);
}
}
}
pub fn gemv_dispatch(
dtype: DType,
data: &[u8],
x: &[f32],
y: &mut [f32],
m: usize,
k: usize,
q8_scratch: Option<(&mut Vec<f32>, &mut Vec<i8>)>,
) {
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
macro_rules! kq_gemv {
($f:path) => {{
match q8_scratch {
Some((scales, quants)) => unsafe { $f(data, x, y, m, k, scales, quants) },
None => {
let mut s = Vec::new();
let mut q = Vec::new();
unsafe { $f(data, x, y, m, k, &mut s, &mut q) }
}
}
return;
}};
}
match dtype {
DType::Q4_0 => {
if let Some((scales, quants)) = q8_scratch {
gemv_q4_0_f32(data, x, y, m, k, scales, quants);
} else {
let mut s = Vec::new();
let mut q = Vec::new();
gemv_q4_0_f32(data, x, y, m, k, &mut s, &mut q);
}
}
DType::Q8_0 => {
if let Some((scales, quants)) = q8_scratch {
gemv_q8_0_f32(data, x, y, m, k, scales, quants);
} else {
let mut s = Vec::new();
let mut q = Vec::new();
gemv_q8_0_f32(data, x, y, m, k, &mut s, &mut q);
}
}
DType::Q4_1 => {
if let Some((scales, quants)) = q8_scratch {
gemv_q4_1_f32(data, x, y, m, k, scales, quants);
} else {
let mut s = Vec::new();
let mut q = Vec::new();
gemv_q4_1_f32(data, x, y, m, k, &mut s, &mut q);
}
}
DType::F32 => gemv_f32(data, x, y, m, k),
DType::F16 => gemv_f16(data, x, y, m, k),
DType::BF16 => gemv_bf16(data, x, y, m, k),
DType::Q6K => {
#[cfg(target_arch = "aarch64")]
kq_gemv!(crate::backend::simd::neon::gemv_q6k_f32_neon);
#[cfg(not(target_arch = "aarch64"))]
{
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
if vnni_int8_available() {
kq_gemv!(crate::backend::simd::avx512_vnni::gemv_q6k_f32);
}
#[cfg(target_arch = "x86_64")]
if avx2_int8_available() {
kq_gemv!(crate::backend::simd::avx2_int8::gemv_q6k_f32);
}
let mut s = Vec::new();
let mut q = Vec::new();
gemv_q6k_f32(data, x, y, m, k, &mut s, &mut q);
}
}
DType::Q4KM => {
#[cfg(target_arch = "aarch64")]
kq_gemv!(crate::backend::simd::neon::gemv_q4k_f32_neon);
#[cfg(not(target_arch = "aarch64"))]
{
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
if vnni_int8_available() {
kq_gemv!(crate::backend::simd::avx512_vnni::gemv_q4k_f32);
}
#[cfg(target_arch = "x86_64")]
if avx2_int8_available() {
kq_gemv!(crate::backend::simd::avx2_int8::gemv_q4k_f32);
}
gemv_q4km_f32(data, x, y, m, k);
}
}
DType::Q5KM => {
#[cfg(target_arch = "aarch64")]
kq_gemv!(crate::backend::simd::neon::gemv_q5k_f32_neon);
#[cfg(not(target_arch = "aarch64"))]
gemv_q5km_f32(data, x, y, m, k);
}
_ => panic!("gemv_dispatch: unsupported dtype {:?}", dtype),
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn rmsnorm_neon(
src_ptr: *const f32,
dst_ptr: *mut f32,
weight_ptr: *const f32,
n: usize,
eps: f32,
) {
use core::arch::aarch64::*;
unsafe {
let mut sum_sq0 = vdupq_n_f64(0.0);
let mut sum_sq1 = vdupq_n_f64(0.0);
let n_chunks = n / 4;
for i in 0..n_chunks {
let v = vld1q_f32(src_ptr.add(i * 4));
let v_lo = vcvt_f64_f32(vget_low_f32(v));
let v_hi = vcvt_f64_f32(vget_high_f32(v));
sum_sq0 = vfmaq_f64(sum_sq0, v_lo, v_lo);
sum_sq1 = vfmaq_f64(sum_sq1, v_hi, v_hi);
}
let mut total_sum_sq = vaddvq_f64(vaddq_f64(sum_sq0, sum_sq1));
for i in (n_chunks * 4)..n {
let v = *src_ptr.add(i) as f64;
total_sum_sq += v * v;
}
let mean = total_sum_sq / n as f64;
let rms = (mean + eps as f64).sqrt();
let inv_rms = (1.0 / rms) as f32;
let v_inv_rms = vdupq_n_f32(inv_rms);
for i in 0..n_chunks {
let s = vld1q_f32(src_ptr.add(i * 4));
let w = vld1q_f32(weight_ptr.add(i * 4));
let scaled = vmulq_f32(vmulq_f32(s, v_inv_rms), w);
vst1q_f32(dst_ptr.add(i * 4), scaled);
}
for i in (n_chunks * 4)..n {
*dst_ptr.add(i) = *src_ptr.add(i) * inv_rms * (*weight_ptr.add(i));
}
}
}
pub fn rmsnorm(x: &mut [f32], weight: &[f32], eps: f32) {
debug_assert_eq!(x.len(), weight.len());
#[cfg(target_arch = "aarch64")]
unsafe {
rmsnorm_neon(x.as_ptr(), x.as_mut_ptr(), weight.as_ptr(), x.len(), eps);
}
#[cfg(not(target_arch = "aarch64"))]
{
let n = x.len();
let mut sum_sq = 0.0f64;
for &v in x.iter() {
sum_sq += (v as f64) * (v as f64);
}
let mean = sum_sq / n as f64;
let rms = (mean + eps as f64).sqrt();
let inv_rms = (1.0 / rms) as f32;
for i in 0..n {
x[i] = x[i] * inv_rms * weight[i];
}
}
}
pub fn rmsnorm_into(src: &[f32], dst: &mut [f32], weight: &[f32], eps: f32) {
debug_assert_eq!(src.len(), weight.len());
debug_assert_eq!(dst.len(), weight.len());
#[cfg(target_arch = "aarch64")]
unsafe {
rmsnorm_neon(
src.as_ptr(),
dst.as_mut_ptr(),
weight.as_ptr(),
src.len(),
eps,
);
}
#[cfg(not(target_arch = "aarch64"))]
{
let n = src.len();
let mut sum_sq = 0.0f64;
for &v in src.iter() {
sum_sq += (v as f64) * (v as f64);
}
let mean = sum_sq / n as f64;
let rms = (mean + eps as f64).sqrt();
let inv_rms = (1.0 / rms) as f32;
for i in 0..n {
dst[i] = src[i] * inv_rms * weight[i];
}
}
}
pub fn rmsnorm_and_quantize_q8_0(
x: &[f32],
weight: &[f32],
eps: f32,
scales: &mut [f32],
quants: &mut [i8],
out_normed: Option<&mut [f32]>,
) {
assert_eq!(x.len(), weight.len());
let n = x.len();
assert!(n.is_multiple_of(32));
let n_blocks = n / 32;
assert!(scales.len() >= n_blocks);
assert!(quants.len() >= n);
if let Some(ref out) = out_normed {
assert!(out.len() >= n);
}
#[cfg(target_arch = "aarch64")]
unsafe {
crate::backend::simd::neon::rmsnorm_and_quantize_q8_0_neon(
x, weight, eps, scales, quants, out_normed,
);
}
#[cfg(not(target_arch = "aarch64"))]
{
let mut sum_sq = 0.0f64;
for &v in x.iter() {
sum_sq += (v as f64) * (v as f64);
}
let mean = sum_sq / n as f64;
let rms = (mean + eps as f64).sqrt();
let inv_rms = (1.0 / rms) as f32;
let mut out_opt = out_normed;
#[allow(clippy::needless_range_loop)]
for b in 0..n_blocks {
let b_offset = b * 32;
let mut amax = 0.0f32;
for i in 0..32 {
let v = x[b_offset + i] * inv_rms * weight[b_offset + i];
if let Some(ref mut out) = out_opt {
out[b_offset + i] = v;
}
let av = v.abs();
if av > amax {
amax = av;
}
}
let (d, id) = if amax == 0.0 {
(0.0, 0.0)
} else {
(amax / 127.0, 127.0 / amax)
};
scales[b] = d;
for i in 0..32 {
let v = x[b_offset + i] * inv_rms * weight[b_offset + i];
let q = (v * id).round().clamp(-128.0, 127.0) as i8;
quants[b_offset + i] = q;
}
}
}
}
#[inline(always)]
pub(crate) fn ggml_expf(x: f32) -> f32 {
const R: f32 = f32::from_bits(0x4B400000); const LOG2E: f32 = f32::from_bits(0x3FB8AA3B); const LN2_HI: f32 = f32::from_bits(0x3F317200); const LN2_LO: f32 = f32::from_bits(0x35BFBE8E); const C1: f32 = f32::from_bits(0x3F7FFFF6); const C2: f32 = f32::from_bits(0x3EFFFEDB); const C3: f32 = f32::from_bits(0x3E2AAF33); const C4: f32 = f32::from_bits(0x3D2B9F17); const C5: f32 = f32::from_bits(0x3C072010);
let z = R + x * LOG2E;
let n = z - R;
let b = x - n * LN2_HI - n * LN2_LO;
let e = z.to_bits().wrapping_shl(23);
let k = f32::from_bits(e.wrapping_add(1.0f32.to_bits()));
let u = b * b;
let j = C1 * b + (C2 + C3 * b + (C4 + C5 * b) * u) * u;
let abs_n = f32::from_bits(n.to_bits() & 0x7FFF_FFFF);
if abs_n <= 126.0 {
k + j * k
} else if abs_n > 192.0 {
if n > 0.0 { f32::INFINITY } else { 0.0 }
} else {
let d = if n <= 0.0 { 0x82000000u32 } else { 0u32 };
let s1 = f32::from_bits(d.wrapping_add(0x7f000000));
let s2 = f32::from_bits(e.wrapping_sub(d));
(s2 + s2 * j) * s1
}
}
pub fn silu_inplace(x: &mut [f32]) {
for v in x.iter_mut() {
*v = *v / (1.0 + ggml_expf(-*v));
}
}
pub fn relu_inplace(x: &mut [f32]) {
for v in x.iter_mut() {
if *v < 0.0 {
*v = 0.0;
}
}
}
pub fn silu_mul_inplace(gate: &mut [f32], up: &[f32]) {
debug_assert_eq!(gate.len(), up.len());
let len = gate.len();
if len >= 1024 {
let chunk_size = 512;
let up_ptr = up.as_ptr() as usize;
par_rows_n(gate, chunk_size, 4, move |(idx, g_chunk)| {
let u_chunk = unsafe {
core::slice::from_raw_parts(
(up_ptr as *const f32).add(idx * chunk_size),
g_chunk.len(),
)
};
for (g, &u) in g_chunk.iter_mut().zip(u_chunk.iter()) {
*g = *g / (1.0 + ggml_expf(-*g)) * u;
}
});
} else {
for (g, &u) in gate.iter_mut().zip(up.iter()) {
*g = *g / (1.0 + ggml_expf(-*g)) * u;
}
}
}
pub fn sigmoid_inplace(x: &mut [f32]) {
for v in x.iter_mut() {
*v = 1.0 / (1.0 + ggml_expf(-*v));
}
}
pub fn glu_split(input: &[f32], output: &mut [f32]) {
debug_assert_eq!(input.len() % 2, 0);
let half = input.len() / 2;
debug_assert_eq!(output.len(), half);
let (a, b) = input.split_at(half);
for i in 0..half {
let gate = 1.0 / (1.0 + ggml_expf(-b[i]));
output[i] = a[i] * gate;
}
}
pub fn softmax_inplace(x: &mut [f32]) {
let max = x.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut sum = 0.0f64;
for v in x.iter_mut() {
*v = ggml_expf(*v - max);
sum += *v as f64;
}
let inv_sum = (1.0 / sum) as f32;
for v in x.iter_mut() {
*v *= inv_sum;
}
}
pub fn layer_norm_inplace(x: &mut [f32], weight: &[f32], bias: &[f32], eps: f32) {
debug_assert_eq!(x.len(), weight.len());
debug_assert_eq!(x.len(), bias.len());
let n = x.len();
if n == 0 {
return;
}
let mut sum = 0.0f64;
for &v in x.iter() {
sum += v as f64;
}
let mean = sum / n as f64;
let mut var_sum = 0.0f64;
for &v in x.iter() {
let d = v as f64 - mean;
var_sum += d * d;
}
let var = var_sum / n as f64;
let inv_std = (1.0 / (var + eps as f64).sqrt()) as f32;
let mean_f32 = mean as f32;
for i in 0..n {
x[i] = (x[i] - mean_f32) * inv_std * weight[i] + bias[i];
}
}
pub fn gelu_erf_inplace(x: &mut [f32]) {
const INV_SQRT_2: f32 = std::f32::consts::FRAC_1_SQRT_2;
for v in x.iter_mut() {
*v = 0.5 * *v * (1.0 + erff(*v * INV_SQRT_2));
}
}
pub fn gelu_inplace(x: &mut [f32]) {
const SQRT_2_OVER_PI: f32 = 0.797_884_6; const COEF: f32 = 0.044_715;
for v in x.iter_mut() {
let xv = *v;
let inner = SQRT_2_OVER_PI * (xv + COEF * xv * xv * xv);
*v = 0.5 * xv * (1.0 + inner.tanh());
}
}
#[inline(always)]
fn erff(x: f32) -> f32 {
const A1: f32 = 0.254_829_6;
const A2: f32 = -0.284_496_7;
const A3: f32 = 1.421_413_7;
const A4: f32 = -1.453_152;
const A5: f32 = 1.061_405_4;
const P: f32 = 0.327_591_1;
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let abs_x = x.abs();
let t = 1.0 / (1.0 + P * abs_x);
let y = 1.0 - (((((A5 * t + A4) * t) + A3) * t + A2) * t + A1) * t * ggml_expf(-abs_x * abs_x);
sign * y
}
#[allow(clippy::too_many_arguments)]
pub fn conv1d(
input: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
output: &mut [f32],
in_channels: usize,
out_channels: usize,
t_in: usize,
kernel_size: usize,
stride: usize,
pad: usize,
groups: usize,
) -> usize {
debug_assert!(stride > 0, "stride must be > 0");
debug_assert!(groups > 0, "groups must be > 0");
debug_assert!(kernel_size > 0, "kernel_size must be > 0");
debug_assert_eq!(input.len(), in_channels * t_in);
debug_assert!(in_channels.is_multiple_of(groups));
debug_assert!(out_channels.is_multiple_of(groups));
if let Some(b) = bias {
debug_assert_eq!(b.len(), out_channels);
}
let in_per_group = in_channels / groups;
let out_per_group = out_channels / groups;
debug_assert_eq!(weight.len(), out_channels * in_per_group * kernel_size);
let padded_t_in = t_in + 2 * pad;
debug_assert!(padded_t_in >= kernel_size, "kernel exceeds padded input");
let t_out = (padded_t_in - kernel_size) / stride + 1;
debug_assert_eq!(output.len(), out_channels * t_out);
for g in 0..groups {
for oc_local in 0..out_per_group {
let oc = g * out_per_group + oc_local;
let bias_v = bias.map_or(0.0, |b| b[oc]);
let out_row_start = oc * t_out;
output[out_row_start..out_row_start + t_out].fill(bias_v);
for ic_local in 0..in_per_group {
let ic = g * in_per_group + ic_local;
let weight_row_start = oc * in_per_group * kernel_size + ic_local * kernel_size;
for ot in 0..t_out {
let mut acc = 0.0f32;
for k in 0..kernel_size {
let padded_pos = ot * stride + k;
if padded_pos >= pad && padded_pos < padded_t_in - pad {
let it = padded_pos - pad;
let w = weight[weight_row_start + k];
let x = input[ic * t_in + it];
acc += w * x;
}
}
output[out_row_start + ot] += acc;
}
}
}
}
t_out
}
#[allow(clippy::too_many_arguments)]
pub fn conv2d(
input: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
output: &mut [f32],
in_channels: usize,
out_channels: usize,
h_in: usize,
w_in: usize,
kh: usize,
kw: usize,
stride_h: usize,
stride_w: usize,
pad_h: usize,
pad_w: usize,
groups: usize,
) -> (usize, usize) {
debug_assert!(stride_h > 0, "stride_h must be > 0");
debug_assert!(stride_w > 0, "stride_w must be > 0");
debug_assert!(groups > 0, "groups must be > 0");
debug_assert!(kh > 0, "kh must be > 0");
debug_assert!(kw > 0, "kw must be > 0");
debug_assert_eq!(input.len(), in_channels * h_in * w_in);
debug_assert!(in_channels.is_multiple_of(groups));
debug_assert!(out_channels.is_multiple_of(groups));
if let Some(b) = bias {
debug_assert_eq!(b.len(), out_channels);
}
let in_per_group = in_channels / groups;
let out_per_group = out_channels / groups;
debug_assert_eq!(weight.len(), out_channels * in_per_group * kh * kw);
let two_pad_h = pad_h
.checked_mul(2)
.expect("conv2d: 2 * pad_h overflowed usize");
let two_pad_w = pad_w
.checked_mul(2)
.expect("conv2d: 2 * pad_w overflowed usize");
let padded_h = h_in
.checked_add(two_pad_h)
.expect("conv2d: h_in + 2 * pad_h overflowed usize");
let padded_w = w_in
.checked_add(two_pad_w)
.expect("conv2d: w_in + 2 * pad_w overflowed usize");
debug_assert!(padded_h >= kh, "kh exceeds padded h_in");
debug_assert!(padded_w >= kw, "kw exceeds padded w_in");
let h_out = (padded_h - kh) / stride_h + 1;
let w_out = (padded_w - kw) / stride_w + 1;
let plane_in = h_in
.checked_mul(w_in)
.expect("conv2d: h_in * w_in overflowed usize");
let plane_out = h_out
.checked_mul(w_out)
.expect("conv2d: h_out * w_out overflowed usize");
let kernel_plane = kh
.checked_mul(kw)
.expect("conv2d: kh * kw overflowed usize");
let total_out = out_channels
.checked_mul(plane_out)
.expect("conv2d: out_channels * plane_out overflowed usize");
debug_assert_eq!(output.len(), total_out);
fn gemm_with_bias_broadcast(
output: &mut [f32],
weight: &[f32],
input: &[f32],
bias: Option<&[f32]>,
m: usize,
n: usize,
k: usize,
) {
#[cfg(has_blas)]
{
crate::backend::blas::sgemm_rowmajor_nn(m, n, k, weight, input, output);
if let Some(b) = bias {
for oc in 0..m {
let bias_v = b[oc];
for v in output[oc * n..(oc + 1) * n].iter_mut() {
*v += bias_v;
}
}
}
}
#[cfg(not(has_blas))]
{
for oc in 0..m {
let bias_v = bias.map_or(0.0, |b| b[oc]);
output[oc * n..(oc + 1) * n].fill(bias_v);
}
matmul_f32(weight, input, output, m, n, k);
}
}
let pointwise = kh == 1
&& kw == 1
&& stride_h == 1
&& stride_w == 1
&& pad_h == 0
&& pad_w == 0
&& groups == 1;
if pointwise {
gemm_with_bias_broadcast(
output,
weight,
input,
bias,
out_channels,
plane_out,
in_channels,
);
return (h_out, w_out);
}
if groups == 1 {
let cols = kernel_plane
.checked_mul(in_channels)
.expect("conv2d: kernel_plane * in_channels overflowed usize");
let im2col_len = cols
.checked_mul(plane_out)
.expect("conv2d: im2col buffer size overflowed usize");
let mut im2col = vec![0.0f32; im2col_len];
for ic in 0..in_channels {
let in_plane = ic * plane_in;
for ki in 0..kh {
for kj in 0..kw {
let row_idx = (ic * kh + ki) * kw + kj;
let im_row_start = row_idx * plane_out;
for oh in 0..h_out {
let pad_row = oh * stride_h + ki;
if pad_row < pad_h || pad_row >= h_in + pad_h {
continue;
}
let ih = pad_row - pad_h;
let in_row_start = in_plane + ih * w_in;
let out_row_start = im_row_start + oh * w_out;
for ow in 0..w_out {
let pad_col = ow * stride_w + kj;
if pad_col >= pad_w && pad_col < w_in + pad_w {
let iw = pad_col - pad_w;
im2col[out_row_start + ow] = input[in_row_start + iw];
}
}
}
}
}
}
gemm_with_bias_broadcast(output, weight, &im2col, bias, out_channels, plane_out, cols);
return (h_out, w_out);
}
let depthwise = groups == in_channels && out_channels == in_channels;
if depthwise {
let im2col_len = kernel_plane
.checked_mul(plane_out)
.expect("conv2d: depthwise im2col buffer size overflowed usize");
let do_channel = |im2col: &mut [f32], ic: usize, out_chunk: &mut [f32]| {
let bias_v = bias.map_or(0.0, |b| b[ic]);
out_chunk.fill(bias_v);
let in_plane = ic * plane_in;
for ki in 0..kh {
for kj in 0..kw {
let row_idx = ki * kw + kj;
let im_row_start = row_idx * plane_out;
for oh in 0..h_out {
let pad_row = oh * stride_h + ki;
if pad_row < pad_h || pad_row >= h_in + pad_h {
continue;
}
let ih = pad_row - pad_h;
let in_row_start = in_plane + ih * w_in;
let out_row_start = im_row_start + oh * w_out;
for ow in 0..w_out {
let pad_col = ow * stride_w + kj;
if pad_col >= pad_w && pad_col < w_in + pad_w {
let iw = pad_col - pad_w;
im2col[out_row_start + ow] = input[in_row_start + iw];
}
}
}
}
}
let w_start = ic * kernel_plane;
matmul_f32(
&weight[w_start..w_start + kernel_plane],
im2col,
out_chunk,
1,
plane_out,
kernel_plane,
);
};
#[cfg(feature = "parallel")]
{
use rayon::prelude::*;
output.par_chunks_mut(plane_out).enumerate().for_each_init(
|| vec![0.0f32; im2col_len],
|im2col, (ic, out_chunk)| do_channel(im2col, ic, out_chunk),
);
}
#[cfg(not(feature = "parallel"))]
{
let mut im2col = vec![0.0f32; im2col_len];
for (ic, out_chunk) in output.chunks_mut(plane_out).enumerate() {
do_channel(&mut im2col, ic, out_chunk);
}
}
return (h_out, w_out);
}
for g in 0..groups {
for oc_local in 0..out_per_group {
let oc = g * out_per_group + oc_local;
let bias_v = bias.map_or(0.0, |b| b[oc]);
let oc_offset = oc * plane_out;
output[oc_offset..oc_offset + plane_out].fill(bias_v);
for ic_local in 0..in_per_group {
let ic = g * in_per_group + ic_local;
let w_oc_ic = (oc * in_per_group + ic_local) * kernel_plane;
let in_plane = ic * plane_in;
for oh in 0..h_out {
for ow in 0..w_out {
let mut acc = 0.0f32;
for ki in 0..kh {
let pad_row = oh * stride_h + ki;
if pad_row < pad_h || pad_row >= h_in + pad_h {
continue;
}
let ih = pad_row - pad_h;
for kj in 0..kw {
let pad_col = ow * stride_w + kj;
if pad_col < pad_w || pad_col >= w_in + pad_w {
continue;
}
let iw = pad_col - pad_w;
let w = weight[w_oc_ic + ki * kw + kj];
let x = input[in_plane + ih * w_in + iw];
acc += w * x;
}
}
output[oc_offset + oh * w_out + ow] += acc;
}
}
}
}
}
(h_out, w_out)
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
pub fn attn_scores(
q_head: &[f32],
k_cache: &[f32],
scores: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
debug_assert!(q_head.len() >= head_dim);
debug_assert!(scores.len() >= seq_len);
if seq_len > 0 {
debug_assert!(k_cache.len() >= (seq_len - 1) * kv_dim + kv_h_offset + head_dim);
}
#[cfg(target_arch = "aarch64")]
if head_dim <= 128 {
unsafe {
attn_scores_neon(
q_head,
k_cache,
scores,
kv_dim,
kv_h_offset,
head_dim,
scale,
seq_len,
);
}
return;
}
#[cfg(target_arch = "x86_64")]
{
if head_dim.is_multiple_of(8)
&& head_dim <= 256
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
{
unsafe {
attn_scores_avx2(
q_head,
k_cache,
scores,
kv_dim,
kv_h_offset,
head_dim,
scale,
seq_len,
);
}
return;
}
}
for t in 0..seq_len {
let mut dot = 0.0f32;
let k_off = t * kv_dim + kv_h_offset;
for d in 0..head_dim {
dot += q_head[d] * k_cache[k_off + d];
}
scores[t] = dot * scale;
}
}
#[allow(clippy::needless_range_loop)]
pub fn attn_values(
scores: &[f32],
v_cache: &[f32],
attn_out: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
seq_len: usize,
) {
debug_assert!(scores.len() >= seq_len);
debug_assert!(attn_out.len() >= head_dim);
if seq_len > 0 {
debug_assert!(v_cache.len() >= (seq_len - 1) * kv_dim + kv_h_offset + head_dim);
}
#[cfg(target_arch = "aarch64")]
if head_dim <= 128 {
unsafe {
attn_values_neon(
scores,
v_cache,
attn_out,
kv_dim,
kv_h_offset,
head_dim,
seq_len,
);
}
return;
}
#[cfg(target_arch = "x86_64")]
{
if head_dim.is_multiple_of(8)
&& head_dim <= 256
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
{
unsafe {
attn_values_avx2(
scores,
v_cache,
attn_out,
kv_dim,
kv_h_offset,
head_dim,
seq_len,
);
}
return;
}
}
attn_out[..head_dim].fill(0.0);
for t in 0..seq_len {
let s = scores[t];
let v_base = t * kv_dim + kv_h_offset;
for d in 0..head_dim {
attn_out[d] += s * v_cache[v_base + d];
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[target_feature(enable = "neon")]
unsafe fn attn_scores_neon(
q_head: &[f32],
k_cache: &[f32],
scores: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
use std::arch::aarch64::*;
unsafe {
let q_ptr = q_head.as_ptr();
let k_ptr = k_cache.as_ptr();
const MAX_Q_VECS: usize = 32;
let n_q_vecs = head_dim / 4;
debug_assert!(n_q_vecs <= MAX_Q_VECS, "head_dim > 128 not supported");
let mut q_vecs = [vdupq_n_f32(0.0); MAX_Q_VECS];
for i in 0..n_q_vecs {
q_vecs[i] = vld1q_f32(q_ptr.add(i * 4));
}
for t in 0..seq_len {
let k_off = t * kv_dim + kv_h_offset;
let mut sum0 = vdupq_n_f32(0.0);
let mut sum1 = vdupq_n_f32(0.0);
let mut d = 0usize;
let mut qi = 0usize;
while d + 8 <= head_dim {
let k0 = vld1q_f32(k_ptr.add(k_off + d));
let k1 = vld1q_f32(k_ptr.add(k_off + d + 4));
sum0 = vfmaq_f32(sum0, q_vecs[qi], k0);
sum1 = vfmaq_f32(sum1, q_vecs[qi + 1], k1);
d += 8;
qi += 2;
}
if d + 4 <= head_dim {
let k0 = vld1q_f32(k_ptr.add(k_off + d));
sum0 = vfmaq_f32(sum0, q_vecs[qi], k0);
d += 4;
}
let mut total = vaddvq_f32(vaddq_f32(sum0, sum1));
while d < head_dim {
total += *q_ptr.add(d) * *k_ptr.add(k_off + d);
d += 1;
}
scores[t] = total * scale;
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::needless_range_loop)]
#[target_feature(enable = "neon")]
unsafe fn attn_values_neon(
scores: &[f32],
v_cache: &[f32],
attn_out: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
seq_len: usize,
) {
use std::arch::aarch64::*;
unsafe {
let v_ptr = v_cache.as_ptr();
let out_ptr = attn_out.as_mut_ptr();
const MAX_ACC_VECS: usize = 32;
let n_vec = head_dim / 4;
let n_tail = head_dim % 4;
debug_assert!(n_vec <= MAX_ACC_VECS, "head_dim > 128 not supported");
let mut acc = [vdupq_n_f32(0.0); MAX_ACC_VECS];
for t in 0..seq_len {
let s = vdupq_n_f32(scores[t]);
let v_base = t * kv_dim + kv_h_offset;
for i in 0..n_vec {
let v = vld1q_f32(v_ptr.add(v_base + i * 4));
acc[i] = vfmaq_f32(acc[i], s, v);
}
}
for i in 0..n_vec {
vst1q_f32(out_ptr.add(i * 4), acc[i]);
}
let tail_start = n_vec * 4;
for dd in 0..n_tail {
let mut val = 0.0f32;
for t in 0..seq_len {
val += scores[t] * *v_ptr.add(t * kv_dim + kv_h_offset + tail_start + dd);
}
*out_ptr.add(tail_start + dd) = val;
}
}
}
#[cfg(target_arch = "x86_64")]
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn hsum256(v: std::arch::x86_64::__m256) -> f32 {
use std::arch::x86_64::*;
let hi = _mm256_extractf128_ps(v, 1);
let lo = _mm256_castps256_ps128(v);
let s128 = _mm_add_ps(lo, hi);
let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
let s32 = _mm_add_ss(s64, _mm_shuffle_ps(s64, s64, 1));
_mm_cvtss_f32(s32)
}
#[cfg(target_arch = "x86_64")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[target_feature(enable = "avx2,fma")]
unsafe fn attn_scores_avx2(
q_head: &[f32],
k_cache: &[f32],
scores: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
use std::arch::x86_64::*;
unsafe {
let q_ptr = q_head.as_ptr();
let k_ptr = k_cache.as_ptr();
const MAX_Q_VECS: usize = 32;
let n_q_vecs = head_dim / 8;
debug_assert!(n_q_vecs <= MAX_Q_VECS, "head_dim > 256 not supported");
let mut q_vecs = [_mm256_setzero_ps(); MAX_Q_VECS];
for i in 0..n_q_vecs {
q_vecs[i] = _mm256_loadu_ps(q_ptr.add(i * 8));
}
for t in 0..seq_len {
let k_off = t * kv_dim + kv_h_offset;
let mut sum0 = _mm256_setzero_ps();
let mut sum1 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 2 <= n_q_vecs {
let k0 = _mm256_loadu_ps(k_ptr.add(k_off + i * 8));
let k1 = _mm256_loadu_ps(k_ptr.add(k_off + (i + 1) * 8));
sum0 = _mm256_fmadd_ps(q_vecs[i], k0, sum0);
sum1 = _mm256_fmadd_ps(q_vecs[i + 1], k1, sum1);
i += 2;
}
if i < n_q_vecs {
let k0 = _mm256_loadu_ps(k_ptr.add(k_off + i * 8));
sum0 = _mm256_fmadd_ps(q_vecs[i], k0, sum0);
}
scores[t] = hsum256(_mm256_add_ps(sum0, sum1)) * scale;
}
}
}
#[cfg(target_arch = "x86_64")]
#[allow(clippy::needless_range_loop)]
#[target_feature(enable = "avx2,fma")]
unsafe fn attn_values_avx2(
scores: &[f32],
v_cache: &[f32],
attn_out: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
seq_len: usize,
) {
use std::arch::x86_64::*;
unsafe {
let v_ptr = v_cache.as_ptr();
let out_ptr = attn_out.as_mut_ptr();
const MAX_ACC_VECS: usize = 32;
let n_vec = head_dim / 8;
debug_assert!(n_vec <= MAX_ACC_VECS, "head_dim > 256 not supported");
let mut acc = [_mm256_setzero_ps(); MAX_ACC_VECS];
for t in 0..seq_len {
let s = _mm256_set1_ps(scores[t]);
let v_base = t * kv_dim + kv_h_offset;
for i in 0..n_vec {
let v = _mm256_loadu_ps(v_ptr.add(v_base + i * 8));
acc[i] = _mm256_fmadd_ps(s, v, acc[i]);
}
}
for i in 0..n_vec {
_mm256_storeu_ps(out_ptr.add(i * 8), acc[i]);
}
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
pub fn attn_scores_f16(
q_head: &[f32],
k_cache: &[u16],
scores: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
debug_assert!(q_head.len() >= head_dim);
debug_assert!(scores.len() >= seq_len);
if seq_len > 0 {
debug_assert!(k_cache.len() >= (seq_len - 1) * kv_dim + kv_h_offset + head_dim);
}
#[cfg(target_arch = "aarch64")]
if head_dim <= 128 && super::cpu_features::cpu_features().fp16 {
unsafe {
attn_scores_f16_neon(
q_head,
k_cache,
scores,
kv_dim,
kv_h_offset,
head_dim,
scale,
seq_len,
);
}
return;
}
#[cfg(target_arch = "x86_64")]
{
if head_dim.is_multiple_of(8)
&& head_dim <= 256
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
&& is_x86_feature_detected!("f16c")
{
unsafe {
attn_scores_f16_avx2(
q_head,
k_cache,
scores,
kv_dim,
kv_h_offset,
head_dim,
scale,
seq_len,
);
}
return;
}
}
for t in 0..seq_len {
let mut dot = 0.0f32;
let k_off = t * kv_dim + kv_h_offset;
for d in 0..head_dim {
dot += q_head[d] * f16_to_f32(k_cache[k_off + d]);
}
scores[t] = dot * scale;
}
}
#[allow(clippy::needless_range_loop)]
pub fn attn_values_f16(
scores: &[f32],
v_cache: &[u16],
attn_out: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
seq_len: usize,
) {
debug_assert!(scores.len() >= seq_len);
debug_assert!(attn_out.len() >= head_dim);
if seq_len > 0 {
debug_assert!(v_cache.len() >= (seq_len - 1) * kv_dim + kv_h_offset + head_dim);
}
#[cfg(target_arch = "aarch64")]
if head_dim <= 128 && super::cpu_features::cpu_features().fp16 {
unsafe {
attn_values_f16_neon(
scores,
v_cache,
attn_out,
kv_dim,
kv_h_offset,
head_dim,
seq_len,
);
}
return;
}
#[cfg(target_arch = "x86_64")]
{
if head_dim.is_multiple_of(8)
&& head_dim <= 256
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
&& is_x86_feature_detected!("f16c")
{
unsafe {
attn_values_f16_avx2(
scores,
v_cache,
attn_out,
kv_dim,
kv_h_offset,
head_dim,
seq_len,
);
}
return;
}
}
attn_out[..head_dim].fill(0.0);
for t in 0..seq_len {
let s = scores[t];
let v_base = t * kv_dim + kv_h_offset;
for d in 0..head_dim {
attn_out[d] += s * f16_to_f32(v_cache[v_base + d]);
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[target_feature(enable = "neon,fp16")]
unsafe fn attn_scores_f16_neon(
q_head: &[f32],
k_cache: &[u16],
scores: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
use std::arch::aarch64::*;
unsafe {
let q_ptr = q_head.as_ptr();
let k_ptr = k_cache.as_ptr();
const MAX_Q_VECS: usize = 32;
let n_q_vecs = head_dim / 4;
debug_assert!(n_q_vecs <= MAX_Q_VECS, "head_dim > 128 not supported");
let mut q_vecs = [vdupq_n_f32(0.0); MAX_Q_VECS];
for i in 0..n_q_vecs {
q_vecs[i] = vld1q_f32(q_ptr.add(i * 4));
}
for t in 0..seq_len {
let k_off = t * kv_dim + kv_h_offset;
let mut sum0 = vdupq_n_f32(0.0);
let mut sum1 = vdupq_n_f32(0.0);
let mut d = 0usize;
let mut qi = 0usize;
while d + 8 <= head_dim {
let k0 = vcvt_f32_f16(vreinterpret_f16_u16(vld1_u16(k_ptr.add(k_off + d))));
let k1 = vcvt_f32_f16(vreinterpret_f16_u16(vld1_u16(k_ptr.add(k_off + d + 4))));
sum0 = vfmaq_f32(sum0, q_vecs[qi], k0);
sum1 = vfmaq_f32(sum1, q_vecs[qi + 1], k1);
d += 8;
qi += 2;
}
if d + 4 <= head_dim {
let k0 = vcvt_f32_f16(vreinterpret_f16_u16(vld1_u16(k_ptr.add(k_off + d))));
sum0 = vfmaq_f32(sum0, q_vecs[qi], k0);
d += 4;
}
let mut total = vaddvq_f32(vaddq_f32(sum0, sum1));
while d < head_dim {
total += *q_ptr.add(d) * f16_to_f32(*k_ptr.add(k_off + d));
d += 1;
}
scores[t] = total * scale;
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::needless_range_loop)]
#[target_feature(enable = "neon,fp16")]
unsafe fn attn_values_f16_neon(
scores: &[f32],
v_cache: &[u16],
attn_out: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
seq_len: usize,
) {
use std::arch::aarch64::*;
unsafe {
let v_ptr = v_cache.as_ptr();
let out_ptr = attn_out.as_mut_ptr();
const MAX_ACC_VECS: usize = 32;
let n_vec = head_dim / 4;
let n_tail = head_dim % 4;
debug_assert!(n_vec <= MAX_ACC_VECS, "head_dim > 128 not supported");
let mut acc = [vdupq_n_f32(0.0); MAX_ACC_VECS];
for t in 0..seq_len {
let s = vdupq_n_f32(scores[t]);
let v_base = t * kv_dim + kv_h_offset;
for i in 0..n_vec {
let v = vcvt_f32_f16(vreinterpret_f16_u16(vld1_u16(v_ptr.add(v_base + i * 4))));
acc[i] = vfmaq_f32(acc[i], s, v);
}
}
for i in 0..n_vec {
vst1q_f32(out_ptr.add(i * 4), acc[i]);
}
let tail_start = n_vec * 4;
for dd in 0..n_tail {
let mut val = 0.0f32;
for t in 0..seq_len {
val +=
scores[t] * f16_to_f32(*v_ptr.add(t * kv_dim + kv_h_offset + tail_start + dd));
}
*out_ptr.add(tail_start + dd) = val;
}
}
}
#[cfg(target_arch = "x86_64")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn attn_scores_f16_avx2(
q_head: &[f32],
k_cache: &[u16],
scores: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
use std::arch::x86_64::*;
unsafe {
let q_ptr = q_head.as_ptr();
let k_ptr = k_cache.as_ptr();
const MAX_Q_VECS: usize = 32;
let n_q_vecs = head_dim / 8;
debug_assert!(n_q_vecs <= MAX_Q_VECS, "head_dim > 256 not supported");
let mut q_vecs = [_mm256_setzero_ps(); MAX_Q_VECS];
for i in 0..n_q_vecs {
q_vecs[i] = _mm256_loadu_ps(q_ptr.add(i * 8));
}
for t in 0..seq_len {
let k_off = t * kv_dim + kv_h_offset;
let mut sum0 = _mm256_setzero_ps();
let mut sum1 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 2 <= n_q_vecs {
let k0 = _mm256_cvtph_ps(_mm_loadu_si128(k_ptr.add(k_off + i * 8).cast()));
let k1 = _mm256_cvtph_ps(_mm_loadu_si128(k_ptr.add(k_off + (i + 1) * 8).cast()));
sum0 = _mm256_fmadd_ps(q_vecs[i], k0, sum0);
sum1 = _mm256_fmadd_ps(q_vecs[i + 1], k1, sum1);
i += 2;
}
if i < n_q_vecs {
let k0 = _mm256_cvtph_ps(_mm_loadu_si128(k_ptr.add(k_off + i * 8).cast()));
sum0 = _mm256_fmadd_ps(q_vecs[i], k0, sum0);
}
scores[t] = hsum256(_mm256_add_ps(sum0, sum1)) * scale;
}
}
}
#[cfg(target_arch = "x86_64")]
#[allow(clippy::needless_range_loop)]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn attn_values_f16_avx2(
scores: &[f32],
v_cache: &[u16],
attn_out: &mut [f32],
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
seq_len: usize,
) {
use std::arch::x86_64::*;
unsafe {
let v_ptr = v_cache.as_ptr();
let out_ptr = attn_out.as_mut_ptr();
const MAX_ACC_VECS: usize = 32;
let n_vec = head_dim / 8;
debug_assert!(n_vec <= MAX_ACC_VECS, "head_dim > 256 not supported");
let mut acc = [_mm256_setzero_ps(); MAX_ACC_VECS];
for t in 0..seq_len {
let s = _mm256_set1_ps(scores[t]);
let v_base = t * kv_dim + kv_h_offset;
for i in 0..n_vec {
let v = _mm256_cvtph_ps(_mm_loadu_si128(v_ptr.add(v_base + i * 8).cast()));
acc[i] = _mm256_fmadd_ps(s, v, acc[i]);
}
}
for i in 0..n_vec {
_mm256_storeu_ps(out_ptr.add(i * 8), acc[i]);
}
}
}
const FLASH_TILE_KV: usize = 32;
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
pub fn flash_attention_gqa_cpu(
q_mat: &[f32],
k_cache: &[f32],
v_cache: &[f32],
out: &mut [f32],
n_heads_start: usize,
group_size: usize,
n_queries: usize,
q_stride: usize,
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
start_pos: usize,
) {
#[cfg(target_arch = "aarch64")]
{
if head_dim.is_multiple_of(4) && head_dim <= 128 {
unsafe {
flash_attention_gqa_neon(
q_mat,
k_cache,
v_cache,
out,
n_heads_start,
group_size,
n_queries,
q_stride,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
}
return;
}
}
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
{
if head_dim.is_multiple_of(16) && head_dim <= 256 && is_x86_feature_detected!("avx512f") {
unsafe {
flash_attention_gqa_avx512(
q_mat,
k_cache,
v_cache,
out,
n_heads_start,
group_size,
n_queries,
q_stride,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
}
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if head_dim.is_multiple_of(8)
&& head_dim <= 256
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
{
unsafe {
flash_attention_gqa_avx2(
q_mat,
k_cache,
v_cache,
out,
n_heads_start,
group_size,
n_queries,
q_stride,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
}
return;
}
}
flash_attention_gqa_scalar(
q_mat,
k_cache,
v_cache,
out,
n_heads_start,
group_size,
n_queries,
q_stride,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
}
#[allow(dead_code, clippy::too_many_arguments, clippy::needless_range_loop)]
fn flash_attention_gqa_scalar(
q_mat: &[f32],
k_cache: &[f32],
v_cache: &[f32],
out: &mut [f32],
n_heads_start: usize,
group_size: usize,
n_queries: usize,
q_stride: usize,
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
start_pos: usize,
) {
assert!(
head_dim <= 256,
"flash_attention_gqa_scalar: head_dim {head_dim} > 256"
);
let mut q_buf = [0.0f32; 256];
let mut acc_buf = [0.0f32; 256];
let q_local = &mut q_buf[..head_dim];
let acc = &mut acc_buf[..head_dim];
let mut tile_scores = [0.0f32; FLASH_TILE_KV];
for g in 0..group_size {
let h = n_heads_start + g;
let h_off = h * head_dim;
for j in 0..n_queries {
let max_kv = start_pos + j + 1;
for d in 0..head_dim {
q_local[d] = q_mat[(h_off + d) * q_stride + j];
}
let mut running_max = f32::NEG_INFINITY;
let mut running_sum = 0.0f64;
acc.fill(0.0);
for kv_start in (0..max_kv).step_by(FLASH_TILE_KV) {
let kv_end = (kv_start + FLASH_TILE_KV).min(max_kv);
let tile_len = kv_end - kv_start;
for ti in 0..tile_len {
let k_off = (kv_start + ti) * kv_dim + kv_h_offset;
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q_local[d] * k_cache[k_off + d];
}
tile_scores[ti] = dot * scale;
}
let tile_max = tile_scores[..tile_len]
.iter()
.fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let new_max = running_max.max(tile_max);
let rescale = if running_max > f32::NEG_INFINITY {
ggml_expf(running_max - new_max)
} else {
0.0
};
let mut tile_sum = 0.0f64;
for ti in 0..tile_len {
tile_scores[ti] = ggml_expf(tile_scores[ti] - new_max);
tile_sum += tile_scores[ti] as f64;
}
for d in 0..head_dim {
acc[d] *= rescale;
}
for ti in 0..tile_len {
let s = tile_scores[ti];
let v_off = (kv_start + ti) * kv_dim + kv_h_offset;
for d in 0..head_dim {
acc[d] += s * v_cache[v_off + d];
}
}
running_sum = running_sum * rescale as f64 + tile_sum;
running_max = new_max;
}
let inv_sum = (1.0 / running_sum) as f32;
let out_off = (g * n_queries + j) * head_dim;
for d in 0..head_dim {
out[out_off + d] = acc[d] * inv_sum;
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
unsafe fn flash_attention_gqa_avx2(
q_mat: &[f32],
k_cache: &[f32],
v_cache: &[f32],
out: &mut [f32],
n_heads_start: usize,
group_size: usize,
n_queries: usize,
q_stride: usize,
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
start_pos: usize,
) {
use std::arch::x86_64::*;
unsafe {
debug_assert!(
q_mat.len() >= ((n_heads_start + group_size) * head_dim - 1) * q_stride + n_queries,
"q_mat too small for the given head range and q_stride"
);
debug_assert!(
(start_pos + n_queries == 0)
|| k_cache.len() >= (start_pos + n_queries - 1) * kv_dim + kv_h_offset + head_dim,
"k_cache too small"
);
debug_assert!(
(start_pos + n_queries == 0)
|| v_cache.len() >= (start_pos + n_queries - 1) * kv_dim + kv_h_offset + head_dim,
"v_cache too small"
);
debug_assert!(
out.len() >= group_size * n_queries * head_dim,
"out buffer too small for contiguous [group_size, n_queries, head_dim] output"
);
const MAX_VECS: usize = 32;
let n_vecs = head_dim / 8;
let q_ptr = q_mat.as_ptr();
let k_ptr = k_cache.as_ptr();
let v_ptr = v_cache.as_ptr();
let out_ptr = out.as_mut_ptr();
let mut q_vecs = [_mm256_setzero_ps(); MAX_VECS];
let mut acc_vecs = [_mm256_setzero_ps(); MAX_VECS];
let mut q_gather = [0.0f32; 256];
let mut tile_scores = [0.0f32; FLASH_TILE_KV];
for g in 0..group_size {
let h = n_heads_start + g;
let h_off = h * head_dim;
for j in 0..n_queries {
let max_kv = start_pos + j + 1;
for d in 0..head_dim {
q_gather[d] = *q_ptr.add((h_off + d) * q_stride + j);
}
for i in 0..n_vecs {
q_vecs[i] = _mm256_loadu_ps(q_gather.as_ptr().add(i * 8));
acc_vecs[i] = _mm256_setzero_ps();
}
let mut running_max = f32::NEG_INFINITY;
let mut running_sum = 0.0f64;
for kv_start in (0..max_kv).step_by(FLASH_TILE_KV) {
let kv_end = (kv_start + FLASH_TILE_KV).min(max_kv);
let tile_len = kv_end - kv_start;
for ti in 0..tile_len {
let k_off = (kv_start + ti) * kv_dim + kv_h_offset;
let mut s0 = _mm256_setzero_ps();
let mut s1 = _mm256_setzero_ps();
let mut i = 0;
while i + 2 <= n_vecs {
let k0 = _mm256_loadu_ps(k_ptr.add(k_off + i * 8));
let k1 = _mm256_loadu_ps(k_ptr.add(k_off + i * 8 + 8));
s0 = _mm256_fmadd_ps(q_vecs[i], k0, s0);
s1 = _mm256_fmadd_ps(q_vecs[i + 1], k1, s1);
i += 2;
}
if i < n_vecs {
let k0 = _mm256_loadu_ps(k_ptr.add(k_off + i * 8));
s0 = _mm256_fmadd_ps(q_vecs[i], k0, s0);
}
tile_scores[ti] = hsum256(_mm256_add_ps(s0, s1)) * scale;
}
let mut tile_max = f32::NEG_INFINITY;
for ti in 0..tile_len {
if tile_scores[ti] > tile_max {
tile_max = tile_scores[ti];
}
}
let new_max = running_max.max(tile_max);
let rescale = if running_max > f32::NEG_INFINITY {
ggml_expf(running_max - new_max)
} else {
0.0
};
let mut tile_sum = 0.0f64;
for ti in 0..tile_len {
tile_scores[ti] = ggml_expf(tile_scores[ti] - new_max);
tile_sum += tile_scores[ti] as f64;
}
let rescale_v = _mm256_set1_ps(rescale);
for i in 0..n_vecs {
acc_vecs[i] = _mm256_mul_ps(acc_vecs[i], rescale_v);
}
for ti in 0..tile_len {
let s = _mm256_set1_ps(tile_scores[ti]);
let v_base = (kv_start + ti) * kv_dim + kv_h_offset;
for i in 0..n_vecs {
let v = _mm256_loadu_ps(v_ptr.add(v_base + i * 8));
acc_vecs[i] = _mm256_fmadd_ps(s, v, acc_vecs[i]);
}
}
running_sum = running_sum * rescale as f64 + tile_sum;
running_max = new_max;
}
let inv_sum = (1.0 / running_sum) as f32;
let inv_sum_v = _mm256_set1_ps(inv_sum);
let out_off = (g * n_queries + j) * head_dim;
for i in 0..n_vecs {
let r = _mm256_mul_ps(acc_vecs[i], inv_sum_v);
_mm256_storeu_ps(out_ptr.add(out_off + i * 8), r);
}
}
}
}
}
#[cfg(all(target_arch = "x86_64", feature = "avx512"))]
#[target_feature(enable = "avx512f")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
unsafe fn flash_attention_gqa_avx512(
q_mat: &[f32],
k_cache: &[f32],
v_cache: &[f32],
out: &mut [f32],
n_heads_start: usize,
group_size: usize,
n_queries: usize,
q_stride: usize,
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
start_pos: usize,
) {
use std::arch::x86_64::*;
unsafe {
debug_assert!(
q_mat.len() >= ((n_heads_start + group_size) * head_dim - 1) * q_stride + n_queries,
"q_mat too small for the given head range and q_stride"
);
debug_assert!(
(start_pos + n_queries == 0)
|| k_cache.len() >= (start_pos + n_queries - 1) * kv_dim + kv_h_offset + head_dim,
"k_cache too small"
);
debug_assert!(
(start_pos + n_queries == 0)
|| v_cache.len() >= (start_pos + n_queries - 1) * kv_dim + kv_h_offset + head_dim,
"v_cache too small"
);
debug_assert!(
out.len() >= group_size * n_queries * head_dim,
"out buffer too small for contiguous [group_size, n_queries, head_dim] output"
);
const MAX_VECS: usize = 16;
let n_vecs = head_dim / 16;
let q_ptr = q_mat.as_ptr();
let k_ptr = k_cache.as_ptr();
let v_ptr = v_cache.as_ptr();
let out_ptr = out.as_mut_ptr();
let mut q_vecs = [_mm512_setzero_ps(); MAX_VECS];
let mut acc_vecs = [_mm512_setzero_ps(); MAX_VECS];
let mut q_gather = [0.0f32; 256];
let mut tile_scores = [0.0f32; FLASH_TILE_KV];
for g in 0..group_size {
let h = n_heads_start + g;
let h_off = h * head_dim;
for j in 0..n_queries {
let max_kv = start_pos + j + 1;
for d in 0..head_dim {
q_gather[d] = *q_ptr.add((h_off + d) * q_stride + j);
}
for i in 0..n_vecs {
q_vecs[i] = _mm512_loadu_ps(q_gather.as_ptr().add(i * 16));
acc_vecs[i] = _mm512_setzero_ps();
}
let mut running_max = f32::NEG_INFINITY;
let mut running_sum = 0.0f64;
for kv_start in (0..max_kv).step_by(FLASH_TILE_KV) {
let kv_end = (kv_start + FLASH_TILE_KV).min(max_kv);
let tile_len = kv_end - kv_start;
for ti in 0..tile_len {
let k_off = (kv_start + ti) * kv_dim + kv_h_offset;
let mut s0 = _mm512_setzero_ps();
let mut s1 = _mm512_setzero_ps();
let mut i = 0;
while i + 2 <= n_vecs {
let k0 = _mm512_loadu_ps(k_ptr.add(k_off + i * 16));
let k1 = _mm512_loadu_ps(k_ptr.add(k_off + i * 16 + 16));
s0 = _mm512_fmadd_ps(q_vecs[i], k0, s0);
s1 = _mm512_fmadd_ps(q_vecs[i + 1], k1, s1);
i += 2;
}
if i < n_vecs {
let k0 = _mm512_loadu_ps(k_ptr.add(k_off + i * 16));
s0 = _mm512_fmadd_ps(q_vecs[i], k0, s0);
}
tile_scores[ti] = _mm512_reduce_add_ps(_mm512_add_ps(s0, s1)) * scale;
}
let mut tile_max = f32::NEG_INFINITY;
for ti in 0..tile_len {
if tile_scores[ti] > tile_max {
tile_max = tile_scores[ti];
}
}
let new_max = running_max.max(tile_max);
let rescale = if running_max > f32::NEG_INFINITY {
ggml_expf(running_max - new_max)
} else {
0.0
};
let mut tile_sum = 0.0f64;
for ti in 0..tile_len {
tile_scores[ti] = ggml_expf(tile_scores[ti] - new_max);
tile_sum += tile_scores[ti] as f64;
}
let rescale_v = _mm512_set1_ps(rescale);
for i in 0..n_vecs {
acc_vecs[i] = _mm512_mul_ps(acc_vecs[i], rescale_v);
}
for ti in 0..tile_len {
let s = _mm512_set1_ps(tile_scores[ti]);
let v_base = (kv_start + ti) * kv_dim + kv_h_offset;
for i in 0..n_vecs {
let v = _mm512_loadu_ps(v_ptr.add(v_base + i * 16));
acc_vecs[i] = _mm512_fmadd_ps(s, v, acc_vecs[i]);
}
}
running_sum = running_sum * rescale as f64 + tile_sum;
running_max = new_max;
}
let inv_sum = (1.0 / running_sum) as f32;
let inv_sum_v = _mm512_set1_ps(inv_sum);
let out_off = (g * n_queries + j) * head_dim;
for i in 0..n_vecs {
let r = _mm512_mul_ps(acc_vecs[i], inv_sum_v);
_mm512_storeu_ps(out_ptr.add(out_off + i * 16), r);
}
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
unsafe fn flash_attention_gqa_neon(
q_mat: &[f32],
k_cache: &[f32],
v_cache: &[f32],
out: &mut [f32],
n_heads_start: usize,
group_size: usize,
n_queries: usize,
q_stride: usize,
kv_dim: usize,
kv_h_offset: usize,
head_dim: usize,
scale: f32,
start_pos: usize,
) {
use std::arch::aarch64::*;
unsafe {
debug_assert!(
(n_queries == 0)
|| (q_stride == n_queries
&& q_mat.len()
>= ((n_heads_start + group_size) * head_dim - 1) * q_stride + n_queries)
|| (q_stride >= head_dim
&& q_mat.len()
>= (n_queries - 1) * q_stride + (n_heads_start + group_size) * head_dim),
"q_mat too small for the given head range and q_stride"
);
debug_assert!(
(start_pos + n_queries == 0)
|| k_cache.len() >= (start_pos + n_queries - 1) * kv_dim + kv_h_offset + head_dim,
"k_cache too small"
);
debug_assert!(
(start_pos + n_queries == 0)
|| v_cache.len() >= (start_pos + n_queries - 1) * kv_dim + kv_h_offset + head_dim,
"v_cache too small"
);
debug_assert!(
out.len() >= group_size * n_queries * head_dim,
"out buffer too small for contiguous [group_size, n_queries, head_dim] output"
);
let q_ptr = q_mat.as_ptr();
let k_ptr = k_cache.as_ptr();
let v_ptr = v_cache.as_ptr();
let out_ptr = out.as_mut_ptr();
let n_vecs = head_dim / 4;
debug_assert!(
head_dim.is_multiple_of(4) && n_vecs <= 32,
"head_dim must be a multiple of 4 and <= 128"
);
const MAX_VECS: usize = 32;
let mut q_vecs = [vdupq_n_f32(0.0); MAX_VECS];
let mut acc_vecs = [vdupq_n_f32(0.0); MAX_VECS];
let mut tile_scores = [0.0f32; FLASH_TILE_KV];
for g in 0..group_size {
let h = n_heads_start + g;
let h_off = h * head_dim;
for j in 0..n_queries {
let max_kv = start_pos + j + 1;
if q_stride == n_queries {
for i in 0..n_vecs {
let d = i * 4;
let q = [
*q_ptr.add((h_off + d) * q_stride + j),
*q_ptr.add((h_off + d + 1) * q_stride + j),
*q_ptr.add((h_off + d + 2) * q_stride + j),
*q_ptr.add((h_off + d + 3) * q_stride + j),
];
q_vecs[i] = vld1q_f32(q.as_ptr());
}
} else {
for i in 0..n_vecs {
q_vecs[i] = vld1q_f32(q_ptr.add(j * q_stride + h_off + i * 4));
}
}
let mut running_max = f32::NEG_INFINITY;
let mut running_sum = 0.0f64;
for i in 0..n_vecs {
acc_vecs[i] = vdupq_n_f32(0.0);
}
for kv_start in (0..max_kv).step_by(FLASH_TILE_KV) {
let kv_end = (kv_start + FLASH_TILE_KV).min(max_kv);
let tile_len = kv_end - kv_start;
for ti in 0..tile_len {
let k_off = (kv_start + ti) * kv_dim + kv_h_offset;
let mut sum0 = vdupq_n_f32(0.0);
let mut sum1 = vdupq_n_f32(0.0);
let mut i = 0;
while i + 2 <= n_vecs {
let k0 = vld1q_f32(k_ptr.add(k_off + i * 4));
let k1 = vld1q_f32(k_ptr.add(k_off + i * 4 + 4));
sum0 = vfmaq_f32(sum0, q_vecs[i], k0);
sum1 = vfmaq_f32(sum1, q_vecs[i + 1], k1);
i += 2;
}
if i < n_vecs {
let k0 = vld1q_f32(k_ptr.add(k_off + i * 4));
sum0 = vfmaq_f32(sum0, q_vecs[i], k0);
}
tile_scores[ti] = vaddvq_f32(vaddq_f32(sum0, sum1)) * scale;
}
let mut tile_max = f32::NEG_INFINITY;
for ti in 0..tile_len {
if tile_scores[ti] > tile_max {
tile_max = tile_scores[ti];
}
}
let new_max = running_max.max(tile_max);
let rescale = if running_max > f32::NEG_INFINITY {
ggml_expf(running_max - new_max)
} else {
0.0
};
let mut tile_sum = 0.0f64;
for ti in 0..tile_len {
tile_scores[ti] = ggml_expf(tile_scores[ti] - new_max);
tile_sum += tile_scores[ti] as f64;
}
let rescale_v = vdupq_n_f32(rescale);
for i in 0..n_vecs {
acc_vecs[i] = vmulq_f32(acc_vecs[i], rescale_v);
}
for ti in 0..tile_len {
let s = vdupq_n_f32(tile_scores[ti]);
let v_base = (kv_start + ti) * kv_dim + kv_h_offset;
for i in 0..n_vecs {
let v = vld1q_f32(v_ptr.add(v_base + i * 4));
acc_vecs[i] = vfmaq_f32(acc_vecs[i], s, v);
}
}
running_sum = running_sum * rescale as f64 + tile_sum;
running_max = new_max;
}
let inv_sum = (1.0 / running_sum) as f32;
let inv_sum_v = vdupq_n_f32(inv_sum);
let out_off = (g * n_queries + j) * head_dim;
for i in 0..n_vecs {
let result = vmulq_f32(acc_vecs[i], inv_sum_v);
vst1q_f32(out_ptr.add(out_off + i * 4), result);
}
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn attn_scores_turboquant_neon(
q_rot_all: &[f32], q_jl_all: &[f32], polar_data: &[u8], jl_data: &[u8], norms_f32: &[f32], residual_norms_f32: &[f32],
q_jl_total_sums: &[f32], group_start: usize, group_size: usize, scores_flat: &mut [f32], head_dim: usize,
centroids: &[f32; 4],
scale: f32,
qjl_scale: f32,
seq_len: usize,
) {
use std::arch::aarch64::*;
unsafe {
let polar_bytes = head_dim / 4;
let jl_bytes = head_dim / 8;
let c_arr = *centroids;
const MAX_VECS: usize = 32;
let n_cent_vecs = polar_bytes; let n_mask_vecs = jl_bytes * 2; debug_assert!(n_cent_vecs <= MAX_VECS);
debug_assert!(n_mask_vecs <= MAX_VECS);
let mut cent_vecs = [vdupq_n_f32(0.0); MAX_VECS];
let mut mask_vecs = [vdupq_n_f32(0.0); MAX_VECS];
for t in 0..seq_len {
let p_base = t * polar_bytes;
let j_base = t * jl_bytes;
let norm = norms_f32[t];
let residual_norm = residual_norms_f32[t];
for (i, cv) in cent_vecs.iter_mut().enumerate().take(n_cent_vecs) {
let b = *polar_data.get_unchecked(p_base + i);
*cv = select_centroids_4(b, &c_arr);
}
for i in 0..jl_bytes {
let b = *jl_data.get_unchecked(j_base + i) as u32;
mask_vecs[i * 2] = bits_to_f32_mask_lo(b);
mask_vecs[i * 2 + 1] = bits_to_f32_mask_hi(b);
}
for g in 0..group_size {
let h = group_start + g;
let q_rot = &q_rot_all[h * head_dim..];
let q_jl = &q_jl_all[h * head_dim..];
let mut dot_acc0 = vdupq_n_f32(0.0);
let mut dot_acc1 = vdupq_n_f32(0.0);
let mut ci = 0usize;
let mut q_off = 0usize;
while ci + 4 <= n_cent_vecs {
let qv0 = vld1q_f32(q_rot.as_ptr().add(q_off));
let qv1 = vld1q_f32(q_rot.as_ptr().add(q_off + 4));
let qv2 = vld1q_f32(q_rot.as_ptr().add(q_off + 8));
let qv3 = vld1q_f32(q_rot.as_ptr().add(q_off + 12));
dot_acc0 = vfmaq_f32(dot_acc0, qv0, cent_vecs[ci]);
dot_acc1 = vfmaq_f32(dot_acc1, qv1, cent_vecs[ci + 1]);
dot_acc0 = vfmaq_f32(dot_acc0, qv2, cent_vecs[ci + 2]);
dot_acc1 = vfmaq_f32(dot_acc1, qv3, cent_vecs[ci + 3]);
ci += 4;
q_off += 16;
}
while ci < n_cent_vecs {
let qv = vld1q_f32(q_rot.as_ptr().add(q_off));
dot_acc0 = vfmaq_f32(dot_acc0, qv, cent_vecs[ci]);
ci += 1;
q_off += 4;
}
let polar_dot = vaddvq_f32(vaddq_f32(dot_acc0, dot_acc1)) * norm;
let total_sum = *q_jl_total_sums.get_unchecked(h);
let mut pos_acc0 = vdupq_n_f32(0.0);
let mut pos_acc1 = vdupq_n_f32(0.0);
let mut mi = 0usize;
let mut jl_q_off = 0usize;
while mi + 4 <= n_mask_vecs {
let q0 = vld1q_f32(q_jl.as_ptr().add(jl_q_off));
let q1 = vld1q_f32(q_jl.as_ptr().add(jl_q_off + 4));
let q2 = vld1q_f32(q_jl.as_ptr().add(jl_q_off + 8));
let q3 = vld1q_f32(q_jl.as_ptr().add(jl_q_off + 12));
pos_acc0 = vfmaq_f32(pos_acc0, q0, mask_vecs[mi]);
pos_acc1 = vfmaq_f32(pos_acc1, q1, mask_vecs[mi + 1]);
pos_acc0 = vfmaq_f32(pos_acc0, q2, mask_vecs[mi + 2]);
pos_acc1 = vfmaq_f32(pos_acc1, q3, mask_vecs[mi + 3]);
mi += 4;
jl_q_off += 16;
}
while mi < n_mask_vecs {
let q = vld1q_f32(q_jl.as_ptr().add(jl_q_off));
pos_acc0 = vfmaq_f32(pos_acc0, q, mask_vecs[mi]);
mi += 1;
jl_q_off += 4;
}
let pos_sum = vaddvq_f32(vaddq_f32(pos_acc0, pos_acc1));
let signed_sum = 2.0 * pos_sum - total_sum;
let correction = norm * residual_norm * qjl_scale * signed_sum;
scores_flat[g * seq_len + t] = (polar_dot + correction) * scale;
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
pub unsafe fn attn_values_turboquant_neon(
polar_data: &[u8], norms_f32: &[f32], scores: &[f32], attn_out: &mut [f32], group_start: usize,
group_size: usize,
head_dim: usize,
seq_len: usize,
centroids: &[f32; 4],
) {
use std::arch::aarch64::*;
unsafe {
let polar_bytes = head_dim / 4;
debug_assert!(
polar_bytes <= 32,
"head_dim > 128 not supported by NEON path"
);
const MAX_VECS: usize = 32;
for g in 0..group_size {
let h = group_start + g;
let head_scores = scores.as_ptr().add(g * seq_len);
let mut acc = [vdupq_n_f32(0.0); MAX_VECS];
for t in 0..seq_len {
let w = *head_scores.add(t) * *norms_f32.get_unchecked(t);
let w_vec = vdupq_n_f32(w);
let base = t * polar_bytes;
for i in 0..polar_bytes {
let b = *polar_data.get_unchecked(base + i);
let c_vec = select_centroids_4(b, centroids);
acc[i] = vfmaq_f32(acc[i], w_vec, c_vec);
}
}
let out_ptr = attn_out.as_mut_ptr().add(h * head_dim);
for i in 0..polar_bytes {
vst1q_f32(out_ptr.add(i * 4), acc[i]);
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn select_centroids_4(byte: u8, c: &[f32; 4]) -> std::arch::aarch64::float32x4_t {
use std::arch::aarch64::*;
unsafe {
let vals: [f32; 4] = [
*c.get_unchecked((byte & 0x03) as usize),
*c.get_unchecked(((byte >> 2) & 0x03) as usize),
*c.get_unchecked(((byte >> 4) & 0x03) as usize),
*c.get_unchecked(((byte >> 6) & 0x03) as usize),
];
vld1q_f32(vals.as_ptr())
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn bits_to_f32_mask_lo(byte: u32) -> std::arch::aarch64::float32x4_t {
use std::arch::aarch64::*;
unsafe {
let vals: [f32; 4] = [
(byte & 1) as f32,
((byte >> 1) & 1) as f32,
((byte >> 2) & 1) as f32,
((byte >> 3) & 1) as f32,
];
vld1q_f32(vals.as_ptr())
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn bits_to_f32_mask_hi(byte: u32) -> std::arch::aarch64::float32x4_t {
use std::arch::aarch64::*;
unsafe {
let vals: [f32; 4] = [
((byte >> 4) & 1) as f32,
((byte >> 5) & 1) as f32,
((byte >> 6) & 1) as f32,
((byte >> 7) & 1) as f32,
];
vld1q_f32(vals.as_ptr())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RopeType {
Neox,
Norm,
}
pub fn rope(
q: &mut [f32],
k: &mut [f32],
pos: usize,
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
freq_base: f32,
) {
debug_assert_eq!(q.len(), n_heads * head_dim);
debug_assert_eq!(k.len(), n_kv_heads * head_dim);
for h in 0..n_heads {
let offset = h * head_dim;
apply_rope_to_head(&mut q[offset..offset + head_dim], pos, head_dim, freq_base);
}
for h in 0..n_kv_heads {
let offset = h * head_dim;
apply_rope_to_head(&mut k[offset..offset + head_dim], pos, head_dim, freq_base);
}
}
pub fn apply_rope_to_head(head: &mut [f32], pos: usize, head_dim: usize, freq_base: f32) {
let half_dim = head_dim / 2;
let theta_scale = freq_base.powf(-2.0 / head_dim as f32);
let mut theta = pos as f32;
for i in 0..half_dim {
let (sin_t, cos_t) = theta.sin_cos();
let x0 = head[i];
let x1 = head[i + half_dim];
head[i] = x0 * cos_t - x1 * sin_t;
head[i + half_dim] = x0 * sin_t + x1 * cos_t;
theta *= theta_scale;
}
}
pub fn apply_rope_delta_to_head(head: &mut [f32], delta_pos: i32, head_dim: usize, freq_base: f32) {
let half_dim = head_dim / 2;
let theta_scale = freq_base.powf(-2.0 / head_dim as f32);
let mut theta = delta_pos as f32;
for i in 0..half_dim {
let (sin_t, cos_t) = theta.sin_cos();
let x0 = head[i];
let x1 = head[i + half_dim];
head[i] = x0 * cos_t - x1 * sin_t;
head[i + half_dim] = x0 * sin_t + x1 * cos_t;
theta *= theta_scale;
}
}
#[allow(clippy::too_many_arguments)]
pub fn rope_norm(
q: &mut [f32],
k: &mut [f32],
pos: usize,
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
freq_base: f32,
freq_factors: Option<&[f32]>,
) {
debug_assert_eq!(q.len(), n_heads * head_dim);
debug_assert_eq!(k.len(), n_kv_heads * head_dim);
for h in 0..n_heads {
let offset = h * head_dim;
apply_rope_norm_to_head(
&mut q[offset..offset + head_dim],
pos,
head_dim,
freq_base,
freq_factors,
);
}
for h in 0..n_kv_heads {
let offset = h * head_dim;
apply_rope_norm_to_head(
&mut k[offset..offset + head_dim],
pos,
head_dim,
freq_base,
freq_factors,
);
}
}
pub fn apply_rope_norm_to_head(
head: &mut [f32],
pos: usize,
head_dim: usize,
freq_base: f32,
freq_factors: Option<&[f32]>,
) {
rope_norm_pairs(head, pos as f32, head_dim, freq_base, freq_factors);
}
pub fn apply_rope_norm_delta_to_head(
head: &mut [f32],
delta_pos: i32,
head_dim: usize,
freq_base: f32,
freq_factors: Option<&[f32]>,
) {
rope_norm_pairs(head, delta_pos as f32, head_dim, freq_base, freq_factors);
}
fn rope_norm_pairs(
head: &mut [f32],
theta_start: f32,
head_dim: usize,
freq_base: f32,
freq_factors: Option<&[f32]>,
) {
let theta_scale = freq_base.powf(-2.0 / head_dim as f32);
let mut theta_base = theta_start;
for (i, pair) in head.as_chunks_mut::<2>().0.iter_mut().enumerate() {
let ff = freq_factors.map_or(1.0, |f| f[i]);
let theta = theta_base / ff;
let (sin_t, cos_t) = theta.sin_cos();
let x0 = pair[0];
let x1 = pair[1];
pair[0] = x0 * cos_t - x1 * sin_t;
pair[1] = x0 * sin_t + x1 * cos_t;
theta_base *= theta_scale;
}
}
pub fn conv1d_depthwise(
input: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
output: &mut [f32],
channels: usize,
kernel_size: usize,
seq_len: usize,
) {
debug_assert_eq!(input.len(), seq_len * channels);
debug_assert_eq!(weight.len(), channels * kernel_size);
debug_assert_eq!(output.len(), seq_len * channels);
let pad = kernel_size / 2;
for t in 0..seq_len {
for c in 0..channels {
let mut sum = if let Some(b) = bias { b[c] } else { 0.0 };
for ki in 0..kernel_size {
let input_t = t as isize + ki as isize - pad as isize;
if input_t >= 0 && (input_t as usize) < seq_len {
sum += input[input_t as usize * channels + c] * weight[c * kernel_size + ki];
}
}
output[t * channels + c] = sum;
}
}
}
pub fn add_inplace(a: &mut [f32], b: &[f32]) {
debug_assert_eq!(a.len(), b.len());
for (a, b) in a.iter_mut().zip(b.iter()) {
*a += *b;
}
}
pub fn mul_inplace(a: &mut [f32], b: &[f32]) {
debug_assert_eq!(a.len(), b.len());
for (a, b) in a.iter_mut().zip(b.iter()) {
*a *= *b;
}
}
pub fn scale_inplace(a: &mut [f32], s: f32) {
for x in a.iter_mut() {
*x *= s;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
#[test]
fn rayon_pool_width_prefers_env_override() {
assert_eq!(rayon_pool_width(Some(3), 8), 3);
assert_eq!(rayon_pool_width(None, 8), 8);
assert_eq!(rayon_pool_width(None, 0), 1);
}
fn lcg(state: &mut u64) -> u32 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(*state >> 33) as u32
}
fn weights(dtype: DType, m: usize, k: usize, st: &mut u64) -> Vec<u8> {
let bb = dtype.block_bytes();
let nb = k / dtype.block_size();
let mut data: Vec<u8> = (0..m * nb * bb).map(|_| (lcg(st) % 256) as u8).collect();
for (bi, blk) in data.chunks_mut(bb).enumerate() {
let d = half::f16::from_f32(0.01 + 0.004 * (bi % 7) as f32);
match dtype {
DType::Q4_0 | DType::Q8_0 | DType::Q4_1 | DType::Q4KM => {
blk[0..2].copy_from_slice(&d.to_bits().to_le_bytes());
}
DType::Q6K => {
let n = blk.len();
blk[n - 2..].copy_from_slice(&d.to_bits().to_le_bytes());
}
_ => unreachable!("dtype without an int8 kernel"),
}
if matches!(dtype, DType::Q4_1 | DType::Q4KM) {
let d2 = half::f16::from_f32(0.02 + 0.003 * (bi % 5) as f32);
blk[2..4].copy_from_slice(&d2.to_bits().to_le_bytes());
}
}
data
}
mod decode_prefill_identity {
use super::*;
#[test]
fn gemv_is_bit_identical_to_gemm_at_n1() {
let (m, k) = (7usize, 512usize);
let mut st = 0x5eed_1234u64;
let x: Vec<f32> = (0..k)
.map(|_| (lcg(&mut st) % 4000) as f32 / 1000.0 - 2.0)
.collect();
let mut b_scales = vec![0.0f32; k / 32];
let mut b_quants = vec![0i8; k];
quantize_f32_to_q8_0_into(&x, &mut b_scales, &mut b_quants);
let mut dtypes: Vec<DType> = vec![DType::Q4_1, DType::Q4KM, DType::Q6K];
#[cfg(target_arch = "aarch64")]
let scalar_remainder_at_n1 = crate::backend::cpu_features::cpu_features().tier
== crate::backend::cpu_features::CpuTier::NeonI8mm;
#[cfg(not(target_arch = "aarch64"))]
let scalar_remainder_at_n1 = false;
if !scalar_remainder_at_n1 {
dtypes.extend([DType::Q4_0, DType::Q8_0]);
}
let mut ran_any = false;
for dtype in dtypes {
let data = weights(dtype, m, k, &mut st);
let mut prefill = vec![0.0f32; m];
if !gemm_preq_dispatch(dtype, &data, &b_scales, &b_quants, &mut prefill, m, 1, k) {
continue;
}
ran_any = true;
let mut decode = vec![0.0f32; m];
gemv_dispatch(dtype, &data, &x, &mut decode, m, k, None);
assert!(
decode.iter().any(|v| *v != 0.0),
"{dtype:?}: decode produced all zeros, so the comparison below \
would pass against an equally empty prefill buffer"
);
for (i, (d, p)) in decode.iter().zip(&prefill).enumerate() {
assert_eq!(
d.to_bits(),
p.to_bits(),
"{dtype:?} row {i}: decode {d:e} vs prefill-at-n=1 {p:e}: \
the decode GEMV is not the prefill GEMM at n=1, so the \
same token yields different logits depending on whether \
it was consumed by prefill or by decode"
);
}
}
assert!(
ran_any || !int8_gemm_available(),
"no dtype reached a batched kernel although this host reports an \
int8 GEMM, so the loop asserted nothing"
);
}
#[test]
#[cfg(target_arch = "aarch64")]
fn gemv_with_preq_matches_gemv_dispatch_for_q4_1() {
let (m, k) = (7usize, 512usize);
let mut st = 0x1dea_5eedu64;
let x: Vec<f32> = (0..k)
.map(|_| (lcg(&mut st) % 4000) as f32 / 1000.0 - 2.0)
.collect();
let data = weights(DType::Q4_1, m, k, &mut st);
let mut scales = vec![0.0f32; k / 32 + 5];
let mut quants = vec![0i8; k + 160];
quantize_f32_to_q8_0_into(&x, &mut scales[..k / 32], &mut quants[..k]);
let mut want = vec![0.0f32; m];
gemv_dispatch(DType::Q4_1, &data, &x, &mut want, m, k, None);
assert!(
want.iter().any(|v| *v != 0.0),
"reference produced all zeros, so the comparison proves nothing"
);
let mut got = vec![0.0f32; m];
gemv_with_preq(DType::Q4_1, &data, &scales, &quants, &x, &mut got, m, k);
for (i, (a, b)) in want.iter().zip(&got).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"Q4_1 row {i}: gemv_dispatch {a:e} vs gemv_with_preq {b:e}, so \
the pre-quantized decode path is not the path everything else \
is pinned against"
);
}
}
}
#[cfg(target_arch = "x86_64")]
mod flash_attention_simd {
use super::*;
fn lcg(state: &mut u64) -> f32 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*state >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
}
fn rel_l2(a: &[f32], b: &[f32]) -> f32 {
let num: f64 = a
.iter()
.zip(b)
.map(|(&x, &y)| ((x - y) as f64).powi(2))
.sum();
let den: f64 = a.iter().map(|&x| (x as f64).powi(2)).sum();
(num.sqrt() / (den.sqrt() + 1e-12)) as f32
}
fn check(head_dim: usize, start_pos: usize, kernel: &str) {
let group_size = 3;
let n_kv_heads = 2;
let n_heads = n_kv_heads * group_size;
let n = 70usize; let kv_len = start_pos + n; let kv_dim = n_kv_heads * head_dim;
let q_dim = n_heads * head_dim;
let scale = 1.0 / (head_dim as f32).sqrt();
let mut st = 0x1234_5678_9abc_def0u64 ^ (head_dim as u64) ^ ((start_pos as u64) << 40);
let q: Vec<f32> = (0..q_dim * n).map(|_| lcg(&mut st)).collect();
let k: Vec<f32> = (0..kv_len * kv_dim).map(|_| lcg(&mut st)).collect();
let v: Vec<f32> = (0..kv_len * kv_dim).map(|_| lcg(&mut st)).collect();
for kv_h in 0..n_kv_heads {
let n_heads_start = kv_h * group_size;
let kv_h_offset = kv_h * head_dim;
let mut out_ref = vec![0.0f32; group_size * n * head_dim];
flash_attention_gqa_scalar(
&q,
&k,
&v,
&mut out_ref,
n_heads_start,
group_size,
n,
n,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
let mut out_simd = vec![0.0f32; group_size * n * head_dim];
match kernel {
"avx2" => unsafe {
flash_attention_gqa_avx2(
&q,
&k,
&v,
&mut out_simd,
n_heads_start,
group_size,
n,
n,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
},
#[cfg(feature = "avx512")]
"avx512" => unsafe {
flash_attention_gqa_avx512(
&q,
&k,
&v,
&mut out_simd,
n_heads_start,
group_size,
n,
n,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
);
},
"dispatch" => flash_attention_gqa_cpu(
&q,
&k,
&v,
&mut out_simd,
n_heads_start,
group_size,
n,
n,
kv_dim,
kv_h_offset,
head_dim,
scale,
start_pos,
),
other => panic!("unknown kernel {other}"),
}
let rel = rel_l2(&out_ref, &out_simd);
assert!(
rel < 1e-4,
"{kernel} head_dim={head_dim} start_pos={start_pos} kv_h={kv_h}: \
relative L2 deviation {rel:e} exceeds 1e-4 vs scalar reference"
);
}
}
const START_POS: [usize; 2] = [0, 37];
#[test]
fn avx2_matches_scalar_across_head_dims() {
if !is_x86_feature_detected!("avx2") || !is_x86_feature_detected!("fma") {
if std::env::var("CERA_REQUIRE_SIMD")
.unwrap_or_default()
.split(',')
.any(|f| f.trim() == "avx2")
{
panic!("CERA_REQUIRE_SIMD=avx2 but avx2/fma not detected");
}
eprintln!("[flash-avx2] SKIP: avx2/fma not detected");
return;
}
for hd in [64usize, 72, 128, 256] {
for sp in START_POS {
check(hd, sp, "avx2");
}
}
}
#[cfg(feature = "avx512")]
#[test]
fn avx512_matches_scalar_across_head_dims() {
if !is_x86_feature_detected!("avx512f") {
if std::env::var("CERA_REQUIRE_SIMD")
.unwrap_or_default()
.split(',')
.any(|f| f.trim() == "avx512")
{
panic!("CERA_REQUIRE_SIMD=avx512 but avx512f not detected");
}
eprintln!("[flash-avx512] SKIP: avx512f not detected");
return;
}
for hd in [64usize, 80, 128, 256] {
for sp in START_POS {
check(hd, sp, "avx512");
}
}
}
#[test]
fn dispatcher_routes_to_a_correct_kernel() {
if !is_x86_feature_detected!("avx2") || !is_x86_feature_detected!("fma") {
if std::env::var("CERA_REQUIRE_SIMD")
.unwrap_or_default()
.split(',')
.any(|f| f.trim() == "avx2")
{
panic!("CERA_REQUIRE_SIMD=avx2 but avx2/fma not detected");
}
eprintln!("[flash-dispatch] SKIP: avx2/fma not detected");
return;
}
for hd in [64usize, 72] {
for sp in START_POS {
check(hd, sp, "dispatch");
}
}
}
}
mod gemm_preq_guards {
use super::*;
use crate::tensor::DType;
fn args() -> (Vec<u8>, Vec<f32>, Vec<i8>, Vec<f32>) {
let (m, n, k) = (2usize, 1usize, 64usize);
let nb = k / 32;
(
vec![0u8; m * nb * DType::Q8_0.block_bytes()],
vec![0.0f32; n * nb],
vec![0i8; n * k],
vec![0.0f32; m * n],
)
}
#[test]
#[should_panic(expected = "out must be exactly")]
fn over_long_out_is_rejected() {
let (data, bs, bq, mut out) = args();
out.push(0.0);
gemm_preq_dispatch(DType::Q8_0, &data, &bs, &bq, &mut out, 2, 1, 64);
}
#[test]
#[should_panic(expected = "out must be exactly")]
fn short_out_is_rejected() {
let (data, bs, bq, mut out) = args();
out.pop();
gemm_preq_dispatch(DType::Q8_0, &data, &bs, &bq, &mut out, 2, 1, 64);
}
#[test]
#[should_panic(expected = "weights are")]
fn short_weights_are_rejected() {
let (mut data, bs, bq, mut out) = args();
data.truncate(DType::Q8_0.block_bytes());
gemm_preq_dispatch(DType::Q8_0, &data, &bs, &bq, &mut out, 2, 1, 64);
}
#[test]
#[should_panic(expected = "not a multiple")]
fn unaligned_k_is_rejected() {
let (data, bs, bq, mut out) = args();
gemm_preq_dispatch(DType::Q8_0, &data, &bs, &bq, &mut out, 2, 1, 60);
}
#[test]
fn well_formed_call_is_accepted() {
let (data, bs, bq, mut out) = args();
gemm_preq_dispatch(DType::Q8_0, &data, &bs, &bq, &mut out, 2, 1, 64);
}
}
#[cfg(all(any(target_arch = "x86_64", target_arch = "aarch64"), not(has_blas)))]
#[test]
fn repacked_q4_0_dispatch_matches_standard_dispatch() {
use crate::tensor::DType;
#[cfg(target_arch = "x86_64")]
if !int8_gemm_available() {
return;
}
#[cfg(target_arch = "aarch64")]
if !crate::backend::simd::neon::k_quant_gemm_available() {
return;
}
let (m, n, k) = (16usize, 13usize, 128usize);
let nb = k / 32;
let mut st = 0xd15e_a5edu64;
let mut lcg = || {
st = st.wrapping_mul(6364136223846793005).wrapping_add(1);
(st >> 33) as u32
};
let mut data = Vec::with_capacity(m * nb * DType::Q4_0.block_bytes());
for _ in 0..m * nb {
let d = half::f16::from_f32(0.01 + 0.04 * (lcg() as f32 / u32::MAX as f32));
data.extend_from_slice(&d.to_bits().to_le_bytes());
for _ in 0..16 {
data.push(lcg() as u8);
}
}
let mut b_scales = vec![0.0f32; n * nb];
let mut b_quants = vec![0i8; n * k];
for j in 0..n {
let col: Vec<f32> = (0..k)
.map(|_| lcg() as f32 / u32::MAX as f32 * 2.0 - 1.0)
.collect();
quantize_f32_to_q8_0_into(
&col,
&mut b_scales[j * nb..(j + 1) * nb],
&mut b_quants[j * k..(j + 1) * k],
);
}
let mut want = vec![0.0f32; m * n];
assert!(gemm_preq_dispatch(
DType::Q4_0,
&data,
&b_scales,
&b_quants,
&mut want,
m,
n,
k
));
let (packed, scales) = repack_q4_0_8x8(&data, m, k);
let mut got = vec![0.0f32; m * n];
assert!(gemm_preq_repacked_q4_0_dispatch(
&packed, &scales, &b_scales, &b_quants, &mut got, m, n, k
));
for (i, (g, w)) in got.iter().zip(&want).enumerate() {
assert!(
(g - w).abs() <= 1e-4 * w.abs().max(1.0),
"repacked vs standard dispatch [{},{}]: {g} vs {w}",
i / n,
i % n,
);
}
}
#[cfg(all(target_arch = "x86_64", not(has_blas)))]
#[test]
fn repacked_q4_k_dispatch_matches_standard_dispatch() {
use crate::tensor::DType;
if !int8_gemm_available() {
return;
}
let (m, n, k) = (16usize, 13usize, 256usize);
let sb = k / 256;
let nb = k / 32;
let mut st = 0x9e37_79b9u64;
let mut lcg = || {
st = st.wrapping_mul(6364136223846793005).wrapping_add(1);
(st >> 33) as u32
};
let bsz = size_of::<crate::quant::BlockQ4KM>();
let mut data = vec![0u8; m * sb * bsz];
for chunk in data.chunks_mut(bsz) {
let d = half::f16::from_f32(0.01 + 0.04 * (lcg() as f32 / u32::MAX as f32));
let dmin = half::f16::from_f32(0.02 + 0.03 * (lcg() as f32 / u32::MAX as f32));
chunk[0..2].copy_from_slice(&d.to_bits().to_le_bytes());
chunk[2..4].copy_from_slice(&dmin.to_bits().to_le_bytes());
for b in chunk[4..].iter_mut() {
*b = lcg() as u8;
}
}
let mut b_scales = vec![0.0f32; n * nb];
let mut b_quants = vec![0i8; n * k];
for j in 0..n {
let col: Vec<f32> = (0..k)
.map(|_| lcg() as f32 / u32::MAX as f32 * 2.0 - 1.0)
.collect();
quantize_f32_to_q8_0_into(
&col,
&mut b_scales[j * nb..(j + 1) * nb],
&mut b_quants[j * k..(j + 1) * k],
);
}
let mut want = vec![0.0f32; m * n];
assert!(gemm_preq_dispatch(
DType::Q4KM,
&data,
&b_scales,
&b_quants,
&mut want,
m,
n,
k
));
let (packed, dsc, dmn) = repack_q4_k_8x8(&data, m, k);
let mut got = vec![0.0f32; m * n];
assert!(gemm_preq_repacked_q4_k_dispatch(
&packed, &dsc, &dmn, &b_scales, &b_quants, &mut got, m, n, k
));
for (i, (g, w)) in got.iter().zip(&want).enumerate() {
assert!(
(g - w).abs() <= 1e-4 * w.abs().max(1.0),
"repacked Q4_K vs standard dispatch [{},{}]: {g} vs {w}",
i / n,
i % n,
);
}
}
#[test]
fn gemv_dispatch_matches_dequantized_reference() {
use crate::tensor::DType;
let (m, k) = (7usize, 256usize);
let mut st = 0x5eed_1234u64;
let x: Vec<f32> = (0..k)
.map(|_| (lcg(&mut st) % 2000) as f32 / 1000.0 - 1.0)
.collect();
for dtype in [
DType::Q4_0,
DType::Q8_0,
DType::Q4_1,
DType::Q4KM,
DType::Q6K,
] {
let data = weights(dtype, m, k, &mut st);
let mut w = vec![0.0f32; m * k];
match dtype {
DType::Q4_0 => crate::quant::dequantize_q4_0_matrix(&data, m, k, &mut w),
DType::Q8_0 => crate::quant::dequantize_q8_0_matrix(&data, m, k, &mut w),
DType::Q4_1 => crate::quant::dequantize_q4_1_matrix(&data, m, k, &mut w),
DType::Q4KM => crate::quant::dequantize_q4_k_m_matrix(&data, m, k, &mut w),
DType::Q6K => crate::quant::dequantize_q6_k_matrix(&data, m, k, &mut w),
_ => unreachable!(),
}
let want: Vec<f32> = (0..m)
.map(|i| (0..k).map(|j| w[i * k + j] * x[j]).sum())
.collect();
let mut got = vec![0.0f32; m];
gemv_dispatch(dtype, &data, &x, &mut got, m, k, None);
let mut scratch_s = vec![7.0f32; 1];
let mut scratch_q = vec![7i8; 1];
let mut got_scratch = vec![0.0f32; m];
gemv_dispatch(
dtype,
&data,
&x,
&mut got_scratch,
m,
k,
Some((&mut scratch_s, &mut scratch_q)),
);
for (i, (a, b)) in got.iter().zip(&got_scratch).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"{dtype:?} row {i}: lending scratch changed the result \
({a} vs {b}) — the two `kq_gemv!` arms have diverged"
);
}
let scale = want.iter().fold(0.0f32, |a, v| a.max(v.abs())).max(1.0);
for (i, (g, wv)) in got.iter().zip(&want).enumerate() {
assert!(
(g - wv).abs() <= 0.02 * scale,
"{dtype:?} row {i}: dispatch gave {g}, dequantized reference {wv} \
— a wrong kernel, not a rounding difference"
);
}
}
}
#[test]
fn test_matmul_f32_identity() {
let a = vec![1.0, 0.0, 0.0, 1.0];
let b = vec![1.0, 2.0, 3.0, 4.0];
let mut c = vec![0.0; 4];
matmul_f32(&a, &b, &mut c, 2, 2, 2);
assert_eq!(c, vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn test_matmul_f32_3x2_times_2x4() {
let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let b = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let mut c = vec![0.0; 12];
matmul_f32(&a, &b, &mut c, 3, 4, 2);
assert_eq!(
c,
vec![
11.0, 14.0, 17.0, 20.0, 23.0, 30.0, 37.0, 44.0, 35.0, 46.0, 57.0, 68.0
]
);
}
#[test]
fn test_rmsnorm() {
let mut x = vec![1.0, 2.0, 3.0, 4.0];
let weight = vec![1.0, 1.0, 1.0, 1.0];
let eps = 1e-5;
let rms = (7.5f32 + eps).sqrt();
let expected: Vec<f32> = vec![1.0 / rms, 2.0 / rms, 3.0 / rms, 4.0 / rms];
rmsnorm(&mut x, &weight, eps);
for (i, (&got, &exp)) in x.iter().zip(expected.iter()).enumerate() {
assert!(
(got - exp).abs() < 1e-5,
"rmsnorm[{i}]: got {got}, expected {exp}"
);
}
}
#[test]
fn test_rmsnorm_with_weight() {
let mut x = vec![2.0, 2.0];
let weight = vec![3.0, 0.5];
let eps = 1e-5;
let rms = (4.0f32 + eps).sqrt(); let inv_rms = 1.0 / rms;
let expected = [2.0 * inv_rms * 3.0, 2.0 * inv_rms * 0.5];
rmsnorm(&mut x, &weight, eps);
for (i, (&got, &exp)) in x.iter().zip(expected.iter()).enumerate() {
assert!(
(got - exp).abs() < 1e-5,
"rmsnorm[{i}]: got {got}, expected {exp}"
);
}
}
#[test]
fn test_silu() {
let mut x = vec![0.0, 1.0, -1.0, 5.0];
silu_inplace(&mut x);
assert!((x[0] - 0.0).abs() < 1e-5);
assert!((x[1] - 0.7311).abs() < 1e-3);
assert!((x[2] - (-0.2689)).abs() < 1e-3);
assert!((x[3] - 4.9665).abs() < 1e-3);
}
#[test]
fn test_silu_mul_inplace() {
let mut gate = vec![0.0, 1.0, -1.0, 5.0];
let up = vec![2.0, 3.0, 0.5, 1.0];
let mut gate_ref = gate.clone();
silu_inplace(&mut gate_ref);
mul_inplace(&mut gate_ref, &up);
silu_mul_inplace(&mut gate, &up);
for (i, (&got, &expected)) in gate.iter().zip(gate_ref.iter()).enumerate() {
assert!(
(got - expected).abs() < 1e-6,
"silu_mul mismatch at {i}: got {got}, expected {expected}"
);
}
}
#[test]
fn test_sigmoid() {
let mut x = vec![0.0f32, 2.0, -2.0, 10.0, -10.0];
sigmoid_inplace(&mut x);
assert!((x[0] - 0.5).abs() < 1e-5, "sigmoid(0) = {}", x[0]);
assert!((x[1] - 0.880_797).abs() < 1e-3, "sigmoid(2) = {}", x[1]);
assert!((x[2] - 0.119_203).abs() < 1e-3, "sigmoid(-2) = {}", x[2]);
assert!(x[3] > 0.999_5, "sigmoid(10) = {}", x[3]);
assert!(x[4] < 5e-4, "sigmoid(-10) = {}", x[4]);
}
#[test]
fn test_glu_split() {
let input = vec![3.0, 7.0, 0.0, 100.0];
let mut output = vec![0.0; 2];
glu_split(&input, &mut output);
assert!((output[0] - 1.5).abs() < 1e-5, "got {}", output[0]); assert!((output[1] - 7.0).abs() < 1e-3, "got {}", output[1]); }
#[test]
fn test_softmax() {
let mut x = vec![1.0, 2.0, 3.0];
softmax_inplace(&mut x);
let sum: f32 = x.iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
assert!(x[0] < x[1]);
assert!(x[1] < x[2]);
assert!((x[0] - 0.0900).abs() < 1e-3);
assert!((x[1] - 0.2447).abs() < 1e-3);
assert!((x[2] - 0.6652).abs() < 1e-3);
}
#[test]
fn test_layer_norm_zero_mean_unit_var_after_norm() {
let mut x = vec![1.0, 2.0, 3.0, 4.0];
let weight = vec![1.0; 4];
let bias = vec![0.0; 4];
layer_norm_inplace(&mut x, &weight, &bias, 1e-5);
let mean: f32 = x.iter().sum::<f32>() / x.len() as f32;
let var: f32 = x.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / x.len() as f32;
assert!(mean.abs() < 1e-5, "mean = {mean}");
assert!((var.sqrt() - 1.0).abs() < 1e-3, "std = {}", var.sqrt());
}
#[test]
fn test_layer_norm_applies_affine() {
let mut x = vec![1.0, 2.0, 3.0, 4.0];
let weight = vec![2.0; 4];
let bias = vec![10.0; 4];
layer_norm_inplace(&mut x, &weight, &bias, 1e-5);
let mean: f32 = x.iter().sum::<f32>() / x.len() as f32;
assert!((mean - 10.0).abs() < 1e-3, "mean = {mean} expected ~10");
}
#[test]
fn test_gelu_erf_known_values() {
let mut x = vec![0.0f32, 1.0, -1.0, 2.0];
gelu_erf_inplace(&mut x);
assert!(x[0].abs() < 1e-4, "gelu(0) = {}", x[0]);
assert!((x[1] - 0.8413).abs() < 5e-3, "gelu(1) = {}", x[1]);
assert!((x[2] + 0.1587).abs() < 5e-3, "gelu(-1) = {}", x[2]);
assert!((x[3] - 1.9545).abs() < 5e-3, "gelu(2) = {}", x[3]);
}
#[test]
fn test_gelu_tanh_known_values() {
let mut x = vec![0.0f32, 1.0, -1.0, 2.0];
gelu_inplace(&mut x);
assert!(x[0].abs() < 1e-6, "gelu_tanh(0) = {}", x[0]);
assert!((x[1] - 0.841_192).abs() < 1e-4, "gelu_tanh(1) = {}", x[1]);
assert!((x[2] + 0.158_808).abs() < 1e-4, "gelu_tanh(-1) = {}", x[2]);
assert!((x[3] - 1.954_598).abs() < 1e-4, "gelu_tanh(2) = {}", x[3]);
}
#[test]
fn test_conv1d_standard_identity_kernel() {
let input = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, ];
let weight = vec![
1.0, 0.0, 0.0, 1.0, ];
let mut output = vec![0.0; 2 * 3];
let t_out = conv1d(&input, &weight, None, &mut output, 2, 2, 3, 1, 1, 0, 1);
assert_eq!(t_out, 3);
assert_eq!(output, input);
}
#[test]
fn test_conv1d_depthwise_per_channel() {
let input = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, ];
let weight = vec![
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, ];
let mut output = vec![0.0; 2 * 4];
let t_out = conv1d(&input, &weight, None, &mut output, 2, 2, 4, 3, 1, 1, 2);
assert_eq!(t_out, 4);
assert_eq!(&output[0..4], &[0.0, 1.0, 2.0, 3.0]);
assert_eq!(&output[4..8], &[6.0, 7.0, 8.0, 0.0]);
}
#[test]
fn test_conv1d_strided() {
let input = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let weight = vec![1.0, 1.0]; let mut output = vec![0.0; 3]; let t_out = conv1d(&input, &weight, None, &mut output, 1, 1, 6, 2, 2, 0, 1);
assert_eq!(t_out, 3);
assert_eq!(output, vec![3.0, 7.0, 11.0]);
}
#[test]
fn test_conv1d_with_bias() {
let input = vec![1.0, 2.0, 3.0];
let weight = vec![1.0]; let bias = vec![5.0];
let mut output = vec![0.0; 3];
conv1d(
&input,
&weight,
Some(&bias),
&mut output,
1,
1,
3,
1,
1,
0,
1,
);
assert_eq!(output, vec![6.0, 7.0, 8.0]);
}
#[test]
fn test_conv2d_pointwise_identity() {
let input = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, ];
let weight = vec![
1.0, 0.0, 0.0, 1.0, ];
let mut output = vec![0.0; 2 * 2 * 3];
let (h_out, w_out) = conv2d(
&input,
&weight,
None,
&mut output,
2,
2,
2,
3,
1,
1,
1,
1,
0,
0,
1,
);
assert_eq!((h_out, w_out), (2, 3));
assert_eq!(output, input);
}
#[test]
fn test_conv2d_strided_with_pad() {
let input: Vec<f32> = (1..=16).map(|v| v as f32).collect();
let weight = vec![1.0 / 9.0; 9];
let mut output = vec![0.0; 2 * 2];
let (h_out, w_out) = conv2d(
&input,
&weight,
None,
&mut output,
1,
1,
4,
4,
3,
3,
2,
2,
1,
1,
1,
);
assert_eq!((h_out, w_out), (2, 2));
assert!((output[0] - 14.0 / 9.0).abs() < 1e-6, "got {}", output[0]);
}
#[test]
fn test_conv2d_depthwise_per_channel_independence() {
let input = vec![
1.0, 2.0, 3.0, 4.0, 10.0, 20.0, 30.0, 40.0, ];
let weight = vec![2.0, 0.5];
let mut output = vec![0.0; 2 * 2 * 2];
let (h_out, w_out) = conv2d(
&input,
&weight,
None,
&mut output,
2,
2,
2,
2,
1,
1,
1,
1,
0,
0,
2,
);
assert_eq!((h_out, w_out), (2, 2));
assert_eq!(&output[0..4], &[2.0, 4.0, 6.0, 8.0]);
assert_eq!(&output[4..8], &[5.0, 10.0, 15.0, 20.0]);
}
#[test]
fn test_conv2d_with_bias() {
let input = vec![1.0, 2.0, 3.0, 4.0]; let weight = vec![1.0]; let bias = vec![10.0];
let mut output = vec![0.0; 4];
let (h_out, w_out) = conv2d(
&input,
&weight,
Some(&bias),
&mut output,
1,
1,
2,
2,
1,
1,
1,
1,
0,
0,
1,
);
assert_eq!((h_out, w_out), (2, 2));
assert_eq!(output, vec![11.0, 12.0, 13.0, 14.0]);
}
#[test]
fn test_conv2d_pad_zero_contribution() {
let input = vec![7.0];
let weight = vec![1.0; 9];
let mut output = vec![0.0; 1];
let (h_out, w_out) = conv2d(
&input,
&weight,
None,
&mut output,
1,
1,
1,
1,
3,
3,
1,
1,
1,
1,
1,
);
assert_eq!((h_out, w_out), (1, 1));
assert_eq!(output[0], 7.0);
}
#[test]
fn test_conv2d_grouped_non_depthwise_fallback() {
let input = vec![1.0, 2.0, 3.0, 4.0];
let weight = vec![
1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0, 1.0, ];
let mut output = vec![0.0; 4];
let (h_out, w_out) = conv2d(
&input,
&weight,
None,
&mut output,
4,
4,
1,
1,
1,
1,
1,
1,
0,
0,
2,
);
assert_eq!((h_out, w_out), (1, 1));
assert_eq!(output, vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn test_ggml_expf() {
let test_vals = [0.0f32, 1.0, -1.0, 2.0, -5.0, -10.0, -50.0, 80.0];
for &x in &test_vals {
let got = ggml_expf(x);
let expected = x.exp();
let rel_err = if expected.abs() > 1e-10 {
((got - expected) / expected).abs()
} else {
(got - expected).abs()
};
assert!(
rel_err < 1e-5,
"ggml_expf({x}) = {got}, expected {expected}, rel_err = {rel_err}"
);
}
assert!(ggml_expf(100.0).is_infinite() || ggml_expf(100.0) > 1e30);
assert!(ggml_expf(-200.0) < 1e-30);
}
#[test]
fn test_softmax_numerical_stability() {
let mut x = vec![1000.0, 1001.0, 1002.0];
softmax_inplace(&mut x);
let sum: f32 = x.iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
assert!(x.iter().all(|v| v.is_finite()));
}
#[test]
fn test_rope_basic() {
let mut q = vec![1.0, 2.0, 3.0, 4.0]; let mut k = vec![5.0, 6.0, 7.0, 8.0];
let q_orig = q.clone();
let k_orig = k.clone();
rope(&mut q, &mut k, 0, 1, 1, 4, 10000.0);
for i in 0..4 {
assert!((q[i] - q_orig[i]).abs() < 1e-5, "q[{i}] changed at pos=0");
assert!((k[i] - k_orig[i]).abs() < 1e-5, "k[{i}] changed at pos=0");
}
}
#[test]
fn test_rope_rotates() {
let mut q = vec![1.0, 0.0, 0.0, 0.0]; let mut k = vec![1.0, 0.0, 0.0, 0.0];
rope(&mut q, &mut k, 10, 1, 1, 4, 10000.0);
assert!((q[0] - 1.0).abs() > 1e-3 || (q[2]).abs() > 1e-3);
}
#[test]
fn test_rope_norm_basic() {
let mut q = vec![1.0, 2.0, 3.0, 4.0];
let mut k = vec![5.0, 6.0, 7.0, 8.0];
let q_orig = q.clone();
let k_orig = k.clone();
rope_norm(&mut q, &mut k, 0, 1, 1, 4, 10000.0, None);
for i in 0..4 {
assert!((q[i] - q_orig[i]).abs() < 1e-5, "q[{i}] changed at pos=0");
assert!((k[i] - k_orig[i]).abs() < 1e-5, "k[{i}] changed at pos=0");
}
}
#[test]
fn test_rope_norm_rotates_adjacent_pairs() {
let head_dim = 4;
let freq_base = 10000.0_f32;
let pos = 1usize;
let mut head = vec![1.0, 0.0, 1.0, 0.0];
apply_rope_norm_to_head(&mut head, pos, head_dim, freq_base, None);
let theta0 = pos as f32; assert!((head[0] - theta0.cos()).abs() < 1e-5);
assert!((head[1] - theta0.sin()).abs() < 1e-5);
let theta1 = pos as f32 * freq_base.powf(-2.0 / head_dim as f32);
assert!((head[2] - theta1.cos()).abs() < 1e-5);
assert!((head[3] - theta1.sin()).abs() < 1e-5);
}
#[test]
fn test_rope_norm_freq_factors_divide_theta() {
let head_dim = 4;
let freq_base = 10000.0_f32;
let pos = 1usize;
let ff = [2.0_f32, 4.0]; let mut head = vec![1.0, 0.0, 1.0, 0.0];
apply_rope_norm_to_head(&mut head, pos, head_dim, freq_base, Some(&ff));
let theta0 = pos as f32 / ff[0];
assert!((head[0] - theta0.cos()).abs() < 1e-5);
assert!((head[1] - theta0.sin()).abs() < 1e-5);
let theta1 = pos as f32 * freq_base.powf(-2.0 / head_dim as f32) / ff[1];
assert!((head[2] - theta1.cos()).abs() < 1e-5);
assert!((head[3] - theta1.sin()).abs() < 1e-5);
}
#[test]
fn test_rope_norm_delta_composition() {
let head_dim = 8;
let freq_base = 1_000_000.0_f32;
let raw = vec![0.3, -1.1, 2.0, 0.5, -0.7, 1.3, 0.9, -0.2];
let mut direct = raw.clone();
apply_rope_norm_to_head(&mut direct, 13, head_dim, freq_base, None);
let mut composed = raw.clone();
apply_rope_norm_to_head(&mut composed, 5, head_dim, freq_base, None);
apply_rope_norm_delta_to_head(&mut composed, 13 - 5, head_dim, freq_base, None);
for i in 0..head_dim {
assert!(
(direct[i] - composed[i]).abs() < 1e-4,
"norm rope delta mismatch at {i}: {} vs {}",
direct[i],
composed[i]
);
}
}
#[test]
fn test_conv1d_depthwise_identity() {
let input = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let weight = vec![0.0, 1.0, 0.0, 0.0, 1.0, 0.0]; let mut output = vec![0.0; 6];
conv1d_depthwise(&input, &weight, None, &mut output, 2, 3, 3);
assert_eq!(output, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
}
#[test]
fn test_conv1d_depthwise_with_bias() {
let input = vec![1.0, 2.0]; let weight = vec![1.0, 1.0]; let bias = vec![10.0, 20.0];
let mut output = vec![0.0; 2];
conv1d_depthwise(&input, &weight, Some(&bias), &mut output, 2, 1, 1);
assert_eq!(output, vec![11.0, 22.0]);
}
#[test]
fn test_add_inplace() {
let mut a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
add_inplace(&mut a, &b);
assert_eq!(a, vec![5.0, 7.0, 9.0]);
}
#[test]
fn test_mul_inplace() {
let mut a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
mul_inplace(&mut a, &b);
assert_eq!(a, vec![4.0, 10.0, 18.0]);
}
#[test]
fn test_par_rows_n_basic() {
let mut out = vec![0.0f32; 6];
par_rows_n(&mut out, 2, 1, |(i, row)| {
row[0] = i as f32;
row[1] = i as f32 * 2.0;
});
assert_eq!(out, vec![0.0, 0.0, 1.0, 2.0, 2.0, 4.0]);
}
#[test]
fn test_par_rows_n_empty() {
let mut out: Vec<f32> = vec![];
par_rows_n(&mut out, 3, 1, |(_i, _row)| {
panic!("should not be called");
});
}
#[allow(clippy::too_many_arguments)]
fn attn_scores_scalar(
q: &[f32],
k_cache: &[f32],
scores: &mut [f32],
kv_dim: usize,
kv_h_off: usize,
head_dim: usize,
scale: f32,
seq_len: usize,
) {
for t in 0..seq_len {
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q[d] * k_cache[t * kv_dim + kv_h_off + d];
}
scores[t] = dot * scale;
}
}
fn attn_values_scalar(
scores: &[f32],
v_cache: &[f32],
out: &mut [f32],
kv_dim: usize,
kv_h_off: usize,
head_dim: usize,
seq_len: usize,
) {
for d in 0..head_dim {
let mut val = 0.0f32;
for t in 0..seq_len {
val += scores[t] * v_cache[t * kv_dim + kv_h_off + d];
}
out[d] = val;
}
}
#[test]
fn test_attn_scores_matches_scalar() {
let head_dim = 64;
let kv_dim = 128; let kv_h_off = 64; let seq_len = 10;
let scale = 1.0 / (head_dim as f32).sqrt();
let q: Vec<f32> = (0..head_dim).map(|i| (i as f32 - 32.0) * 0.05).collect();
let k_cache: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 7 + 3) % 31) as f32 * 0.04 - 0.6)
.collect();
let mut expected = vec![0.0f32; seq_len];
attn_scores_scalar(
&q,
&k_cache,
&mut expected,
kv_dim,
kv_h_off,
head_dim,
scale,
seq_len,
);
let mut actual = vec![0.0f32; seq_len];
attn_scores(
&q,
&k_cache,
&mut actual,
kv_dim,
kv_h_off,
head_dim,
scale,
seq_len,
);
for t in 0..seq_len {
let diff = (expected[t] - actual[t]).abs();
assert!(
diff < 1e-5,
"attn_scores mismatch at t={t}: expected={}, actual={}, diff={diff}",
expected[t],
actual[t]
);
}
}
#[test]
fn test_attn_values_matches_scalar() {
let head_dim = 64;
let kv_dim = 128;
let kv_h_off = 0;
let seq_len = 10;
let scores: Vec<f32> = (0..seq_len)
.map(|i| (i as f32 + 1.0) / seq_len as f32)
.collect();
let v_cache: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 11 + 5) % 29) as f32 * 0.03 - 0.4)
.collect();
let mut expected = vec![0.0f32; head_dim];
attn_values_scalar(
&scores,
&v_cache,
&mut expected,
kv_dim,
kv_h_off,
head_dim,
seq_len,
);
let mut actual = vec![0.0f32; head_dim];
attn_values(
&scores,
&v_cache,
&mut actual,
kv_dim,
kv_h_off,
head_dim,
seq_len,
);
for d in 0..head_dim {
let diff = (expected[d] - actual[d]).abs();
assert!(
diff < 1e-4,
"attn_values mismatch at d={d}: expected={}, actual={}, diff={diff}",
expected[d],
actual[d]
);
}
}
#[test]
fn test_attn_scores_seq_len_zero() {
let mut scores = vec![];
attn_scores(&[0.0; 64], &[], &mut scores, 64, 0, 64, 0.125, 0);
assert!(scores.is_empty());
}
fn check_attn_f16_kernels(head_dim: usize) {
let n_kv_heads = 2;
let kv_dim = n_kv_heads * head_dim;
let kv_h_off = head_dim; let seq_len = 13;
let scale = 1.0 / (head_dim as f32).sqrt();
let q: Vec<f32> = (0..head_dim).map(|i| (i as f32 - 8.0) * 0.05).collect();
let k32: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 7 + 3) % 31) as f32 * 0.04 - 0.6)
.collect();
let v32: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 11 + 5) % 29) as f32 * 0.03 - 0.4)
.collect();
let k16: Vec<u16> = k32
.iter()
.map(|&x| half::f16::from_f32(x).to_bits())
.collect();
let v16: Vec<u16> = v32
.iter()
.map(|&x| half::f16::from_f32(x).to_bits())
.collect();
let k_ref: Vec<f32> = k16.iter().map(|&b| f16_to_f32(b)).collect();
let v_ref: Vec<f32> = v16.iter().map(|&b| f16_to_f32(b)).collect();
let mut expected = vec![0.0f32; seq_len];
attn_scores_scalar(
&q,
&k_ref,
&mut expected,
kv_dim,
kv_h_off,
head_dim,
scale,
seq_len,
);
let mut actual = vec![0.0f32; seq_len];
attn_scores_f16(
&q,
&k16,
&mut actual,
kv_dim,
kv_h_off,
head_dim,
scale,
seq_len,
);
for t in 0..seq_len {
let diff = (expected[t] - actual[t]).abs();
assert!(
diff < 1e-4,
"attn_scores_f16 mismatch (head_dim={head_dim}) at t={t}: \
expected={}, actual={}, diff={diff}",
expected[t],
actual[t]
);
}
let scores: Vec<f32> = (0..seq_len)
.map(|i| (i as f32 + 1.0) / seq_len as f32)
.collect();
let mut vexp = vec![0.0f32; head_dim];
attn_values_scalar(
&scores, &v_ref, &mut vexp, kv_dim, kv_h_off, head_dim, seq_len,
);
let mut vact = vec![0.0f32; head_dim];
attn_values_f16(
&scores, &v16, &mut vact, kv_dim, kv_h_off, head_dim, seq_len,
);
for d in 0..head_dim {
let diff = (vexp[d] - vact[d]).abs();
assert!(
diff < 1e-4,
"attn_values_f16 mismatch (head_dim={head_dim}) at d={d}: \
expected={}, actual={}, diff={diff}",
vexp[d],
vact[d]
);
}
}
#[test]
fn test_attn_f16_matches_scalar_over_widened() {
check_attn_f16_kernels(64); check_attn_f16_kernels(12); check_attn_f16_kernels(10); check_attn_f16_kernels(6); }
fn check_attn_f32_kernels(head_dim: usize) {
let n_kv_heads = 2;
let kv_dim = n_kv_heads * head_dim;
let kv_h_off = head_dim; let seq_len = 13;
let scale = 1.0 / (head_dim as f32).sqrt();
let q: Vec<f32> = (0..head_dim).map(|i| (i as f32 - 8.0) * 0.05).collect();
let k: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 7 + 3) % 31) as f32 * 0.04 - 0.6)
.collect();
let v: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 11 + 5) % 29) as f32 * 0.03 - 0.4)
.collect();
let mut expected = vec![0.0f32; seq_len];
attn_scores_scalar(
&q,
&k,
&mut expected,
kv_dim,
kv_h_off,
head_dim,
scale,
seq_len,
);
let mut actual = vec![0.0f32; seq_len];
attn_scores(
&q,
&k,
&mut actual,
kv_dim,
kv_h_off,
head_dim,
scale,
seq_len,
);
for t in 0..seq_len {
let diff = (expected[t] - actual[t]).abs();
assert!(
diff < 1e-5,
"attn_scores mismatch (head_dim={head_dim}) at t={t}: \
expected={}, actual={}, diff={diff}",
expected[t],
actual[t]
);
}
let scores: Vec<f32> = (0..seq_len)
.map(|i| (i as f32 + 1.0) / seq_len as f32)
.collect();
let mut vexp = vec![0.0f32; head_dim];
attn_values_scalar(&scores, &v, &mut vexp, kv_dim, kv_h_off, head_dim, seq_len);
let mut vact = vec![0.0f32; head_dim];
attn_values(&scores, &v, &mut vact, kv_dim, kv_h_off, head_dim, seq_len);
for d in 0..head_dim {
let diff = (vexp[d] - vact[d]).abs();
assert!(
diff < 1e-5,
"attn_values mismatch (head_dim={head_dim}) at d={d}: \
expected={}, actual={}, diff={diff}",
vexp[d],
vact[d]
);
}
}
#[test]
fn test_attn_kernels_across_head_dims() {
for &hd in &[8usize, 16, 24, 32, 40, 48, 56, 64, 72, 128] {
check_attn_f32_kernels(hd);
check_attn_f16_kernels(hd);
}
check_attn_f32_kernels(256);
check_attn_f16_kernels(256);
}
#[test]
fn test_flash_attention_matches_naive() {
let n_heads = 4;
let n_kv_heads = 2;
let group_size = n_heads / n_kv_heads;
let head_dim = 64;
let hs = n_heads * head_dim;
let kv_dim = n_kv_heads * head_dim;
let n = 8;
let start_pos = 4;
let scale = 1.0 / (head_dim as f32).sqrt();
let total_seq = start_pos + n;
let mut q_mat = vec![0.0f32; hs * n];
let mut seed: u64 = 0xCAFE_BABE;
for v in q_mat.iter_mut() {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
*v = ((seed >> 33) as i32 as f32) * 1e-9;
}
let mut k_cache = vec![0.0f32; total_seq * kv_dim];
for v in k_cache.iter_mut() {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
*v = ((seed >> 33) as i32 as f32) * 1e-9;
}
let mut v_cache = vec![0.0f32; total_seq * kv_dim];
for v in v_cache.iter_mut() {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
*v = ((seed >> 33) as i32 as f32) * 1e-9;
}
let chunk_size = group_size * n * head_dim;
let mut flash_raw = vec![0.0f32; n_kv_heads * chunk_size];
for kv_h in 0..n_kv_heads {
let chunk = &mut flash_raw[kv_h * chunk_size..(kv_h + 1) * chunk_size];
flash_attention_gqa_cpu(
&q_mat,
&k_cache,
&v_cache,
chunk,
kv_h * group_size,
group_size,
n,
n, kv_dim,
kv_h * head_dim,
head_dim,
scale,
start_pos,
);
}
let mut flash_out = vec![0.0f32; hs * n];
for kv_h in 0..n_kv_heads {
for g in 0..group_size {
let h = kv_h * group_size + g;
let src_base = kv_h * chunk_size + g * n * head_dim;
for j in 0..n {
for d in 0..head_dim {
flash_out[(h * head_dim + d) * n + j] =
flash_raw[src_base + j * head_dim + d];
}
}
}
}
let mut naive_out = vec![0.0f32; hs * n];
for j in 0..n {
let seq_len = start_pos + j + 1; for h in 0..n_heads {
let kv_h = h / group_size;
let kv_h_offset = kv_h * head_dim;
let mut q_head = vec![0.0f32; head_dim];
for d in 0..head_dim {
q_head[d] = q_mat[(h * head_dim + d) * n + j];
}
let mut scores = vec![0.0f32; seq_len];
attn_scores(
&q_head,
&k_cache,
&mut scores,
kv_dim,
kv_h_offset,
head_dim,
scale,
seq_len,
);
softmax_inplace(&mut scores);
let mut attn_out = vec![0.0f32; head_dim];
attn_values(
&scores,
&v_cache,
&mut attn_out,
kv_dim,
kv_h_offset,
head_dim,
seq_len,
);
for d in 0..head_dim {
naive_out[(h * head_dim + d) * n + j] = attn_out[d];
}
}
}
let mut max_diff = 0.0f32;
for i in 0..hs * n {
max_diff = max_diff.max((flash_out[i] - naive_out[i]).abs());
}
assert!(
max_diff < 1e-4,
"flash vs naive max_diff = {max_diff} (expected < 1e-4)"
);
}
#[cfg(all(feature = "parallel", not(target_arch = "wasm32")))]
#[test]
#[ignore]
fn microbench_dispatch() {
use std::time::Instant;
let threads = crate::backend::threadpool::RowPool::decode().num_threads();
let mut warm = vec![0.0f32; 4096];
for _ in 0..100 {
par_rows(&mut warm, gemv_min_rows(), |(i, v)| *v = i as f32);
}
eprintln!("\n=== RowPool dispatch cost ({threads} decode threads) ===");
for rows in [512usize, 2048, 8192] {
let mut y = vec![0.0f32; rows];
let iters = 2000;
let t = Instant::now();
for _ in 0..iters {
par_rows(&mut y, gemv_min_rows(), |(_i, v)| *v += 1.0);
}
let per = t.elapsed().as_secs_f64() / iters as f64;
std::hint::black_box(&y);
eprintln!(
" rows={rows:<5} {:>7.1} us/dispatch -> {:>6.2} ms/token at 113 dispatches",
per * 1e6,
per * 1e3 * 113.0
);
}
}
#[cfg(all(target_arch = "aarch64", feature = "parallel"))]
#[test]
#[ignore]
fn microbench_gemv_q4_0() {
use std::time::Instant;
let n_threads = crate::backend::threadpool::RowPool::decode().num_threads();
let m = 6912; let k = 2048; let iters = 200;
let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * size_of::<crate::quant::BlockQ4_0>();
let mut weight = vec![0u8; m * row_bytes];
let mut s: u64 = 0xdead_beef;
for b in weight.iter_mut() {
s = s.wrapping_mul(6364136223846793005).wrapping_add(1);
*b = (s >> 33) as u8;
}
let x: Vec<f32> = (0..k)
.map(|i| ((i * 31) % 127) as f32 * 0.01 - 0.5)
.collect();
let (x_scales, x_quants) = quantize_f32_to_q8_0(&x);
let mut y = vec![0.0f32; m];
gemv_q4_0_with_q8(&weight, &x_scales, &x_quants, &mut y, m, k);
let t0 = Instant::now();
for _ in 0..iters {
gemv_q4_0_with_q8(&weight, &x_scales, &x_quants, &mut y, m, k);
}
let elapsed = t0.elapsed().as_secs_f64();
let per_call = elapsed / iters as f64;
let weight_bytes = m * row_bytes;
let input_bytes = x_scales.len() * 4 + x_quants.len();
let total_bytes = weight_bytes + input_bytes;
let bw_gbps = (total_bytes as f64 / per_call) / 1e9;
eprintln!("\n=== GEMV Q4_0×Q8_0 microbench (m={m}, k={k}) ===");
eprintln!(" per-call: {:.1} µs", per_call * 1e6);
eprintln!(" weight: {:.2} MB", weight_bytes as f64 / 1e6);
eprintln!(" bandwidth: {:.1} GB/s", bw_gbps);
eprintln!(" decode pool threads: {n_threads}");
let m_large = 65536;
let mut weight_large = vec![0u8; m_large * row_bytes];
for b in weight_large.iter_mut() {
s = s.wrapping_mul(6364136223846793005).wrapping_add(1);
*b = (s >> 33) as u8;
}
let mut y_large = vec![0.0f32; m_large];
gemv_q4_0_with_q8(
&weight_large,
&x_scales,
&x_quants,
&mut y_large,
m_large,
k,
);
let t0 = Instant::now();
for _ in 0..20 {
gemv_q4_0_with_q8(
&weight_large,
&x_scales,
&x_quants,
&mut y_large,
m_large,
k,
);
}
let elapsed = t0.elapsed().as_secs_f64();
let per_call = elapsed / 20.0;
let weight_bytes_large = m_large * row_bytes;
let bw_large = ((weight_bytes_large + input_bytes) as f64 / per_call) / 1e9;
eprintln!("\n=== GEMV Q4_0×Q8_0 large (m={m_large}, k={k}) ===");
eprintln!(" per-call: {:.1} µs", per_call * 1e6);
eprintln!(" weight: {:.2} MB", weight_bytes_large as f64 / 1e6);
eprintln!(" bandwidth: {:.1} GB/s", bw_large);
}
}
#[cfg(test)]
mod f16_gemv_tests {
use super::*;
const SHAPES: &[(usize, usize)] =
&[(7, 12), (5, 64), (3, 70), (512, 64), (512, 70), (64, 2048)];
fn check_widened_gemv(
label: &str,
seed: u32,
narrow: fn(f32) -> u16,
widen: fn(u16) -> f32,
kernel: fn(&[u8], &[f32], &mut [f32], usize, usize),
) {
for &(m, k) in SHAPES {
let mut st = seed;
let mut next = || {
st = st.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((st >> 8) as f32 / 8_388_608.0) - 1.0
};
let halves: Vec<u16> = (0..m * k).map(|_| narrow(next())).collect();
let x: Vec<f32> = (0..k).map(|_| next()).collect();
let mut got = vec![0.0f32; m];
kernel(bytemuck::cast_slice(&halves), &x, &mut got, m, k);
let widened: Vec<f32> = halves.iter().map(|&h| widen(h)).collect();
let mut want = vec![0.0f32; m];
gemv_f32(bytemuck::cast_slice(&widened), &x, &mut want, m, k);
for (i, (g, w)) in got.iter().zip(&want).enumerate() {
let sum_abs: f32 = widened[i * k..(i + 1) * k]
.iter()
.zip(&x)
.map(|(a, b)| (a * b).abs())
.sum();
let tol = 1e-5 * (1.0 + w.abs()) + 2e-6 * sum_abs;
assert!(
(g - w).abs() <= tol,
"{label} m={m} k={k} row {i}: got {g} want {w} (tol {tol})"
);
}
}
}
#[test]
fn gemv_bf16_matches_widened_f32() {
check_widened_gemv(
"bf16",
0x1357_9BDF,
|v| half::bf16::from_f32(v).to_bits(),
|h| half::bf16::from_bits(h).to_f32(),
gemv_bf16,
);
}
#[test]
fn gemv_f16_matches_widened_f32() {
check_widened_gemv(
"f16",
0x9E37_79B9,
|v| half::f16::from_f32(v).to_bits(),
|h| half::f16::from_bits(h).to_f32(),
gemv_f16,
);
}
#[test]
fn dot_f32_various_lengths() {
for len in [
0, 1, 2, 3, 7, 8, 15, 16, 17, 31, 32, 33, 63, 64, 65, 127, 128, 255, 256, 513, 1024,
1025,
] {
let a: Vec<f32> = (0..len).map(|i| (i as f32 + 1.0) * 0.05).collect();
let b: Vec<f32> = (0..len).map(|i| (i as f32 + 2.0) * 0.025).collect();
let expected: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
let actual = dot_f32(&a, &b);
let diff = (actual - expected).abs();
let rel_diff = diff / expected.abs().max(1.0);
assert!(
rel_diff <= 1e-4,
"len {len}: expected {expected}, got {actual} (diff {diff}, rel_diff {rel_diff})"
);
}
}
}