1use crate::matrix::DenseMatrix;
2use std::fmt;
3use std::sync::Arc;
4
5#[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 pub dense_ops: u128,
136 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 pub sparse_candidate_products: u128,
144 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
200pub 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
285pub 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 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 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 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}