fn normalize_scalar<M>(x: u32) -> u32where
M: Montgomery32NttModulus,Examples found in repository?
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 141)
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}More examples
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx512.rs (line 17)
4unsafe fn normalize_avx512<M>(a: &mut [u32])
5where
6 M: Montgomery32NttModulus,
7{
8 let mod_vec = _mm512_set1_epi32(M::MOD as i32);
9 let mut i = 0;
10 while i + 16 <= a.len() {
11 let x = _mm512_loadu_si512(a.as_ptr().add(i) as *const __m512i);
12 let y = _mm512_min_epu32(x, _mm512_sub_epi32(x, mod_vec));
13 _mm512_storeu_si512(a.as_mut_ptr().add(i) as *mut __m512i, y);
14 i += 16;
15 }
16 while i < a.len() {
17 a[i] = normalize_scalar::<M>(a[i]);
18 i += 1;
19 }
20}
21
22unsafe fn add_vec_avx512<M>(a: __m512i, b: __m512i, mod_vec: __m512i, mod2_vec: __m512i) -> __m512i
23where
24 M: Montgomery32NttModulus,
25{
26 if M::MOD < LAZY_THRESHOLD {
27 montgomery_simd::montgomery_add_512(a, b, mod2_vec)
28 } else {
29 montgomery_simd::add_mod_512(a, b, mod_vec)
30 }
31}
32
33unsafe fn sub_vec_avx512<M>(a: __m512i, b: __m512i, mod_vec: __m512i, mod2_vec: __m512i) -> __m512i
34where
35 M: Montgomery32NttModulus,
36{
37 if M::MOD < LAZY_THRESHOLD {
38 montgomery_simd::montgomery_sub_512(a, b, mod2_vec)
39 } else {
40 montgomery_simd::sub_mod_512(a, b, mod_vec)
41 }
42}
43
44unsafe fn mul_vec_avx512<M>(a: __m512i, b: __m512i, r_vec: __m512i, mod_vec: __m512i) -> __m512i
45where
46 M: Montgomery32NttModulus,
47{
48 if M::MOD < LAZY_THRESHOLD {
49 montgomery_simd::montgomery_mul_512(a, b, r_vec, mod_vec)
50 } else {
51 montgomery_simd::montgomery_mul_512_canon(a, b, r_vec, mod_vec)
52 }
53}
54
55#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
56pub unsafe fn pointwise_multiply_avx512<M>(f: &mut [MInt<M>], g: &[MInt<M>])
57where
58 M: Montgomery32NttModulus,
59{
60 let r_vec = _mm512_set1_epi32(M::R as i32);
61 let mod_vec = _mm512_set1_epi32(M::MOD as i32);
62 let mut i = 0;
63 while i + 16 <= f.len() {
64 let a = _mm512_loadu_si512(f.as_ptr().add(i) as *const __m512i);
65 let b = _mm512_loadu_si512(g.as_ptr().add(i) as *const __m512i);
66 let x = montgomery_simd::montgomery_mul_512_canon(a, b, r_vec, mod_vec);
67 _mm512_storeu_si512(f.as_mut_ptr().add(i) as *mut __m512i, x);
68 i += 16;
69 }
70 while i < f.len() {
71 f[i] *= g[i];
72 i += 1;
73 }
74}
75
76#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
77pub unsafe fn pointwise_multiply_add_avx512<M>(sum: &mut [MInt<M>], f: &[MInt<M>], g: &[MInt<M>])
78where
79 M: Montgomery32NttModulus,
80{
81 let r_vec = _mm512_set1_epi32(M::R as i32);
82 let mod_vec = _mm512_set1_epi32(M::MOD as i32);
83 let mut i = 0;
84 while i + 16 <= sum.len() {
85 let s = _mm512_loadu_si512(sum.as_ptr().add(i).cast());
86 let f = _mm512_loadu_si512(f.as_ptr().add(i).cast());
87 let g = _mm512_loadu_si512(g.as_ptr().add(i).cast());
88 let product = montgomery_simd::montgomery_mul_512_canon(f, g, r_vec, mod_vec);
89 _mm512_storeu_si512(
90 sum.as_mut_ptr().add(i).cast(),
91 montgomery_simd::add_mod_512(s, product, mod_vec),
92 );
93 i += 16;
94 }
95 while i < sum.len() {
96 sum[i] += f[i] * g[i];
97 i += 1;
98 }
99}
100
101#[inline]
102#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
103pub unsafe fn ntt_batch_avx512<M>(a: &mut [MInt<M>], width: usize)
104where
105 M: Montgomery32NttModulus,
106{
107 let n = a.len() / width;
108 if n <= 1 {
109 return;
110 }
111 let ptr = a.as_mut_ptr() as *mut u32;
112 let a = std::slice::from_raw_parts_mut(ptr, a.len());
113 let mod_vec = _mm512_set1_epi32(M::MOD as i32);
114 let mod2_vec = _mm512_set1_epi32(M::MOD.wrapping_add(M::MOD) as i32);
115 let r_vec = _mm512_set1_epi32(M::R as i32);
116 let imag = M::INFO.root[2];
117 let imag_vec = _mm512_set1_epi32(imag as i32);
118
119 let mut v = n / 2;
120 if n.trailing_zeros() & 1 == 1 {
121 let half = v * width;
122 let mut i = 0;
123 while i + 16 <= half {
124 let x0 = _mm512_loadu_si512(a.as_ptr().add(i) as *const __m512i);
125 let x1 = _mm512_loadu_si512(a.as_ptr().add(half + i) as *const __m512i);
126 let y0 = add_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
127 let y1 = sub_vec_avx512::<M>(x0, x1, mod_vec, mod2_vec);
128 _mm512_storeu_si512(a.as_mut_ptr().add(i) as *mut __m512i, y0);
129 _mm512_storeu_si512(a.as_mut_ptr().add(half + i) as *mut __m512i, y1);
130 i += 16;
131 }
132 while i < half {
133 let x0 = a[i];
134 let x1 = a[half + i];
135 a[i] = M::mod_add(x0, x1);
136 a[half + i] = M::mod_sub(x0, x1);
137 i += 1;
138 }
139 v >>= 1;
140 }
141 while v > 1 {
142 if width == 1 && v == 2 && a.len() >= 8 && M::MOD < LAZY_THRESHOLD {
143 ntt_avx2::ntt_four_avx2::<M, false>(a);
144 break;
145 }
146 let half = (v >> 1) * width;
147 let mut w1 = M::N1;
148 for (s, block) in a.chunks_exact_mut((v << 1) * width).enumerate() {
149 let base = block.as_mut_ptr();
150 let ll = base;
151 let lr = base.add(half);
152 let rl = base.add(v * width);
153 let rr = base.add(v * width + half);
154 let w2 = M::mod_mul(w1, w1);
155 let w3 = M::mod_mul(w2, w1);
156 let w1v = _mm512_set1_epi32(w1 as i32);
157 let w2v = _mm512_set1_epi32(w2 as i32);
158 let w3v = _mm512_set1_epi32(w3 as i32);
159
160 let mut i = 0;
161 while i + 16 <= half {
162 let x0 = _mm512_loadu_si512(ll.add(i) as *const __m512i);
163 let x1 = _mm512_loadu_si512(lr.add(i) as *const __m512i);
164 let x2 = _mm512_loadu_si512(rl.add(i) as *const __m512i);
165 let x3 = _mm512_loadu_si512(rr.add(i) as *const __m512i);
166
167 let (a1, a2, a3) = if s == 0 {
168 (x1, x2, x3)
169 } else {
170 (
171 mul_vec_avx512::<M>(x1, w1v, r_vec, mod_vec),
172 mul_vec_avx512::<M>(x2, w2v, r_vec, mod_vec),
173 mul_vec_avx512::<M>(x3, w3v, r_vec, mod_vec),
174 )
175 };
176
177 let a0pa2 = add_vec_avx512::<M>(x0, a2, mod_vec, mod2_vec);
178 let a0na2 = sub_vec_avx512::<M>(x0, a2, mod_vec, mod2_vec);
179 let a1pa3 = add_vec_avx512::<M>(a1, a3, mod_vec, mod2_vec);
180 let a1na3 = sub_vec_avx512::<M>(a1, a3, mod_vec, mod2_vec);
181 let a1na3imag = mul_vec_avx512::<M>(a1na3, imag_vec, r_vec, mod_vec);
182
183 let y0 = add_vec_avx512::<M>(a0pa2, a1pa3, mod_vec, mod2_vec);
184 let y1 = sub_vec_avx512::<M>(a0pa2, a1pa3, mod_vec, mod2_vec);
185 let y2 = add_vec_avx512::<M>(a0na2, a1na3imag, mod_vec, mod2_vec);
186 let y3 = sub_vec_avx512::<M>(a0na2, a1na3imag, mod_vec, mod2_vec);
187
188 _mm512_storeu_si512(ll.add(i) as *mut __m512i, y0);
189 _mm512_storeu_si512(lr.add(i) as *mut __m512i, y1);
190 _mm512_storeu_si512(rl.add(i) as *mut __m512i, y2);
191 _mm512_storeu_si512(rr.add(i) as *mut __m512i, y3);
192 i += 16;
193 }
194 while i < half {
195 let a0 = normalize_scalar::<M>(*ll.add(i));
196 let a1 = M::mod_mul(normalize_scalar::<M>(*lr.add(i)), w1);
197 let a2 = M::mod_mul(normalize_scalar::<M>(*rl.add(i)), w2);
198 let a3 = M::mod_mul(normalize_scalar::<M>(*rr.add(i)), w3);
199 let a0pa2 = M::mod_add(a0, a2);
200 let a0na2 = M::mod_sub(a0, a2);
201 let a1pa3 = M::mod_add(a1, a3);
202 let a1na3 = M::mod_sub(a1, a3);
203 let a1na3imag = M::mod_mul(a1na3, imag);
204 *ll.add(i) = M::mod_add(a0pa2, a1pa3);
205 *lr.add(i) = M::mod_sub(a0pa2, a1pa3);
206 *rl.add(i) = M::mod_add(a0na2, a1na3imag);
207 *rr.add(i) = M::mod_sub(a0na2, a1na3imag);
208 i += 1;
209 }
210 w1 = M::mod_mul(w1, M::INFO.rate3[s.trailing_ones() as usize]);
211 }
212 v >>= 2;
213 }
214 normalize_avx512::<M>(a);
215}