Skip to main content

montgomery_add_256

Function montgomery_add_256 

Source
pub unsafe fn montgomery_add_256(
    a: __m256i,
    b: __m256i,
    mod2_vec: __m256i,
) -> __m256i
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 95)
23pub unsafe fn ntt_four_avx2<M, const INVERSE: bool>(a: &mut [u32])
24where
25    M: Montgomery32NttModulus,
26{
27    let roots = if INVERSE {
28        &M::INFO.inv_root
29    } else {
30        &M::INFO.root
31    };
32    let rates = &const {
33        let info = M::INFO;
34        let (roots, rates) = if INVERSE {
35            (info.inv_root, info.inv_rate3_packed)
36        } else {
37            (info.root, info.rate3_packed)
38        };
39        let mut result = [[0u32; 4]; 32];
40        let mut i = 0;
41        while i < 32 {
42            let r = mod_mul(rates[i][1], roots[3], M::MOD, M::R);
43            let r2 = mod_mul(r, r, M::MOD, M::R);
44            result[i] = [M::N1, r, r2, mod_mul(r2, r, M::MOD, M::R)];
45            i += 1;
46        }
47        result
48    };
49    let one = _mm256_set1_epi32(M::N1 as i32);
50    let modulus = _mm256_set1_epi32(M::MOD as i32);
51    let modulus2 = _mm256_set1_epi32((M::MOD * 2) as i32);
52    let r = _mm256_set1_epi32(M::R as i32);
53    let root3 = M::mod_mul(roots[2], roots[3]) as i32;
54    let step = _mm256_setr_epi32(
55        M::N1 as i32,
56        roots[3] as i32,
57        roots[2] as i32,
58        root3,
59        M::N1 as i32,
60        roots[3] as i32,
61        roots[2] as i32,
62        root3,
63    );
64    let mut twiddle = _mm256_blend_epi32::<0xf0>(one, step);
65    let imag = if INVERSE {
66        _mm256_setr_epi32(
67            M::N1 as i32,
68            roots[2] as i32,
69            M::N1 as i32,
70            roots[2] as i32,
71            M::N1 as i32,
72            roots[2] as i32,
73            M::N1 as i32,
74            roots[2] as i32,
75        )
76    } else {
77        _mm256_setr_epi32(
78            M::N1 as i32,
79            M::N1 as i32,
80            roots[2] as i32,
81            roots[2] as i32,
82            M::N1 as i32,
83            M::N1 as i32,
84            roots[2] as i32,
85            roots[2] as i32,
86        )
87    };
88    for (s, a) in a.as_chunks_mut::<8>().0.iter_mut().enumerate() {
89        let mut x = _mm256_loadu_si256(a.as_ptr().cast());
90        if !INVERSE {
91            x = montgomery_simd::montgomery_mul_256(x, twiddle, r, modulus);
92        }
93        let pair = if INVERSE {
94            let y = _mm256_shuffle_epi32::<0xb1>(x);
95            let sum = montgomery_simd::montgomery_add_256(x, y, modulus2);
96            let diff = montgomery_simd::montgomery_sub_256(x, y, modulus2);
97            let sum = _mm256_shuffle_epi32::<0x88>(sum);
98            let diff = _mm256_shuffle_epi32::<0x88>(diff);
99            let diff = montgomery_simd::montgomery_mul_256(diff, imag, r, modulus);
100            _mm256_unpacklo_epi64(sum, diff)
101        } else {
102            let y = _mm256_shuffle_epi32::<0x4e>(x);
103            let sum = montgomery_simd::montgomery_add_256(x, y, modulus2);
104            let diff = montgomery_simd::montgomery_sub_256(x, y, modulus2);
105            _mm256_unpacklo_epi64(sum, diff)
106        };
107        let left = _mm256_shuffle_epi32::<0xa0>(pair);
108        let mut right = _mm256_shuffle_epi32::<0xf5>(pair);
109        if !INVERSE {
110            right = montgomery_simd::montgomery_mul_256(right, imag, r, modulus);
111        }
112        let sum = montgomery_simd::montgomery_add_256(left, right, modulus2);
113        let diff = montgomery_simd::montgomery_sub_256(left, right, modulus2);
114        let mut value = _mm256_blend_epi32::<0xaa>(sum, diff);
115        if INVERSE {
116            value = _mm256_shuffle_epi32::<0xd8>(value);
117            value = montgomery_simd::montgomery_mul_256(value, twiddle, r, modulus);
118        }
119        _mm256_storeu_si256(a.as_mut_ptr().cast(), value);
120        let rate = _mm256_broadcastsi128_si256(_mm_loadu_si128(
121            rates[s.trailing_ones() as usize + 1].as_ptr().cast(),
122        ));
123        twiddle = montgomery_simd::montgomery_mul_256(twiddle, rate, r, modulus);
124    }
125}
126
127#[target_feature(enable = "avx2")]
128unsafe fn normalize_avx2<M>(a: &mut [u32])
129where
130    M: Montgomery32NttModulus,
131{
132    let mod_vec = _mm256_set1_epi32(M::MOD as i32);
133    let mut i = 0;
134    while i + 8 <= a.len() {
135        let x = _mm256_loadu_si256(a.as_ptr().add(i) as *const __m256i);
136        let y = _mm256_min_epu32(x, _mm256_sub_epi32(x, mod_vec));
137        _mm256_storeu_si256(a.as_mut_ptr().add(i) as *mut __m256i, y);
138        i += 8;
139    }
140    while i < a.len() {
141        a[i] = normalize_scalar::<M>(a[i]);
142        i += 1;
143    }
144}
145
146pub unsafe fn add_vec_avx2<M>(
147    a: __m256i,
148    b: __m256i,
149    mod_vec: __m256i,
150    mod2_vec: __m256i,
151) -> __m256i
152where
153    M: Montgomery32NttModulus,
154{
155    if M::MOD < LAZY_THRESHOLD {
156        montgomery_simd::montgomery_add_256(a, b, mod2_vec)
157    } else {
158        montgomery_simd::add_mod_256(a, b, mod_vec)
159    }
160}