pub unsafe fn montgomery_mul_256(
a: __m256i,
b: __m256i,
r_vec: __m256i,
mod_vec: __m256i,
) -> __m256iExamples found in repository?
More examples
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 91)
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}
161
162pub unsafe fn sub_vec_avx2<M>(
163 a: __m256i,
164 b: __m256i,
165 mod_vec: __m256i,
166 mod2_vec: __m256i,
167) -> __m256i
168where
169 M: Montgomery32NttModulus,
170{
171 if M::MOD < LAZY_THRESHOLD {
172 montgomery_simd::montgomery_sub_256(a, b, mod2_vec)
173 } else {
174 montgomery_simd::sub_mod_256(a, b, mod_vec)
175 }
176}
177
178unsafe fn mul_vec_avx2<M>(a: __m256i, b: __m256i, r_vec: __m256i, mod_vec: __m256i) -> __m256i
179where
180 M: Montgomery32NttModulus,
181{
182 if M::MOD < LAZY_THRESHOLD {
183 montgomery_simd::montgomery_mul_256(a, b, r_vec, mod_vec)
184 } else {
185 montgomery_simd::montgomery_mul_256_canon(a, b, r_vec, mod_vec)
186 }
187}