Skip to main content

leaf_avx512

Function leaf_avx512 

Source
unsafe fn leaf_avx512<const FIXED: bool>(
    a: *const u32,
    b: *const u32,
    c: *mut u32,
    shape: (usize, usize, usize),
    modulus: u32,
    inverse: u32,
)
Examples found in repository?
crates/competitive/src/num/mint/simd_matrix.rs (line 148)
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}