Skip to main content

montgomery_mul_256_fixed

Function montgomery_mul_256_fixed 

Source
pub unsafe fn montgomery_mul_256_fixed(
    a: __m256i,
    b: __m256i,
    b_r: __m256i,
    mod_vec: __m256i,
) -> __m256i
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform/ntt_simd/convolution_avx2.rs (lines 139-144)
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            // fw < 5 * MOD / 4, so the eight-product sum still reduces below 7 * MOD / 2.
562            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}