hermes_simd/dispatch/
modular.rs1use hermes_simd_core::view::SimdError;
9
10#[inline]
23pub fn ntt_butterfly_stage_u64(
24 data: &mut [u64],
25 stage_len: usize,
26 twiddles: &[u64],
27 modulus: u64,
28) -> Result<(), SimdError> {
29 if stage_len == 0
30 || !stage_len.is_multiple_of(2)
31 || !data.len().is_multiple_of(stage_len)
32 || twiddles.len() != stage_len / 2
33 {
34 return Err(SimdError::LengthMismatch);
35 }
36
37 let half = stage_len / 2;
38 for chunk in data.chunks_mut(stage_len) {
39 let (left, right) = chunk.split_at_mut(half);
40 for index in 0..half {
41 let lhs = left[index];
42 let rhs = mod_mul_u64(right[index], twiddles[index], modulus);
43 left[index] = mod_add_u64(lhs, rhs, modulus);
44 right[index] = mod_sub_u64(lhs, rhs, modulus);
45 }
46 }
47 Ok(())
48}
49
50#[inline]
51fn mod_mul_u64(lhs: u64, rhs: u64, modulus: u64) -> u64 {
52 ((lhs as u128 * rhs as u128) % modulus as u128) as u64
53}
54
55#[inline]
56fn mod_add_u64(lhs: u64, rhs: u64, modulus: u64) -> u64 {
57 ((lhs as u128 + rhs as u128) % modulus as u128) as u64
58}
59
60#[inline]
61fn mod_sub_u64(lhs: u64, rhs: u64, modulus: u64) -> u64 {
62 if lhs >= rhs {
63 lhs - rhs
64 } else {
65 modulus - (rhs - lhs)
66 }
67}