proof_engine/symbolic/
matrix.rs1use super::expr::Expr;
4
5#[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 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 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 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 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 pub fn minor(&self, row: usize, col: usize) -> Expr {
89 let sub = self.submatrix(row, col);
90 sub.determinant()
91 }
92
93 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 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 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 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}