1use super::*;
2
3#[inline]
4#[target_feature(enable = "avx2")]
5unsafe fn load_block_avx2(a: *const u32, i: usize) -> __m256i {
6 _mm256_loadu_si256(a.add(i << 3) as *const __m256i)
7}
8
9#[inline]
10#[target_feature(enable = "avx2")]
11unsafe fn store_block_avx2(a: *mut u32, i: usize, x: __m256i) {
12 _mm256_storeu_si256(a.add(i << 3) as *mut __m256i, x);
13}
14
15#[inline]
16#[target_feature(enable = "avx2")]
17unsafe fn shrink_avx2(x: __m256i, modulus: __m256i) -> __m256i {
18 _mm256_min_epu32(x, _mm256_sub_epi32(x, modulus))
19}
20
21#[inline]
22#[target_feature(enable = "avx2")]
23unsafe fn normalize_avx2(x: __m256i, modulus: __m256i, modulus2: __m256i) -> __m256i {
24 shrink_avx2(shrink_avx2(x, modulus2), modulus)
25}
26
27#[inline]
28#[target_feature(enable = "avx2")]
29unsafe fn add_mod_avx2(x: __m256i, y: __m256i, modulus: __m256i) -> __m256i {
30 shrink_avx2(_mm256_add_epi32(x, y), modulus)
31}
32
33#[inline]
34#[target_feature(enable = "avx2")]
35unsafe fn sub_mod_avx2(x: __m256i, y: __m256i, modulus: __m256i) -> __m256i {
36 let d = _mm256_sub_epi32(x, y);
37 _mm256_min_epu32(d, _mm256_add_epi32(d, modulus))
38}
39
40#[inline]
41#[target_feature(enable = "avx2")]
42unsafe fn lazy_sub_avx2(x: __m256i, y: __m256i, modulus: __m256i) -> __m256i {
43 _mm256_add_epi32(x, _mm256_sub_epi32(modulus, y))
44}
45
46#[inline]
47#[target_feature(enable = "avx2")]
48unsafe fn montgomery_mul_even_avx2(
49 x: __m256i,
50 y: __m256i,
51 r: __m256i,
52 modulus: __m256i,
53) -> __m256i {
54 let x_odd = _mm256_bsrli_epi128::<4>(x);
55 let t_even = _mm256_mul_epu32(x, y);
56 let t_odd = _mm256_mul_epu32(x_odd, y);
57 let m_even = _mm256_mul_epu32(t_even, r);
58 let m_odd = _mm256_mul_epu32(t_odd, r);
59 let u_even = _mm256_add_epi64(t_even, _mm256_mul_epu32(m_even, modulus));
60 let u_odd = _mm256_add_epi64(t_odd, _mm256_mul_epu32(m_odd, modulus));
61 _mm256_or_si256(_mm256_bsrli_epi128::<4>(u_even), u_odd)
62}
63
64#[inline]
65#[target_feature(enable = "avx2")]
66unsafe fn update_root_avx2(root: __m256i, rate: __m256i, modulus: __m256i) -> __m256i {
67 let correction = _mm256_mul_epu32(root, rate);
68 let product = _mm256_mul_epu32(root, _mm256_srli_epi64::<32>(rate));
69 let correction = _mm256_mul_epu32(correction, modulus);
70 shrink_avx2(
71 _mm256_srli_epi64::<32>(_mm256_add_epi64(product, correction)),
72 modulus,
73 )
74}
75
76#[inline]
77#[target_feature(enable = "avx2")]
78unsafe fn packed_rate_avx2<M>(index: usize, inverse: bool) -> __m256i
79where
80 M: Montgomery32NttModulus,
81{
82 let rate = if inverse {
83 &M::INFO.inv_rate3_packed[index]
84 } else {
85 &M::INFO.rate3_packed[index]
86 };
87 _mm256_loadu_si256(rate.as_ptr().cast())
88}
89
90#[target_feature(enable = "avx2")]
91unsafe fn ntt_blocks_avx2<M>(a: *mut u32, n: usize)
92where
93 M: Montgomery32NttModulus,
94{
95 let modulus = _mm256_set1_epi32(M::MOD as i32);
96 let modulus2 = _mm256_set1_epi32(M::MOD.wrapping_mul(2) as i32);
97 let r = _mm256_set1_epi32(M::R as i32);
98 let imag = M::INFO.root[2];
99 let imag_r = _mm256_set1_epi32(imag.wrapping_mul(M::R) as i32);
100 let imag = _mm256_set1_epi32(imag as i32);
101 let root_indices = _mm256_setr_epi32(0, 2, 0, 4, 0, 2, 0, 4);
102 let root3 = M::INFO.root[3];
103 let root2 = M::INFO.root[2];
104 let initial_root = _mm256_setr_epi32(
105 root3 as i32,
106 0,
107 root2 as i32,
108 0,
109 (M::MOD - M::mod_mul(root2, root3)) as i32,
110 0,
111 0,
112 0,
113 );
114 let log_n = n.trailing_zeros() as usize;
115 let mut roots = [initial_root; 16];
116 let nn = n >> (log_n & 1);
117 let tile_len = n.min(64);
118
119 if nn != n {
120 let mut i = 0;
121 while i < nn {
122 let x0 = load_block_avx2(a, i);
123 let x1 = load_block_avx2(a, nn + i);
124 store_block_avx2(a, i, add_mod_avx2(x0, x1, modulus2));
125 store_block_avx2(a, nn + i, lazy_sub_avx2(x0, x1, modulus2));
126 i += 1;
127 }
128 }
129
130 let mut size = nn >> 2;
131 while size > 0 {
132 let final_stage = size == 1;
133 let mut i = 0;
134 while i < size {
135 let x0 = load_block_avx2(a, i);
136 let x1 = load_block_avx2(a, size + i);
137 let x2 = load_block_avx2(a, size * 2 + i);
138 let x3 = load_block_avx2(a, size * 3 + i);
139 let g3 = montgomery_simd::montgomery_mul_256_fixed(
140 lazy_sub_avx2(x1, x3, modulus2),
141 imag,
142 imag_r,
143 modulus,
144 );
145 let g1 = add_mod_avx2(x1, x3, modulus2);
146 let g0 = add_mod_avx2(x0, x2, modulus2);
147 let g2 = sub_mod_avx2(x0, x2, modulus2);
148 let mut y0 = add_mod_avx2(g0, g1, modulus2);
149 let mut y1 = lazy_sub_avx2(g0, g1, modulus2);
150 let mut y2 = _mm256_add_epi32(g2, g3);
151 let mut y3 = lazy_sub_avx2(g2, g3, modulus2);
152 if final_stage {
153 y0 = normalize_avx2(y0, modulus, modulus2);
154 y1 = normalize_avx2(y1, modulus, modulus2);
155 y2 = normalize_avx2(y2, modulus, modulus2);
156 y3 = normalize_avx2(y3, modulus, modulus2);
157 }
158 store_block_avx2(a, i, y0);
159 store_block_avx2(a, size + i, y1);
160 store_block_avx2(a, size * 2 + i, y2);
161 store_block_avx2(a, size * 3 + i, y3);
162 i += 1;
163 }
164 size >>= 2;
165 }
166
167 let mut tile = 0;
168 let mut stage_log = log_n.min(6) & !1;
169 let mut root_slot = (stage_log - 2) >> 1;
170 while tile < n {
171 let base = a.add(tile << 3);
172 let mut group_len = 1usize << stage_log;
173 let mut quarter = group_len >> 2;
174 while quarter > 1 {
175 let mut root = roots[root_slot];
176 let mut i = if tile == 0 { group_len } else { 0 };
177 let mut group = (tile + i) >> stage_log;
178 while i < tile_len {
179 let r1 = _mm256_permutevar8x32_epi32(root, root_indices);
180 let r1_r = _mm256_permutevar8x32_epi32(_mm256_mul_epu32(root, r), root_indices);
181 root = update_root_avx2(
182 root,
183 packed_rate_avx2::<M>((!group).trailing_zeros() as usize, false),
184 modulus,
185 );
186 let r2 = _mm256_shuffle_epi32::<0x55>(r1);
187 let nr3 = _mm256_shuffle_epi32::<0xff>(r1);
188 let r2_r = _mm256_shuffle_epi32::<0x55>(r1_r);
189 let nr3_r = _mm256_shuffle_epi32::<0xff>(r1_r);
190 let mut j = 0;
191 while j < quarter {
192 let p0 = (i + j) << 3;
193 let x0 = _mm256_loadu_si256(base.add(p0).cast());
194 let x1 = _mm256_loadu_si256(base.add(p0 + (quarter << 3)).cast());
195 let x2 = _mm256_loadu_si256(base.add(p0 + (quarter << 4)).cast());
196 let x3 = _mm256_loadu_si256(base.add(p0 + quarter * 24).cast());
197 let g1 = montgomery_simd::montgomery_mul_256_fixed(x1, r1, r1_r, modulus);
198 let ng3 = montgomery_simd::montgomery_mul_256_fixed(x3, nr3, nr3_r, modulus);
199 let g2 = montgomery_simd::montgomery_mul_256_fixed(x2, r2, r2_r, modulus);
200 let g0 = shrink_avx2(x0, modulus2);
201 let h3 = montgomery_simd::montgomery_mul_256_fixed(
202 _mm256_add_epi32(g1, ng3),
203 imag,
204 imag_r,
205 modulus,
206 );
207 let h1 = sub_mod_avx2(g1, ng3, modulus2);
208 let h0 = add_mod_avx2(g0, g2, modulus2);
209 let h2 = sub_mod_avx2(g0, g2, modulus2);
210 _mm256_storeu_si256(base.add(p0).cast(), _mm256_add_epi32(h0, h1));
211 _mm256_storeu_si256(
212 base.add(p0 + (quarter << 3)).cast(),
213 lazy_sub_avx2(h0, h1, modulus2),
214 );
215 _mm256_storeu_si256(
216 base.add(p0 + (quarter << 4)).cast(),
217 _mm256_add_epi32(h2, h3),
218 );
219 _mm256_storeu_si256(
220 base.add(p0 + quarter * 24).cast(),
221 lazy_sub_avx2(h2, h3, modulus2),
222 );
223 j += 1;
224 }
225 i += group_len;
226 group += 1;
227 }
228 roots[root_slot] = root;
229 group_len = quarter;
230 quarter >>= 2;
231 stage_log -= 2;
232 root_slot -= 1;
233 }
234
235 let mut root = roots[0];
236 let mut i = tile + if tile == 0 { 4 } else { 0 };
237 while i < tile + tile_len {
238 let r1 = _mm256_permutevar8x32_epi32(root, root_indices);
239 root = update_root_avx2(
240 root,
241 packed_rate_avx2::<M>((!(i >> 2)).trailing_zeros() as usize, false),
242 modulus,
243 );
244 let r2 = _mm256_shuffle_epi32::<0x55>(r1);
245 let nr3 = _mm256_shuffle_epi32::<0xff>(r1);
246 let x0 = load_block_avx2(a, i);
247 let x1 = load_block_avx2(a, i + 1);
248 let x2 = load_block_avx2(a, i + 2);
249 let x3 = load_block_avx2(a, i + 3);
250 let g1 = montgomery_mul_even_avx2(x1, r1, r, modulus);
251 let ng3 = montgomery_mul_even_avx2(x3, nr3, r, modulus);
252 let g2 = montgomery_mul_even_avx2(x2, r2, r, modulus);
253 let g0 = shrink_avx2(x0, modulus2);
254 let h3 = montgomery_simd::montgomery_mul_256_fixed(
255 _mm256_add_epi32(g1, ng3),
256 imag,
257 imag_r,
258 modulus,
259 );
260 let h1 = sub_mod_avx2(g1, ng3, modulus2);
261 let h0 = add_mod_avx2(g0, g2, modulus2);
262 let h2 = sub_mod_avx2(g0, g2, modulus2);
263 store_block_avx2(
264 a,
265 i,
266 normalize_avx2(_mm256_add_epi32(h0, h1), modulus, modulus2),
267 );
268 store_block_avx2(
269 a,
270 i + 1,
271 normalize_avx2(lazy_sub_avx2(h0, h1, modulus2), modulus, modulus2),
272 );
273 store_block_avx2(
274 a,
275 i + 2,
276 normalize_avx2(_mm256_add_epi32(h2, h3), modulus, modulus2),
277 );
278 store_block_avx2(
279 a,
280 i + 3,
281 normalize_avx2(lazy_sub_avx2(h2, h3, modulus2), modulus, modulus2),
282 );
283 i += 4;
284 }
285 roots[0] = root;
286
287 tile += tile_len;
288 if tile < n {
289 stage_log = tile.trailing_zeros() as usize & !1;
290 root_slot = (stage_log - 2) >> 1;
291 }
292 }
293}
294
295#[target_feature(enable = "avx2")]
296unsafe fn intt_blocks_avx2<M>(a: *mut u32, n: usize)
297where
298 M: Montgomery32NttModulus,
299{
300 let inv = M::mod_inv(<M as MIntConvert<u32>>::from(n as u32));
301 let modulus = _mm256_set1_epi32(M::MOD as i32);
302 let modulus2 = _mm256_set1_epi32(M::MOD.wrapping_mul(2) as i32);
303 let r = _mm256_set1_epi32(M::R as i32);
304 let imag = _mm256_set1_epi32(M::INFO.root[2] as i32);
305 let imag_r = _mm256_set1_epi32(M::INFO.root[2].wrapping_mul(M::R) as i32);
306 let root_indices = _mm256_setr_epi32(0, 2, 0, 4, 0, 2, 0, 4);
307 let root3 = M::INFO.inv_root[3];
308 let root2 = M::INFO.inv_root[2];
309 let initial_root = _mm256_setr_epi32(
310 root3 as i32,
311 0,
312 root2 as i32,
313 0,
314 M::mod_mul(root2, root3) as i32,
315 0,
316 0,
317 0,
318 );
319 let log_n = n.trailing_zeros() as usize;
320 let mut roots = [initial_root; 16];
321 let nn = n >> (log_n & 1);
322 let tile_len = n.min(64);
323 let inv_vec = _mm256_set1_epi32(inv as i32);
324 let inv_r = _mm256_set1_epi32(inv.wrapping_mul(M::R) as i32);
325 roots[0] = inv_vec;
326
327 let mut tile = 0;
328 while tile < n {
329 let max_stage_log = (tile + tile_len).trailing_zeros() as usize;
330 let mut stage_log = 4usize;
331 let mut root_slot = 1usize;
332
333 let mut root = roots[0];
334 let mut i = tile;
335 while i < tile + tile_len {
336 let r1 = _mm256_permutevar8x32_epi32(root, root_indices);
337 root = update_root_avx2(
338 root,
339 packed_rate_avx2::<M>((!(i >> 2)).trailing_zeros() as usize, true),
340 modulus,
341 );
342 let r2 = _mm256_shuffle_epi32::<0x55>(r1);
343 let r3 = _mm256_shuffle_epi32::<0xff>(r1);
344 let x0 = load_block_avx2(a, i);
345 let x1 = load_block_avx2(a, i + 1);
346 let x2 = load_block_avx2(a, i + 2);
347 let x3 = load_block_avx2(a, i + 3);
348 let g3 = montgomery_simd::montgomery_mul_256_fixed(
349 lazy_sub_avx2(x3, x2, modulus2),
350 imag,
351 imag_r,
352 modulus,
353 );
354 let g2 = add_mod_avx2(x2, x3, modulus2);
355 let g0 = add_mod_avx2(x0, x1, modulus2);
356 let g1 = sub_mod_avx2(x0, x1, modulus2);
357 let h2 = lazy_sub_avx2(g0, g2, modulus2);
358 let h3 = lazy_sub_avx2(g1, g3, modulus2);
359 let h0 = _mm256_add_epi32(g0, g2);
360 let h1 = _mm256_add_epi32(g1, g3);
361 if inv == M::N1 {
362 store_block_avx2(a, i, shrink_avx2(h0, modulus2));
363 } else {
364 store_block_avx2(
365 a,
366 i,
367 montgomery_simd::montgomery_mul_256_fixed(h0, inv_vec, inv_r, modulus),
368 );
369 }
370 store_block_avx2(a, i + 1, montgomery_mul_even_avx2(h1, r1, r, modulus));
371 store_block_avx2(a, i + 2, montgomery_mul_even_avx2(h2, r2, r, modulus));
372 store_block_avx2(a, i + 3, montgomery_mul_even_avx2(h3, r3, r, modulus));
373 i += 4;
374 }
375 roots[0] = root;
376
377 let mut group_len = 16usize;
378 let mut quarter = 4usize;
379 while stage_log <= max_stage_log {
380 let offset = tile + tile_len - group_len.max(tile_len);
381 let base = a.add(offset << 3);
382 let mut i = 0;
383 let mut root = roots[root_slot];
384 if offset == 0 {
385 let final_stage = group_len == n;
386 while i < quarter {
387 let x0 = load_block_avx2(a, i);
388 let x1 = load_block_avx2(a, quarter + i);
389 let x2 = load_block_avx2(a, quarter * 2 + i);
390 let x3 = load_block_avx2(a, quarter * 3 + i);
391 let g3 = montgomery_simd::montgomery_mul_256_fixed(
392 lazy_sub_avx2(x3, x2, modulus2),
393 imag,
394 imag_r,
395 modulus,
396 );
397 let g2 = add_mod_avx2(x2, x3, modulus2);
398 let g0 = add_mod_avx2(x0, x1, modulus2);
399 let g1 = sub_mod_avx2(x0, x1, modulus2);
400 let mut y0 = _mm256_add_epi32(g0, g2);
401 let mut y1 = _mm256_add_epi32(g1, g3);
402 let mut y2 = sub_mod_avx2(g0, g2, modulus2);
403 let mut y3 = sub_mod_avx2(g1, g3, modulus2);
404 if final_stage {
405 y0 = shrink_avx2(y0, modulus);
406 y1 = shrink_avx2(y1, modulus);
407 y2 = shrink_avx2(y2, modulus);
408 y3 = shrink_avx2(y3, modulus);
409 } else {
410 y0 = shrink_avx2(y0, modulus2);
411 y1 = shrink_avx2(y1, modulus2);
412 }
413 store_block_avx2(a, i, y0);
414 store_block_avx2(a, quarter + i, y1);
415 store_block_avx2(a, quarter * 2 + i, y2);
416 store_block_avx2(a, quarter * 3 + i, y3);
417 i += 1;
418 }
419 i = group_len;
420 }
421
422 let mut group = (tile + i) >> stage_log;
423 while i < tile_len {
424 let r1 = _mm256_permutevar8x32_epi32(root, root_indices);
425 let r1_r = _mm256_permutevar8x32_epi32(_mm256_mul_epu32(root, r), root_indices);
426 root = update_root_avx2(
427 root,
428 packed_rate_avx2::<M>((!group).trailing_zeros() as usize, true),
429 modulus,
430 );
431 let r2 = _mm256_shuffle_epi32::<0x55>(r1);
432 let r3 = _mm256_shuffle_epi32::<0xff>(r1);
433 let r2_r = _mm256_shuffle_epi32::<0x55>(r1_r);
434 let r3_r = _mm256_shuffle_epi32::<0xff>(r1_r);
435 let mut j = 0;
436 while j < quarter {
437 let p0 = (i + j) << 3;
438 let x0 = _mm256_loadu_si256(base.add(p0).cast());
439 let x1 = _mm256_loadu_si256(base.add(p0 + (quarter << 3)).cast());
440 let x2 = _mm256_loadu_si256(base.add(p0 + (quarter << 4)).cast());
441 let x3 = _mm256_loadu_si256(base.add(p0 + quarter * 24).cast());
442 let g3 = montgomery_simd::montgomery_mul_256_fixed(
443 lazy_sub_avx2(x3, x2, modulus2),
444 imag,
445 imag_r,
446 modulus,
447 );
448 let g2 = add_mod_avx2(x2, x3, modulus2);
449 let g0 = add_mod_avx2(x0, x1, modulus2);
450 let g1 = sub_mod_avx2(x0, x1, modulus2);
451 let h2 = lazy_sub_avx2(g0, g2, modulus2);
452 let h3 = lazy_sub_avx2(g1, g3, modulus2);
453 let h0 = _mm256_add_epi32(g0, g2);
454 let h1 = _mm256_add_epi32(g1, g3);
455 _mm256_storeu_si256(base.add(p0).cast(), shrink_avx2(h0, modulus2));
456 _mm256_storeu_si256(
457 base.add(p0 + (quarter << 3)).cast(),
458 montgomery_simd::montgomery_mul_256_fixed(h1, r1, r1_r, modulus),
459 );
460 _mm256_storeu_si256(
461 base.add(p0 + (quarter << 4)).cast(),
462 montgomery_simd::montgomery_mul_256_fixed(h2, r2, r2_r, modulus),
463 );
464 _mm256_storeu_si256(
465 base.add(p0 + quarter * 24).cast(),
466 montgomery_simd::montgomery_mul_256_fixed(h3, r3, r3_r, modulus),
467 );
468 j += 1;
469 }
470 i += group_len;
471 group += 1;
472 }
473 roots[root_slot] = root;
474 quarter = group_len;
475 group_len <<= 2;
476 stage_log += 2;
477 root_slot += 1;
478 }
479 tile += tile_len;
480 }
481
482 if nn != n {
483 let mut i = 0;
484 while i < nn {
485 let x0 = load_block_avx2(a, i);
486 let x1 = load_block_avx2(a, nn + i);
487 store_block_avx2(
488 a,
489 i,
490 shrink_avx2(
491 shrink_avx2(add_mod_avx2(x0, x1, modulus2), modulus),
492 modulus,
493 ),
494 );
495 store_block_avx2(
496 a,
497 nn + i,
498 shrink_avx2(
499 shrink_avx2(sub_mod_avx2(x0, x1, modulus2), modulus),
500 modulus,
501 ),
502 );
503 i += 1;
504 }
505 } else {
506 let mut i = 0;
507 while i < n {
508 store_block_avx2(
509 a,
510 i,
511 shrink_avx2(shrink_avx2(load_block_avx2(a, i), modulus), modulus),
512 );
513 i += 1;
514 }
515 }
516}
517
518#[inline]
519#[target_feature(enable = "avx2")]
520unsafe fn reduce_sum_avx2(
521 even: __m256i,
522 odd: __m256i,
523 r_vec: __m256i,
524 mod_vec: __m256i,
525) -> __m256i {
526 let even_m = _mm256_mul_epu32(even, r_vec);
527 let odd_m = _mm256_mul_epu32(odd, r_vec);
528 let even = _mm256_add_epi64(even, _mm256_mul_epu32(even_m, mod_vec));
529 let odd = _mm256_add_epi64(odd, _mm256_mul_epu32(odd_m, mod_vec));
530 _mm256_or_si256(_mm256_bsrli_epi128::<4>(even), odd)
531}
532
533#[target_feature(enable = "avx2")]
534unsafe fn convolve_8_avx2<M>(f: *mut u32, g: *const u32, n: usize)
535where
536 M: Montgomery32NttModulus,
537{
538 #[repr(C, align(32))]
539 struct AlignedWork([u32; 64]);
540
541 let mod_vec = _mm256_set1_epi32(M::MOD as i32);
542 let mod2_vec = _mm256_set1_epi32(M::MOD.wrapping_add(M::MOD) as i32);
543 let r_vec = _mm256_set1_epi32(M::R as i32);
544 let mut rr = M::N1;
545 let mut i = 0;
546 while i < n {
547 let rr_i = M::mod_mul(rr, M::INFO.root[2]);
548 let mut work = std::mem::MaybeUninit::<AlignedWork>::uninit();
549 let work = work.as_mut_ptr().cast::<u32>();
550 for (j, ww) in [
551 rr,
552 M::MOD.wrapping_mul(2) - rr,
553 rr_i,
554 M::MOD.wrapping_mul(2) - rr_i,
555 ]
556 .into_iter()
557 .enumerate()
558 {
559 let k = i + j;
560 let ff = load_block_avx2(f, k);
561 let fw = shrink_avx2(
563 montgomery_simd::montgomery_mul_256_fixed(
564 ff,
565 _mm256_set1_epi32(ww as i32),
566 _mm256_set1_epi32(ww.wrapping_mul(M::R) as i32),
567 mod_vec,
568 ),
569 mod_vec,
570 );
571 _mm256_store_si256(work.add(j << 4).cast(), fw);
572 _mm256_store_si256(work.add((j << 4) + 8).cast(), ff);
573 }
574 let mut even = [_mm256_setzero_si256(); 4];
575 let mut odd = [_mm256_setzero_si256(); 4];
576 let mut l = 0;
577 while l < 8 {
578 let mut j = 0;
579 while j < 4 {
580 let x = _mm256_loadu_si256(work.add((j << 4) + 8 - l).cast());
581 let y = _mm256_set1_epi32(*g.add(((i + j) << 3) + l) as i32);
582 even[j] = _mm256_add_epi64(even[j], _mm256_mul_epu32(x, y));
583 odd[j] = _mm256_add_epi64(odd[j], _mm256_mul_epu32(_mm256_bsrli_epi128::<4>(x), y));
584 j += 1;
585 }
586 l += 1;
587 }
588 let mut j = 0;
589 while j < 4 {
590 let x = reduce_sum_avx2(even[j], odd[j], r_vec, mod_vec);
591 store_block_avx2(f, i + j, _mm256_min_epu32(x, _mm256_sub_epi32(x, mod2_vec)));
592 j += 1;
593 }
594 i += 4;
595 rr = M::mod_mul(rr, M::INFO.rate3[(i >> 2).trailing_zeros() as usize]);
596 }
597}
598
599#[target_feature(enable = "avx2")]
600pub unsafe fn transform_blocks_avx2<M>(f: &mut [MInt<M>])
601where
602 M: Montgomery32NttModulus,
603{
604 let n = f.len() >> 3;
605 let f = f.as_mut_ptr() as *mut u32;
606 ntt_blocks_avx2::<M>(f, n);
607}
608
609#[target_feature(enable = "avx2")]
610pub unsafe fn multiply_blocks_avx2<M>(f: &mut [MInt<M>], g: &[MInt<M>])
611where
612 M: Montgomery32NttModulus,
613{
614 let n = f.len() >> 3;
615 let f = f.as_mut_ptr() as *mut u32;
616 let g = g.as_ptr() as *const u32;
617 convolve_8_avx2::<M>(f, g, n);
618}
619
620#[target_feature(enable = "avx2")]
621pub unsafe fn inverse_transform_blocks_avx2<M>(f: &mut [MInt<M>])
622where
623 M: Montgomery32NttModulus,
624{
625 let n = f.len() >> 3;
626 let f = f.as_mut_ptr() as *mut u32;
627 intt_blocks_avx2::<M>(f, n);
628}
629
630#[target_feature(enable = "avx2")]
631pub unsafe fn convolve_blocks_avx2<M>(f: &mut [MInt<M>], g: &mut [MInt<M>], same: bool)
632where
633 M: Montgomery32NttModulus,
634{
635 let n = f.len() >> 3;
636 let f = f.as_mut_ptr().cast();
637 let g = g.as_mut_ptr().cast();
638 ntt_blocks_avx2::<M>(f, n);
639 if same {
640 std::ptr::copy_nonoverlapping(f, g, n << 3);
641 } else {
642 ntt_blocks_avx2::<M>(g, n);
643 }
644 convolve_8_avx2::<M>(f, g, n);
645 intt_blocks_avx2::<M>(f, n);
646}