#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
#[cfg(target_arch = "x86_64")]
use crate::ntt::ZETAS;
use crate::params::N;
#[cfg(any(target_arch = "x86_64", test))]
use crate::params::Q;
#[cfg(target_arch = "x86_64")]
const QINV32: i32 = crate::params::QINV as i32;
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn montgomery_mul_avx2(zeta: __m256i, y: __m256i) -> __m256i {
let q_v = _mm256_set1_epi32(Q);
let qinv_v = _mm256_set1_epi32(QINV32);
let a_even = _mm256_mul_epi32(zeta, y);
let a_even_lo = _mm256_mul_epi32(
a_even, qinv_v, );
let t_even_lo_32 = a_even_lo; let tq_even = _mm256_mul_epi32(t_even_lo_32, q_v); let r_even_64 = _mm256_sub_epi64(a_even, tq_even);
let r_even_32 = _mm256_srli_epi64::<32>(r_even_64);
let zeta_odd = _mm256_srli_epi64::<32>(zeta);
let y_odd = _mm256_srli_epi64::<32>(y);
let a_odd = _mm256_mul_epi32(zeta_odd, y_odd);
let a_odd_lo = _mm256_mul_epi32(a_odd, qinv_v);
let tq_odd = _mm256_mul_epi32(a_odd_lo, q_v);
let r_odd_64 = _mm256_sub_epi64(a_odd, tq_odd);
let r_odd_hi = _mm256_and_si256(r_odd_64, _mm256_set1_epi64x(-4294967296i64));
_mm256_or_si256(r_even_32, r_odd_hi)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn butterfly_avx2(a: &mut [i32; N], j: usize, len: usize, zeta: i32) {
let zeta_v = _mm256_set1_epi32(zeta);
let x = _mm256_loadu_si256(a.as_ptr().add(j) as *const __m256i);
let y = _mm256_loadu_si256(a.as_ptr().add(j + len) as *const __m256i);
let t = montgomery_mul_avx2(zeta_v, y);
_mm256_storeu_si256(
a.as_mut_ptr().add(j) as *mut __m256i,
_mm256_add_epi32(x, t),
);
_mm256_storeu_si256(
a.as_mut_ptr().add(j + len) as *mut __m256i,
_mm256_sub_epi32(x, t),
);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn inv_butterfly_avx2(a: &mut [i32; N], j: usize, len: usize, zeta_neg: i32) {
let zeta_v = _mm256_set1_epi32(zeta_neg);
let x = _mm256_loadu_si256(a.as_ptr().add(j) as *const __m256i);
let y = _mm256_loadu_si256(a.as_ptr().add(j + len) as *const __m256i);
let sum = _mm256_add_epi32(x, y);
let diff = _mm256_sub_epi32(x, y);
let reduced = montgomery_mul_avx2(zeta_v, diff);
_mm256_storeu_si256(a.as_mut_ptr().add(j) as *mut __m256i, sum);
_mm256_storeu_si256(a.as_mut_ptr().add(j + len) as *mut __m256i, reduced);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub unsafe fn ntt_avx2(a: &mut [i32; N]) {
let mut k: usize = 0;
let mut len = 128;
while len > 0 {
let mut start = 0;
while start < N {
k += 1;
let zeta = ZETAS[k];
if len >= 8 {
let mut j = start;
while j + 8 <= start + len {
butterfly_avx2(a, j, len, zeta);
j += 8;
}
} else {
for j in start..start + len {
let t = crate::reduce::montgomery_reduce(zeta as i64 * a[j + len] as i64);
a[j + len] = a[j] - t;
a[j] += t;
}
}
start += 2 * len;
}
len >>= 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub unsafe fn invntt_avx2(a: &mut [i32; N]) {
let f: i32 = 41978; let mut k: usize = 256;
let mut len = 1;
while len < N {
let mut start = 0;
while start < N {
k -= 1;
let zeta = -ZETAS[k];
if len >= 8 {
let mut j = start;
while j + 8 <= start + len {
inv_butterfly_avx2(a, j, len, zeta);
j += 8;
}
} else {
for j in start..start + len {
let t = a[j];
a[j] = t + a[j + len];
a[j + len] = t - a[j + len];
a[j + len] = crate::reduce::montgomery_reduce(zeta as i64 * a[j + len] as i64);
}
}
start += 2 * len;
}
len <<= 1;
}
let f_v = _mm256_set1_epi32(f);
let mut j = 0;
while j + 8 <= N {
let v = _mm256_loadu_si256(a.as_ptr().add(j) as *const __m256i);
let scaled = montgomery_mul_avx2(f_v, v);
_mm256_storeu_si256(a.as_mut_ptr().add(j) as *mut __m256i, scaled);
j += 8;
}
}
pub fn ntt_simd(a: &mut [i32; N]) {
#[cfg(all(target_arch = "x86_64", feature = "std"))]
{
if is_x86_feature_detected!("avx2") {
unsafe {
ntt_avx2(a);
}
return;
}
}
#[cfg(all(target_arch = "x86_64", not(feature = "std"), target_feature = "avx2"))]
{
unsafe {
ntt_avx2(a);
}
return;
}
#[allow(unreachable_code)]
crate::ntt::ntt(a);
}
pub fn invntt_simd(a: &mut [i32; N]) {
#[cfg(all(target_arch = "x86_64", feature = "std"))]
{
if is_x86_feature_detected!("avx2") {
unsafe {
invntt_avx2(a);
}
return;
}
}
#[cfg(all(target_arch = "x86_64", not(feature = "std"), target_feature = "avx2"))]
{
unsafe {
invntt_avx2(a);
}
return;
}
#[allow(unreachable_code)]
crate::ntt::invntt_tomont(a);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_ntt_simd_matches_scalar() {
let mut a_scalar = [0i32; N];
let mut a_simd = [0i32; N];
for i in 0..N {
let v = (i as i32 * 37 + 11) % Q;
a_scalar[i] = v;
a_simd[i] = v;
}
crate::ntt::ntt(&mut a_scalar);
ntt_simd(&mut a_simd);
assert_eq!(a_scalar, a_simd, "SIMD NTT diverged from scalar");
}
#[test]
fn test_invntt_simd_matches_scalar() {
let mut a_scalar = [0i32; N];
let mut a_simd = [0i32; N];
for i in 0..N {
let v = (i as i32 * 37 + 11) % Q;
a_scalar[i] = v;
a_simd[i] = v;
}
crate::ntt::ntt(&mut a_scalar);
crate::ntt::ntt(&mut a_simd);
crate::ntt::invntt_tomont(&mut a_scalar);
invntt_simd(&mut a_simd);
assert_eq!(a_scalar, a_simd, "SIMD INVNTT diverged from scalar");
}
#[test]
fn test_ntt_roundtrip_simd() {
let mut a = [0i32; N];
for i in 0..N {
let v = (i as i32 * 13 + 5) % Q;
a[i] = v;
}
let before = a;
ntt_simd(&mut a);
assert_ne!(a, before, "NTT did not transform input");
invntt_simd(&mut a);
let mut b = before;
ntt_simd(&mut b);
invntt_simd(&mut b);
assert_eq!(a, b, "SIMD NTT round-trip not deterministic");
}
}