use hermes_simd_core::view::SimdError;
#[inline]
pub fn ntt_butterfly_stage_u64(
data: &mut [u64],
stage_len: usize,
twiddles: &[u64],
modulus: u64,
) -> Result<(), SimdError> {
if stage_len == 0
|| !stage_len.is_multiple_of(2)
|| !data.len().is_multiple_of(stage_len)
|| twiddles.len() != stage_len / 2
{
return Err(SimdError::LengthMismatch);
}
let half = stage_len / 2;
for chunk in data.chunks_mut(stage_len) {
let (left, right) = chunk.split_at_mut(half);
for index in 0..half {
let lhs = left[index];
let rhs = mod_mul_u64(right[index], twiddles[index], modulus);
left[index] = mod_add_u64(lhs, rhs, modulus);
right[index] = mod_sub_u64(lhs, rhs, modulus);
}
}
Ok(())
}
#[inline]
fn mod_mul_u64(lhs: u64, rhs: u64, modulus: u64) -> u64 {
((lhs as u128 * rhs as u128) % modulus as u128) as u64
}
#[inline]
fn mod_add_u64(lhs: u64, rhs: u64, modulus: u64) -> u64 {
((lhs as u128 + rhs as u128) % modulus as u128) as u64
}
#[inline]
fn mod_sub_u64(lhs: u64, rhs: u64, modulus: u64) -> u64 {
if lhs >= rhs {
lhs - rhs
} else {
modulus - (rhs - lhs)
}
}