Skip to main content

blocks

Function blocks 

Source
fn blocks(
    rows: usize,
    cols: usize,
    depth: usize,
) -> impl Iterator<Item = (usize, usize, usize)>
Examples found in repository?
crates/competitive/src/num/mint/simd_matrix.rs (line 250)
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    }