use spiral_rs::{arith::*, params::*, poly::*};
use crate::server::ToU64;
use super::server::ToM512;
#[cfg(feature = "explicit_avx512")]
use std::arch::x86_64::*;
#[cfg(feature = "rayon")]
use rayon::prelude::*;
#[cfg(feature = "rayon")]
#[derive(Copy, Clone)]
struct SendPtr<T>(*const T);
#[cfg(feature = "rayon")]
unsafe impl<T> Send for SendPtr<T> {}
#[cfg(feature = "rayon")]
unsafe impl<T> Sync for SendPtr<T> {}
#[cfg(feature = "rayon")]
impl<T> SendPtr<T> {
#[inline(always)]
fn get(self) -> *const T {
self.0
}
}
#[inline(always)]
fn split_a<const K: usize>(a: &[u64]) -> [&[u64]; K] {
assert!(K > 0, "split_a requires K > 0");
assert_eq!(
a.len() % K,
0,
"split_a: a.len() ({}) must be a multiple of K ({})",
a.len(),
K
);
let mut out: [&[u64]; K] = [&[]; K];
for (slot, chunk) in out.iter_mut().zip(a.chunks_exact(a.len() / K)) {
*slot = chunk;
}
out
}
#[inline(always)]
fn writeback(params: &Params, c_cell: &mut u64, sum_lo: u64, sum_hi: u64) {
let lo = barrett_coeff_u64(params, sum_lo, 0);
let hi = barrett_coeff_u64(params, sum_hi, 1);
let res = params.crt_compose_2(lo, hi);
*c_cell = barrett_u64(params, *c_cell + res);
}
#[cfg(feature = "explicit_avx512")]
#[inline(always)]
unsafe fn writeback_avx512(
params: &Params,
c_cell: &mut u64,
sum_lo: __m512i,
sum_hi: __m512i,
) {
let mut vl = [0u64; 8];
let mut vh = [0u64; 8];
_mm512_storeu_si512(vl.as_mut_ptr() as *mut _, sum_lo);
_mm512_storeu_si512(vh.as_mut_ptr() as *mut _, sum_hi);
writeback(params, c_cell, vl.iter().sum(), vh.iter().sum());
}
pub fn fast_batched_dot_product<const K: usize, T: Copy>(
params: &Params,
c: &mut [u64],
a: &[u64],
a_elems: usize,
b_t: &[T], b_rows: usize,
b_cols: usize,
) where
*const T: ToM512 + ToU64,
{
#[cfg(feature = "explicit_avx512")]
fast_batched_dot_product_explicit_avx512::<K, T>(params, c, a, a_elems, b_t, b_rows, b_cols);
#[cfg(not(feature = "explicit_avx512"))]
fast_batched_dot_product_implicit::<K, T>(params, c, a, a_elems, b_t, b_rows, b_cols);
}
#[cfg(not(feature = "explicit_avx512"))]
pub fn fast_batched_dot_product_explicit_avx512<const K: usize, T: Copy>(
_params: &Params,
_c: &mut [u64],
_a: &[u64],
_a_elems: usize,
_b_t: &[T], _b_rows: usize,
_b_cols: usize,
) where
*const T: ToM512 + ToU64,
{
panic!("explicit_avx512 not enabled");
}
#[cfg(feature = "explicit_avx512")]
pub fn fast_batched_dot_product_explicit_avx512<const K: usize, T: Copy>(
params: &Params,
c: &mut [u64],
a: &[u64],
a_elems: usize,
b_t: &[T], b_rows: usize,
b_cols: usize,
) where
*const T: ToM512 + ToU64,
{
assert_eq!(a_elems, b_rows);
let simd_width = 8;
let chunk_size = (8192 / K.next_power_of_two()).min(a_elems / simd_width);
let num_chunks = (a_elems / simd_width) / chunk_size;
#[cfg(feature = "rayon")]
if K == 1 {
let b_send = SendPtr(b_t.as_ptr());
let a_send = SendPtr(a.as_ptr());
let a_len = a.len();
let num_threads = rayon::current_num_threads();
let cols_per_chunk = ((b_cols + num_threads - 1) / num_threads).max(1);
c.par_chunks_mut(cols_per_chunk)
.enumerate()
.for_each(move |(chunk_idx, c_chunk)| {
let j_start = chunk_idx * cols_per_chunk;
unsafe {
let b_ptr = b_send.get();
let a_slice = std::slice::from_raw_parts(a_send.get(), a_len);
let a_slcs = split_a::<K>(a_slice);
for k_outer in 0..num_chunks {
for (j_local, c_cell) in c_chunk.iter_mut().enumerate() {
let j = j_start + j_local;
debug_assert!(
j < b_cols,
"par_chunks_mut partition invariant broken: j={j} >= b_cols={b_cols}"
);
let mut total_sum_lo = [_mm512_setzero_si512(); K];
let mut total_sum_hi = [_mm512_setzero_si512(); K];
let mut tmp = [_mm512_setzero_si512(); K];
for k_inner in 0..chunk_size {
let k = simd_width * (k_outer * chunk_size + k_inner);
let b_val_simd = b_ptr.add(j * b_rows + k).to_m512();
for batch in 0..K {
tmp[batch] = _mm512_load_si512(
a_slcs[batch].as_ptr().add(k) as *const _,
);
}
for batch in 0..K {
let a_val_lo = tmp[batch];
let a_val_hi = _mm512_srli_epi64(tmp[batch], 32);
total_sum_lo[batch] = _mm512_add_epi64(
total_sum_lo[batch],
_mm512_mul_epu32(a_val_lo, b_val_simd),
);
total_sum_hi[batch] = _mm512_add_epi64(
total_sum_hi[batch],
_mm512_mul_epu32(a_val_hi, b_val_simd),
);
}
}
writeback_avx512(params, c_cell, total_sum_lo[0], total_sum_hi[0]);
}
}
}
});
return;
}
unsafe {
let a_slcs = split_a::<K>(a);
let b_ptr = b_t.as_ptr();
for k_outer in 0..num_chunks {
for j in 0..b_cols {
let mut total_sum_lo = [_mm512_setzero_si512(); K];
let mut total_sum_hi = [_mm512_setzero_si512(); K];
let mut tmp = [_mm512_setzero_si512(); K];
for k_inner in 0..chunk_size {
let k = simd_width * (k_outer * chunk_size + k_inner);
let b_val_simd = b_ptr.add(j * b_rows + k).to_m512();
for batch in 0..K {
tmp[batch] = _mm512_load_si512(a_slcs[batch].as_ptr().add(k) as *const _);
}
for batch in 0..K {
let a_val_lo = tmp[batch];
let a_val_hi = _mm512_srli_epi64(tmp[batch], 32);
total_sum_lo[batch] = _mm512_add_epi64(
total_sum_lo[batch],
_mm512_mul_epu32(a_val_lo, b_val_simd),
);
total_sum_hi[batch] = _mm512_add_epi64(
total_sum_hi[batch],
_mm512_mul_epu32(a_val_hi, b_val_simd),
);
}
}
for (batch, c_row) in c.chunks_exact_mut(c.len() / K).enumerate() {
writeback_avx512(params, &mut c_row[j], total_sum_lo[batch], total_sum_hi[batch]);
}
}
}
}
}
pub fn fast_batched_dot_product_implicit<const K: usize, T: Copy>(
params: &Params,
c: &mut [u64],
a: &[u64],
a_elems: usize,
b_t: &[T], b_rows: usize,
b_cols: usize,
) where
*const T: ToM512 + ToU64,
{
assert_eq!(a_elems, b_rows);
assert_eq!(K, 1);
let simd_width = 1;
let chunk_size = (65536 / K.next_power_of_two()).min(a_elems / simd_width);
let num_chunks = (a_elems / simd_width) / chunk_size;
#[cfg(feature = "rayon")]
{
let b_send = SendPtr(b_t.as_ptr());
let a_send = SendPtr(a.as_ptr());
let a_len = a.len();
let num_threads = rayon::current_num_threads();
let cols_per_chunk = ((b_cols + num_threads - 1) / num_threads).max(1);
c.par_chunks_mut(cols_per_chunk)
.enumerate()
.for_each(move |(chunk_idx, c_chunk)| {
let j_start = chunk_idx * cols_per_chunk;
unsafe {
let b_ptr = b_send.get();
let a_slice = std::slice::from_raw_parts(a_send.get(), a_len);
let a_slcs = split_a::<K>(a_slice);
for k_outer in 0..num_chunks {
for (j_local, c_cell) in c_chunk.iter_mut().enumerate() {
let j = j_start + j_local;
debug_assert!(
j < b_cols,
"par_chunks_mut partition invariant broken: j={j} >= b_cols={b_cols}"
);
let mut total_sum_lo = [0u64; K];
let mut total_sum_hi = [0u64; K];
let mut tmp = [0u64; K];
for k_inner in 0..chunk_size {
let k = simd_width * (k_outer * chunk_size + k_inner);
let b_val_simd = (b_ptr.add(j * b_rows + k)).to_u64();
for batch in 0..K {
tmp[batch] = *(a_slcs[batch].as_ptr().add(k));
}
for batch in 0..K {
let a_val_lo = (tmp[batch] as u32) as u64;
let a_val_hi = ((tmp[batch] >> 32) as u32) as u64;
total_sum_lo[batch] += a_val_lo * b_val_simd;
total_sum_hi[batch] += a_val_hi * b_val_simd;
}
}
writeback(params, c_cell, total_sum_lo[0], total_sum_hi[0]);
}
}
}
});
return;
}
#[cfg(not(feature = "rayon"))]
unsafe {
let a_slcs = split_a::<K>(a);
let b_ptr = b_t.as_ptr();
for k_outer in 0..num_chunks {
for j in 0..b_cols {
let mut total_sum_lo = [0u64; K];
let mut total_sum_hi = [0u64; K];
let mut tmp = [0u64; K];
for k_inner in 0..chunk_size {
let k = simd_width * (k_outer * chunk_size + k_inner);
let b_val_simd = (b_ptr.add(j * b_rows + k) as *const T).to_u64();
for batch in 0..K {
tmp[batch] = *(a_slcs[batch].as_ptr().add(k));
}
for batch in 0..K {
let a_val_lo = (tmp[batch] as u32) as u64;
let a_val_hi = ((tmp[batch] >> 32) as u32) as u64;
total_sum_lo[batch] += a_val_lo * b_val_simd;
total_sum_hi[batch] += a_val_hi * b_val_simd;
}
}
for (batch, c_row) in c.chunks_exact_mut(c.len() / K).enumerate() {
writeback(params, &mut c_row[j], total_sum_lo[batch], total_sum_hi[batch]);
}
}
}
}
}
pub fn scalar_multiply_avx(res: &mut PolyMatrixNTT, a: &PolyMatrixNTT, b: &PolyMatrixNTT) {
assert_eq!(a.rows, 1);
assert_eq!(a.cols, 1);
let params = res.params;
let pol2 = a.get_poly(0, 0);
for i in 0..b.rows {
for j in 0..b.cols {
let res_poly = res.get_poly_mut(i, j);
let pol1 = b.get_poly(i, j);
crate::packing::multiply_poly_avx(params, res_poly, pol1, pol2);
}
}
}
pub fn multiply_matrices_raw_not_transposed<T>(
params: &Params,
a: &[u64],
a_rows: usize,
a_cols: usize,
b: &[T], b_rows: usize,
b_cols: usize,
) -> Vec<u64>
where
T: ToU64 + Copy,
{
assert_eq!(a_cols, b_rows);
let mut result = vec![0u128; a_rows * b_cols];
for i in 0..a_rows {
for k in 0..a_cols {
for j in 0..b_cols {
let a_idx = i * a_cols + k;
let b_idx = k * b_cols + j;
let res_idx = i * b_cols + j;
unsafe {
let a_val = *a.get_unchecked(a_idx);
let b_val = (*b.get_unchecked(b_idx)).to_u64();
let prod = a_val as u128 * b_val as u128;
result[res_idx] += prod;
}
}
}
}
let mut result_u64 = vec![0u64; a_rows * b_cols];
for i in 0..result.len() {
result_u64[i] = barrett_reduction_u128(params, result[i]);
}
result_u64
}
#[cfg(test)]
mod test {
use std::time::Instant;
use log::debug;
use spiral_rs::aligned_memory::AlignedMemory64;
use spiral_rs::poly::*;
use super::super::util::test_params;
use super::*;
use crate::{transpose::*, util::*};
use test_log::test;
fn test_fast_batched_dot_product(use_explicit: bool) {
let params = test_params();
const A_ROWS: usize = 1;
let a_cols = 65536;
let b_rows = a_cols;
let b_cols = 32768;
let a = PolyMatrixRaw::random(¶ms, A_ROWS, a_cols);
let mut b = AlignedMemory64::new(b_rows * b_cols);
let mut c = AlignedMemory64::new(A_ROWS * b_cols);
let trials = 10;
let mut sum = 0u64;
let mut sum_time = 0;
for _ in 0..trials {
for i in 0..b.len() {
b[i] = fastrand::u64(..);
}
let b_u16_slc =
unsafe { std::slice::from_raw_parts(b.as_ptr() as *const u16, b.len() * 4) };
let now = Instant::now();
if use_explicit {
fast_batched_dot_product_explicit_avx512::<A_ROWS, _>(
¶ms,
c.as_mut_slice(),
a.as_slice(),
a_cols,
b_u16_slc,
b_rows,
b_cols,
);
} else {
fast_batched_dot_product_implicit::<A_ROWS, _>(
¶ms,
c.as_mut_slice(),
a.as_slice(),
a_cols,
b_u16_slc,
b_rows,
b_cols,
);
}
sum_time += now.elapsed().as_micros();
sum += c.as_slice()[fastrand::usize(..c.len())];
}
debug!(
"fast_matmul_avx512 in {} us ({}: {} x {})",
sum_time, trials, a_cols, b_cols
);
debug!("");
debug!("{}", sum);
}
#[cfg(feature = "explicit_avx512")]
#[test]
#[ignore]
fn test_fast_batched_dot_product_explicit() {
test_fast_batched_dot_product(true);
}
#[test]
#[ignore]
fn test_fast_batched_dot_product_implicit() {
test_fast_batched_dot_product(false);
}
#[test]
fn test_negacyclic_mul_db_col() {
let params = test_params();
let pol_a = PolyMatrixRaw::random(¶ms, 1, 1);
let pol_b: PolyMatrixRaw<'_> = PolyMatrixRaw::random(¶ms, 1, 1);
let a = pol_a.get_poly(0, 0);
let b = pol_b.get_poly(0, 0);
let negacylic_a = negacyclic_matrix(&a, params.modulus);
let negacyclic_a_t = transpose_generic(&negacylic_a, params.poly_len, params.poly_len);
assert_eq!(negacylic_a[0], a[0]);
assert_eq!(
negacylic_a[params.poly_len],
(params.modulus - a[params.poly_len - 1]) % params.modulus
);
let prod = multiply_matrices_raw_not_transposed(
¶ms,
b,
1,
params.poly_len,
&negacyclic_a_t,
params.poly_len,
params.poly_len,
);
let transformed_a = negacyclic_perm(a, 0, params.modulus);
let mut pol_a_transformed = PolyMatrixRaw::zero(¶ms, 1, 1);
pol_a_transformed
.data
.as_mut_slice()
.copy_from_slice(&transformed_a);
let pol_c = (&pol_a_transformed.ntt() * &pol_b.ntt()).raw();
let c = pol_c.get_poly(0, 0);
for i in 0..params.poly_len {
assert_eq!(prod[i] % params.modulus, c[i] % params.modulus, "i = {}", i);
}
}
#[test]
fn test_negacyclic_mul() {
let params = test_params();
let pol_a = PolyMatrixRaw::random(¶ms, 1, 1);
let pol_b = PolyMatrixRaw::random(¶ms, 1, 1);
let a = pol_a.get_poly(0, 0);
let b = pol_b.get_poly(0, 0);
let negacylic_a = negacyclic_matrix(&a, params.modulus);
assert_eq!(negacylic_a[0], a[0]);
assert_eq!(
negacylic_a[params.poly_len],
(params.modulus - a[params.poly_len - 1]) % params.modulus
);
let prod = multiply_matrices_raw_not_transposed(
¶ms,
b,
1,
params.poly_len,
&negacylic_a,
params.poly_len,
params.poly_len,
);
let pol_c = (&pol_a.ntt() * &pol_b.ntt()).raw();
let c = pol_c.get_poly(0, 0);
for i in 0..params.poly_len {
assert_eq!(prod[i] % params.modulus, c[i] % params.modulus, "i = {}", i);
}
}
fn reference_dot_product_transposed_u16(
params: &Params,
c: &mut [u64],
a: &[u64],
a_elems: usize,
b_t: &[u16],
b_rows: usize,
b_cols: usize,
) {
assert_eq!(a_elems, b_rows);
for j in 0..b_cols {
let mut sum_lo = 0u64;
let mut sum_hi = 0u64;
for k in 0..a_elems {
let a_val = a[k];
let a_lo = (a_val as u32) as u64;
let a_hi = ((a_val >> 32) as u32) as u64;
let b_val = b_t[j * b_rows + k] as u64;
sum_lo += a_lo * b_val;
sum_hi += a_hi * b_val;
}
let (lo, hi) = (
barrett_coeff_u64(params, sum_lo, 0),
barrett_coeff_u64(params, sum_hi, 1),
);
let res = params.crt_compose_2(lo, hi);
c[j] = barrett_u64(params, c[j] + res);
}
}
fn random_bounded_aligned(len: usize, bound: u64) -> AlignedMemory64 {
let mut mem = AlignedMemory64::new(len);
for i in 0..len {
mem[i] = fastrand::u64(..) % bound;
}
mem
}
fn random_u16_vec(len: usize) -> Vec<u16> {
(0..len).map(|_| fastrand::u16(..)).collect()
}
#[test]
fn test_implicit_matches_reference() {
let params = test_params();
let a_elems = 65536;
let b_cols = 1024;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_impl = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_impl.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_impl.as_slice()[j], c_ref[j],
"mismatch at column {j}: implicit={} ref={}",
c_impl.as_slice()[j], c_ref[j]
);
}
}
#[test]
fn test_implicit_accumulates_into_c() {
let params = test_params();
let a_elems = 65536;
let b_cols = 128;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_once = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_once.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_twice = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_twice.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_twice.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
let expected = barrett_u64(
¶ms,
c_once.as_slice()[j] + c_once.as_slice()[j],
);
assert_eq!(
c_twice.as_slice()[j], expected,
"accumulation mismatch at column {j}"
);
}
}
#[test]
fn test_implicit_single_column() {
let params = test_params();
let a_elems = 65536;
let b_cols = 1;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems);
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_impl = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_impl.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
assert_eq!(c_impl.as_slice()[0], c_ref[0], "single column mismatch");
}
#[test]
fn test_implicit_zero_input() {
let params = test_params();
let a_elems = 65536;
let b_cols = 64;
let a = AlignedMemory64::new(a_elems); let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(c.as_slice()[j], 0, "expected zero at column {j}");
}
}
#[test]
fn test_implicit_modulus_boundary_values() {
let params = test_params();
let a_elems = 65536;
let b_cols = 32;
let mut a = AlignedMemory64::new(a_elems);
for i in 0..a_elems {
a[i] = params.modulus - 1;
}
let b_t_u16: Vec<u16> = vec![u16::MAX; a_elems * b_cols];
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_impl = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_impl.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_impl.as_slice()[j], c_ref[j],
"modulus-boundary mismatch at column {j}"
);
}
}
#[test]
fn test_implicit_small_dimension() {
let params = test_params();
let a_elems = 65536;
let b_cols = 2;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_impl = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_impl.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_impl.as_slice()[j], c_ref[j],
"small dimension mismatch at column {j}"
);
}
}
#[cfg(feature = "rayon")]
mod rayon_tests {
use super::*;
use test_log::test;
fn assert_rayon_matches_reference(a_elems: usize, b_cols: usize) {
let params = test_params();
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_rayon = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_rayon.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_rayon.as_slice()[j], c_ref[j],
"rayon mismatch at col {j} (a_elems={a_elems}, b_cols={b_cols})"
);
}
}
#[test]
fn test_rayon_implicit_standard() {
assert_rayon_matches_reference(65536, 1024);
}
#[test]
fn test_rayon_implicit_single_column() {
assert_rayon_matches_reference(65536, 1);
}
#[test]
fn test_rayon_implicit_two_columns() {
assert_rayon_matches_reference(65536, 2);
}
#[test]
fn test_rayon_implicit_cols_less_than_threads() {
let num_threads = rayon::current_num_threads();
if num_threads > 3 {
assert_rayon_matches_reference(65536, 3);
}
}
#[test]
fn test_rayon_implicit_non_power_of_two_cols() {
assert_rayon_matches_reference(65536, 100);
}
#[test]
fn test_rayon_implicit_large() {
assert_rayon_matches_reference(65536, 32768);
}
#[test]
fn test_rayon_implicit_b_cols_zero() {
let params = test_params();
let a_elems = 65536;
let b_cols = 0;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16: Vec<u16> = Vec::new();
let mut c = AlignedMemory64::new(b_cols.max(1));
let c_before = c.as_slice()[0];
fast_batched_dot_product_implicit::<1, _>(
¶ms,
&mut c.as_mut_slice()[..b_cols],
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
assert_eq!(c.as_slice()[0], c_before);
}
#[test]
fn test_rayon_implicit_b_cols_equal_num_threads() {
let num_threads = rayon::current_num_threads();
assert_rayon_matches_reference(65536, num_threads);
}
#[test]
fn test_rayon_implicit_accumulates() {
let params = test_params();
let a_elems = 65536;
let b_cols = 256;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_once = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_once.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_twice = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_twice.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_twice.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
let expected = barrett_u64(
¶ms,
c_once.as_slice()[j] + c_once.as_slice()[j],
);
assert_eq!(
c_twice.as_slice()[j], expected,
"rayon accumulation mismatch at column {j}"
);
}
}
#[test]
fn test_rayon_implicit_zero_a() {
let params = test_params();
let a_elems = 65536;
let b_cols = 64;
let a = AlignedMemory64::new(a_elems);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(c.as_slice()[j], 0, "expected zero at column {j}");
}
}
#[test]
fn test_rayon_implicit_modulus_boundary() {
let params = test_params();
let a_elems = 65536;
let b_cols = 32;
let mut a = AlignedMemory64::new(a_elems);
for i in 0..a_elems {
a[i] = params.modulus - 1;
}
let b_t_u16: Vec<u16> = vec![u16::MAX; a_elems * b_cols];
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_rayon = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_rayon.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_rayon.as_slice()[j], c_ref[j],
"rayon modulus-boundary mismatch at column {j}"
);
}
}
#[test]
fn test_rayon_deterministic() {
let params = test_params();
let a_elems = 65536;
let b_cols = 512;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c1 = AlignedMemory64::new(b_cols);
let mut c2 = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c1.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c2.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
assert_eq!(c1.as_slice(), c2.as_slice(), "rayon results not deterministic");
}
#[test]
fn test_rayon_concurrent_invocations() {
use std::sync::Arc;
use std::thread;
let params = Arc::new(test_params());
let a_elems = 65536;
let b_cols = 1024;
let num_drivers = 8;
let inputs: Vec<_> = (0..num_drivers)
.map(|seed| {
fastrand::seed(seed as u64);
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
(a.as_slice().to_vec(), b_t_u16, c_ref)
})
.collect();
let handles: Vec<_> = inputs
.into_iter()
.enumerate()
.map(|(idx, (a_vec, b_t_u16, c_ref))| {
let params = Arc::clone(¶ms);
thread::spawn(move || {
let mut a = AlignedMemory64::new(a_vec.len());
a.as_mut_slice().copy_from_slice(&a_vec);
let mut c = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c.as_slice()[j], c_ref[j],
"driver {idx}: concurrent mismatch at col {j}"
);
}
})
})
.collect();
for h in handles {
h.join().expect("concurrent kernel driver panicked");
}
}
}
#[cfg(feature = "explicit_avx512")]
mod avx512_tests {
use super::*;
use test_log::test;
fn assert_avx512_matches_reference(a_elems: usize, b_cols: usize) {
let params = test_params();
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_avx = AlignedMemory64::new(b_cols);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c_avx.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_avx.as_slice()[j], c_ref[j],
"AVX-512 vs reference mismatch at col {j} \
(a_elems={a_elems}, b_cols={b_cols})"
);
}
}
fn assert_avx512_matches_implicit(a_elems: usize, b_cols: usize) {
let params = test_params();
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_avx = AlignedMemory64::new(b_cols);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c_avx.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_scalar = AlignedMemory64::new(b_cols);
fast_batched_dot_product_implicit::<1, _>(
¶ms,
c_scalar.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
assert_eq!(
c_avx.as_slice(),
c_scalar.as_slice(),
"AVX-512 and scalar kernels disagree (a_elems={a_elems}, b_cols={b_cols})"
);
}
#[test]
fn test_avx512_standard() {
assert_avx512_matches_reference(65536, 1024);
}
#[test]
fn test_avx512_single_column() {
assert_avx512_matches_reference(65536, 1);
}
#[test]
fn test_avx512_two_columns() {
assert_avx512_matches_reference(65536, 2);
}
#[test]
fn test_avx512_non_power_of_two_cols() {
assert_avx512_matches_reference(65536, 100);
}
#[test]
fn test_avx512_large() {
assert_avx512_matches_reference(65536, 32768);
}
#[test]
fn test_avx512_b_cols_zero() {
let params = test_params();
let a_elems = 65536;
let b_cols = 0;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16: Vec<u16> = Vec::new();
let mut c = AlignedMemory64::new(b_cols.max(1));
let c_before = c.as_slice()[0];
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
&mut c.as_mut_slice()[..b_cols],
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
assert_eq!(c.as_slice()[0], c_before);
}
#[test]
fn test_avx512_accumulates() {
let params = test_params();
let a_elems = 65536;
let b_cols = 256;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_once = AlignedMemory64::new(b_cols);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c_once.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_twice = AlignedMemory64::new(b_cols);
for _ in 0..2 {
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c_twice.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
}
for j in 0..b_cols {
let expected = barrett_u64(
¶ms,
c_once.as_slice()[j] + c_once.as_slice()[j],
);
assert_eq!(
c_twice.as_slice()[j], expected,
"AVX-512 accumulation mismatch at column {j}"
);
}
}
#[test]
fn test_avx512_zero_a() {
let params = test_params();
let a_elems = 65536;
let b_cols = 64;
let a = AlignedMemory64::new(a_elems);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c = AlignedMemory64::new(b_cols);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(c.as_slice()[j], 0, "expected zero at column {j}");
}
}
#[test]
fn test_avx512_modulus_boundary() {
let params = test_params();
let a_elems = 65536;
let b_cols = 32;
let mut a = AlignedMemory64::new(a_elems);
for i in 0..a_elems {
a[i] = params.modulus - 1;
}
let b_t_u16: Vec<u16> = vec![u16::MAX; a_elems * b_cols];
let mut c_ref = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref,
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_avx = AlignedMemory64::new(b_cols);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c_avx.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_avx.as_slice()[j], c_ref[j],
"AVX-512 modulus-boundary mismatch at column {j}"
);
}
}
#[test]
fn test_avx512_deterministic() {
let params = test_params();
let a_elems = 65536;
let b_cols = 512;
let a = random_bounded_aligned(a_elems, params.modulus);
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c1 = AlignedMemory64::new(b_cols);
let mut c2 = AlignedMemory64::new(b_cols);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c1.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
fast_batched_dot_product_explicit_avx512::<1, _>(
¶ms,
c2.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
assert_eq!(
c1.as_slice(),
c2.as_slice(),
"AVX-512 results not deterministic"
);
}
#[test]
fn test_avx512_matches_scalar_standard() {
assert_avx512_matches_implicit(65536, 1024);
}
#[test]
fn test_avx512_matches_scalar_single_column() {
assert_avx512_matches_implicit(65536, 1);
}
#[test]
fn test_avx512_matches_scalar_non_power_of_two() {
assert_avx512_matches_implicit(65536, 100);
}
#[test]
fn test_avx512_k2_sequential() {
let params = test_params();
let a_elems = 65536;
let b_cols = 256;
const K: usize = 2;
let a0 = random_bounded_aligned(a_elems, params.modulus);
let a1 = random_bounded_aligned(a_elems, params.modulus);
let mut a = AlignedMemory64::new(K * a_elems);
a.as_mut_slice()[..a_elems].copy_from_slice(a0.as_slice());
a.as_mut_slice()[a_elems..].copy_from_slice(a1.as_slice());
let b_t_u16 = random_u16_vec(a_elems * b_cols);
let mut c_ref0 = vec![0u64; b_cols];
let mut c_ref1 = vec![0u64; b_cols];
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref0,
a0.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
reference_dot_product_transposed_u16(
¶ms,
&mut c_ref1,
a1.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
let mut c_avx = AlignedMemory64::new(K * b_cols);
fast_batched_dot_product_explicit_avx512::<K, _>(
¶ms,
c_avx.as_mut_slice(),
a.as_slice(),
a_elems,
&b_t_u16,
a_elems,
b_cols,
);
for j in 0..b_cols {
assert_eq!(
c_avx.as_slice()[j], c_ref0[j],
"AVX-512 K=2 batch 0 mismatch at column {j}"
);
assert_eq!(
c_avx.as_slice()[b_cols + j], c_ref1[j],
"AVX-512 K=2 batch 1 mismatch at column {j}"
);
}
}
}
}