Skip to main content

solver/linalg/
csr.rs

1/// Compressed Sparse Row Matrix
2///
3/// # Examples
4///
5/// ```
6/// use solver::linalg::csr::CsrMatrix;
7///
8/// // Create a 3x3 identity matrix
9/// let mut mat = CsrMatrix::new(3, 3);
10///
11/// // Row 0
12/// mat.add_entry(0, 1.0);
13/// mat.finish_row();
14///
15/// // Row 1
16/// mat.add_entry(1, 1.0);
17/// mat.finish_row();
18///
19/// // Row 2
20/// mat.add_entry(2, 1.0);
21/// mat.finish_row();
22///
23/// assert_eq!(mat.values, vec![1.0, 1.0, 1.0]);
24/// assert_eq!(mat.col_indices, vec![0, 1, 2]);
25/// assert_eq!(mat.row_ptr, vec![0, 1, 2, 3]);
26/// ```
27pub struct CsrMatrix {
28    pub values: Vec<f64>,
29    pub col_indices: Vec<usize>,
30    pub row_ptr: Vec<usize>,
31    pub size: usize,
32    pub diag_indices: Vec<usize>,
33}
34
35impl CsrMatrix {
36    pub fn new(size: usize, estimated_entries: usize) -> Self {
37        let mut row_ptr = Vec::with_capacity(size + 1);
38        row_ptr.push(0);
39        Self {
40            values: Vec::with_capacity(estimated_entries),
41            col_indices: Vec::with_capacity(estimated_entries),
42            row_ptr,
43            size,
44            diag_indices: Vec::with_capacity(size),
45        }
46    }
47
48    #[inline]
49    pub fn add_entry(&mut self, col: usize, val: f64) {
50        self.col_indices.push(col);
51        self.values.push(val);
52    }
53
54    #[inline]
55    pub fn finish_row(&mut self) {
56        // Find diagonal index for the current row (which is self.row_ptr.len() - 1)
57        let row_idx = self.row_ptr.len() - 1;
58        let start = self.row_ptr[row_idx];
59        let end = self.values.len();
60
61        let mut diag_found = false;
62        for idx in start..end {
63            if self.col_indices[idx] == row_idx {
64                self.diag_indices.push(idx);
65                diag_found = true;
66                break;
67            }
68        }
69
70        if !diag_found {
71            // If diagonal not found, push a dummy value or handle it.
72            // For now, we assume diagonal exists or we push usize::MAX to indicate missing.
73            // However, solve_sor assumes diagonal exists.
74            // Let's push usize::MAX and handle it in solve_sor if needed,
75            // but for this specific solver, diagonal is always added.
76            self.diag_indices.push(usize::MAX);
77        }
78
79        self.row_ptr.push(self.values.len());
80    }
81
82    /// Solves the linear system Ax = b using LAPACK's DGTSV (General Tridiagonal Solve).
83    /// Assumes the matrix is tridiagonal.
84    ///
85    /// # Examples
86    ///
87    /// ```
88    /// use solver::linalg::csr::CsrMatrix;
89    ///
90    /// // Solve A x = b where A is:
91    /// // [ 2 -1  0 ]
92    /// // [-1  2 -1 ]
93    /// // [ 0 -1  2 ]
94    /// // and b = [1, 0, 1]
95    /// // Solution should be x = [1, 1, 1]
96    ///
97    /// let mut mat = CsrMatrix::new(3, 9);
98    ///
99    /// // Row 0
100    /// mat.add_entry(0, 2.0);
101    /// mat.add_entry(1, -1.0);
102    /// mat.finish_row();
103    ///
104    /// // Row 1
105    /// mat.add_entry(0, -1.0);
106    /// mat.add_entry(1, 2.0);
107    /// mat.add_entry(2, -1.0);
108    /// mat.finish_row();
109    ///
110    /// // Row 2
111    /// mat.add_entry(1, -1.0);
112    /// mat.add_entry(2, 2.0);
113    /// mat.finish_row();
114    ///
115    /// let b = vec![1.0, 0.0, 1.0];
116    /// let x = mat.solve_tridiagonal_system(&b);
117    ///
118    /// assert!((x[0] - 1.0).abs() < 1e-6);
119    /// assert!((x[1] - 1.0).abs() < 1e-6);
120    /// assert!((x[2] - 1.0).abs() < 1e-6);
121    /// ```
122    pub fn solve_tridiagonal_system(&self, b: &[f64]) -> Vec<f64> {
123        use lapack::dgtsv;
124
125        let n = self.size;
126        let mut dl = vec![0.0; n - 1];
127        let mut d = vec![0.0; n];
128        let mut du = vec![0.0; n - 1];
129
130        // Extract diagonals
131        for i in 0..n {
132            let row_start = self.row_ptr[i];
133            let row_end = self.row_ptr[i + 1];
134
135            for idx in row_start..row_end {
136                let col = self.col_indices[idx];
137                let val = self.values[idx];
138
139                if col == i {
140                    d[i] = val;
141                } else if i > 0 && col == i - 1 {
142                    dl[i - 1] = val;
143                } else if col == i + 1 {
144                    du[i] = val;
145                }
146            }
147        }
148
149        let mut x = b.to_vec();
150        let n_i32 = n as i32;
151        let nrhs = 1;
152        let ldb = n_i32;
153        let mut info = 0;
154
155        unsafe {
156            dgtsv(
157                n_i32, nrhs, &mut dl, &mut d, &mut du, &mut x, ldb, &mut info,
158            );
159        }
160
161        if info != 0 {
162            panic!("LAPACK dgtsv failed with info = {}", info);
163        }
164
165        x
166    }
167}