use trueno::Vector;
#[must_use]
pub fn dot(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "dot product requires equal lengths");
let va = Vector::from_slice(a);
let vb = Vector::from_slice(b);
va.dot(&vb).unwrap_or_else(|_| dot_scalar(a, b))
}
#[inline]
#[must_use]
pub fn dot_nalloc(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "dot product requires equal lengths");
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("fma") && is_x86_feature_detected!("avx2") {
return unsafe { dot_fma_avx2_public(a, b) };
}
}
dot_scalar(a, b)
}
#[inline]
#[must_use]
pub fn dot_scalar(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
pub(crate) unsafe fn dot_fma_avx2_public(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::{
__m256, _mm256_add_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps,
_mm256_storeu_ps,
};
let n = a.len();
let mut i = 0;
let mut acc0: __m256;
let mut acc1: __m256;
let mut acc2: __m256;
let mut acc3: __m256;
unsafe {
acc0 = _mm256_setzero_ps();
acc1 = _mm256_setzero_ps();
acc2 = _mm256_setzero_ps();
acc3 = _mm256_setzero_ps();
while i + 32 <= n {
let a0 = _mm256_loadu_ps(a.as_ptr().add(i));
let b0 = _mm256_loadu_ps(b.as_ptr().add(i));
acc0 = _mm256_fmadd_ps(a0, b0, acc0);
let a1 = _mm256_loadu_ps(a.as_ptr().add(i + 8));
let b1 = _mm256_loadu_ps(b.as_ptr().add(i + 8));
acc1 = _mm256_fmadd_ps(a1, b1, acc1);
let a2 = _mm256_loadu_ps(a.as_ptr().add(i + 16));
let b2 = _mm256_loadu_ps(b.as_ptr().add(i + 16));
acc2 = _mm256_fmadd_ps(a2, b2, acc2);
let a3 = _mm256_loadu_ps(a.as_ptr().add(i + 24));
let b3 = _mm256_loadu_ps(b.as_ptr().add(i + 24));
acc3 = _mm256_fmadd_ps(a3, b3, acc3);
i += 32;
}
while i + 8 <= n {
let av = _mm256_loadu_ps(a.as_ptr().add(i));
let bv = _mm256_loadu_ps(b.as_ptr().add(i));
acc0 = _mm256_fmadd_ps(av, bv, acc0);
i += 8;
}
acc0 = _mm256_add_ps(acc0, acc1);
acc2 = _mm256_add_ps(acc2, acc3);
acc0 = _mm256_add_ps(acc0, acc2);
let mut buf = [0.0_f32; 8];
_mm256_storeu_ps(buf.as_mut_ptr(), acc0);
let mut sum = buf[0] + buf[1] + buf[2] + buf[3] + buf[4] + buf[5] + buf[6] + buf[7];
while i < n {
sum += a[i] * b[i];
i += 1;
}
sum
}
}
pub fn softmax_online_inplace(scores: &[f32], weights: &mut [f32]) {
assert_eq!(scores.len(), weights.len());
if scores.is_empty() {
return;
}
let mut max_val = scores[0];
let mut sum_exp = 1.0_f32;
for &s in &scores[1..] {
if s > max_val {
sum_exp = sum_exp * (max_val - s).exp() + 1.0;
max_val = s;
} else {
sum_exp += (s - max_val).exp();
}
}
let inv_sum = 1.0 / sum_exp;
for (w, &s) in weights.iter_mut().zip(scores.iter()) {
*w = (s - max_val).exp() * inv_sum;
}
}
#[must_use]
pub fn add(a: &[f32], b: &[f32]) -> Vec<f32> {
assert_eq!(a.len(), b.len(), "addition requires equal lengths");
let va = Vector::from_slice(a);
let vb = Vector::from_slice(b);
va.add(&vb)
.map_or_else(|_| vec![0.0; a.len()], |v| v.as_slice().to_vec())
}
#[must_use]
pub fn sub(a: &[f32], b: &[f32]) -> Vec<f32> {
assert_eq!(a.len(), b.len(), "subtraction requires equal lengths");
let va = Vector::from_slice(a);
let vb = Vector::from_slice(b);
va.sub(&vb)
.map_or_else(|_| vec![0.0; a.len()], |v| v.as_slice().to_vec())
}
#[must_use]
pub fn mul(a: &[f32], b: &[f32]) -> Vec<f32> {
assert_eq!(a.len(), b.len(), "multiplication requires equal lengths");
let va = Vector::from_slice(a);
let vb = Vector::from_slice(b);
va.mul(&vb)
.map_or_else(|_| vec![0.0; a.len()], |v| v.as_slice().to_vec())
}
#[must_use]
pub fn scale(a: &[f32], s: f32) -> Vec<f32> {
let va = Vector::from_slice(a);
va.scale(s)
.map_or_else(|_| vec![0.0; a.len()], |v| v.as_slice().to_vec())
}
#[must_use]
pub fn sum(a: &[f32]) -> f32 {
let va = Vector::from_slice(a);
va.sum().unwrap_or_else(|_| a.iter().sum())
}
#[must_use]
pub fn mean(a: &[f32]) -> f32 {
if a.is_empty() {
return 0.0;
}
sum(a) / a.len() as f32
}
#[must_use]
pub fn variance(a: &[f32]) -> f32 {
if a.is_empty() {
return 0.0;
}
let va = Vector::from_slice(a);
va.variance().unwrap_or_else(|_| {
let mean = a.iter().sum::<f32>() / a.len() as f32;
a.iter().map(|&x| (x - mean) * (x - mean)).sum::<f32>() / a.len() as f32
})
}
#[must_use]
pub fn std_dev(a: &[f32]) -> f32 {
variance(a).sqrt()
}
#[must_use]
pub fn max(a: &[f32]) -> f32 {
if a.is_empty() {
return f32::NEG_INFINITY;
}
let va = Vector::from_slice(a);
va.max().unwrap_or(f32::NEG_INFINITY)
}
#[must_use]
pub fn min(a: &[f32]) -> f32 {
if a.is_empty() {
return f32::INFINITY;
}
let va = Vector::from_slice(a);
va.min().unwrap_or(f32::INFINITY)
}
#[must_use]
pub fn argmax(a: &[f32]) -> usize {
if a.is_empty() {
return 0;
}
let va = Vector::from_slice(a);
va.argmax().unwrap_or(0)
}
#[must_use]
#[inline]
pub fn max_element(a: &[f32]) -> f32 {
max(a)
}
pub fn scale_inplace(a: &mut [f32], s: f32) {
for x in a.iter_mut() {
*x *= s;
}
}
pub fn axpy(a: f32, x: &[f32], y: &mut [f32]) {
assert_eq!(x.len(), y.len(), "axpy requires equal lengths");
for (yi, &xi) in y.iter_mut().zip(x.iter()) {
*yi += a * xi;
}
}
pub fn add_inplace(x: &[f32], y: &mut [f32]) {
assert_eq!(x.len(), y.len(), "add_inplace requires equal lengths");
for (yi, &xi) in y.iter_mut().zip(x.iter()) {
*yi += xi;
}
}
pub fn broadcast_add_inplace(matrix: &mut [f32], vec: &[f32], rows: usize, cols: usize) {
assert_eq!(matrix.len(), rows * cols, "matrix dimensions mismatch");
assert_eq!(vec.len(), cols, "vector dimension mismatch");
for row in 0..rows {
let row_start = row * cols;
add_inplace(vec, &mut matrix[row_start..row_start + cols]);
}
}
pub fn dequant_f16_row(f16_data: &[u16], out: &mut [f32]) {
debug_assert!(
out.len() >= f16_data.len(),
"output buffer too small: {} < {}",
out.len(),
f16_data.len()
);
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("avx") {
unsafe {
dequant_f16_row_f16c(f16_data, out);
}
return;
}
}
for (o, &bits) in out.iter_mut().zip(f16_data.iter()) {
*o = half::f16::from_bits(bits).to_f32();
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "f16c", enable = "avx")]
unsafe fn dequant_f16_row_f16c(f16_data: &[u16], out: &mut [f32]) {
use std::arch::x86_64::{_mm256_cvtph_ps, _mm256_storeu_ps, _mm_loadu_si128};
let n = f16_data.len();
let chunks = n / 8;
let remainder = n % 8;
let src = f16_data.as_ptr();
let dst = out.as_mut_ptr();
for i in 0..chunks {
let offset = i * 8;
unsafe {
let half8 = _mm_loadu_si128(src.add(offset).cast());
let float8 = _mm256_cvtph_ps(half8);
_mm256_storeu_ps(dst.add(offset), float8);
}
}
let base = chunks * 8;
for j in 0..remainder {
unsafe {
*dst.add(base + j) = half::f16::from_bits(*src.add(base + j)).to_f32();
}
}
}
#[must_use]
pub fn dot_f16(a_f16: &[u16], b: &[f32], buf: &mut [f32]) -> f32 {
assert_eq!(a_f16.len(), b.len(), "dot_f16 requires equal lengths");
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("f16c")
&& is_x86_feature_detected!("avx")
&& is_x86_feature_detected!("fma")
{
return unsafe { dot_f16_fused_f16c(a_f16, b) };
}
}
dequant_f16_row(a_f16, buf);
dot(&buf[..a_f16.len()], b)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "f16c", enable = "avx", enable = "fma")]
pub unsafe fn dot_f16_fused_f16c(a_f16: &[u16], b: &[f32]) -> f32 {
use std::arch::x86_64::{
__m256, _mm256_add_ps, _mm256_cvtph_ps, _mm256_fmadd_ps, _mm256_loadu_ps,
_mm256_setzero_ps, _mm256_storeu_ps, _mm_loadu_si128,
};
let n = a_f16.len();
let a_ptr = a_f16.as_ptr();
let b_ptr = b.as_ptr();
let mut acc0: __m256 = _mm256_setzero_ps();
let mut acc1: __m256 = _mm256_setzero_ps();
let mut acc2: __m256 = _mm256_setzero_ps();
let mut acc3: __m256 = _mm256_setzero_ps();
let mut i = 0;
while i + 32 <= n {
unsafe {
let h0 = _mm_loadu_si128(a_ptr.add(i).cast());
let a0 = _mm256_cvtph_ps(h0);
let b0 = _mm256_loadu_ps(b_ptr.add(i));
acc0 = _mm256_fmadd_ps(a0, b0, acc0);
let h1 = _mm_loadu_si128(a_ptr.add(i + 8).cast());
let a1 = _mm256_cvtph_ps(h1);
let b1 = _mm256_loadu_ps(b_ptr.add(i + 8));
acc1 = _mm256_fmadd_ps(a1, b1, acc1);
let h2 = _mm_loadu_si128(a_ptr.add(i + 16).cast());
let a2 = _mm256_cvtph_ps(h2);
let b2 = _mm256_loadu_ps(b_ptr.add(i + 16));
acc2 = _mm256_fmadd_ps(a2, b2, acc2);
let h3 = _mm_loadu_si128(a_ptr.add(i + 24).cast());
let a3 = _mm256_cvtph_ps(h3);
let b3 = _mm256_loadu_ps(b_ptr.add(i + 24));
acc3 = _mm256_fmadd_ps(a3, b3, acc3);
}
i += 32;
}
while i + 8 <= n {
unsafe {
let h0 = _mm_loadu_si128(a_ptr.add(i).cast());
let a0 = _mm256_cvtph_ps(h0);
let b0 = _mm256_loadu_ps(b_ptr.add(i));
acc0 = _mm256_fmadd_ps(a0, b0, acc0);
}
i += 8;
}
unsafe {
acc0 = _mm256_add_ps(acc0, acc1);
acc2 = _mm256_add_ps(acc2, acc3);
acc0 = _mm256_add_ps(acc0, acc2);
let mut sum_buf = [0.0_f32; 8];
_mm256_storeu_ps(sum_buf.as_mut_ptr(), acc0);
let mut result = sum_buf[0]
+ sum_buf[1]
+ sum_buf[2]
+ sum_buf[3]
+ sum_buf[4]
+ sum_buf[5]
+ sum_buf[6]
+ sum_buf[7];
while i < n {
let a_val = half::f16::from_bits(*a_ptr.add(i)).to_f32();
let b_val = *b_ptr.add(i);
result += a_val * b_val;
i += 1;
}
result
}
}
pub fn quant_f32_row_to_i8(row: &[f32]) -> (Vec<i8>, f32) {
let abs_max = row.iter().fold(0.0_f32, |m, &v| m.max(v.abs()));
if abs_max == 0.0 {
return (vec![0i8; row.len()], 0.0);
}
let scale = abs_max / 127.0;
let inv_scale = 127.0 / abs_max;
let quantized: Vec<i8> = row
.iter()
.map(|&v| (v * inv_scale).round().clamp(-127.0, 127.0) as i8)
.collect();
(quantized, scale)
}
#[must_use]
pub fn dot_i8(a_i8: &[i8], b: &[f32], scale: f32) -> f32 {
assert_eq!(a_i8.len(), b.len(), "dot_i8 requires equal lengths");
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return unsafe { dot_i8_avx2(a_i8, b, scale) };
}
}
dot_i8_scalar(a_i8, b, scale)
}
fn dot_i8_scalar(a_i8: &[i8], b: &[f32], scale: f32) -> f32 {
let mut sum = 0.0_f32;
for (&a, &x) in a_i8.iter().zip(b.iter()) {
sum += (a as f32) * x;
}
sum * scale
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn dot_i8_avx2(a_i8: &[i8], b: &[f32], scale: f32) -> f32 {
use std::arch::x86_64::{
_mm256_add_ps, _mm256_cvtepi32_ps, _mm256_cvtepi8_epi32, _mm256_fmadd_ps, _mm256_loadu_ps,
_mm256_setzero_ps, _mm256_storeu_ps, _mm_loadl_epi64,
};
let n = a_i8.len();
let mut i = 0;
unsafe {
let mut acc0 = _mm256_setzero_ps();
let mut acc1 = _mm256_setzero_ps();
while i + 16 <= n {
let i8_0 = _mm_loadl_epi64(a_i8.as_ptr().add(i).cast());
let i32_0 = _mm256_cvtepi8_epi32(i8_0);
let f32_0 = _mm256_cvtepi32_ps(i32_0);
let b0 = _mm256_loadu_ps(b.as_ptr().add(i));
acc0 = _mm256_fmadd_ps(f32_0, b0, acc0);
let i8_1 = _mm_loadl_epi64(a_i8.as_ptr().add(i + 8).cast());
let i32_1 = _mm256_cvtepi8_epi32(i8_1);
let f32_1 = _mm256_cvtepi32_ps(i32_1);
let b1 = _mm256_loadu_ps(b.as_ptr().add(i + 8));
acc1 = _mm256_fmadd_ps(f32_1, b1, acc1);
i += 16;
}
while i + 8 <= n {
let i8_r = _mm_loadl_epi64(a_i8.as_ptr().add(i).cast());
let i32_r = _mm256_cvtepi8_epi32(i8_r);
let f32_r = _mm256_cvtepi32_ps(i32_r);
let br = _mm256_loadu_ps(b.as_ptr().add(i));
acc0 = _mm256_fmadd_ps(f32_r, br, acc0);
i += 8;
}
acc0 = _mm256_add_ps(acc0, acc1);
let mut buf = [0.0_f32; 8];
_mm256_storeu_ps(buf.as_mut_ptr(), acc0);
let mut sum = buf[0] + buf[1] + buf[2] + buf[3] + buf[4] + buf[5] + buf[6] + buf[7];
while i < n {
sum += (a_i8[i] as f32) * b[i];
i += 1;
}
sum * scale
}
}
pub fn quant_f32_row_to_i4(row: &[f32], group_size: usize) -> (Vec<u8>, Vec<f32>) {
assert_eq!(
row.len() % group_size,
0,
"row len must be multiple of group_size"
);
let mut packed = Vec::with_capacity(row.len() / 2);
let mut scales = Vec::with_capacity(row.len() / group_size);
for chunk in row.chunks(group_size) {
let abs_max = chunk.iter().fold(0.0_f32, |m, &v| m.max(v.abs()));
let scale = if abs_max == 0.0 { 0.0 } else { abs_max / 7.0 };
let inv_scale = if scale == 0.0 { 0.0 } else { 1.0 / scale };
scales.push(scale);
for pair in chunk.chunks(2) {
let v0 = pair[0];
let v1 = if pair.len() > 1 { pair[1] } else { 0.0 };
let q0 = (v0 * inv_scale).round().clamp(-8.0, 7.0) as i8;
let q1 = (v1 * inv_scale).round().clamp(-8.0, 7.0) as i8;
let packed_byte = ((q0 as u8) & 0x0F) | (((q1 as u8) & 0x0F) << 4);
packed.push(packed_byte);
}
}
(packed, scales)
}
#[must_use]
pub fn dot_i4(a_i4: &[u8], scales: &[f32], b: &[f32], group_size: usize) -> f32 {
assert_eq!(a_i4.len() * 2, b.len(), "dot_i4 requires equal lengths");
assert_eq!(
b.len() % group_size,
0,
"b len must be multiple of group_size"
);
let mut sum = 0.0;
let mut i8_buf = vec![0i8; group_size];
for (group_idx, (b_chunk, &scale)) in b.chunks(group_size).zip(scales.iter()).enumerate() {
let a_chunk = &a_i4[group_idx * (group_size / 2)..(group_idx + 1) * (group_size / 2)];
for (i, &byte) in a_chunk.iter().enumerate() {
let q0 = ((byte << 4) as i8) >> 4;
let q1 = (byte as i8) >> 4;
i8_buf[i * 2] = q0;
i8_buf[i * 2 + 1] = q1;
}
sum += dot_i8(&i8_buf, b_chunk, scale);
}
sum
}
#[must_use]
pub fn quant_f32_to_f16(f32_data: &[f32]) -> Vec<u16> {
f32_data
.iter()
.map(|&v| half::f16::from_f32(v).to_bits())
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
const EPSILON: f32 = 1e-4;
fn approx_eq(a: f32, b: f32) -> bool {
(a - b).abs() < EPSILON
}
fn vec_approx_eq(a: &[f32], b: &[f32]) -> bool {
a.len() == b.len() && a.iter().zip(b).all(|(x, y)| approx_eq(*x, *y))
}
#[test]
fn test_dot_product() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![5.0, 6.0, 7.0, 8.0];
let result = dot(&a, &b);
assert!(approx_eq(result, 70.0)); }
#[test]
fn test_add() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
let result = add(&a, &b);
assert!(vec_approx_eq(&result, &[5.0, 7.0, 9.0]));
}
#[test]
fn test_sub() {
let a = vec![5.0, 7.0, 9.0];
let b = vec![1.0, 2.0, 3.0];
let result = sub(&a, &b);
assert!(vec_approx_eq(&result, &[4.0, 5.0, 6.0]));
}
#[test]
fn test_mul() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
let result = mul(&a, &b);
assert!(vec_approx_eq(&result, &[4.0, 10.0, 18.0]));
}
#[test]
fn test_scale() {
let a = vec![1.0, 2.0, 3.0];
let result = scale(&a, 2.0);
assert!(vec_approx_eq(&result, &[2.0, 4.0, 6.0]));
}
#[test]
fn test_sum() {
let a = vec![1.0, 2.0, 3.0, 4.0];
assert!(approx_eq(sum(&a), 10.0));
}
#[test]
fn test_mean() {
let a = vec![1.0, 2.0, 3.0, 4.0];
assert!(approx_eq(mean(&a), 2.5));
}
#[test]
fn test_mean_empty() {
let a: Vec<f32> = vec![];
assert!(approx_eq(mean(&a), 0.0));
}
#[test]
fn test_variance() {
let a = vec![1.0, 2.0, 3.0, 4.0, 5.0];
assert!(approx_eq(variance(&a), 2.0));
}
#[test]
fn test_std_dev() {
let a = vec![1.0, 2.0, 3.0, 4.0, 5.0];
assert!(approx_eq(std_dev(&a), 2.0_f32.sqrt()));
}
#[test]
fn test_max() {
let a = vec![1.0, 5.0, 3.0, 2.0];
assert!(approx_eq(max(&a), 5.0));
}
#[test]
fn test_min() {
let a = vec![1.0, 5.0, 3.0, 2.0];
assert!(approx_eq(min(&a), 1.0));
}
#[test]
fn test_argmax() {
let a = vec![1.0, 5.0, 3.0, 2.0];
assert_eq!(argmax(&a), 1);
}
#[test]
fn test_scale_inplace() {
let mut a = vec![1.0, 2.0, 3.0, 4.0];
scale_inplace(&mut a, 2.0);
assert!(vec_approx_eq(&a, &[2.0, 4.0, 6.0, 8.0]));
}
#[test]
fn test_scale_inplace_zero() {
let mut a = vec![1.0, 2.0, 3.0];
scale_inplace(&mut a, 0.0);
assert!(vec_approx_eq(&a, &[0.0, 0.0, 0.0]));
}
#[test]
fn test_axpy() {
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![10.0, 20.0, 30.0];
axpy(2.0, &x, &mut y);
assert!(vec_approx_eq(&y, &[12.0, 24.0, 36.0]));
}
#[test]
fn test_axpy_zero_scalar() {
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![10.0, 20.0, 30.0];
axpy(0.0, &x, &mut y);
assert!(vec_approx_eq(&y, &[10.0, 20.0, 30.0]));
}
#[test]
fn test_add_inplace() {
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![10.0, 20.0, 30.0];
add_inplace(&x, &mut y);
assert!(vec_approx_eq(&y, &[11.0, 22.0, 33.0]));
}
#[test]
fn test_broadcast_add_inplace() {
let mut matrix = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let vec = vec![10.0, 20.0, 30.0];
broadcast_add_inplace(&mut matrix, &vec, 2, 3);
assert!(vec_approx_eq(
&matrix,
&[11.0, 22.0, 33.0, 14.0, 25.0, 36.0]
));
}
#[test]
fn test_max_element() {
let a = vec![1.0, 5.0, 3.0, 2.0];
assert!(approx_eq(max_element(&a), 5.0));
}
#[test]
fn test_max_empty() {
let a: Vec<f32> = vec![];
assert_eq!(max(&a), f32::NEG_INFINITY);
}
#[test]
fn test_min_empty() {
let a: Vec<f32> = vec![];
assert_eq!(min(&a), f32::INFINITY);
}
#[test]
fn test_argmax_empty() {
let a: Vec<f32> = vec![];
assert_eq!(argmax(&a), 0);
}
#[test]
fn test_variance_empty() {
let a: Vec<f32> = vec![];
assert!(approx_eq(variance(&a), 0.0));
}
#[test]
fn test_dot_empty() {
let a: Vec<f32> = vec![];
let b: Vec<f32> = vec![];
assert!(approx_eq(dot(&a, &b), 0.0));
}
#[test]
fn test_sum_empty() {
let a: Vec<f32> = vec![];
assert!(approx_eq(sum(&a), 0.0));
}
#[test]
fn test_scale_empty() {
let a: Vec<f32> = vec![];
let result = scale(&a, 2.0);
assert!(result.is_empty());
}
#[test]
fn test_add_empty() {
let a: Vec<f32> = vec![];
let b: Vec<f32> = vec![];
let result = add(&a, &b);
assert!(result.is_empty());
}
#[test]
fn test_sub_empty() {
let a: Vec<f32> = vec![];
let b: Vec<f32> = vec![];
let result = sub(&a, &b);
assert!(result.is_empty());
}
#[test]
fn test_mul_empty() {
let a: Vec<f32> = vec![];
let b: Vec<f32> = vec![];
let result = mul(&a, &b);
assert!(result.is_empty());
}
#[test]
fn test_dequant_f16_row() {
let f32_vals = [1.0_f32, 2.0, 3.0, 4.0];
let f16_bits: Vec<u16> = f32_vals
.iter()
.map(|&v| half::f16::from_f32(v).to_bits())
.collect();
let mut out = vec![0.0_f32; 4];
dequant_f16_row(&f16_bits, &mut out);
for (a, &b) in out.iter().zip(f32_vals.iter()) {
assert!(approx_eq(*a, b));
}
}
#[test]
fn test_dot_f16() {
let a_f32 = [1.0_f32, 2.0, 3.0, 4.0];
let b = [5.0_f32, 6.0, 7.0, 8.0];
let a_f16: Vec<u16> = a_f32
.iter()
.map(|&v| half::f16::from_f32(v).to_bits())
.collect();
let mut buf = vec![0.0_f32; 4];
let result = dot_f16(&a_f16, &b, &mut buf);
assert!(approx_eq(result, 70.0));
}
#[test]
fn test_quant_f32_to_f16_roundtrip() {
let original = vec![0.0_f32, 1.0, -1.0, 0.5, 65504.0];
let f16_bits = quant_f32_to_f16(&original);
assert_eq!(f16_bits.len(), original.len());
let mut recovered = vec![0.0_f32; original.len()];
dequant_f16_row(&f16_bits, &mut recovered);
for (a, &b) in recovered.iter().zip(original.iter()) {
assert!(approx_eq(*a, b));
}
}
#[test]
fn test_quant_f32_to_f16_empty() {
let result = quant_f32_to_f16(&[]);
assert!(result.is_empty());
}
#[test]
fn pv_dot_fma_scalar_equivalence() {
use std::f32::consts::PI;
for len in [1, 7, 8, 15, 16, 31, 32, 63, 64, 128, 384, 1536] {
let a: Vec<f32> = (0..len).map(|i| (i as f32 * 0.01 * PI).sin()).collect();
let b: Vec<f32> = (0..len).map(|i| (i as f32 * 0.017 + 0.3).cos()).collect();
let scalar = dot_scalar(&a, &b);
let nalloc = dot_nalloc(&a, &b);
let diff = (scalar - nalloc).abs();
let tol = len as f32 * f32::EPSILON * scalar.abs().max(1.0);
assert!(
diff < tol,
"len={len}: scalar={scalar}, nalloc={nalloc}, diff={diff}, tol={tol}"
);
}
}
#[test]
fn pv_dot_empty_and_unit() {
assert_eq!(dot_nalloc(&[], &[]), 0.0);
assert_eq!(dot_nalloc(&[3.0], &[7.0]), 21.0);
}
#[test]
fn pv_dot_commutativity() {
let a: Vec<f32> = (0..384).map(|i| (i as f32 * 0.1).sin()).collect();
let b: Vec<f32> = (0..384).map(|i| (i as f32 * 0.2).cos()).collect();
let ab = dot_nalloc(&a, &b);
let ba = dot_nalloc(&b, &a);
assert!((ab - ba).abs() < 1e-6, "ab={ab}, ba={ba}");
}
#[test]
fn pv_dot_self_non_negative() {
let a: Vec<f32> = (0..384).map(|i| (i as f32 - 192.0) * 0.01).collect();
assert!(dot_nalloc(&a, &a) >= 0.0);
}
#[test]
fn pv_dot_nan_propagation() {
let a = [f32::NAN, 1.0];
let b = [1.0, 1.0];
assert!(dot_nalloc(&a, &b).is_nan());
}
fn softmax_standard(scores: &[f32]) -> Vec<f32> {
let max_s = scores.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let exps: Vec<f32> = scores.iter().map(|&s| (s - max_s).exp()).collect();
let sum: f32 = exps.iter().sum();
exps.iter().map(|&e| e / sum).collect()
}
#[test]
fn pv_softmax_online_matches_standard() {
for len in [1, 2, 6, 64, 384, 448, 1500] {
let scores: Vec<f32> = (0..len)
.map(|i| i as f32 * 0.1 - len as f32 * 0.05)
.collect();
let reference = softmax_standard(&scores);
let mut online = vec![0.0_f32; len];
softmax_online_inplace(&scores, &mut online);
for (i, (&r, &o)) in reference.iter().zip(online.iter()).enumerate() {
let diff = (r - o).abs();
assert!(
diff < 1e-5,
"len={len}, i={i}: ref={r}, online={o}, diff={diff}"
);
}
}
}
#[test]
fn pv_softmax_sum_to_one() {
for len in [1, 6, 64, 1500] {
let scores: Vec<f32> = (0..len).map(|i| i as f32 * 0.3 - 5.0).collect();
let mut weights = vec![0.0_f32; len];
softmax_online_inplace(&scores, &mut weights);
let sum: f32 = weights.iter().sum();
assert!((sum - 1.0).abs() < 1e-6, "len={len}: sum={sum}");
}
}
#[test]
fn pv_softmax_positivity() {
let scores = [-20.0_f32, -10.0, 0.0, 10.0, 20.0];
let mut weights = vec![0.0_f32; 5];
softmax_online_inplace(&scores, &mut weights);
for (i, &w) in weights.iter().enumerate() {
assert!(w > 0.0, "i={i}: weight={w} should be positive");
}
}
#[test]
fn pv_softmax_order_preservation() {
let scores = [1.0_f32, 3.0, 2.0, 5.0, 4.0];
let mut weights = vec![0.0_f32; 5];
softmax_online_inplace(&scores, &mut weights);
assert!(weights[3] > weights[1]); assert!(weights[1] > weights[2]); assert!(weights[2] > weights[0]); }
#[test]
fn pv_softmax_shift_invariance() {
let scores = [1.0_f32, 2.0, 3.0, 4.0];
let shifted: Vec<f32> = scores.iter().map(|&s| s + 1000.0).collect();
let mut w1 = vec![0.0_f32; 4];
let mut w2 = vec![0.0_f32; 4];
softmax_online_inplace(&scores, &mut w1);
softmax_online_inplace(&shifted, &mut w2);
for (i, (&a, &b)) in w1.iter().zip(w2.iter()).enumerate() {
assert!((a - b).abs() < 1e-6, "i={i}: w1={a}, w2={b}");
}
}
#[test]
fn pv_softmax_single_element() {
let mut w = [0.0_f32];
softmax_online_inplace(&[42.0], &mut w);
assert_eq!(w[0], 1.0);
}
#[test]
fn pv_i8q_roundtrip_accuracy() {
let row: Vec<f32> = (0..384).map(|i| (i as f32 * 0.01).sin() * 0.5).collect();
let (q, scale) = quant_f32_row_to_i8(&row);
for (i, (&orig, &qi)) in row.iter().zip(q.iter()).enumerate() {
let recovered = qi as f32 * scale;
let diff = (orig - recovered).abs();
assert!(
diff < scale + 1e-6,
"i={i}: orig={orig}, recovered={recovered}, diff={diff}"
);
}
}
#[test]
fn pv_i8q_zero_row() {
let row = vec![0.0_f32; 64];
let (q, scale) = quant_f32_row_to_i8(&row);
assert_eq!(scale, 0.0);
assert!(q.iter().all(|&v| v == 0));
}
#[test]
fn pv_i8q_range_bounded() {
let row: Vec<f32> = (0..384).map(|i| i as f32 * 0.1 - 19.2).collect();
let (q, _scale) = quant_f32_row_to_i8(&row);
for &v in &q {
assert!((-127..=127).contains(&v), "i8 value {v} out of range");
}
}
#[test]
fn pv_i8q_dot_scalar_equivalence() {
use std::f32::consts::PI;
for len in [1, 8, 16, 64, 384, 1536] {
let weights: Vec<f32> = (0..len)
.map(|i| (i as f32 * 0.01 * PI).sin() * 0.3)
.collect();
let input: Vec<f32> = (0..len).map(|i| (i as f32 * 0.017 + 0.3).cos()).collect();
let ref_dot: f32 = weights.iter().zip(input.iter()).map(|(w, x)| w * x).sum();
let (q, scale) = quant_f32_row_to_i8(&weights);
let i8_dot = dot_i8(&q, &input, scale);
let diff = (ref_dot - i8_dot).abs();
let tol = len as f32 * scale * 0.5 + 1e-4;
assert!(
diff < tol,
"len={len}: ref={ref_dot}, i8={i8_dot}, diff={diff}, tol={tol}"
);
}
}
#[test]
fn pv_i8q_dot_empty() {
assert_eq!(dot_i8(&[], &[], 1.0), 0.0);
}
#[test]
fn pv_i8q_scale_positive() {
let row: Vec<f32> = (0..64).map(|i| (i as f32 - 32.0) * 0.1).collect();
let (_q, scale) = quant_f32_row_to_i8(&row);
assert!(scale > 0.0, "scale should be positive for non-zero row");
}
}