Skip to main content

competitive/num/mint/
simd_matrix.rs

1use super::{MInt, MIntBase, advise_huge_pages, avx512_enabled};
2use std::arch::x86_64::*;
3
4struct Kernel {
5    modulus: u32,
6    inverse: u32,
7    avx512: bool,
8}
9
10#[target_feature(enable = "avx2")]
11unsafe fn leaf_avx2(
12    a: *const u32,
13    b: *const u32,
14    c: *mut u32,
15    shape: (usize, usize, usize),
16    modulus: u32,
17    inverse: u32,
18) {
19    let (n, m, p) = shape;
20    unsafe {
21        // Eight products fit in u64 for modulus < 2^30. Subtracting 2*modulus*2^32
22        // keeps each accumulator below that bound without changing its residue.
23        let bound = _mm256_set1_epi64x(((2 * modulus as u64) << 32) as i64);
24        let mv = _mm256_set1_epi32(modulus as i32);
25        let rv = _mm256_set1_epi32(inverse as i32);
26        for i in (0..n).step_by(4) {
27            for j in (0..p).step_by(8) {
28                let mut lo = [_mm256_setzero_si256(); 4];
29                let mut hi = lo;
30                for first in (0..m).step_by(8) {
31                    for k in first..first + 8 {
32                        let b0 = _mm256_loadu_si256(b.add(k * p + j).cast());
33                        let b1 = _mm256_shuffle_epi32::<0xf5>(b0);
34                        for t in 0..4 {
35                            let factor = _mm256_set1_epi32(*a.add((i + t) * m + k) as i32);
36                            lo[t] = _mm256_add_epi64(lo[t], _mm256_mul_epu32(factor, b0));
37                            hi[t] = _mm256_add_epi64(hi[t], _mm256_mul_epu32(factor, b1));
38                        }
39                    }
40                    for t in 0..4 {
41                        lo[t] = _mm256_min_epu32(lo[t], _mm256_sub_epi32(lo[t], bound));
42                        hi[t] = _mm256_min_epu32(hi[t], _mm256_sub_epi32(hi[t], bound));
43                    }
44                }
45                for t in 0..4 {
46                    let x =
47                        _mm256_add_epi64(lo[t], _mm256_mul_epu32(_mm256_mul_epu32(lo[t], rv), mv));
48                    let y =
49                        _mm256_add_epi64(hi[t], _mm256_mul_epu32(_mm256_mul_epu32(hi[t], rv), mv));
50                    let x = _mm256_or_si256(_mm256_srli_epi64::<32>(x), y);
51                    let x = _mm256_min_epu32(x, _mm256_sub_epi32(x, mv));
52                    let x = _mm256_min_epu32(x, _mm256_sub_epi32(x, mv));
53                    _mm256_storeu_si256(c.add((i + t) * p + j).cast(), x);
54                }
55            }
56        }
57    }
58}
59
60#[target_feature(enable = "avx512f")]
61unsafe fn leaf_avx512<const FIXED: bool>(
62    a: *const u32,
63    b: *const u32,
64    c: *mut u32,
65    shape: (usize, usize, usize),
66    modulus: u32,
67    inverse: u32,
68) {
69    let (n, m, p) = if FIXED { (64, 64, 64) } else { shape };
70    unsafe {
71        let bound = _mm512_set1_epi64(((2 * modulus as u64) << 32) as i64);
72        let mv = _mm512_set1_epi32(modulus as i32);
73        let rv = _mm512_set1_epi32(inverse as i32);
74        for i in (0..n).step_by(8) {
75            for j in (0..p).step_by(16) {
76                let mask = if j + 16 <= p { 0xffff } else { 0xff };
77                let mut lo = [_mm512_setzero_si512(); 8];
78                let mut hi = lo;
79                for first in (0..m).step_by(8) {
80                    for k in first..first + 8 {
81                        let b0 = _mm512_maskz_loadu_epi32(mask, b.add(k * p + j).cast());
82                        let b1 = _mm512_srli_epi64::<32>(b0);
83                        for t in 0..8 {
84                            let factor = _mm512_set1_epi32(*a.add((i + t) * m + k) as i32);
85                            lo[t] = _mm512_add_epi64(lo[t], _mm512_mul_epu32(factor, b0));
86                            hi[t] = _mm512_add_epi64(hi[t], _mm512_mul_epu32(factor, b1));
87                        }
88                    }
89                    for t in 0..8 {
90                        lo[t] = _mm512_min_epu32(lo[t], _mm512_sub_epi32(lo[t], bound));
91                        hi[t] = _mm512_min_epu32(hi[t], _mm512_sub_epi32(hi[t], bound));
92                    }
93                }
94                for t in 0..8 {
95                    let x =
96                        _mm512_add_epi64(lo[t], _mm512_mul_epu32(_mm512_mul_epu32(lo[t], rv), mv));
97                    let y =
98                        _mm512_add_epi64(hi[t], _mm512_mul_epu32(_mm512_mul_epu32(hi[t], rv), mv));
99                    let x = _mm512_or_si512(_mm512_srli_epi64::<32>(x), y);
100                    let x = _mm512_min_epu32(x, _mm512_sub_epi32(x, mv));
101                    let x = _mm512_min_epu32(x, _mm512_sub_epi32(x, mv));
102                    _mm512_mask_storeu_epi32(c.add((i + t) * p + j).cast(), mask, x);
103                }
104            }
105        }
106    }
107}
108
109#[target_feature(enable = "avx2")]
110unsafe fn combine<const SUB: bool>(
111    a: *const u32,
112    b: *const u32,
113    c: *mut u32,
114    len: usize,
115    modulus: u32,
116) {
117    unsafe {
118        let modulus = _mm256_set1_epi32(modulus as i32);
119        for i in (0..len).step_by(8) {
120            let a = _mm256_loadu_si256(a.add(i).cast());
121            let b = _mm256_loadu_si256(b.add(i).cast());
122            let v = if SUB {
123                _mm256_sub_epi32(_mm256_add_epi32(a, modulus), b)
124            } else {
125                _mm256_add_epi32(a, b)
126            };
127            let v = _mm256_min_epu32(v, _mm256_sub_epi32(v, modulus));
128            _mm256_storeu_si256(c.add(i).cast(), v);
129        }
130    }
131}
132
133#[target_feature(enable = "avx2")]
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    }
304}