Skip to main content

proof_engine/symbolic/
matrix.rs

1//! Symbolic matrix operations — determinant, inverse, eigenvalues.
2
3use super::expr::Expr;
4
5/// A symbolic matrix of expressions.
6#[derive(Debug, Clone)]
7pub struct SymMatrix {
8    pub rows: usize,
9    pub cols: usize,
10    pub data: Vec<Vec<Expr>>,
11}
12
13impl SymMatrix {
14    pub fn new(rows: usize, cols: usize) -> Self {
15        Self { rows, cols, data: vec![vec![Expr::zero(); cols]; rows] }
16    }
17
18    pub fn identity(n: usize) -> Self {
19        let mut m = Self::new(n, n);
20        for i in 0..n { m.data[i][i] = Expr::one(); }
21        m
22    }
23
24    pub fn from_f64(data: &[&[f64]]) -> Self {
25        let rows = data.len();
26        let cols = if rows > 0 { data[0].len() } else { 0 };
27        let mut m = Self::new(rows, cols);
28        for i in 0..rows {
29            for j in 0..cols {
30                m.data[i][j] = Expr::c(data[i][j]);
31            }
32        }
33        m
34    }
35
36    pub fn get(&self, r: usize, c: usize) -> &Expr { &self.data[r][c] }
37    pub fn set(&mut self, r: usize, c: usize, val: Expr) { self.data[r][c] = val; }
38
39    /// Matrix multiplication.
40    pub fn mul(&self, other: &SymMatrix) -> SymMatrix {
41        assert_eq!(self.cols, other.rows);
42        let mut result = SymMatrix::new(self.rows, other.cols);
43        for i in 0..self.rows {
44            for j in 0..other.cols {
45                let mut sum = Expr::zero();
46                for k in 0..self.cols {
47                    sum = sum.add(self.data[i][k].clone().mul(other.data[k][j].clone()));
48                }
49                result.data[i][j] = sum;
50            }
51        }
52        result
53    }
54
55    /// Transpose.
56    pub fn transpose(&self) -> SymMatrix {
57        let mut result = SymMatrix::new(self.cols, self.rows);
58        for i in 0..self.rows {
59            for j in 0..self.cols {
60                result.data[j][i] = self.data[i][j].clone();
61            }
62        }
63        result
64    }
65
66    /// Determinant (recursive cofactor expansion).
67    pub fn determinant(&self) -> Expr {
68        assert_eq!(self.rows, self.cols);
69        let n = self.rows;
70        if n == 1 { return self.data[0][0].clone(); }
71        if n == 2 {
72            let a = self.data[0][0].clone().mul(self.data[1][1].clone());
73            let b = self.data[0][1].clone().mul(self.data[1][0].clone());
74            return a.sub(b);
75        }
76        let mut det = Expr::zero();
77        // Laplace expansion along row 0. The cofactor already carries the
78        // (-1)^(i+j) sign; alternating add/sub as well applied it twice, so
79        // every 3x3 and larger determinant had wrong-signed odd terms.
80        for j in 0..n {
81            let cofactor = self.cofactor(0, j);
82            det = det.add(self.data[0][j].clone().mul(cofactor));
83        }
84        det
85    }
86
87    /// Minor: determinant of the submatrix with row i and col j removed.
88    pub fn minor(&self, row: usize, col: usize) -> Expr {
89        let sub = self.submatrix(row, col);
90        sub.determinant()
91    }
92
93    /// Cofactor: (-1)^(i+j) * minor(i,j).
94    pub fn cofactor(&self, row: usize, col: usize) -> Expr {
95        let m = self.minor(row, col);
96        if (row + col) % 2 == 0 { m } else { m.neg() }
97    }
98
99    /// Remove row i and column j.
100    pub fn submatrix(&self, row: usize, col: usize) -> SymMatrix {
101        let mut result = SymMatrix::new(self.rows - 1, self.cols - 1);
102        let mut ri = 0;
103        for i in 0..self.rows {
104            if i == row { continue; }
105            let mut ci = 0;
106            for j in 0..self.cols {
107                if j == col { continue; }
108                result.data[ri][ci] = self.data[i][j].clone();
109                ci += 1;
110            }
111            ri += 1;
112        }
113        result
114    }
115
116    /// Trace: sum of diagonal elements.
117    pub fn trace(&self) -> Expr {
118        let mut sum = Expr::zero();
119        for i in 0..self.rows.min(self.cols) {
120            sum = sum.add(self.data[i][i].clone());
121        }
122        sum
123    }
124
125    /// Numerical eigenvalues for a 2x2 matrix.
126    pub fn eigenvalues_2x2(&self) -> Option<(f64, f64)> {
127        if self.rows != 2 || self.cols != 2 { return None; }
128        let vars = std::collections::HashMap::new();
129        let a = self.data[0][0].eval(&vars);
130        let b = self.data[0][1].eval(&vars);
131        let c = self.data[1][0].eval(&vars);
132        let d = self.data[1][1].eval(&vars);
133
134        let trace = a + d;
135        let det = a * d - b * c;
136        let disc = trace * trace - 4.0 * det;
137        if disc < 0.0 { return None; }
138        let sqrt_disc = disc.sqrt();
139        Some(((trace + sqrt_disc) / 2.0, (trace - sqrt_disc) / 2.0))
140    }
141}
142
143#[cfg(test)]
144mod tests {
145    use super::*;
146    use std::collections::HashMap;
147
148    #[test]
149    fn det_2x2() {
150        let m = SymMatrix::from_f64(&[&[1.0, 2.0], &[3.0, 4.0]]);
151        let det = m.determinant();
152        let val = det.eval(&HashMap::new());
153        assert!((val - (-2.0)).abs() < 1e-10);
154    }
155
156    #[test]
157    fn det_3x3() {
158        let m = SymMatrix::from_f64(&[&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 0.0]]);
159        let det = m.determinant();
160        let val = det.eval(&HashMap::new());
161        assert!((val - 27.0).abs() < 1e-8);
162    }
163
164    #[test]
165    fn identity_det_is_one() {
166        let m = SymMatrix::identity(3);
167        let det = m.determinant();
168        let val = det.eval(&HashMap::new());
169        assert!((val - 1.0).abs() < 1e-10);
170    }
171
172    #[test]
173    fn eigenvalues_diagonal() {
174        let m = SymMatrix::from_f64(&[&[3.0, 0.0], &[0.0, 5.0]]);
175        let (e1, e2) = m.eigenvalues_2x2().unwrap();
176        assert!((e1 - 5.0).abs() < 1e-10);
177        assert!((e2 - 3.0).abs() < 1e-10);
178    }
179
180    #[test]
181    fn transpose() {
182        let m = SymMatrix::from_f64(&[&[1.0, 2.0], &[3.0, 4.0]]);
183        let t = m.transpose();
184        let val = t.data[1][0].eval(&HashMap::new());
185        assert!((val - 2.0).abs() < 1e-10);
186    }
187}