1use crate::error::InterpolateError;
7use scirs2_core::ndarray::{Array2, Array3, ArrayD, IxDyn};
8
9#[derive(Debug, Clone)]
16pub struct TensorTrain {
17 pub cores: Vec<Array3<f64>>,
19 pub shape: Vec<usize>,
21 pub ranks: Vec<usize>,
23}
24
25impl TensorTrain {
26 pub fn new(cores: Vec<Array3<f64>>) -> Result<Self, InterpolateError> {
30 if cores.is_empty() {
31 return Err(InterpolateError::InvalidInput {
32 message: "TensorTrain requires at least one core".into(),
33 });
34 }
35 let d = cores.len();
36 let mut shape = Vec::with_capacity(d);
37 let mut ranks = Vec::with_capacity(d + 1);
38 ranks.push(cores[0].shape()[0]);
39 for (k, core) in cores.iter().enumerate() {
40 let s = core.shape();
41 if s.len() != 3 {
42 return Err(InterpolateError::InvalidInput {
43 message: format!("Core {k} must be a 3-D array, got {}D", s.len()),
44 });
45 }
46 let prev_rank = *ranks.last().ok_or_else(|| InterpolateError::InvalidInput {
47 message: "Internal rank mismatch".into(),
48 })?;
49 if k > 0 && s[0] != prev_rank {
50 return Err(InterpolateError::InvalidInput {
51 message: format!(
52 "Left rank of core {k} ({}) does not match right rank of core {} ({})",
53 s[0],
54 k - 1,
55 prev_rank,
56 ),
57 });
58 }
59 shape.push(s[1]);
60 ranks.push(s[2]);
61 }
62 if ranks[0] != 1
63 || *ranks.last().ok_or_else(|| InterpolateError::InvalidInput {
64 message: "Empty ranks".into(),
65 })? != 1
66 {
67 return Err(InterpolateError::InvalidInput {
68 message: format!(
69 "Boundary ranks must be 1, got r_0={} and r_d={}",
70 ranks[0], ranks[d]
71 ),
72 });
73 }
74 Ok(Self {
75 cores,
76 shape,
77 ranks,
78 })
79 }
80
81 pub fn eval(&self, idx: &[usize]) -> Result<f64, InterpolateError> {
86 let d = self.cores.len();
87 if idx.len() != d {
88 return Err(InterpolateError::DimensionMismatch(format!(
89 "idx has length {}, expected {d}",
90 idx.len()
91 )));
92 }
93 for (k, (&ik, &nk)) in idx.iter().zip(self.shape.iter()).enumerate() {
94 if ik >= nk {
95 return Err(InterpolateError::OutOfBounds(format!(
96 "Index {ik} out of range [0, {nk}) in dimension {k}"
97 )));
98 }
99 }
100
101 let mut v = vec![1.0f64];
103 for (k, &ik) in idx.iter().enumerate() {
104 let core = &self.cores[k];
105 let (r_left, _n, r_right) = (core.shape()[0], core.shape()[1], core.shape()[2]);
106 let mut new_v = vec![0.0f64; r_right];
108 for j in 0..r_right {
109 let mut s = 0.0f64;
110 for i in 0..r_left {
111 s += v[i] * core[[i, ik, j]];
112 }
113 new_v[j] = s;
114 }
115 v = new_v;
116 }
117 Ok(v[0])
119 }
120
121 pub fn to_dense(&self) -> Result<ArrayD<f64>, InterpolateError> {
125 let d = self.cores.len();
126 let total: usize = self.shape.iter().product();
127 let mut data = vec![0.0f64; total];
128
129 let mut idx = vec![0usize; d];
131 loop {
132 let flat = row_major_index(&idx, &self.shape);
133 data[flat] = self.eval(&idx)?;
134
135 let mut carry = true;
137 for k in (0..d).rev() {
138 if carry {
139 idx[k] += 1;
140 if idx[k] >= self.shape[k] {
141 idx[k] = 0;
142 } else {
143 carry = false;
144 }
145 }
146 }
147 if carry {
148 break; }
150 }
151
152 ArrayD::from_shape_vec(IxDyn(&self.shape), data)
153 .map_err(|e| InterpolateError::ComputationError(format!("to_dense shape error: {e}")))
154 }
155
156 pub fn norm(&self) -> f64 {
158 let d = self.cores.len();
159 if d == 0 {
160 return 0.0;
161 }
162 let mut gram = Array2::<f64>::eye(1);
164
165 for core in &self.cores {
166 let (r_left, n, r_right) = (core.shape()[0], core.shape()[1], core.shape()[2]);
167 let mut new_gram = Array2::<f64>::zeros((r_right, r_right));
168 for ik in 0..n {
169 for beta1 in 0..r_right {
170 for beta2 in 0..r_right {
171 let mut s = 0.0f64;
172 for alpha1 in 0..r_left {
173 for alpha2 in 0..r_left {
174 s += core[[alpha1, ik, beta1]]
175 * gram[[alpha1, alpha2]]
176 * core[[alpha2, ik, beta2]];
177 }
178 }
179 new_gram[[beta1, beta2]] += s;
180 }
181 }
182 }
183 gram = new_gram;
184 }
185 gram[[0, 0]].max(0.0).sqrt()
187 }
188
189 pub fn n_params(&self) -> usize {
191 self.cores.iter().map(|c| c.len()).sum()
192 }
193
194 pub fn from_dense(
196 tensor: &ArrayD<f64>,
197 max_rank: usize,
198 tol: f64,
199 ) -> Result<Self, InterpolateError> {
200 tt_svd(tensor, max_rank, tol)
201 }
202}
203
204fn row_major_index(idx: &[usize], shape: &[usize]) -> usize {
207 let mut flat = 0usize;
208 let mut stride = 1usize;
209 for k in (0..idx.len()).rev() {
210 flat += idx[k] * stride;
211 stride *= shape[k];
212 }
213 flat
214}
215
216pub fn truncated_svd(
225 a: &Array2<f64>,
226 max_rank: usize,
227 tol: f64,
228) -> Result<(Array2<f64>, Vec<f64>, Array2<f64>), InterpolateError> {
229 let m = a.nrows();
230 let n = a.ncols();
231 if m == 0 || n == 0 {
232 return Err(InterpolateError::InvalidInput {
233 message: "truncated_svd: matrix must have positive dimensions".into(),
234 });
235 }
236 let r_max = max_rank.min(m).min(n);
237 let mat_data: Vec<f64> = a.iter().copied().collect();
238 let (u_data, s_vals, vt_data) = svd_deflation(&mat_data, m, n, r_max, tol);
239
240 let r = s_vals.len();
241 let u_arr = Array2::from_shape_vec((m, r), u_data)
242 .map_err(|e| InterpolateError::ComputationError(format!("SVD U shape error: {e}")))?;
243 let vt_arr = Array2::from_shape_vec((r, n), vt_data)
244 .map_err(|e| InterpolateError::ComputationError(format!("SVD Vt shape error: {e}")))?;
245 Ok((u_arr, s_vals, vt_arr))
246}
247
248fn svd_deflation(
252 a: &[f64],
253 m: usize,
254 n: usize,
255 max_rank: usize,
256 tol: f64,
257) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
258 let frob_sq: f64 = a.iter().map(|x| x * x).sum();
260 let sigma_ref = frob_sq.sqrt();
261 if sigma_ref < 1e-300 {
262 let mut u0 = vec![0.0f64; m];
263 let mut v0 = vec![0.0f64; n];
264 if m > 0 {
265 u0[0] = 1.0;
266 }
267 if n > 0 {
268 v0[0] = 1.0;
269 }
270 return (u0, vec![0.0], v0);
271 }
272 let threshold = tol * sigma_ref;
273
274 let mut residual = a.to_vec();
275 let mut u_cols: Vec<Vec<f64>> = Vec::new();
276 let mut s_vals: Vec<f64> = Vec::new();
277 let mut v_rows: Vec<Vec<f64>> = Vec::new();
278
279 for _rank in 0..max_rank {
280 let mut vk: Vec<f64> = vec![0.0f64; n];
282 let mut best_norm = 0.0f64;
283 for i in 0..m {
284 let row_norm: f64 = (0..n)
285 .map(|j| residual[i * n + j] * residual[i * n + j])
286 .sum::<f64>()
287 .sqrt();
288 if row_norm > best_norm {
289 best_norm = row_norm;
290 for j in 0..n {
291 vk[j] = residual[i * n + j];
292 }
293 }
294 }
295 let vnorm: f64 = vk.iter().map(|x| x * x).sum::<f64>().sqrt();
296 if vnorm < 1e-300 {
297 break;
298 }
299 for x in vk.iter_mut() {
300 *x /= vnorm;
301 }
302
303 let mut uk = vec![0.0f64; m];
305 for _iter in 0..20 {
306 for i in 0..m {
308 let mut s = 0.0f64;
309 for j in 0..n {
310 s += residual[i * n + j] * vk[j];
311 }
312 uk[i] = s;
313 }
314 let unorm: f64 = uk.iter().map(|x| x * x).sum::<f64>().sqrt();
315 if unorm < 1e-300 {
316 break;
317 }
318 for x in uk.iter_mut() {
319 *x /= unorm;
320 }
321
322 let mut new_vk = vec![0.0f64; n];
324 for j in 0..n {
325 let mut s = 0.0f64;
326 for i in 0..m {
327 s += residual[i * n + j] * uk[i];
328 }
329 new_vk[j] = s;
330 }
331 let new_vnorm: f64 = new_vk.iter().map(|x| x * x).sum::<f64>().sqrt();
332 if new_vnorm < 1e-300 {
333 break;
334 }
335 let diff: f64 = new_vk
337 .iter()
338 .zip(vk.iter())
339 .map(|(a, b)| (a / new_vnorm - b).powi(2))
340 .sum::<f64>()
341 .sqrt();
342 for j in 0..n {
343 vk[j] = new_vk[j] / new_vnorm;
344 }
345 if diff < 1e-12 {
346 break;
347 }
348 }
349
350 let mut uk_final = vec![0.0f64; m];
352 for i in 0..m {
353 let mut s = 0.0f64;
354 for j in 0..n {
355 s += residual[i * n + j] * vk[j];
356 }
357 uk_final[i] = s;
358 }
359 let sigma: f64 = uk_final.iter().map(|x| x * x).sum::<f64>().sqrt();
360 if sigma < threshold {
361 break;
362 }
363 for x in uk_final.iter_mut() {
364 *x /= sigma;
365 }
366
367 for i in 0..m {
369 for j in 0..n {
370 residual[i * n + j] -= sigma * uk_final[i] * vk[j];
371 }
372 }
373
374 u_cols.push(uk_final);
375 s_vals.push(sigma);
376 v_rows.push(vk);
377 }
378
379 if s_vals.is_empty() {
380 let mut u0 = vec![0.0f64; m];
381 let mut v0 = vec![0.0f64; n];
382 if m > 0 {
383 u0[0] = 1.0;
384 }
385 if n > 0 {
386 v0[0] = 1.0;
387 }
388 return (u0, vec![threshold.max(1e-300)], v0);
389 }
390
391 let r = s_vals.len();
392 let mut u_data = vec![0.0f64; m * r];
394 for i in 0..m {
395 for k in 0..r {
396 u_data[i * r + k] = u_cols[k][i];
397 }
398 }
399 let vt_data: Vec<f64> = v_rows.into_iter().flatten().collect();
400 (u_data, s_vals, vt_data)
401}
402
403pub fn tt_svd(
418 tensor: &ArrayD<f64>,
419 max_rank: usize,
420 tol: f64,
421) -> Result<TensorTrain, InterpolateError> {
422 let shape: Vec<usize> = tensor.shape().to_vec();
423 let d = shape.len();
424
425 if d == 0 {
426 return Err(InterpolateError::InvalidInput {
427 message: "tt_svd: tensor must have at least one dimension".into(),
428 });
429 }
430 if max_rank == 0 {
431 return Err(InterpolateError::InvalidInput {
432 message: "tt_svd: max_rank must be >= 1".into(),
433 });
434 }
435
436 let mut cores = Vec::with_capacity(d);
437 let mut r_left = 1usize;
438
439 let mut remainder: Vec<f64> = tensor.iter().copied().collect();
441
442 for k in 0..d {
443 let n_k = shape[k];
444 let n_right: usize = shape[k + 1..].iter().product::<usize>().max(1);
446
447 let rows = r_left * n_k;
448 let cols = n_right;
449
450 let mat = Array2::from_shape_vec((rows, cols), remainder.clone()).map_err(|e| {
452 InterpolateError::ComputationError(format!("tt_svd reshape error at k={k}: {e}"))
453 })?;
454
455 if k < d - 1 {
456 let (u, s, vt) = truncated_svd(&mat, max_rank, tol)?;
457 let r_right = s.len();
458
459 let u_flat: Vec<f64> = u.iter().copied().collect();
462 let core = Array3::from_shape_vec((r_left, n_k, r_right), u_flat).map_err(|e| {
463 InterpolateError::ComputationError(format!("tt_svd core shape error k={k}: {e}"))
464 })?;
465 cores.push(core);
466
467 let mut new_rem = vec![0.0f64; r_right * cols];
469 for i in 0..r_right {
470 let si = s[i];
471 for j in 0..cols {
472 new_rem[i * cols + j] = si * vt[[i, j]];
473 }
474 }
475 remainder = new_rem;
476 r_left = r_right;
477 } else {
478 let mat_flat: Vec<f64> = mat.iter().copied().collect();
481 let core = Array3::from_shape_vec((r_left, n_k, 1), mat_flat).map_err(|e| {
482 InterpolateError::ComputationError(format!(
483 "tt_svd last core shape error k={k}: {e}"
484 ))
485 })?;
486 cores.push(core);
487 }
488 }
489
490 TensorTrain::new(cores)
491}
492
493#[cfg(test)]
496mod tests {
497 use super::*;
498 use scirs2_core::ndarray::IxDyn;
499
500 fn make_rank1_tt(shape: &[usize]) -> TensorTrain {
501 let cores: Vec<Array3<f64>> = shape.iter().map(|&n| Array3::ones((1, n, 1))).collect();
503 TensorTrain::new(cores).expect("rank-1 TT valid")
504 }
505
506 #[test]
507 fn test_tt_eval_correct() {
508 let core0 = Array3::from_shape_fn((1, 3, 1), |(_, i, _)| (i + 1) as f64);
511 let core1 = Array3::from_shape_fn((1, 4, 1), |(_, j, _)| (j + 1) as f64);
512 let tt = TensorTrain::new(vec![core0, core1]).expect("valid TT");
513
514 for i in 0..3 {
515 for j in 0..4 {
516 let val = tt.eval(&[i, j]).expect("eval ok");
517 let expected = ((i + 1) * (j + 1)) as f64;
518 assert!(
519 (val - expected).abs() < 1e-12,
520 "TT[{i},{j}] expected {expected} got {val}"
521 );
522 }
523 }
524 }
525
526 #[test]
527 fn test_tt_norm() {
528 let tt = make_rank1_tt(&[2, 2]);
530 let norm = tt.norm();
531 assert!((norm - 2.0).abs() < 1e-10, "norm={norm}");
532 }
533
534 #[test]
535 fn test_tt_n_params() {
536 let tt = make_rank1_tt(&[3, 4]);
538 assert_eq!(tt.n_params(), 7);
539 }
540
541 #[test]
542 fn test_tt_svd_2d() {
543 let a = [1.0, 2.0, 3.0f64];
545 let b = [1.0, -1.0, 2.0, -2.0f64];
546 let data: Vec<f64> = a
547 .iter()
548 .flat_map(|&ai| b.iter().map(move |&bj| ai * bj))
549 .collect();
550 let tensor = ArrayD::from_shape_vec(IxDyn(&[3, 4]), data).expect("valid");
551
552 let tt = tt_svd(&tensor, 4, 1e-10).expect("TT-SVD ok");
553 assert_eq!(tt.shape, vec![3, 4]);
554 for i in 0..3 {
556 for j in 0..4 {
557 let val = tt.eval(&[i, j]).expect("eval ok");
558 let expected = a[i] * b[j];
559 assert!(
560 (val - expected).abs() < 1e-7,
561 "TT-SVD[{i},{j}] expected {expected:.6} got {val:.6}"
562 );
563 }
564 }
565 }
566
567 #[test]
568 fn test_tt_from_dense_rank_compression() {
569 let tensor = ArrayD::ones(IxDyn(&[2, 2, 2]));
571 let tt = TensorTrain::from_dense(&tensor, 4, 1e-8).expect("from_dense ok");
572 assert!(tt.n_params() <= 8, "n_params={}", tt.n_params());
574 for i in 0..2 {
576 for j in 0..2 {
577 for k in 0..2 {
578 let val = tt.eval(&[i, j, k]).expect("eval ok");
579 assert!((val - 1.0).abs() < 1e-6, "val={val}");
580 }
581 }
582 }
583 }
584
585 #[test]
586 fn test_tt_to_dense() {
587 let core0 = Array3::from_shape_fn((1, 2, 1), |(_, i, _)| (i + 1) as f64);
588 let core1 = Array3::from_shape_fn((1, 2, 1), |(_, j, _)| (j + 1) as f64);
589 let tt = TensorTrain::new(vec![core0, core1]).expect("valid");
590 let dense = tt.to_dense().expect("to_dense ok");
591 assert_eq!(dense.shape(), &[2, 2]);
592 assert!((dense[[0, 0]] - 1.0).abs() < 1e-12);
594 assert!((dense[[0, 1]] - 2.0).abs() < 1e-12);
595 assert!((dense[[1, 0]] - 2.0).abs() < 1e-12);
596 assert!((dense[[1, 1]] - 4.0).abs() < 1e-12);
597 }
598}