Skip to main content

competitive/num/mint/
montgomery_simd.rs

1#![allow(unsafe_op_in_unsafe_fn)] // SIMD intrinsics and raw pointers are confined here
2use std::arch::x86_64::*;
3
4#[inline]
5#[target_feature(enable = "avx2")]
6pub unsafe fn montgomery_mul_256(
7    a: __m256i,
8    b: __m256i,
9    r_vec: __m256i,
10    mod_vec: __m256i,
11) -> __m256i {
12    let a13 = _mm256_bsrli_epi128::<4>(a);
13    let b13 = _mm256_bsrli_epi128::<4>(b);
14    let t02 = _mm256_mul_epu32(a, b);
15    let t13 = _mm256_mul_epu32(a13, b13);
16    let m02 = _mm256_mul_epu32(t02, r_vec);
17    let m13 = _mm256_mul_epu32(t13, r_vec);
18    let u02 = _mm256_add_epi64(t02, _mm256_mul_epu32(m02, mod_vec));
19    let u13 = _mm256_add_epi64(t13, _mm256_mul_epu32(m13, mod_vec));
20    _mm256_or_si256(_mm256_bsrli_epi128::<4>(u02), u13)
21}
22
23#[inline]
24#[target_feature(enable = "avx2")]
25pub unsafe fn montgomery_mul_256_fixed(
26    a: __m256i,
27    b: __m256i,
28    b_r: __m256i,
29    mod_vec: __m256i,
30) -> __m256i {
31    let a13 = _mm256_bsrli_epi128::<4>(a);
32    let t02 = _mm256_mul_epu32(a, b);
33    let t13 = _mm256_mul_epu32(a13, b);
34    let m02 = _mm256_mul_epu32(a, b_r);
35    let m13 = _mm256_mul_epu32(a13, b_r);
36    let u02 = _mm256_add_epi64(t02, _mm256_mul_epu32(m02, mod_vec));
37    let u13 = _mm256_add_epi64(t13, _mm256_mul_epu32(m13, mod_vec));
38    _mm256_or_si256(_mm256_bsrli_epi128::<4>(u02), u13)
39}
40
41#[inline]
42#[target_feature(enable = "avx2")]
43pub unsafe fn add_mod_256(a: __m256i, b: __m256i, mod_vec: __m256i) -> __m256i {
44    let sum = _mm256_add_epi32(a, b);
45    _mm256_min_epu32(sum, _mm256_sub_epi32(sum, mod_vec))
46}
47
48#[inline]
49#[target_feature(enable = "avx2")]
50pub unsafe fn sub_mod_256(a: __m256i, b: __m256i, mod_vec: __m256i) -> __m256i {
51    let diff = _mm256_sub_epi32(_mm256_add_epi32(a, mod_vec), b);
52    _mm256_min_epu32(diff, _mm256_sub_epi32(diff, mod_vec))
53}
54
55#[inline]
56#[target_feature(enable = "avx2")]
57pub unsafe fn montgomery_mul_256_canon(
58    a: __m256i,
59    b: __m256i,
60    r_vec: __m256i,
61    mod_vec: __m256i,
62) -> __m256i {
63    let x = montgomery_mul_256(a, b, r_vec, mod_vec);
64    _mm256_min_epu32(x, _mm256_sub_epi32(x, mod_vec))
65}
66
67#[inline]
68#[target_feature(enable = "avx2")]
69pub unsafe fn montgomery_add_256(a: __m256i, b: __m256i, mod2_vec: __m256i) -> __m256i {
70    let sum = _mm256_add_epi32(a, b);
71    _mm256_min_epu32(sum, _mm256_sub_epi32(sum, mod2_vec))
72}
73
74#[inline]
75#[target_feature(enable = "avx2")]
76pub unsafe fn montgomery_sub_256(a: __m256i, b: __m256i, mod2_vec: __m256i) -> __m256i {
77    let diff = _mm256_sub_epi32(_mm256_add_epi32(a, mod2_vec), b);
78    _mm256_min_epu32(diff, _mm256_sub_epi32(diff, mod2_vec))
79}
80
81#[inline]
82#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
83pub unsafe fn montgomery_mul_512(
84    a: __m512i,
85    b: __m512i,
86    r_vec: __m512i,
87    mod_vec: __m512i,
88) -> __m512i {
89    let a13 = _mm512_srli_epi64::<32>(a);
90    let b13 = _mm512_srli_epi64::<32>(b);
91    let t02 = _mm512_mul_epu32(a, b);
92    let t13 = _mm512_mul_epu32(a13, b13);
93    let m02 = _mm512_mul_epu32(t02, r_vec);
94    let m13 = _mm512_mul_epu32(t13, r_vec);
95    let u02 = _mm512_add_epi64(t02, _mm512_mul_epu32(m02, mod_vec));
96    let u13 = _mm512_add_epi64(t13, _mm512_mul_epu32(m13, mod_vec));
97    _mm512_or_si512(_mm512_srli_epi64::<32>(u02), u13)
98}
99
100#[inline]
101#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
102pub unsafe fn add_mod_512(a: __m512i, b: __m512i, mod_vec: __m512i) -> __m512i {
103    let sum = _mm512_add_epi32(a, b);
104    _mm512_min_epu32(sum, _mm512_sub_epi32(sum, mod_vec))
105}
106
107#[inline]
108#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
109pub unsafe fn sub_mod_512(a: __m512i, b: __m512i, mod_vec: __m512i) -> __m512i {
110    let diff = _mm512_sub_epi32(_mm512_add_epi32(a, mod_vec), b);
111    _mm512_min_epu32(diff, _mm512_sub_epi32(diff, mod_vec))
112}
113
114#[inline]
115#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
116pub unsafe fn montgomery_mul_512_canon(
117    a: __m512i,
118    b: __m512i,
119    r_vec: __m512i,
120    mod_vec: __m512i,
121) -> __m512i {
122    let x = montgomery_mul_512(a, b, r_vec, mod_vec);
123    _mm512_min_epu32(x, _mm512_sub_epi32(x, mod_vec))
124}
125
126#[inline]
127#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
128pub unsafe fn montgomery_add_512(a: __m512i, b: __m512i, mod2_vec: __m512i) -> __m512i {
129    let sum = _mm512_add_epi32(a, b);
130    _mm512_min_epu32(sum, _mm512_sub_epi32(sum, mod2_vec))
131}
132
133#[inline]
134#[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
135pub unsafe fn montgomery_sub_512(a: __m512i, b: __m512i, mod2_vec: __m512i) -> __m512i {
136    let diff = _mm512_sub_epi32(_mm512_add_epi32(a, mod2_vec), b);
137    _mm512_min_epu32(diff, _mm512_sub_epi32(diff, mod2_vec))
138}