const fn mod_mul(x: u32, y: u32, p: u32, r: u32) -> u32Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 71)
68const fn mod_pow(mut x: u32, mut y: u32, p: u32, r: u32, mut z: u32) -> u32 {
69 while y > 0 {
70 if y & 1 == 1 {
71 z = mod_mul(z, x, p, r);
72 }
73 x = mod_mul(x, x, p, r);
74 y >>= 1;
75 }
76 z
77}
78
79pub trait Montgomery32NttModulus: Sized + MontgomeryReduction32 {
80 const PRIMITIVE_ROOT: u32 = {
81 let mut g = 3u32;
82 loop {
83 let mut ok = true;
84 let mut d = 1u32;
85 while d * d < Self::MOD {
86 if (Self::MOD - 1) % d == 0 {
87 let ds = [d, (Self::MOD - 1) / d];
88 let mut i = 0;
89 while i < 2 {
90 ok &= ds[i] == Self::MOD - 1
91 || mod_pow(
92 reduce(g as u64 * Self::N2 as u64, Self::MOD, Self::R),
93 ds[i],
94 Self::MOD,
95 Self::R,
96 Self::N1,
97 ) != Self::N1;
98 i += 1;
99 }
100 }
101 d += 1;
102 }
103 if ok {
104 break;
105 }
106 g += 2;
107 }
108 g
109 };
110 const RANK: u32 = (Self::MOD - 1).trailing_zeros();
111 const INFO: NttInfo = NttInfo::new::<Self>();
112}
113
114#[derive(Debug, PartialEq)]
115pub struct NttInfo {
116 root: [u32; 32],
117 inv_root: [u32; 32],
118 rate3: [u32; 32],
119 rate3_packed: [[u32; 8]; 32],
120 inv_rate3_packed: [[u32; 8]; 32],
121}
122impl NttInfo {
123 const fn new<M>() -> Self
124 where
125 M: Montgomery32NttModulus,
126 {
127 let mut root = [0; 32];
128 let mut inv_root = [0; 32];
129 let mut rate3_values = [0; 32];
130 let mut rate3_packed = [[0; 8]; 32];
131 let mut inv_rate3_packed = [[0; 8]; 32];
132 let rank = M::RANK as usize;
133
134 let g = reduce(M::PRIMITIVE_ROOT as u64 * M::N2 as u64, M::MOD, M::R);
135 root[rank] = mod_pow(g, (M::MOD - 1) >> rank, M::MOD, M::R, M::N1);
136 inv_root[rank] = mod_pow(root[rank], M::MOD - 2, M::MOD, M::R, M::N1);
137 let mut i = rank - 1;
138 loop {
139 root[i] = mod_mul(root[i + 1], root[i + 1], M::MOD, M::R);
140 inv_root[i] = mod_mul(inv_root[i + 1], inv_root[i + 1], M::MOD, M::R);
141 if i == 0 {
142 break;
143 }
144 i -= 1;
145 }
146
147 let (mut i, mut prod, mut inv_prod) = (0, M::N1, M::N1);
148 while i < rank - 2 {
149 let rate3 = mod_mul(root[i + 3], prod, M::MOD, M::R);
150 rate3_values[i] = rate3;
151 let rate3_2 = mod_mul(rate3, rate3, M::MOD, M::R);
152 let rate3_3 = mod_mul(rate3_2, rate3, M::MOD, M::R);
153 let inv_rate3 = mod_mul(inv_root[i + 3], inv_prod, M::MOD, M::R);
154 let inv_rate3_2 = mod_mul(inv_rate3, inv_rate3, M::MOD, M::R);
155 let inv_rate3_3 = mod_mul(inv_rate3_2, inv_rate3, M::MOD, M::R);
156 rate3_packed[i] = [
157 rate3.wrapping_mul(M::R),
158 rate3,
159 rate3_2.wrapping_mul(M::R),
160 rate3_2,
161 rate3_3.wrapping_mul(M::R),
162 rate3_3,
163 0,
164 0,
165 ];
166 inv_rate3_packed[i] = [
167 inv_rate3.wrapping_mul(M::R),
168 inv_rate3,
169 inv_rate3_2.wrapping_mul(M::R),
170 inv_rate3_2,
171 inv_rate3_3.wrapping_mul(M::R),
172 inv_rate3_3,
173 0,
174 0,
175 ];
176 prod = mod_mul(prod, inv_root[i + 3], M::MOD, M::R);
177 inv_prod = mod_mul(inv_prod, root[i + 3], M::MOD, M::R);
178 i += 1;
179 }
180
181 NttInfo {
182 root,
183 inv_root,
184 rate3: rate3_values,
185 rate3_packed,
186 inv_rate3_packed,
187 }
188 }More examples
crates/competitive/src/math/number_theoretic_transform/ntt_simd/ntt_avx2.rs (line 42)
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}