1use std::collections::HashMap;
43
44use serde::{Deserialize, Serialize};
45
46use crate::error::{KernelError, Result};
47use crate::types::Kernel;
48
49#[derive(Clone, Debug, Serialize, Deserialize)]
67pub struct SparseKernelMatrix {
68 size: usize,
70 row_ptr: Vec<usize>,
72 col_idx: Vec<usize>,
74 values: Vec<f64>,
76 #[serde(skip)]
78 temp_map: HashMap<(usize, usize), f64>,
79}
80
81impl SparseKernelMatrix {
82 pub fn new(size: usize) -> Self {
84 Self {
85 size,
86 row_ptr: vec![0; size + 1],
87 col_idx: Vec::new(),
88 values: Vec::new(),
89 temp_map: HashMap::new(),
90 }
91 }
92
93 pub fn set(&mut self, row: usize, col: usize, value: f64) {
95 if row >= self.size || col >= self.size {
96 return;
97 }
98
99 if value.abs() < 1e-10 {
100 self.temp_map.remove(&(row, col));
102 } else {
103 self.temp_map.insert((row, col), value);
104 }
105 }
106
107 pub fn get(&self, row: usize, col: usize) -> Option<f64> {
109 if row >= self.size || col >= self.size {
110 return None;
111 }
112
113 if let Some(&value) = self.temp_map.get(&(row, col)) {
115 return Some(value);
116 }
117
118 let start = self.row_ptr[row];
120 let end = self.row_ptr[row + 1];
121
122 for i in start..end {
123 if self.col_idx[i] == col {
124 return Some(self.values[i]);
125 }
126 }
127
128 None
129 }
130
131 pub fn finalize(&mut self) {
133 if self.temp_map.is_empty() {
134 return;
135 }
136
137 self.col_idx.clear();
139 self.values.clear();
140 self.row_ptr = vec![0; self.size + 1];
141
142 let mut entries: Vec<_> = self.temp_map.iter().collect();
144 entries.sort_by_key(|&((row, col), _)| (*row, *col));
145
146 let mut current_row = 0;
148 for (&(row, col), &value) in &entries {
149 while current_row < row {
151 current_row += 1;
152 self.row_ptr[current_row] = self.col_idx.len();
153 }
154
155 self.col_idx.push(col);
156 self.values.push(value);
157 }
158
159 while current_row < self.size {
161 current_row += 1;
162 self.row_ptr[current_row] = self.col_idx.len();
163 }
164
165 self.temp_map.clear();
167 }
168
169 pub fn nnz(&self) -> usize {
171 self.values.len() + self.temp_map.len()
172 }
173
174 pub fn size(&self) -> usize {
176 self.size
177 }
178
179 pub fn density(&self) -> f64 {
181 let total = self.size * self.size;
182 if total == 0 {
183 0.0
184 } else {
185 self.nnz() as f64 / total as f64
186 }
187 }
188
189 #[allow(clippy::needless_range_loop)]
191 pub fn to_dense(&mut self) -> Vec<Vec<f64>> {
192 self.finalize();
193
194 let mut dense = vec![vec![0.0; self.size]; self.size];
195
196 for row in 0..self.size {
197 let start = self.row_ptr[row];
198 let end = self.row_ptr[row + 1];
199
200 for i in start..end {
201 let col = self.col_idx[i];
202 let value = self.values[i];
203 dense[row][col] = value;
204 }
205 }
206
207 dense
208 }
209
210 pub fn from_kernel_with_threshold(
212 data: &[Vec<f64>],
213 kernel: &dyn Kernel,
214 threshold: f64,
215 ) -> Result<Self> {
216 let n = data.len();
217 let mut matrix = Self::new(n);
218
219 for i in 0..n {
220 for j in 0..n {
221 let value = kernel.compute(&data[i], &data[j])?;
222 if value.abs() >= threshold {
223 matrix.set(i, j, value);
224 }
225 }
226 }
227
228 matrix.finalize();
229 Ok(matrix)
230 }
231
232 pub fn row(&mut self, row_idx: usize) -> Option<Vec<(usize, f64)>> {
234 if row_idx >= self.size {
235 return None;
236 }
237
238 self.finalize();
239
240 let start = self.row_ptr[row_idx];
241 let end = self.row_ptr[row_idx + 1];
242
243 let mut row_data = Vec::new();
244 for i in start..end {
245 row_data.push((self.col_idx[i], self.values[i]));
246 }
247
248 Some(row_data)
249 }
250}
251
252pub struct SparseKernelMatrixBuilder {
254 threshold: f64,
256 max_entries_per_row: Option<usize>,
258}
259
260impl SparseKernelMatrixBuilder {
261 pub fn new() -> Self {
263 Self {
264 threshold: 1e-10,
265 max_entries_per_row: None,
266 }
267 }
268
269 pub fn with_threshold(mut self, threshold: f64) -> Result<Self> {
271 if threshold < 0.0 {
272 return Err(KernelError::InvalidParameter {
273 parameter: "threshold".to_string(),
274 value: threshold.to_string(),
275 reason: "must be non-negative".to_string(),
276 });
277 }
278 self.threshold = threshold;
279 Ok(self)
280 }
281
282 pub fn with_max_entries_per_row(mut self, max_entries: usize) -> Result<Self> {
284 if max_entries == 0 {
285 return Err(KernelError::InvalidParameter {
286 parameter: "max_entries_per_row".to_string(),
287 value: max_entries.to_string(),
288 reason: "must be positive".to_string(),
289 });
290 }
291 self.max_entries_per_row = Some(max_entries);
292 Ok(self)
293 }
294
295 pub fn build(&self, data: &[Vec<f64>], kernel: &dyn Kernel) -> Result<SparseKernelMatrix> {
297 let n = data.len();
298 let mut matrix = SparseKernelMatrix::new(n);
299
300 for i in 0..n {
301 let mut row_entries = Vec::new();
302
303 for j in 0..n {
305 let value = kernel.compute(&data[i], &data[j])?;
306 if value.abs() >= self.threshold {
307 row_entries.push((j, value));
308 }
309 }
310
311 if let Some(max_entries) = self.max_entries_per_row {
313 if row_entries.len() > max_entries {
314 row_entries.sort_by(|(_, a), (_, b)| {
316 b.abs()
317 .partial_cmp(&a.abs())
318 .unwrap_or(std::cmp::Ordering::Equal)
319 });
320 row_entries.truncate(max_entries);
321 }
322 }
323
324 for (j, value) in row_entries {
326 matrix.set(i, j, value);
327 }
328 }
329
330 matrix.finalize();
331 Ok(matrix)
332 }
333}
334
335impl Default for SparseKernelMatrixBuilder {
336 fn default() -> Self {
337 Self::new()
338 }
339}
340
341impl SparseKernelMatrix {
343 pub fn spmv(&mut self, x: &[f64]) -> Result<Vec<f64>> {
345 if x.len() != self.size {
346 return Err(KernelError::InvalidParameter {
347 parameter: "x".to_string(),
348 value: x.len().to_string(),
349 reason: format!("vector length must match matrix size {}", self.size),
350 });
351 }
352
353 self.finalize();
354
355 let mut y = vec![0.0; self.size];
356
357 for (row, y_elem) in y.iter_mut().enumerate() {
358 let start = self.row_ptr[row];
359 let end = self.row_ptr[row + 1];
360
361 let mut sum = 0.0;
362 for i in start..end {
363 let col = self.col_idx[i];
364 let value = self.values[i];
365 sum += value * x[col];
366 }
367 *y_elem = sum;
368 }
369
370 Ok(y)
371 }
372
373 pub fn transpose(&self) -> Result<Self> {
375 let mut transposed = Self::new(self.size);
376
377 for row in 0..self.size {
378 let start = self.row_ptr[row];
379 let end = self.row_ptr[row + 1];
380
381 for i in start..end {
382 let col = self.col_idx[i];
383 let value = self.values[i];
384 transposed.set(col, row, value);
385 }
386 }
387
388 transposed.finalize();
389 Ok(transposed)
390 }
391
392 pub fn add(&mut self, other: &Self) -> Result<Self> {
394 if self.size != other.size {
395 return Err(KernelError::InvalidParameter {
396 parameter: "other".to_string(),
397 value: other.size.to_string(),
398 reason: format!("matrix sizes must match: {} vs {}", self.size, other.size),
399 });
400 }
401
402 self.finalize();
403
404 let mut other_finalized = other.clone();
406 other_finalized.finalize();
407
408 let mut result = Self::new(self.size);
409
410 for row in 0..self.size {
412 let start = self.row_ptr[row];
413 let end = self.row_ptr[row + 1];
414
415 for i in start..end {
416 let col = self.col_idx[i];
417 let value = self.values[i];
418 result.set(row, col, value);
419 }
420 }
421
422 for row in 0..other_finalized.size {
424 let start = other_finalized.row_ptr[row];
425 let end = other_finalized.row_ptr[row + 1];
426
427 for i in start..end {
428 let col = other_finalized.col_idx[i];
429 let value = other_finalized.values[i];
430 let existing = result.get(row, col).unwrap_or(0.0);
431 result.set(row, col, existing + value);
432 }
433 }
434
435 result.finalize();
436 Ok(result)
437 }
438
439 pub fn frobenius_norm(&self) -> f64 {
441 let mut sum_squares = 0.0;
442
443 for row in 0..self.size {
444 let start = self.row_ptr[row];
445 let end = self.row_ptr[row + 1];
446
447 for i in start..end {
448 let value = self.values[i];
449 sum_squares += value * value;
450 }
451 }
452
453 sum_squares.sqrt()
454 }
455
456 pub fn iter_nonzeros(&mut self) -> SparseMatrixIterator<'_> {
458 self.finalize();
459 SparseMatrixIterator {
460 matrix: self,
461 current_row: 0,
462 current_idx: 0,
463 }
464 }
465
466 pub fn scale(&mut self, scalar: f64) {
468 for value in &mut self.values {
469 *value *= scalar;
470 }
471
472 for value in self.temp_map.values_mut() {
473 *value *= scalar;
474 }
475 }
476}
477
478pub struct SparseMatrixIterator<'a> {
480 matrix: &'a SparseKernelMatrix,
481 current_row: usize,
482 current_idx: usize,
483}
484
485impl<'a> Iterator for SparseMatrixIterator<'a> {
486 type Item = (usize, usize, f64);
487
488 fn next(&mut self) -> Option<Self::Item> {
489 while self.current_row < self.matrix.size {
490 let row_end = self.matrix.row_ptr[self.current_row + 1];
491
492 if self.current_idx < row_end {
493 let col = self.matrix.col_idx[self.current_idx];
494 let value = self.matrix.values[self.current_idx];
495 self.current_idx += 1;
496 return Some((self.current_row, col, value));
497 }
498
499 self.current_row += 1;
500 self.current_idx = self
501 .matrix
502 .row_ptr
503 .get(self.current_row)
504 .copied()
505 .unwrap_or(0);
506 }
507
508 None
509 }
510}
511
512impl SparseKernelMatrixBuilder {
514 pub fn build_parallel(
516 &self,
517 data: &[Vec<f64>],
518 kernel: &dyn Kernel,
519 ) -> Result<SparseKernelMatrix> {
520 use rayon::prelude::*;
521
522 let n = data.len();
523 let mut matrix = SparseKernelMatrix::new(n);
524
525 let row_data: Vec<Vec<(usize, f64)>> = (0..n)
527 .into_par_iter()
528 .map(|i| {
529 let mut row_entries = Vec::new();
530
531 for j in 0..n {
532 match kernel.compute(&data[i], &data[j]) {
533 Ok(value) => {
534 if value.abs() >= self.threshold {
535 row_entries.push((j, value));
536 }
537 }
538 Err(_) => continue,
539 }
540 }
541
542 if let Some(max_entries) = self.max_entries_per_row {
544 if row_entries.len() > max_entries {
545 row_entries.sort_by(|(_, a), (_, b)| {
546 b.abs()
547 .partial_cmp(&a.abs())
548 .unwrap_or(std::cmp::Ordering::Equal)
549 });
550 row_entries.truncate(max_entries);
551 }
552 }
553
554 row_entries
555 })
556 .collect();
557
558 for (i, row_entries) in row_data.into_iter().enumerate() {
560 for (j, value) in row_entries {
561 matrix.set(i, j, value);
562 }
563 }
564
565 matrix.finalize();
566 Ok(matrix)
567 }
568}
569
570#[cfg(test)]
571mod tests {
572 use super::*;
573 use crate::tensor_kernels::LinearKernel;
574
575 #[test]
576 fn test_sparse_matrix_creation() {
577 let matrix = SparseKernelMatrix::new(3);
578 assert_eq!(matrix.size(), 3);
579 assert_eq!(matrix.nnz(), 0);
580 }
581
582 #[test]
583 fn test_sparse_matrix_set_get() {
584 let mut matrix = SparseKernelMatrix::new(3);
585 matrix.set(0, 1, 0.8);
586 matrix.set(1, 2, 0.6);
587
588 assert_eq!(matrix.get(0, 1), Some(0.8));
589 assert_eq!(matrix.get(1, 2), Some(0.6));
590 assert_eq!(matrix.get(0, 2), None);
591 }
592
593 #[test]
594 fn test_sparse_matrix_finalize() {
595 let mut matrix = SparseKernelMatrix::new(3);
596 matrix.set(0, 1, 0.8);
597 matrix.set(1, 2, 0.6);
598 matrix.set(2, 0, 0.4);
599
600 matrix.finalize();
601
602 assert_eq!(matrix.get(0, 1), Some(0.8));
603 assert_eq!(matrix.get(1, 2), Some(0.6));
604 assert_eq!(matrix.get(2, 0), Some(0.4));
605 }
606
607 #[test]
608 fn test_sparse_matrix_nnz() {
609 let mut matrix = SparseKernelMatrix::new(3);
610 matrix.set(0, 1, 0.8);
611 matrix.set(1, 2, 0.6);
612
613 assert_eq!(matrix.nnz(), 2);
614 }
615
616 #[test]
617 fn test_sparse_matrix_density() {
618 let mut matrix = SparseKernelMatrix::new(3);
619 matrix.set(0, 1, 0.8);
620 matrix.set(1, 2, 0.6);
621
622 let density = matrix.density();
623 assert!((density - 2.0 / 9.0).abs() < 1e-10);
624 }
625
626 #[test]
627 fn test_sparse_matrix_to_dense() {
628 let mut matrix = SparseKernelMatrix::new(3);
629 matrix.set(0, 1, 0.8);
630 matrix.set(1, 2, 0.6);
631
632 let dense = matrix.to_dense();
633 assert_eq!(dense.len(), 3);
634 assert!((dense[0][1] - 0.8).abs() < 1e-10);
635 assert!((dense[1][2] - 0.6).abs() < 1e-10);
636 assert!(dense[0][0].abs() < 1e-10);
637 }
638
639 #[test]
640 fn test_sparse_matrix_from_kernel() {
641 let kernel = LinearKernel::new();
642 let data = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![0.5, 0.5]];
643
644 let mut matrix =
645 SparseKernelMatrix::from_kernel_with_threshold(&data, &kernel, 0.1).expect("unwrap");
646
647 assert!(matrix.nnz() > 0);
648 let dense = matrix.to_dense();
649 assert_eq!(dense.len(), 3);
650 }
651
652 #[test]
653 fn test_sparse_matrix_row() {
654 let mut matrix = SparseKernelMatrix::new(3);
655 matrix.set(0, 1, 0.8);
656 matrix.set(0, 2, 0.6);
657
658 let row = matrix.row(0).expect("unwrap");
659 assert_eq!(row.len(), 2);
660 assert!(row.contains(&(1, 0.8)));
661 assert!(row.contains(&(2, 0.6)));
662 }
663
664 #[test]
665 fn test_sparse_matrix_builder() {
666 let builder = SparseKernelMatrixBuilder::new();
667 let kernel = LinearKernel::new();
668 let data = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
669
670 let matrix = builder.build(&data, &kernel).expect("unwrap");
671 assert!(matrix.nnz() > 0);
672 }
673
674 #[test]
675 fn test_sparse_matrix_builder_with_threshold() {
676 let builder = SparseKernelMatrixBuilder::new()
677 .with_threshold(0.5)
678 .expect("unwrap");
679 let kernel = LinearKernel::new();
680 let data = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
681
682 let matrix = builder.build(&data, &kernel).expect("unwrap");
683 assert!(matrix.nnz() > 0);
684 }
685
686 #[test]
687 fn test_sparse_matrix_builder_invalid_threshold() {
688 let result = SparseKernelMatrixBuilder::new().with_threshold(-0.1);
689 assert!(result.is_err());
690 }
691
692 #[test]
693 fn test_sparse_matrix_builder_max_entries() {
694 let builder = SparseKernelMatrixBuilder::new()
695 .with_max_entries_per_row(2)
696 .expect("unwrap");
697 let kernel = LinearKernel::new();
698 let data = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![0.5, 0.5]];
699
700 let matrix = builder.build(&data, &kernel).expect("unwrap");
701 for i in 0..matrix.size() {
703 let mut temp_matrix = matrix.clone();
704 let row = temp_matrix.row(i).expect("unwrap");
705 assert!(row.len() <= 2);
706 }
707 }
708
709 #[test]
710 fn test_sparse_matrix_builder_invalid_max_entries() {
711 let result = SparseKernelMatrixBuilder::new().with_max_entries_per_row(0);
712 assert!(result.is_err());
713 }
714
715 #[test]
716 fn test_sparse_matrix_zero_threshold() {
717 let mut matrix = SparseKernelMatrix::new(3);
718 matrix.set(0, 1, 1e-11); matrix.finalize();
720
721 assert_eq!(matrix.nnz(), 0);
723 }
724
725 #[test]
726 fn test_sparse_matrix_spmv() {
727 let mut matrix = SparseKernelMatrix::new(3);
728 matrix.set(0, 0, 2.0);
729 matrix.set(0, 2, 1.0);
730 matrix.set(1, 1, 3.0);
731 matrix.set(2, 0, 1.0);
732 matrix.set(2, 2, 2.0);
733
734 let x = vec![1.0, 2.0, 3.0];
735 let y = matrix.spmv(&x).expect("unwrap");
736
737 assert_eq!(y.len(), 3);
738 assert!((y[0] - 5.0).abs() < 1e-10); assert!((y[1] - 6.0).abs() < 1e-10); assert!((y[2] - 7.0).abs() < 1e-10); }
742
743 #[test]
744 fn test_sparse_matrix_spmv_invalid_size() {
745 let mut matrix = SparseKernelMatrix::new(3);
746 matrix.set(0, 0, 1.0);
747
748 let x = vec![1.0, 2.0]; let result = matrix.spmv(&x);
750 assert!(result.is_err());
751 }
752
753 #[test]
754 fn test_sparse_matrix_transpose() {
755 let mut matrix = SparseKernelMatrix::new(3);
756 matrix.set(0, 1, 0.8);
757 matrix.set(1, 2, 0.6);
758 matrix.set(2, 0, 0.4);
759 matrix.finalize();
760
761 let transposed = matrix.transpose().expect("unwrap");
762
763 assert_eq!(transposed.get(1, 0), Some(0.8));
764 assert_eq!(transposed.get(2, 1), Some(0.6));
765 assert_eq!(transposed.get(0, 2), Some(0.4));
766 }
767
768 #[test]
769 fn test_sparse_matrix_add() {
770 let mut matrix1 = SparseKernelMatrix::new(3);
771 matrix1.set(0, 0, 1.0);
772 matrix1.set(0, 1, 2.0);
773 matrix1.set(1, 1, 3.0);
774
775 let mut matrix2 = SparseKernelMatrix::new(3);
776 matrix2.set(0, 1, 1.0);
777 matrix2.set(1, 2, 4.0);
778 matrix2.set(2, 2, 5.0);
779
780 let result = matrix1.add(&matrix2).expect("unwrap");
781
782 assert_eq!(result.get(0, 0), Some(1.0));
783 assert_eq!(result.get(0, 1), Some(3.0)); assert_eq!(result.get(1, 1), Some(3.0));
785 assert_eq!(result.get(1, 2), Some(4.0));
786 assert_eq!(result.get(2, 2), Some(5.0));
787 }
788
789 #[test]
790 fn test_sparse_matrix_add_invalid_size() {
791 let mut matrix1 = SparseKernelMatrix::new(3);
792 matrix1.set(0, 0, 1.0);
793
794 let matrix2 = SparseKernelMatrix::new(2);
795 let result = matrix1.add(&matrix2);
796 assert!(result.is_err());
797 }
798
799 #[test]
800 fn test_sparse_matrix_frobenius_norm() {
801 let mut matrix = SparseKernelMatrix::new(3);
802 matrix.set(0, 0, 3.0);
803 matrix.set(1, 1, 4.0);
804 matrix.finalize();
805
806 let norm = matrix.frobenius_norm();
807 assert!((norm - 5.0).abs() < 1e-10); }
809
810 #[test]
811 fn test_sparse_matrix_iterator() {
812 let mut matrix = SparseKernelMatrix::new(3);
813 matrix.set(0, 1, 0.8);
814 matrix.set(1, 2, 0.6);
815 matrix.set(2, 0, 0.4);
816
817 let entries: Vec<_> = matrix.iter_nonzeros().collect();
818
819 assert_eq!(entries.len(), 3);
820 assert!(entries.contains(&(0, 1, 0.8)));
821 assert!(entries.contains(&(1, 2, 0.6)));
822 assert!(entries.contains(&(2, 0, 0.4)));
823 }
824
825 #[test]
826 fn test_sparse_matrix_scale() {
827 let mut matrix = SparseKernelMatrix::new(3);
828 matrix.set(0, 0, 2.0);
829 matrix.set(1, 1, 4.0);
830 matrix.finalize();
831
832 matrix.scale(0.5);
833
834 assert_eq!(matrix.get(0, 0), Some(1.0));
835 assert_eq!(matrix.get(1, 1), Some(2.0));
836 }
837
838 #[test]
839 fn test_sparse_matrix_builder_parallel() {
840 let builder = SparseKernelMatrixBuilder::new();
841 let kernel = LinearKernel::new();
842 let data = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![0.5, 0.5]];
843
844 let matrix = builder.build_parallel(&data, &kernel).expect("unwrap");
845 assert!(matrix.nnz() > 0);
846
847 let matrix_seq = builder.build(&data, &kernel).expect("unwrap");
849 assert_eq!(matrix.nnz(), matrix_seq.nnz());
850 }
851
852 #[test]
853 fn test_sparse_matrix_parallel_with_threshold() {
854 let builder = SparseKernelMatrixBuilder::new()
855 .with_threshold(0.5)
856 .expect("unwrap");
857 let kernel = LinearKernel::new();
858 let data = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![0.5, 0.5]];
859
860 let matrix = builder.build_parallel(&data, &kernel).expect("unwrap");
861 assert!(matrix.nnz() > 0);
862 }
863}