1use crate::error::{Error, require_nonzero_size, require_threshold};
4use crate::matrix::Mat;
5use crate::poly::PiecewiseLegendrePolyVector;
6
7#[derive(Debug, Clone)]
9pub struct SVEResult {
10 pub(crate) u: PiecewiseLegendrePolyVector,
12 pub(crate) s: Vec<f64>,
14 pub(crate) v: PiecewiseLegendrePolyVector,
16 pub(crate) epsilon: f64,
18}
19
20impl SVEResult {
21 pub fn u(&self) -> &PiecewiseLegendrePolyVector {
23 &self.u
24 }
25
26 pub fn s(&self) -> &[f64] {
28 &self.s
29 }
30
31 pub fn v(&self) -> &PiecewiseLegendrePolyVector {
33 &self.v
34 }
35
36 pub fn epsilon(&self) -> f64 {
38 self.epsilon
39 }
40
41 pub fn from_discretized_matrix<T: crate::numeric::CustomNumeric + 'static>(
61 matrix: &Mat<T>,
62 gauss_x: &crate::gauss::Rule<T>,
63 gauss_y: &crate::gauss::Rule<T>,
64 segments_x: &[f64],
65 segments_y: &[f64],
66 n_gauss: usize,
67 epsilon: f64,
68 ) -> Result<Self, Error> {
69 require_gauss_point_count(matrix.shape().0, segments_x, n_gauss, "rows")?;
70 require_gauss_point_count(matrix.shape().1, segments_y, n_gauss, "columns")?;
71
72 let (u, s, v) = crate::tsvd::compute_svd_dtensor(matrix)?;
73
74 let u_unweighted = crate::sve::utils::remove_weights(&u, gauss_x.w.as_slice(), true);
77 let v_unweighted = crate::sve::utils::remove_weights(&v, gauss_y.w.as_slice(), true);
78
79 let u_f64 = Mat::<f64>::from_fn(u_unweighted.dims(), |idx| u_unweighted[idx].to_f64());
80 let v_f64 = Mat::<f64>::from_fn(v_unweighted.dims(), |idx| v_unweighted[idx].to_f64());
81
82 let gauss_rule_f64 = crate::gauss::legendre::<f64>(n_gauss);
83 let u_polys =
84 crate::sve::utils::svd_to_polynomials(&u_f64, segments_x, &gauss_rule_f64, n_gauss)?;
85 let v_polys =
86 crate::sve::utils::svd_to_polynomials(&v_f64, segments_y, &gauss_rule_f64, n_gauss)?;
87
88 let s_f64: Vec<f64> = s.iter().map(|sv| sv.to_f64()).collect();
89
90 Self::new(
94 PiecewiseLegendrePolyVector::from_polys_unchecked(u_polys),
95 s_f64,
96 PiecewiseLegendrePolyVector::from_polys_unchecked(v_polys),
97 epsilon,
98 )
99 }
100
101 #[allow(clippy::too_many_arguments)]
118 pub fn from_discretized_matrices_centrosymmetric(
119 even: &Mat<f64>,
120 odd: &Mat<f64>,
121 gauss_x: &crate::gauss::Rule<f64>,
122 gauss_y: &crate::gauss::Rule<f64>,
123 segments_x: &[f64],
124 segments_y: &[f64],
125 n_gauss: usize,
126 xmax: f64,
127 ymax: f64,
128 epsilon: f64,
129 ) -> Result<Self, Error> {
130 use crate::kernel::SymmetryType;
131 use crate::sve::utils::{extend_to_full_domain, merge_results, svd_to_polynomials};
132
133 for matrix in [even, odd] {
134 require_gauss_point_count(matrix.shape().0, segments_x, n_gauss, "rows")?;
135 require_gauss_point_count(matrix.shape().1, segments_y, n_gauss, "columns")?;
136 }
137
138 let gauss_rule_f64 = crate::gauss::legendre::<f64>(n_gauss);
139 let block = |matrix: &Mat<f64>, symmetry: SymmetryType| {
140 let (u, s, v) = crate::tsvd::compute_svd_dtensor(matrix)?;
141 let u_unweighted = crate::sve::utils::remove_weights(&u, gauss_x.w.as_slice(), true);
142 let v_unweighted = crate::sve::utils::remove_weights(&v, gauss_y.w.as_slice(), true);
143 let u_polys = svd_to_polynomials(&u_unweighted, segments_x, &gauss_rule_f64, n_gauss)?;
144 let v_polys = svd_to_polynomials(&v_unweighted, segments_y, &gauss_rule_f64, n_gauss)?;
145 let u_full = extend_to_full_domain(u_polys, symmetry, xmax)?;
146 let v_full = extend_to_full_domain(v_polys, symmetry, ymax)?;
147 Ok::<_, Error>((
148 PiecewiseLegendrePolyVector::from_polys_unchecked(u_full),
149 s,
150 PiecewiseLegendrePolyVector::from_polys_unchecked(v_full),
151 ))
152 };
153
154 let result_even = block(even, SymmetryType::Even)?;
155 let result_odd = block(odd, SymmetryType::Odd)?;
156 merge_results(result_even, result_odd, epsilon)
157 }
158
159 pub fn new(
169 u: PiecewiseLegendrePolyVector,
170 s: Vec<f64>,
171 v: PiecewiseLegendrePolyVector,
172 epsilon: f64,
173 ) -> Result<Self, Error> {
174 let result = Self { u, s, v, epsilon };
175 result.check()?;
176 Ok(result)
177 }
178
179 fn check(&self) -> Result<(), Error> {
182 let n = self.s.len();
183 if n == 0 {
184 return Err(Error::EmptyInput { name: "s" });
185 }
186 for (name, funcs) in [("u", &self.u), ("v", &self.v)] {
187 let len = funcs.get_polys().len();
188 if len != n {
189 return Err(Error::InvalidParameter {
190 name,
191 value: format!("{len} functions"),
192 reason: format!("must have one function per singular value ({n})"),
193 });
194 }
195 }
196 for (i, &x) in self.s.iter().enumerate() {
197 if !x.is_finite() {
198 return Err(Error::NonFiniteInput {
199 name: "s",
200 index: vec![i],
201 value: x,
202 });
203 }
204 if x <= 0.0 {
205 return Err(Error::InvalidParameter {
206 name: "s",
207 value: format!("{x:?} at index {i}"),
208 reason: "singular values must be positive".to_string(),
209 });
210 }
211 if i > 0 && x > self.s[i - 1] {
212 return Err(Error::InvalidParameter {
213 name: "s",
214 value: format!("{x:?} at index {i}, after {:?}", self.s[i - 1]),
215 reason: "singular values must be non-increasing".to_string(),
216 });
217 }
218 }
219 require_threshold("epsilon", Some(self.epsilon))
220 }
221
222 pub fn part(
240 &self,
241 eps: Option<f64>,
242 max_size: Option<usize>,
243 ) -> Result<
244 (
245 PiecewiseLegendrePolyVector,
246 Vec<f64>,
247 PiecewiseLegendrePolyVector,
248 ),
249 Error,
250 > {
251 self.check()?;
252 require_threshold("eps", eps)?;
253 require_nonzero_size("max_size", max_size)?;
254 let eps = eps.unwrap_or(self.epsilon);
255 let threshold = eps * self.s[0];
256
257 let mut cut = 0;
258 for &val in self.s.iter() {
259 if val >= threshold {
260 cut += 1;
261 } else {
262 break;
263 }
264 }
265
266 if let Some(max) = max_size {
267 cut = cut.min(max);
268 }
269
270 let u_part = PiecewiseLegendrePolyVector::new(self.u.get_polys()[..cut].to_vec())?;
272 let s_part = self.s[..cut].to_vec();
273 let v_part = PiecewiseLegendrePolyVector::new(self.v.get_polys()[..cut].to_vec())?;
274
275 Ok((u_part, s_part, v_part))
276 }
277}
278
279fn require_gauss_point_count(
286 len: usize,
287 segments: &[f64],
288 n_gauss: usize,
289 what: &'static str,
290) -> Result<(), Error> {
291 if segments.len() < 2 {
292 return Err(Error::InvalidParameter {
293 name: "segments",
294 value: format!("{} boundaries", segments.len()),
295 reason: "must have at least 2 entries".to_string(),
296 });
297 }
298 let expected = n_gauss * (segments.len() - 1);
299 if len != expected {
300 return Err(Error::InvalidParameter {
301 name: "matrix",
302 value: format!("{len} {what}"),
303 reason: format!("must have one per Gauss point of the segments ({expected})"),
304 });
305 }
306 Ok(())
307}