const fn reduce(z: u64, p: u32, r: u32) -> u32Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 66)
65const fn mod_mul(x: u32, y: u32, p: u32, r: u32) -> u32 {
66 reduce(x as u64 * y as u64, p, r)
67}
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 }