unsafe fn multiply(
a: *const u32,
b: *const u32,
c: *mut u32,
shape: (usize, usize, usize),
work: *mut u32,
kernel: &Kernel,
)Examples found in repository?
crates/competitive/src/num/mint/simd_matrix.rs (line 163)
134unsafe fn multiply(
135 a: *const u32,
136 b: *const u32,
137 c: *mut u32,
138 shape: (usize, usize, usize),
139 work: *mut u32,
140 kernel: &Kernel,
141) {
142 unsafe {
143 let (n, m, p) = shape;
144 let (modulus, inverse) = (kernel.modulus, kernel.inverse);
145 if n.min(m).min(p) <= 64 || n % 16 != 0 || m % 16 != 0 || p % 16 != 0 {
146 if kernel.avx512 {
147 if shape == (64, 64, 64) {
148 leaf_avx512::<true>(a, b, c, shape, modulus, inverse);
149 } else {
150 leaf_avx512::<false>(a, b, c, shape, modulus, inverse);
151 }
152 } else {
153 leaf_avx2(a, b, c, shape, modulus, inverse);
154 }
155 return;
156 }
157 let (n, m, p) = (n / 2, m / 2, p / 2);
158 let (s, t, q) = (work, work.add(n * m), work.add(n * m + m * p));
159 let work = q.add(n * p);
160 let (a00, a01, a10, a11) = (a, a.add(n * m), a.add(2 * n * m), a.add(3 * n * m));
161 let (b00, b01, b10, b11) = (b, b.add(m * p), b.add(2 * m * p), b.add(3 * m * p));
162 let (c00, c01, c10, c11) = (c, c.add(n * p), c.add(2 * n * p), c.add(3 * n * p));
163 multiply(a00, b00, c11, (n, m, p), work, kernel);
164 multiply(a01, b10, c00, (n, m, p), work, kernel);
165 combine::<false>(c00, c11, c00, n * p, modulus);
166 combine::<false>(a10, a11, s, n * m, modulus);
167 combine::<true>(b01, b00, t, m * p, modulus);
168 multiply(s, t, c01, (n, m, p), work, kernel);
169 combine::<true>(s, a00, s, n * m, modulus);
170 combine::<true>(b11, t, t, m * p, modulus);
171 multiply(s, t, c10, (n, m, p), work, kernel);
172 combine::<false>(c11, c10, c10, n * p, modulus);
173 combine::<true>(a01, s, s, n * m, modulus);
174 multiply(s, b11, q, (n, m, p), work, kernel);
175 combine::<false>(c10, c01, c11, n * p, modulus);
176 combine::<false>(c11, q, c01, n * p, modulus);
177 combine::<true>(t, b10, t, m * p, modulus);
178 multiply(a11, t, q, (n, m, p), work, kernel);
179 combine::<true>(c10, q, c10, n * p, modulus);
180 combine::<true>(a00, a10, s, n * m, modulus);
181 combine::<true>(b11, b01, t, m * p, modulus);
182 multiply(s, t, q, (n, m, p), work, kernel);
183 combine::<false>(c10, q, c10, n * p, modulus);
184 combine::<false>(c11, q, c11, n * p, modulus);
185 }
186}
187
188fn blocks(rows: usize, cols: usize, depth: usize) -> impl Iterator<Item = (usize, usize, usize)> {
189 (0..1usize << (2 * depth)).map(move |block| {
190 let (mut row, mut col) = (0, 0);
191 for bit in 0..depth {
192 row |= (block >> (2 * bit + 1) & 1) << bit;
193 col |= (block >> (2 * bit) & 1) << bit;
194 }
195 (
196 block * (rows >> depth) * (cols >> depth),
197 row * (rows >> depth),
198 col * (cols >> depth),
199 )
200 })
201}
202
203impl<M> MInt<M>
204where
205 M: MIntBase<Inner = u32>,
206{
207 /// # Safety
208 /// AVX2 must be available. The modulus must be odd, greater than one and below 2^30.
209 /// `scale` must be below the modulus; entries must have canonical raw representations.
210 #[target_feature(enable = "avx2")]
211 pub unsafe fn matrix_product_avx2(
212 a: &[Vec<Self>],
213 b: &[Vec<Self>],
214 scale: u32,
215 ) -> Vec<Vec<Self>> {
216 let (n, m, p) = (a.len(), b.len(), b.first().map_or(0, Vec::len));
217 assert!(a.iter().all(|row| row.len() == m));
218 assert!(b.iter().all(|row| row.len() == p));
219 let modulus = M::get_mod();
220 let alignment = if n.min(m).min(p) <= 64 { 8 } else { 32 };
221 let (nn, mm, pp) = (
222 n.div_ceil(alignment) * alignment,
223 m.div_ceil(alignment) * alignment,
224 p.div_ceil(alignment) * alignment,
225 );
226 let mut depth = 0;
227 let (mut x, mut y, mut z) = (nn, mm, pp);
228 while x.min(y).min(z) > 64 && x % 16 == 0 && y % 16 == 0 && z % 16 == 0 {
229 depth += 1;
230 x /= 2;
231 y /= 2;
232 z /= 2;
233 }
234 let entries = nn * mm + mm * pp + nn * pp;
235 // A recursive level uses one quarter of its parent's storage; siblings reuse it.
236 let mut data = if entries + entries / 3 >= 1 << 20 {
237 let mut data = Vec::with_capacity(entries + entries / 3);
238 advise_huge_pages(&mut data);
239 data.resize(entries + entries / 3, 0u32);
240 data
241 } else {
242 vec![0u32; entries + entries / 3]
243 };
244 let quotient = (((scale as u64) << 32) / modulus as u64) as u32;
245 let mut inverse = 1u32;
246 for _ in 0..5 {
247 inverse = inverse.wrapping_mul(2u32.wrapping_sub(modulus.wrapping_mul(inverse)));
248 }
249 let inverse = inverse.wrapping_neg();
250 for (offset, row, col) in blocks(nn, mm, depth) {
251 let (nr, nc) = (nn >> depth, mm >> depth);
252 for i in row..(row + nr).min(n) {
253 // SAFETY: MInt is transparent over u32. Reading raw words preserves Montgomery encoding.
254 let values: &[u32] = unsafe { std::slice::from_raw_parts(a[i].as_ptr().cast(), m) };
255 for j in col..(col + nc).min(m) {
256 let x = values[j];
257 let q = ((x as u64 * quotient as u64) >> 32) as u32;
258 let x = x.wrapping_mul(scale).wrapping_sub(q.wrapping_mul(modulus));
259 data[offset + (i - row) * nc + j - col] = x.min(x.wrapping_sub(modulus));
260 }
261 }
262 }
263 for (offset, row, col) in blocks(mm, pp, depth) {
264 let (nr, nc) = (mm >> depth, pp >> depth);
265 for i in row..(row + nr).min(m) {
266 // SAFETY: MInt is transparent over u32 and row lengths were checked above.
267 let values: &[u32] = unsafe { std::slice::from_raw_parts(b[i].as_ptr().cast(), p) };
268 for j in col..(col + nc).min(p) {
269 data[nn * mm + offset + (i - row) * nc + j - col] = values[j];
270 }
271 }
272 }
273 let kernel = Kernel {
274 modulus,
275 inverse,
276 avx512: avx512_enabled() && is_x86_feature_detected!("avx512f"),
277 };
278 // SAFETY: padding keeps every leaf dimension divisible by eight. The three matrices
279 // and the geometric scratch space are disjoint parts of the allocated buffer.
280 unsafe {
281 let ptr = data.as_mut_ptr();
282 multiply(
283 ptr,
284 ptr.add(nn * mm),
285 ptr.add(nn * mm + mm * pp),
286 (nn, mm, pp),
287 ptr.add(entries),
288 &kernel,
289 );
290 }
291 let mut result = vec![vec![MInt::new_unchecked(M::mod_zero()); p]; n];
292 for (offset, row, col) in blocks(nn, pp, depth) {
293 let (nr, nc) = (nn >> depth, pp >> depth);
294 for i in row..(row + nr).min(n) {
295 for j in col..(col + nc).min(p) {
296 result[i][j] = MInt::new_unchecked(
297 data[nn * mm + mm * pp + offset + (i - row) * nc + j - col],
298 );
299 }
300 }
301 }
302 result
303 }