Skip to main content

scirs2_interpolate/tensor_train/
tt_decomp.rs

1//! TT-SVD decomposition for tensor-train format.
2//!
3//! Implements the TT-SVD algorithm from Oseledets 2011:
4//! "Tensor-Train Decomposition", SIAM J. Sci. Comput., 33(5), 2295-2317.
5
6use crate::error::InterpolateError;
7use scirs2_core::ndarray::{Array2, Array3, ArrayD, IxDyn};
8
9/// Tensor-train representation of a d-dimensional array.
10///
11/// `A[i1,...,id] ≈ G1[i1] * G2[i2] * ... * Gd[id]`
12///
13/// where each core `Gk` has shape `[r_{k-1}, n_k, r_k]` and the boundary ranks
14/// satisfy `r_0 = r_d = 1`.
15#[derive(Debug, Clone)]
16pub struct TensorTrain {
17    /// Cores: `cores[k]` has shape `[r_{k-1}, n_k, r_k]`.
18    pub cores: Vec<Array3<f64>>,
19    /// Mode sizes `n_1, ..., n_d`.
20    pub shape: Vec<usize>,
21    /// Ranks `r_0=1, r_1, ..., r_{d-1}, r_d=1`.
22    pub ranks: Vec<usize>,
23}
24
25impl TensorTrain {
26    /// Create a TT from pre-built cores.
27    ///
28    /// Validates that shapes are consistent.
29    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    /// Evaluate the TT at a multi-index `idx = [i1, ..., id]`.
82    ///
83    /// Returns the scalar `G1[:,i1,:] * G2[:,i2,:] * ... * Gd[:,id,:]` contracted
84    /// to a single `f64`.
85    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        // Carry a row-vector of shape [1, r_k]; start with [1.0] (r_0=1).
102        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            // Slice: G_k[:, ik, :] => shape [r_left, r_right]
107            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        // v should have length 1 (r_d = 1)
118        Ok(v[0])
119    }
120
121    /// Reconstruct the full dense tensor.
122    ///
123    /// **Warning**: exponential cost in `d` — only practical for small tensors.
124    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        // Iterate over all multi-indices
130        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            // Increment multi-index (last dimension varies fastest)
136            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; // all combinations exhausted
149            }
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    /// Compute the Frobenius norm of the TT tensor via a transfer-matrix approach.
157    pub fn norm(&self) -> f64 {
158        let d = self.cores.len();
159        if d == 0 {
160            return 0.0;
161        }
162        // Build Gram matrix G of shape [1,1] = [[1]], then transfer left to right.
163        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 is now 1x1
186        gram[[0, 0]].max(0.0).sqrt()
187    }
188
189    /// Number of parameters (stored f64 elements) in the TT format.
190    pub fn n_params(&self) -> usize {
191        self.cores.iter().map(|c| c.len()).sum()
192    }
193
194    /// Convenience: build from a dense tensor with TT-SVD.
195    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
204// ─── Row-major flat index ─────────────────────────────────────────────────────
205
206fn 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
216// ─── Truncated SVD via deflation + power iteration ────────────────────────────
217
218/// Compute a rank-`r` truncated SVD of an `m x n` matrix.
219///
220/// Uses sequential deflation with power-iteration to extract singular triplets.
221/// Returns `(U, s, Vt)` where `U` is `m x r`, `s` is `r`, `Vt` is `r x n`.
222/// All singular values are non-negative; only those >= `tol * sigma_max` are kept
223/// (and at most `max_rank`).
224pub 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
248/// Extract singular triplets one by one via power-iteration + deflation.
249///
250/// Guarantees non-negative singular values and consistent U/Vt.
251fn 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    // Compute Frobenius norm to determine threshold
259    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        // Initialise v as the row of residual with largest norm
281        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        // Power iteration: alternate between A*v and A^T*u
304        let mut uk = vec![0.0f64; m];
305        for _iter in 0..20 {
306            // uk = A * vk
307            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            // vk = A^T * uk
323            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            // Check convergence
336            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        // Compute sigma = ||A * vk||
351        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        // Deflate: residual -= sigma * uk * vk^T
368        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    // Pack U in row-major: u_data[i*r + k] = u_cols[k][i]
393    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
403// ─── TT-SVD ──────────────────────────────────────────────────────────────────
404
405/// TT-SVD algorithm (Oseledets 2011).
406///
407/// Sequentially unfolds the tensor into matrices and performs truncated SVD
408/// to obtain the TT cores.
409///
410/// # Parameters
411/// - `tensor`:   dense input tensor of any shape
412/// - `max_rank`: maximum TT rank per bond
413/// - `tol`:      relative truncation tolerance (`tol * sigma_max` threshold)
414///
415/// # Returns
416/// A [`TensorTrain`] approximating `tensor`.
417pub 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    // Start with the flat data copy
440    let mut remainder: Vec<f64> = tensor.iter().copied().collect();
441
442    for k in 0..d {
443        let n_k = shape[k];
444        // Remaining dimensions product
445        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        // Build matrix of shape [r_left * n_k, n_right]
451        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            // Core k has shape [r_left, n_k, r_right]
460            // U has shape [r_left * n_k, r_right]; reshape to [r_left, n_k, r_right]
461            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            // New remainder = diag(s) * Vt  (shape [r_right, n_right])
468            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            // Last core: the remainder is already [r_left * n_d, 1]
479            // => shape [r_left, n_d, 1]
480            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// ─── Tests ────────────────────────────────────────────────────────────────────
494
495#[cfg(test)]
496mod tests {
497    use super::*;
498    use scirs2_core::ndarray::IxDyn;
499
500    fn make_rank1_tt(shape: &[usize]) -> TensorTrain {
501        // Build a TT of all ones: each core is all-1s of shape [1, n_k, 1]
502        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        // 2D: TT where core0[0, i, 0] = i+1  and core1[0, j, 0] = j+1
509        // => TT[i,j] = (i+1)*(j+1)
510        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        // Rank-1 TT of shape [2,2]: each entry = 1, Frobenius norm = 2
529        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        // shape [3, 4], ranks [1,1,1]: params = 1*3*1 + 1*4*1 = 7
537        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        // TT-SVD of a rank-1 2D tensor: outer product a x b
544        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        // Recover values — should match original tensor
555        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        // A 2x2x2 tensor of all 1s has TT rank 1
570        let tensor = ArrayD::ones(IxDyn(&[2, 2, 2]));
571        let tt = TensorTrain::from_dense(&tensor, 4, 1e-8).expect("from_dense ok");
572        // Rank-1 TT of shape [2,2,2] has 2+2+2=6 params; max allowed 8
573        assert!(tt.n_params() <= 8, "n_params={}", tt.n_params());
574        // Values should be recoverable
575        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        // [1*1, 1*2; 2*1, 2*2] = [[1,2],[2,4]]
593        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}