Skip to main content

sketch_spgemm/
rect.rs

1use crate::matrix::DenseMatrix;
2use std::fmt;
3use std::sync::Arc;
4
5/// Rectangular multiplication policy for the compressed product (H*A)(B*G^T).
6#[derive(Clone, Copy, Debug, PartialEq, Eq)]
7pub enum RectangularPolicy {
8    Auto,
9    Dense,
10    SparseLeft,
11    SparseRight,
12    SparseSparse,
13}
14
15impl Default for RectangularPolicy {
16    fn default() -> Self {
17        Self::Auto
18    }
19}
20
21impl std::str::FromStr for RectangularPolicy {
22    type Err = String;
23
24    fn from_str(s: &str) -> Result<Self, Self::Err> {
25        match s {
26            "auto" => Ok(Self::Auto),
27            "dense" => Ok(Self::Dense),
28            "sparse-left" | "left" => Ok(Self::SparseLeft),
29            "sparse-right" | "right" => Ok(Self::SparseRight),
30            "sparse-sparse" | "sparse" => Ok(Self::SparseSparse),
31            _ => Err(format!(
32                "unknown rectangular policy {s:?}; expected auto|dense|sparse-left|sparse-right|sparse-sparse"
33            )),
34        }
35    }
36}
37
38impl fmt::Display for RectangularPolicy {
39    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
40        let s = match self {
41            Self::Auto => "auto",
42            Self::Dense => "dense",
43            Self::SparseLeft => "sparse-left",
44            Self::SparseRight => "sparse-right",
45            Self::SparseSparse => "sparse-sparse",
46        };
47        f.write_str(s)
48    }
49}
50
51#[derive(Clone, Copy, Debug, PartialEq, Eq)]
52pub enum RectangularKernel {
53    DenseBlocked,
54    SparseLeft,
55    SparseRight,
56    SparseSparse,
57}
58
59impl fmt::Display for RectangularKernel {
60    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
61        let s = match self {
62            Self::DenseBlocked => "dense-blocked",
63            Self::SparseLeft => "sparse-left",
64            Self::SparseRight => "sparse-right",
65            Self::SparseSparse => "sparse-sparse",
66        };
67        f.write_str(s)
68    }
69}
70
71#[derive(Clone, Debug)]
72pub struct PreparedFactor {
73    pub dense: Arc<DenseMatrix>,
74    pub sparse_rows: Arc<Vec<Vec<(usize, i64)>>>,
75    pub nnz: usize,
76    pub row_nnz: Arc<Vec<usize>>,
77    pub col_nnz: Arc<Vec<usize>>,
78}
79
80impl PreparedFactor {
81    pub fn new(matrix: DenseMatrix) -> Self {
82        Self::from_arc(Arc::new(matrix))
83    }
84
85    pub fn from_arc(dense: Arc<DenseMatrix>) -> Self {
86        let mut rows = Vec::with_capacity(dense.rows);
87        let mut row_nnz = vec![0usize; dense.rows];
88        let mut col_nnz = vec![0usize; dense.cols];
89        let mut nnz = 0usize;
90        for i in 0..dense.rows {
91            let mut row = Vec::new();
92            let base = i * dense.cols;
93            for j in 0..dense.cols {
94                let v = dense.data[base + j];
95                if v != 0 {
96                    row.push((j, v));
97                    row_nnz[i] += 1;
98                    col_nnz[j] += 1;
99                    nnz += 1;
100                }
101            }
102            rows.push(row);
103        }
104        Self {
105            dense,
106            sparse_rows: Arc::new(rows),
107            nnz,
108            row_nnz: Arc::new(row_nnz),
109            col_nnz: Arc::new(col_nnz),
110        }
111    }
112
113    #[inline]
114    pub fn rows(&self) -> usize {
115        self.dense.rows
116    }
117    #[inline]
118    pub fn cols(&self) -> usize {
119        self.dense.cols
120    }
121    #[inline]
122    pub fn density(&self) -> f64 {
123        self.nnz as f64 / self.rows().saturating_mul(self.cols()).max(1) as f64
124    }
125}
126
127#[derive(Clone, Debug, Default)]
128pub struct RectangularStats {
129    pub kernel: Option<RectangularKernel>,
130    pub a_nnz: usize,
131    pub b_nnz: usize,
132    pub a_density: f64,
133    pub b_density: f64,
134    /// m*n*g, i.e. multiply count of a truly dense GEMM.
135    pub dense_ops: u128,
136    /// Cost estimates used by auto dispatch. Sparse estimates include scans
137    /// needed to create their compact row views.
138    pub dense_estimated_cost: u128,
139    pub sparse_left_estimated_cost: u128,
140    pub sparse_right_estimated_cost: u128,
141    pub sparse_sparse_estimated_cost: u128,
142    /// `sum_k nnz(A[:, k]) * nnz(B[k, :])`.
143    pub sparse_candidate_products: u128,
144    /// Scalar multiplies actually issued by the selected kernel.
145    pub scalar_multiplications: u128,
146}
147
148#[derive(Clone, Debug)]
149struct FactorCounts {
150    a_nnz: usize,
151    b_nnz: usize,
152    a_col_nnz: Vec<usize>,
153    b_row_nnz: Vec<usize>,
154}
155
156impl FactorCounts {
157    fn build(a: &DenseMatrix, b: &DenseMatrix) -> Self {
158        assert_eq!(a.cols, b.rows);
159        let mut a_col_nnz = vec![0usize; a.cols];
160        let mut a_nnz = 0usize;
161        for i in 0..a.rows {
162            let base = i * a.cols;
163            for k in 0..a.cols {
164                if a.data[base + k] != 0 {
165                    a_col_nnz[k] += 1;
166                    a_nnz += 1;
167                }
168            }
169        }
170
171        let mut b_row_nnz = vec![0usize; b.rows];
172        let mut b_nnz = 0usize;
173        for (k, row_count) in b_row_nnz.iter_mut().enumerate() {
174            let base = k * b.cols;
175            for j in 0..b.cols {
176                if b.data[base + j] != 0 {
177                    *row_count += 1;
178                    b_nnz += 1;
179                }
180            }
181        }
182
183        Self {
184            a_nnz,
185            b_nnz,
186            a_col_nnz,
187            b_row_nnz,
188        }
189    }
190
191    fn sparse_candidate_products(&self) -> u128 {
192        self.a_col_nnz
193            .iter()
194            .zip(&self.b_row_nnz)
195            .map(|(&x, &y)| (x as u128) * (y as u128))
196            .sum()
197    }
198}
199
200/// Adaptive dense-output rectangular multiplication.
201///
202/// The outer sparse-recovery decoder wants a dense measurement matrix W, but
203/// its factors can be highly sparse. This routine keeps the output dense while
204/// selecting how to traverse the two input factors.
205pub fn adaptive_matmul(
206    a: &DenseMatrix,
207    b: &DenseMatrix,
208    policy: RectangularPolicy,
209) -> (DenseMatrix, RectangularStats) {
210    assert_eq!(a.cols, b.rows, "incompatible matrix dimensions");
211
212    let counts = FactorCounts::build(a, b);
213    let m = a.rows as u128;
214    let n = a.cols as u128;
215    let g = b.cols as u128;
216    let a_cells = m.saturating_mul(n);
217    let b_cells = n.saturating_mul(g);
218
219    let dense_ops = m.saturating_mul(n).saturating_mul(g);
220    let sparse_left_ops = (counts.a_nnz as u128).saturating_mul(g);
221    let sparse_right_ops = m.saturating_mul(counts.b_nnz as u128);
222    let sparse_sparse_ops = counts.sparse_candidate_products();
223
224    let dense_cost = dense_ops;
225    let sparse_left_cost = a_cells.saturating_add(sparse_left_ops);
226    let sparse_right_cost = b_cells.saturating_add(sparse_right_ops);
227    let sparse_sparse_cost = a_cells
228        .saturating_add(b_cells)
229        .saturating_add(sparse_sparse_ops);
230
231    let kernel = match policy {
232        RectangularPolicy::Dense => RectangularKernel::DenseBlocked,
233        RectangularPolicy::SparseLeft => RectangularKernel::SparseLeft,
234        RectangularPolicy::SparseRight => RectangularKernel::SparseRight,
235        RectangularPolicy::SparseSparse => RectangularKernel::SparseSparse,
236        RectangularPolicy::Auto => [
237            (dense_cost, RectangularKernel::DenseBlocked),
238            (sparse_left_cost, RectangularKernel::SparseLeft),
239            (sparse_right_cost, RectangularKernel::SparseRight),
240            (sparse_sparse_cost, RectangularKernel::SparseSparse),
241        ]
242        .into_iter()
243        .min_by_key(|(cost, _)| *cost)
244        .map(|(_, kernel)| kernel)
245        .unwrap(),
246    };
247
248    let (out, scalar_multiplications) = match kernel {
249        RectangularKernel::DenseBlocked => dense_blocked(a, b),
250        RectangularKernel::SparseLeft => {
251            let a_rows = sparse_rows(a);
252            sparse_left(a, b, &a_rows)
253        }
254        RectangularKernel::SparseRight => {
255            let b_rows = sparse_rows(b);
256            sparse_right(a, b, &b_rows)
257        }
258        RectangularKernel::SparseSparse => {
259            let a_rows = sparse_rows(a);
260            let b_rows = sparse_rows(b);
261            sparse_sparse(a, b, &a_rows, &b_rows)
262        }
263    };
264
265    let a_total = a.rows.saturating_mul(a.cols).max(1);
266    let b_total = b.rows.saturating_mul(b.cols).max(1);
267    let stats = RectangularStats {
268        kernel: Some(kernel),
269        a_nnz: counts.a_nnz,
270        b_nnz: counts.b_nnz,
271        a_density: counts.a_nnz as f64 / a_total as f64,
272        b_density: counts.b_nnz as f64 / b_total as f64,
273        dense_ops,
274        dense_estimated_cost: dense_cost,
275        sparse_left_estimated_cost: sparse_left_cost,
276        sparse_right_estimated_cost: sparse_right_cost,
277        sparse_sparse_estimated_cost: sparse_sparse_cost,
278        sparse_candidate_products: sparse_sparse_ops,
279        scalar_multiplications,
280    };
281
282    (out, stats)
283}
284
285/// Adaptive multiplication when dense factors and their sparse row views are
286/// already prepared/cached. Unlike `adaptive_matmul`, the auto cost model does
287/// not charge another full factor scan for sparse traversal. A small empirical
288/// penalty is applied to sparse×sparse pointer chasing so nearly-dense left
289/// factors tend to use the more cache-friendly sparse-right kernel.
290pub fn adaptive_matmul_prepared(
291    a: &PreparedFactor,
292    b: &PreparedFactor,
293    policy: RectangularPolicy,
294) -> (DenseMatrix, RectangularStats) {
295    assert_eq!(a.cols(), b.rows(), "incompatible matrix dimensions");
296    let m = a.rows() as u128;
297    let n = a.cols() as u128;
298    let g = b.cols() as u128;
299    let dense_ops = m.saturating_mul(n).saturating_mul(g);
300    let sparse_left_ops = (a.nnz as u128).saturating_mul(g);
301    let sparse_right_ops = m.saturating_mul(b.nnz as u128);
302    let sparse_sparse_ops: u128 = a
303        .col_nnz
304        .iter()
305        .zip(b.row_nnz.iter())
306        .map(|(&x, &y)| (x as u128) * (y as u128))
307        .sum();
308
309    let dense_cost = dense_ops;
310    let sparse_left_cost = sparse_left_ops;
311    let sparse_right_cost = sparse_right_ops;
312    // Prepared sparse views remove scan cost, but sparse×sparse has more
313    // irregular indirection. A 9/8 penalty is intentionally conservative and
314    // can still be beaten easily when both factors are truly sparse.
315    let sparse_sparse_cost = sparse_sparse_ops.saturating_mul(9) / 8;
316
317    let kernel = match policy {
318        RectangularPolicy::Dense => RectangularKernel::DenseBlocked,
319        RectangularPolicy::SparseLeft => RectangularKernel::SparseLeft,
320        RectangularPolicy::SparseRight => RectangularKernel::SparseRight,
321        RectangularPolicy::SparseSparse => RectangularKernel::SparseSparse,
322        RectangularPolicy::Auto => [
323            (dense_cost, RectangularKernel::DenseBlocked),
324            (sparse_left_cost, RectangularKernel::SparseLeft),
325            (sparse_right_cost, RectangularKernel::SparseRight),
326            (sparse_sparse_cost, RectangularKernel::SparseSparse),
327        ]
328        .into_iter()
329        .min_by_key(|(cost, _)| *cost)
330        .map(|(_, kernel)| kernel)
331        .unwrap(),
332    };
333
334    let (out, scalar_multiplications) = match kernel {
335        RectangularKernel::DenseBlocked => dense_blocked(&a.dense, &b.dense),
336        RectangularKernel::SparseLeft => sparse_left(&a.dense, &b.dense, &a.sparse_rows),
337        RectangularKernel::SparseRight => sparse_right(&a.dense, &b.dense, &b.sparse_rows),
338        RectangularKernel::SparseSparse => {
339            sparse_sparse(&a.dense, &b.dense, &a.sparse_rows, &b.sparse_rows)
340        }
341    };
342
343    let stats = RectangularStats {
344        kernel: Some(kernel),
345        a_nnz: a.nnz,
346        b_nnz: b.nnz,
347        a_density: a.density(),
348        b_density: b.density(),
349        dense_ops,
350        dense_estimated_cost: dense_cost,
351        sparse_left_estimated_cost: sparse_left_cost,
352        sparse_right_estimated_cost: sparse_right_cost,
353        sparse_sparse_estimated_cost: sparse_sparse_cost,
354        sparse_candidate_products: sparse_sparse_ops,
355        scalar_multiplications,
356    };
357    (out, stats)
358}
359
360fn sparse_rows(m: &DenseMatrix) -> Vec<Vec<(usize, i64)>> {
361    let mut rows = Vec::with_capacity(m.rows);
362    for i in 0..m.rows {
363        let mut row = Vec::new();
364        let base = i * m.cols;
365        for j in 0..m.cols {
366            let v = m.data[base + j];
367            if v != 0 {
368                row.push((j, v));
369            }
370        }
371        rows.push(row);
372    }
373    rows
374}
375
376fn dense_blocked(a: &DenseMatrix, b: &DenseMatrix) -> (DenseMatrix, u128) {
377    let mut out = DenseMatrix::zeros(a.rows, b.cols);
378    const BI: usize = 24;
379    const BK: usize = 32;
380    const BJ: usize = 64;
381
382    let mut ii = 0usize;
383    while ii < a.rows {
384        let i_end = (ii + BI).min(a.rows);
385        let mut kk = 0usize;
386        while kk < a.cols {
387            let k_end = (kk + BK).min(a.cols);
388            let mut jj = 0usize;
389            while jj < b.cols {
390                let j_end = (jj + BJ).min(b.cols);
391                for i in ii..i_end {
392                    let abase = i * a.cols;
393                    let obase = i * out.cols;
394                    for k in kk..k_end {
395                        let av = a.data[abase + k];
396                        let bbase = k * b.cols;
397                        for j in jj..j_end {
398                            out.data[obase + j] += av * b.data[bbase + j];
399                        }
400                    }
401                }
402                jj = j_end;
403            }
404            kk = k_end;
405        }
406        ii = i_end;
407    }
408
409    let ops = (a.rows as u128)
410        .saturating_mul(a.cols as u128)
411        .saturating_mul(b.cols as u128);
412    (out, ops)
413}
414
415fn sparse_left(
416    a: &DenseMatrix,
417    b: &DenseMatrix,
418    a_rows: &[Vec<(usize, i64)>],
419) -> (DenseMatrix, u128) {
420    let mut out = DenseMatrix::zeros(a.rows, b.cols);
421    let mut ops = 0u128;
422    for (i, row) in a_rows.iter().enumerate() {
423        let obase = i * out.cols;
424        for &(k, av) in row {
425            let bbase = k * b.cols;
426            for j in 0..b.cols {
427                out.data[obase + j] += av * b.data[bbase + j];
428                ops += 1;
429            }
430        }
431    }
432    (out, ops)
433}
434
435fn sparse_right(
436    a: &DenseMatrix,
437    b: &DenseMatrix,
438    b_rows: &[Vec<(usize, i64)>],
439) -> (DenseMatrix, u128) {
440    let mut out = DenseMatrix::zeros(a.rows, b.cols);
441    let mut ops = 0u128;
442    for i in 0..a.rows {
443        let abase = i * a.cols;
444        let obase = i * out.cols;
445        for (k, row) in b_rows.iter().enumerate() {
446            let av = a.data[abase + k];
447            for &(j, bv) in row {
448                out.data[obase + j] += av * bv;
449                ops += 1;
450            }
451        }
452    }
453    (out, ops)
454}
455
456fn sparse_sparse(
457    a: &DenseMatrix,
458    b: &DenseMatrix,
459    a_rows: &[Vec<(usize, i64)>],
460    b_rows: &[Vec<(usize, i64)>],
461) -> (DenseMatrix, u128) {
462    let mut out = DenseMatrix::zeros(a.rows, b.cols);
463    let mut ops = 0u128;
464    for (i, row) in a_rows.iter().enumerate() {
465        let obase = i * out.cols;
466        for &(k, av) in row {
467            for &(j, bv) in &b_rows[k] {
468                out.data[obase + j] += av * bv;
469                ops += 1;
470            }
471        }
472    }
473    (out, ops)
474}
475
476#[cfg(test)]
477mod tests {
478    use super::*;
479
480    fn sample() -> (DenseMatrix, DenseMatrix) {
481        let mut a = DenseMatrix::zeros(3, 4);
482        a[(0, 0)] = 2;
483        a[(0, 3)] = 1;
484        a[(1, 1)] = -3;
485        a[(2, 2)] = 5;
486
487        let mut b = DenseMatrix::zeros(4, 3);
488        b[(0, 0)] = 7;
489        b[(1, 1)] = 11;
490        b[(2, 2)] = 13;
491        b[(3, 0)] = 17;
492        b[(3, 2)] = 19;
493        (a, b)
494    }
495
496    #[test]
497    fn all_kernels_are_exactly_equivalent() {
498        let (a, b) = sample();
499        let (reference, _) = adaptive_matmul(&a, &b, RectangularPolicy::Dense);
500        for policy in [
501            RectangularPolicy::SparseLeft,
502            RectangularPolicy::SparseRight,
503            RectangularPolicy::SparseSparse,
504            RectangularPolicy::Auto,
505        ] {
506            let (actual, _) = adaptive_matmul(&a, &b, policy);
507            assert_eq!(actual, reference, "policy={policy}");
508        }
509    }
510
511    #[test]
512    fn auto_accounts_for_sparse_view_build_cost() {
513        let mut a = DenseMatrix::zeros(128, 128);
514        let mut b = DenseMatrix::zeros(128, 128);
515        for i in 0..8 {
516            a[(i, i)] = 1;
517            b[(i, i)] = 1;
518        }
519
520        // Sparse×sparse would issue only 8 multiplies, but adaptive_matmul must
521        // first scan both dense factor buffers to construct sparse row views.
522        // Sparse-left scans only A and then performs 8*128 = 1024 multiplies,
523        // which is cheaper under the current one-shot cost model:
524        //   sparse-left  = 128*128 + 8*128       = 17_408
525        //   sparse-sparse= 128*128 + 128*128 + 8 = 32_776
526        let (_, stats) = adaptive_matmul(&a, &b, RectangularPolicy::Auto);
527        assert_eq!(stats.kernel, Some(RectangularKernel::SparseLeft));
528        assert_eq!(stats.sparse_candidate_products, 8);
529        assert_eq!(stats.scalar_multiplications, 8 * 128);
530        assert_eq!(stats.sparse_left_estimated_cost, 17_408);
531        assert_eq!(stats.sparse_sparse_estimated_cost, 32_776);
532        assert!(stats.sparse_left_estimated_cost < stats.sparse_sparse_estimated_cost);
533
534        // The forced sparse×sparse kernel is still exact and still performs the
535        // minimum 8 scalar products; it just is not the cheapest one-shot path
536        // after sparse-view construction is included.
537        let (_, forced) = adaptive_matmul(&a, &b, RectangularPolicy::SparseSparse);
538        assert_eq!(forced.scalar_multiplications, 8);
539    }
540
541    #[test]
542    fn auto_dispatches_all_density_extremes() {
543        let mut dense_a = DenseMatrix::zeros(32, 32);
544        let mut dense_b = DenseMatrix::zeros(32, 32);
545        for i in 0..32 {
546            for j in 0..32 {
547                dense_a[(i, j)] = 1;
548                dense_b[(i, j)] = 1;
549            }
550        }
551        let (_, dense_stats) = adaptive_matmul(&dense_a, &dense_b, RectangularPolicy::Auto);
552        assert_eq!(dense_stats.kernel, Some(RectangularKernel::DenseBlocked));
553
554        let mut sparse_a = DenseMatrix::zeros(128, 128);
555        let mut dense_b = DenseMatrix::zeros(128, 128);
556        for i in 0..8 {
557            sparse_a[(i, i)] = 1;
558        }
559        for i in 0..128 {
560            for j in 0..128 {
561                dense_b[(i, j)] = 1;
562            }
563        }
564        let (_, left_stats) = adaptive_matmul(&sparse_a, &dense_b, RectangularPolicy::Auto);
565        assert_eq!(left_stats.kernel, Some(RectangularKernel::SparseLeft));
566
567        let mut dense_a = DenseMatrix::zeros(128, 128);
568        let mut sparse_b = DenseMatrix::zeros(128, 128);
569        for i in 0..128 {
570            for j in 0..128 {
571                dense_a[(i, j)] = 1;
572            }
573        }
574        for i in 0..8 {
575            sparse_b[(i, i)] = 1;
576        }
577        let (_, right_stats) = adaptive_matmul(&dense_a, &sparse_b, RectangularPolicy::Auto);
578        assert_eq!(right_stats.kernel, Some(RectangularKernel::SparseRight));
579    }
580
581    #[test]
582    fn forced_kernels_report_expected_multiplication_counts() {
583        let (a, b) = sample();
584        let (_, left) = adaptive_matmul(&a, &b, RectangularPolicy::SparseLeft);
585        let (_, right) = adaptive_matmul(&a, &b, RectangularPolicy::SparseRight);
586        let (_, ss) = adaptive_matmul(&a, &b, RectangularPolicy::SparseSparse);
587        assert_eq!(left.scalar_multiplications, (left.a_nnz * b.cols) as u128);
588        assert_eq!(right.scalar_multiplications, (a.rows * right.b_nnz) as u128);
589        assert_eq!(ss.scalar_multiplications, ss.sparse_candidate_products);
590    }
591
592    #[test]
593    fn prepared_factors_avoid_scan_cost_and_keep_exactness() {
594        let (a, b) = sample();
595        let pa = PreparedFactor::new(a.clone());
596        let pb = PreparedFactor::new(b.clone());
597        let (prepared, stats) = adaptive_matmul_prepared(&pa, &pb, RectangularPolicy::Auto);
598        let (reference, _) = adaptive_matmul(&a, &b, RectangularPolicy::Dense);
599        assert_eq!(prepared, reference);
600        assert_eq!(stats.a_nnz, a.nnz());
601        assert_eq!(stats.b_nnz, b.nnz());
602    }
603}