Skip to main content

tensorlogic_sklears_kernels/
sparse.rs

1//! Sparse kernel matrix support for large-scale problems.
2//!
3//! Provides efficient storage and operations for sparse kernel matrices using
4//! Compressed Sparse Row (CSR) format for memory-efficient representation.
5//!
6//! # Features
7//!
8//! - **Efficient Storage**: CSR format for sparse matrices with configurable thresholds
9//! - **Matrix Operations**: SpMV, transpose, addition, scaling, Frobenius norm
10//! - **Parallel Construction**: Multi-threaded matrix building with rayon
11//! - **Iterator Support**: Efficient iteration over non-zero entries
12//! - **Flexible Builders**: Configurable threshold and max entries per row
13//!
14//! # Example
15//!
16//! ```rust
17//! use tensorlogic_sklears_kernels::{SparseKernelMatrix, SparseKernelMatrixBuilder};
18//! use tensorlogic_sklears_kernels::tensor_kernels::LinearKernel;
19//!
20//! // Build a sparse kernel matrix with parallel computation
21//! let builder = SparseKernelMatrixBuilder::new()
22//!     .with_threshold(0.1).expect("unwrap")
23//!     .with_max_entries_per_row(100).expect("unwrap");
24//!
25//! let kernel = LinearKernel::new();
26//! let data = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
27//! let matrix = builder.build_parallel(&data, &kernel).expect("unwrap");
28//!
29//! // Sparse matrix-vector multiplication
30//! let mut matrix = SparseKernelMatrix::new(3);
31//! matrix.set(0, 0, 2.0);
32//! matrix.set(1, 1, 3.0);
33//! let x = vec![1.0, 2.0, 0.0];
34//! let y = matrix.spmv(&x).expect("unwrap");
35//!
36//! // Iterate over non-zero entries
37//! for (row, col, value) in matrix.iter_nonzeros() {
38//!     println!("({}, {}) = {}", row, col, value);
39//! }
40//! ```
41
42use std::collections::HashMap;
43
44use serde::{Deserialize, Serialize};
45
46use crate::error::{KernelError, Result};
47use crate::types::Kernel;
48
49/// Sparse kernel matrix using Compressed Sparse Row (CSR) format
50///
51/// Stores only non-zero entries for efficient memory usage.
52///
53/// # Example
54///
55/// ```rust
56/// use tensorlogic_sklears_kernels::SparseKernelMatrix;
57///
58/// let mut matrix = SparseKernelMatrix::new(3);
59/// matrix.set(0, 1, 0.8);
60/// matrix.set(1, 2, 0.6);
61///
62/// assert_eq!(matrix.get(0, 1), Some(0.8));
63/// assert_eq!(matrix.get(0, 2), None);
64/// assert_eq!(matrix.nnz(), 2);
65/// ```
66#[derive(Clone, Debug, Serialize, Deserialize)]
67pub struct SparseKernelMatrix {
68    /// Number of rows/columns (square matrix)
69    size: usize,
70    /// Row pointers for CSR format
71    row_ptr: Vec<usize>,
72    /// Column indices
73    col_idx: Vec<usize>,
74    /// Non-zero values
75    values: Vec<f64>,
76    /// Temporary map for construction (not serialized)
77    #[serde(skip)]
78    temp_map: HashMap<(usize, usize), f64>,
79}
80
81impl SparseKernelMatrix {
82    /// Create a new sparse kernel matrix
83    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    /// Set a value in the matrix
94    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            // Remove near-zero values
101            self.temp_map.remove(&(row, col));
102        } else {
103            self.temp_map.insert((row, col), value);
104        }
105    }
106
107    /// Get a value from the matrix
108    pub fn get(&self, row: usize, col: usize) -> Option<f64> {
109        if row >= self.size || col >= self.size {
110            return None;
111        }
112
113        // Check temp map first
114        if let Some(&value) = self.temp_map.get(&(row, col)) {
115            return Some(value);
116        }
117
118        // Search in CSR format
119        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    /// Finalize the matrix (convert temp map to CSR format)
132    pub fn finalize(&mut self) {
133        if self.temp_map.is_empty() {
134            return;
135        }
136
137        // Clear existing CSR data
138        self.col_idx.clear();
139        self.values.clear();
140        self.row_ptr = vec![0; self.size + 1];
141
142        // Sort entries by row, then column
143        let mut entries: Vec<_> = self.temp_map.iter().collect();
144        entries.sort_by_key(|&((row, col), _)| (*row, *col));
145
146        // Build CSR format
147        let mut current_row = 0;
148        for (&(row, col), &value) in &entries {
149            // Update row pointers
150            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        // Finalize row pointers
160        while current_row < self.size {
161            current_row += 1;
162            self.row_ptr[current_row] = self.col_idx.len();
163        }
164
165        // Clear temp map
166        self.temp_map.clear();
167    }
168
169    /// Get number of non-zero entries
170    pub fn nnz(&self) -> usize {
171        self.values.len() + self.temp_map.len()
172    }
173
174    /// Get matrix size
175    pub fn size(&self) -> usize {
176        self.size
177    }
178
179    /// Get density (fraction of non-zero entries)
180    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    /// Convert to dense matrix
190    #[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    /// Compute sparse kernel matrix from data with threshold
211    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    /// Get row as sparse vector
233    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
252/// Sparse kernel matrix builder with configuration
253pub struct SparseKernelMatrixBuilder {
254    /// Sparsity threshold (values below this are treated as zero)
255    threshold: f64,
256    /// Maximum entries per row (for controlled sparsity)
257    max_entries_per_row: Option<usize>,
258}
259
260impl SparseKernelMatrixBuilder {
261    /// Create a new builder
262    pub fn new() -> Self {
263        Self {
264            threshold: 1e-10,
265            max_entries_per_row: None,
266        }
267    }
268
269    /// Set sparsity threshold
270    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    /// Set maximum entries per row
283    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    /// Build sparse kernel matrix from data
296    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            // Compute all values for this row
304            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 max_entries_per_row is set, keep only top-k entries
312            if let Some(max_entries) = self.max_entries_per_row {
313                if row_entries.len() > max_entries {
314                    // Sort by absolute value (descending)
315                    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            // Add entries to matrix
325            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
341/// Advanced sparse matrix operations
342impl SparseKernelMatrix {
343    /// Sparse matrix-vector multiplication: y = A * x
344    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    /// Sparse matrix transpose
374    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    /// Add two sparse matrices element-wise
393    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        // Clone and finalize other to ensure all values are in CSR format
405        let mut other_finalized = other.clone();
406        other_finalized.finalize();
407
408        let mut result = Self::new(self.size);
409
410        // Add values from self
411        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        // Add values from other
423        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    /// Frobenius norm of the sparse matrix
440    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    /// Iterator over non-zero entries (row, col, value)
457    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    /// Scale the matrix by a scalar
467    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
478/// Iterator for sparse matrix non-zero entries
479pub 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
512/// Parallel sparse kernel matrix builder
513impl SparseKernelMatrixBuilder {
514    /// Build sparse kernel matrix with parallel computation
515    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        // Compute rows in parallel
526        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 max_entries_per_row is set, keep only top-k entries
543                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        // Sequentially insert into matrix
559        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        // Each row should have at most 2 entries
702        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); // Very small value (below 1e-10 threshold)
719        matrix.finalize();
720
721        // Should be treated as zero and filtered out
722        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); // 2*1 + 1*3
739        assert!((y[1] - 6.0).abs() < 1e-10); // 3*2
740        assert!((y[2] - 7.0).abs() < 1e-10); // 1*1 + 2*3
741    }
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]; // Wrong size
749        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)); // 2.0 + 1.0
784        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); // sqrt(3^2 + 4^2) = 5
808    }
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        // Compare with sequential build
848        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}