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}