qualia_core_db/solvers/linear_algebra/
svd.rs1use crate::solvers::linear_algebra::eigen::symmetric_eigen;
12use crate::solvers::SolversError;
13
14#[derive(Debug, Clone)]
19pub struct Svd {
20 pub singular_values: Vec<f64>,
21 pub u: Vec<f64>,
22 pub v: Vec<f64>,
23}
24
25pub fn svd(m: usize, n: usize, data: &[f64]) -> Result<Svd, SolversError> {
29 if m == 0 || n == 0 || data.len() != m * n {
30 return Err(SolversError::InvalidDimension);
31 }
32
33 let mut ata = vec![0.0_f64; n * n];
38 super::gemm::gemm(
39 super::gemm::Transpose::Yes,
40 super::gemm::Transpose::No,
41 n,
42 n,
43 m,
44 1.0,
45 data,
46 data,
47 0.0,
48 &mut ata,
49 )?;
50
51 let mut eigvecs = vec![0.0_f64; n * n];
53 symmetric_eigen(n, &mut ata, &mut eigvecs)?;
54 let eigvals: Vec<f64> = (0..n).map(|i| ata[i * n + i]).collect();
55
56 let mut order: Vec<usize> = (0..n).collect();
58 order.sort_by(|&i, &j| {
59 eigvals[j]
60 .partial_cmp(&eigvals[i])
61 .unwrap_or(core::cmp::Ordering::Equal)
62 });
63
64 let mut singular_values = vec![0.0_f64; n];
65 let mut v = vec![0.0_f64; n * n];
66 for (new_col, &old_col) in order.iter().enumerate() {
67 singular_values[new_col] = eigvals[old_col].max(0.0).sqrt();
68 for row in 0..n {
69 v[row * n + new_col] = eigvecs[row * n + old_col];
70 }
71 }
72
73 let mut av = vec![0.0_f64; m * n];
77 super::gemm::matmul(m, n, n, data, &v, &mut av)?;
78 let mut u = vec![0.0_f64; m * n];
79 let smax = singular_values.first().copied().unwrap_or(0.0).max(1.0);
80 for k in 0..n {
81 let sigma = singular_values[k];
82 if sigma <= 1e-12 * smax {
83 continue;
84 }
85 for i in 0..m {
86 u[i * n + k] = av[i * n + k] / sigma;
87 }
88 }
89
90 Ok(Svd {
91 singular_values,
92 u,
93 v,
94 })
95}
96
97#[cfg(test)]
98mod tests {
99 use super::*;
100
101 #[test]
102 fn reconstructs_square() {
103 let m = 3;
105 let n = 3;
106 let a = [4.0, 0.0, 0.0, 0.0, 3.0, 0.0, 0.0, 0.0, 5.0];
107 let s = svd(m, n, &a).unwrap();
108 assert!((s.singular_values[0] - 5.0).abs() < 1e-9);
110 assert!((s.singular_values[1] - 4.0).abs() < 1e-9);
111 assert!((s.singular_values[2] - 3.0).abs() < 1e-9);
112 for i in 0..m {
113 for j in 0..n {
114 let mut acc = 0.0;
115 for k in 0..n {
116 acc += s.u[i * n + k] * s.singular_values[k] * s.v[j * n + k];
117 }
118 assert!(
119 (acc - a[i * n + j]).abs() < 1e-6,
120 "({i},{j}) {acc} != {}",
121 a[i * n + j]
122 );
123 }
124 }
125 }
126
127 #[test]
128 fn reconstructs_tall_rectangular() {
129 let m = 4;
130 let n = 2;
131 let a = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
132 let s = svd(m, n, &a).unwrap();
133 for i in 0..m {
134 for j in 0..n {
135 let mut acc = 0.0;
136 for k in 0..n {
137 acc += s.u[i * n + k] * s.singular_values[k] * s.v[j * n + k];
138 }
139 assert!(
140 (acc - a[i * n + j]).abs() < 1e-6,
141 "({i},{j}) {acc} != {}",
142 a[i * n + j]
143 );
144 }
145 }
146 assert!(s.singular_values[0] >= s.singular_values[1]);
148 }
149
150 #[test]
151 fn rejects_bad_dims() {
152 assert!(matches!(
153 svd(2, 2, &[1.0, 2.0, 3.0]),
154 Err(SolversError::InvalidDimension)
155 ));
156 }
157}