Skip to main content

hermes_simd/dispatch/
modular.rs

1//! Modular-arithmetic kernels for transform workloads.
2//!
3//! These kernels keep exact residue-field arithmetic in Hermes instead of
4//! forcing downstream crates to duplicate butterfly loops around SIMD-facing
5//! provider boundaries. Multiplication uses `u128` widening because the
6//! numerical contract is exact modular arithmetic, not wrapping arithmetic.
7
8use hermes_simd_core::view::SimdError;
9
10/// Executes one radix-2 NTT butterfly stage in place.
11///
12/// `data` is partitioned into chunks of `stage_len`; `twiddles` contains one
13/// stage twiddle per element in the right half of each chunk. Each butterfly
14/// computes:
15///
16/// ```text
17/// left'  = left + right * twiddle (mod modulus)
18/// right' = left - right * twiddle (mod modulus)
19/// ```
20///
21/// Returns [`SimdError::LengthMismatch`] when the stage shape is invalid.
22#[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}