Skip to main content

solver/linalg/
sor.rs

1use super::csr::CsrMatrix;
2
3impl CsrMatrix {
4    #[doc = include_str!("../../docs/math/sor.md")]
5    ///
6    /// # Examples
7    ///
8    /// ```
9    /// use solver::linalg::csr::CsrMatrix;
10    ///
11    /// // Solve A x = b where A is:
12    /// // [ 4 -1 ]
13    /// // [-1  4 ]
14    /// // and b = [3, 3]
15    /// // Solution should be x = [1, 1]
16    ///
17    /// let mut mat = CsrMatrix::new(2, 4);
18    ///
19    /// // Row 0
20    /// mat.add_entry(0, 4.0);
21    /// mat.add_entry(1, -1.0);
22    /// mat.finish_row();
23    ///
24    /// // Row 1
25    /// mat.add_entry(0, -1.0);
26    /// mat.add_entry(1, 4.0);
27    /// mat.finish_row();
28    ///
29    /// let b = vec![3.0, 3.0];
30    /// let mut x = vec![0.0, 0.0]; // Initial guess
31    ///
32    /// let iters = mat.solve_sor(&b, &mut x, 1e-6, 100, 1.0);
33    ///
34    /// assert!((x[0] - 1.0).abs() < 1e-5);
35    /// assert!((x[1] - 1.0).abs() < 1e-5);
36    /// ```
37    pub fn solve_sor(
38        &self,
39        b: &[f64],
40        x: &mut [f64],
41        tol: f64,
42        max_iter: usize,
43        omega: f64,
44    ) -> usize {
45        let mut iter_count = 0;
46        for _ in 0..max_iter {
47            iter_count += 1;
48            let mut max_error = 0.0;
49
50            for i in 0..self.size {
51                let row_start = self.row_ptr[i];
52                let row_end = self.row_ptr[i + 1];
53
54                let diag_idx = self.diag_indices[i];
55                if diag_idx == usize::MAX {
56                    continue;
57                } // Should not happen in well-formed HJB matrices
58
59                let diag = self.values[diag_idx];
60                let mut sigma = 0.0;
61
62                for idx in row_start..row_end {
63                    if idx != diag_idx {
64                        sigma += self.values[idx] * x[self.col_indices[idx]];
65                    }
66                }
67
68                if diag == 0.0 {
69                    continue;
70                }
71
72                let val_new = (1.0 - omega) * x[i] + (omega / diag) * (b[i] - sigma);
73
74                let diff = (val_new - x[i]).abs();
75                if diff > max_error {
76                    max_error = diff;
77                }
78
79                x[i] = val_new;
80            }
81
82            if max_error < tol {
83                return iter_count;
84            }
85        }
86        iter_count
87    }
88}